{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing\nimport tensorflow as tf\nfrom glob import glob\nimport re\nimport ast\nimport cv2\nimport csv\nimport ast\nimport os\nimport urllib\nfrom PIL import Image, ImageDraw\nimport matplotlib.pyplot as plt\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-14T09:24:00.403477Z","iopub.execute_input":"2022-12-14T09:24:00.403955Z","iopub.status.idle":"2022-12-14T09:24:00.412239Z","shell.execute_reply.started":"2022-12-14T09:24:00.403918Z","shell.execute_reply":"2022-12-14T09:24:00.410585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Chargement data quickdraw-doodle-recognition\nINPUT_DIR = '/kaggle/input/quickdraw-doodle-recognition/'\n# Clean names \nclasses_path = os.listdir(INPUT_DIR + 'train_simplified/')\nclasses_path = sorted(classes_path, key=lambda s: s.lower())\nclass_dict = {x[:-4].replace(\" \", \"_\"):i for i, x in enumerate(classes_path)}\nlabels = {x[:-4].replace(\" \", \"_\") for i, x in enumerate(classes_path)}","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:00.639195Z","iopub.execute_input":"2022-12-14T09:24:00.639706Z","iopub.status.idle":"2022-12-14T09:24:00.649092Z","shell.execute_reply.started":"2022-12-14T09:24:00.63966Z","shell.execute_reply":"2022-12-14T09:24:00.647685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_labels = len(labels)\nprint(\"Number of labels: {}\".format(n_labels))\n\nfileList = glob(INPUT_DIR + \"train_simplified/*.csv\")     \n\nn_files = n_labels #number of csv files same as labels.\n\n#time is sacred HARDCODED FOR THE COMP\nn_records = 49707919\nsize = 128","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:00.759167Z","iopub.execute_input":"2022-12-14T09:24:00.760461Z","iopub.status.idle":"2022-12-14T09:24:00.771085Z","shell.execute_reply.started":"2022-12-14T09:24:00.760412Z","shell.execute_reply":"2022-12-14T09:24:00.769401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# to image from stroke\ndef drawing_to_np(drawing, shape=(size, size)):\n    drawing = eval(drawing)\n    fig, ax = plt.subplots()\n    for x,y in drawing:\n        ax.plot(x, y, marker='.')\n        ax.axis('off')\n    fig.canvas.draw()\n    # Convert images to numpy arrat\n    np_drawing = np.array(fig.canvas.renderer._renderer)\n    plt.close(fig)\n    img = cv2.resize(np_drawing, shape)\n    img_gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    img_expanded = img_gray[:, :, np.newaxis]\n    return img_expanded\n\n## Return img array from cv2 and normalize value\ndef draw_cv2_reshape_normalized(raw_strokes, size=size, lw=6):\n    img = np.zeros((BASE_SIZE, BASE_SIZE), np.uint8)\n    for stroke in raw_strokes:\n        for i in range(len(stroke[0]) - 1):\n            _ = cv2.line(img, (stroke[0][i], stroke[1][i]), (stroke[0][i + 1], stroke[1][i + 1]), 255, lw)\n\n    img = cv2.resize(img, (size, size))\n    img = img / 255.\n    img = img[:, :, np.newaxis]\n    return img\n\n## Return img array from cv2\ndef draw_cv2(raw_strokes, size=256, lw=6):\n    img = np.zeros((BASE_SIZE, BASE_SIZE), np.uint8)\n    for stroke in raw_strokes:\n        for i in range(len(stroke[0]) - 1):\n            _ = cv2.line(img, (stroke[0][i], stroke[1][i]), (stroke[0][i + 1], stroke[1][i + 1]), 255, lw)\n    if size != BASE_SIZE:\n        return cv2.resize(img, (size, size))\n    else:\n        return img\n\n## Take df (pandas Dataframe) and convert to image with pixel value between 0 and 1   \ndef df_to_image_array(df, size=size, lw=6):\n    df['drawing'] = df['drawing'].apply(ast.literal_eval)\n    x = np.zeros((len(df), size, size))\n    for i, raw_strokes in enumerate(df.drawing.values):\n        x[i] = draw_cv2(raw_strokes, size=size, lw=lw)\n    x = x / 255.\n    x = x.reshape((len(df), 1, size, size)).astype(np.float32)\n    return x\n\n## Pour avoir un tensor avec des valeurs comprises entre -1 et 1\ndef tensor_1_1(x):\n    return (x*2)-1 ","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:00.958045Z","iopub.execute_input":"2022-12-14T09:24:00.958881Z","iopub.status.idle":"2022-12-14T09:24:00.977388Z","shell.execute_reply.started":"2022-12-14T09:24:00.958832Z","shell.execute_reply":"2022-12-14T09:24:00.976079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_SIZE = 256\n# Load Airplane and convert to array\nvalid_car_df = pd.read_csv('/kaggle/input/quickdraw-doodle-recognition/train_simplified/car.csv', nrows=10000)\nx_valid_df = df_to_image_array(valid_car_df)","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:01.247266Z","iopub.execute_input":"2022-12-14T09:24:01.247778Z","iopub.status.idle":"2022-12-14T09:24:07.181883Z","shell.execute_reply.started":"2022-12-14T09:24:01.247739Z","shell.execute_reply":"2022-12-14T09:24:07.180505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#np.shape(x_valid_df[0].T)\nplt.imshow(x_valid_df[0, 0, :, :]*255)\n#plt.imshow(tensor_1_1(x_valid_df)[0, :, :, 0])","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.184858Z","iopub.execute_input":"2022-12-14T09:24:07.185406Z","iopub.status.idle":"2022-12-14T09:24:07.402057Z","shell.execute_reply.started":"2022-12-14T09:24:07.185353Z","shell.execute_reply":"2022-12-14T09:24:07.400607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Construction du Modèle de Diffusion**","metadata":{}},{"cell_type":"markdown","source":"# 1. Forward Process : Noise Scheduler\n\nCe processus, qui constitue la première étape de l'entraînement du modèle de diffusion, consiste à transformer une image de dimension **n*m** en une image entièrement bruitée de dimension **n*m** en plusieurs étapes. L'image obtenue à la fin est la représentation uniforme d'un bruit, générée par une gaussienne isotropique (la variance est la même dans toutes les direcions) dont la moyenne vaut 0.\n> On définit les fonctions permettant d'ajouter du bruit à une image de façon non séquentielle pour alléger le temps de calucl. Pour cela on calcule l'état de l'image à un instant t indiduellement grâce aux formules ci-dessous.","metadata":{}},{"cell_type":"markdown","source":"$\\alpha_t = 1-\\beta_t$\n\n$q(x_t|x_0) = N(x_t;\\sqrt[2](\\bar\\alpha_t)x_0;(1-\\bar\\alpha_t)I)$\n\n$x_t = \\sqrt[2](\\bar\\alpha_t)x_0+\\sqrt[2](1-\\bar\\alpha_t)\\epsilon$","metadata":{}},{"cell_type":"markdown","source":"En outre, la variance est augmentée linéairement d'un instant t à un autre, afin que l'image générée converge plus rapidement vers une image entièrement composée de bruit. On fixe par ailleurs une valeur limite de $beta_t$.","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\nimport torch","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.403303Z","iopub.execute_input":"2022-12-14T09:24:07.404326Z","iopub.status.idle":"2022-12-14T09:24:07.410932Z","shell.execute_reply.started":"2022-12-14T09:24:07.404281Z","shell.execute_reply":"2022-12-14T09:24:07.409457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def linear_beta_schedule(timesteps, start=0.0001, end=0.02):\n    \"\"\"\"\n    Fonction qui à timesteps instants, et pour une valeur initiale de beta_t\n    et une valeur seuil de beta_t, retourne une interpolation entre ces deux valeurs\n    avec timesteps éléments. \"\"\" \n    return torch.linspace(start, end, timesteps)\n\ndef get_index_from_list(vals, t, x_shape):\n    \"\"\"\n    Retourne un index t de la liste vals en prenant en compte la dimension du batch\n    \"\"\"\n    batch_size = t.shape[0]\n    out = vals.gather(-1, t.cpu())\n    return out.reshape(batch_size, *((1,) * (len(x_shape) - 1))).to(t.device)\n\ndef forward_diffusion_sample(x_0, t, device=\"cpu\"):\n    \"\"\"\n    Prend l'image initiale x_0 et un instant t et retourne l'image x_t avec du bruit obtenue à\n    l'instant t\n    \"\"\"\n    noise = torch.randn_like(x_0)\n    sqrt_alphas_cumprod_t = get_index_from_list(sqrt_alphas_cumprod, t, x_0.shape) # On récupére les index de la list t\n    sqrt_one_minus_alphas_cumprod_t = get_index_from_list(sqrt_one_minus_alphas_cumprod, t, x_0.shape)\n    # moyenne + variance\n    return sqrt_alphas_cumprod_t.to(device) * x_0.to(device) + sqrt_one_minus_alphas_cumprod_t.to(device) * noise.to(device), noise.to(device)\n\n## Constantes ##\n# Definition beta schedule\nBATCH_SIZE = 128\nT = 300 # Nombre d'itération totale\nbetas = linear_beta_schedule(timesteps=T, start=0.00001, end=0.002)\n\n# Calcule des différents terme pour l'application la formule ensuite\nalphas = 1. - betas\nalphas_cumprod = torch.cumprod(alphas, axis=0) # Produit cummulé des différents alphas\nalphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value=1.0)\nsqrt_recip_alphas = torch.sqrt(1.0 / alphas)\nsqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)\nsqrt_one_minus_alphas_cumprod = torch.sqrt(1. - alphas_cumprod)\nposterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod)","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.414271Z","iopub.execute_input":"2022-12-14T09:24:07.415124Z","iopub.status.idle":"2022-12-14T09:24:07.43092Z","shell.execute_reply.started":"2022-12-14T09:24:07.415061Z","shell.execute_reply":"2022-12-14T09:24:07.429236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Convert np.ndarray to tensor\ntensor_car = torch.from_numpy(tensor_1_1(x_valid_df))","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.433152Z","iopub.execute_input":"2022-12-14T09:24:07.433777Z","iopub.status.idle":"2022-12-14T09:24:07.708826Z","shell.execute_reply.started":"2022-12-14T09:24:07.433685Z","shell.execute_reply":"2022-12-14T09:24:07.707302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensor_car[0].shape","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.710501Z","iopub.execute_input":"2022-12-14T09:24:07.711143Z","iopub.status.idle":"2022-12-14T09:24:07.720643Z","shell.execute_reply.started":"2022-12-14T09:24:07.711093Z","shell.execute_reply":"2022-12-14T09:24:07.719117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms\nfrom torch.utils.data import DataLoader\n\ndef show_tensor_image(image):\n    reverse_transforms = transforms.Compose([\n        transforms.Lambda(lambda t: (t+1)/2),\n        transforms.Lambda(lambda t: t.permute(1, 2, 0)),\n        transforms.Lambda(lambda t: t*255),\n        transforms.Lambda(lambda t: t.numpy().astype(np.uint8)),\n        transforms.ToPILImage(),\n    ])\n    #Take first image of batch\n    if (len(image.shape)==4):\n        image = image[0, :, :, :]\n    plt.imshow(reverse_transforms(image))\n\ndataloader = DataLoader(tensor_car[0:10000, :, :, :], batch_size=BATCH_SIZE, shuffle=True, drop_last=True)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.722296Z","iopub.execute_input":"2022-12-14T09:24:07.722953Z","iopub.status.idle":"2022-12-14T09:24:07.736172Z","shell.execute_reply.started":"2022-12-14T09:24:07.722901Z","shell.execute_reply":"2022-12-14T09:24:07.734715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for step, batch in enumerate(dataloader):\n    if (step>5):\n        break\n    print(batch.shape)","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.738217Z","iopub.execute_input":"2022-12-14T09:24:07.738647Z","iopub.status.idle":"2022-12-14T09:24:07.771505Z","shell.execute_reply.started":"2022-12-14T09:24:07.738613Z","shell.execute_reply":"2022-12-14T09:24:07.769896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image = next(iter(dataloader))\nt = torch.Tensor([3]).type(torch.int64) # On s'arrête à 200 ?\nimage, noise = forward_diffusion_sample(image, t)\nshow_tensor_image(image)","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:07.775876Z","iopub.execute_input":"2022-12-14T09:24:07.776354Z","iopub.status.idle":"2022-12-14T09:24:08.036962Z","shell.execute_reply.started":"2022-12-14T09:24:07.776308Z","shell.execute_reply":"2022-12-14T09:24:08.035722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"plt.figure(figsize=(15,15))\nplt.axis('off')\nnum_images = 10\nstepsize = int(T/num_images)\n\nfor idx in range(0, 10):\n    t = torch.Tensor([idx]).type(torch.int64)\n    plt.subplot(1, num_images+1, idx + 1)\n    image, noise = forward_diffusion_sample(tensor_car[0, :, :, 0], t)\n    show_tensor_image(image)\"\"\"\n\nimage = next(iter(dataloader))\n\nplt.figure(figsize=(15,15))\nplt.axis('off')\nnum_images = 10\nstepsize = int(T/num_images)\n\nfor idx in range(0, 30, 3):\n    t = torch.Tensor([idx]).type(torch.int64)\n    plt.subplot(1, num_images+1, int(idx/3) + 1)\n    image, noise = forward_diffusion_sample(image, t)\n    show_tensor_image(image)","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:08.040122Z","iopub.execute_input":"2022-12-14T09:24:08.040644Z","iopub.status.idle":"2022-12-14T09:24:09.232032Z","shell.execute_reply.started":"2022-12-14T09:24:08.040604Z","shell.execute_reply":"2022-12-14T09:24:09.230697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"On remarque que l'image est détruite après seulement 10 étapes... Trop rapide !","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nimport math\n\n#Chaque block applique 2 layers de convolution (conv1, conv2)\nclass Block(nn.Module):\n    def __init__(self, in_ch, out_ch, time_emb_dim, up=False):\n        super().__init__()\n        self.time_mlp =  nn.Linear(time_emb_dim, out_ch)\n        if up:\n            #On augmente la dimension du tenseur\n            self.conv1 = nn.Conv2d(2*in_ch, out_ch, 3, padding=1)\n            self.transform = nn.ConvTranspose2d(out_ch, out_ch, 4, 2, 1)\n        else:\n            #On la réduit\n            self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)\n            self.transform = nn.Conv2d(out_ch, out_ch, 4, 2, 1)\n        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)\n        self.bnorm1 = nn.BatchNorm2d(out_ch)\n        self.bnorm2 = nn.BatchNorm2d(out_ch)\n        self.relu  = nn.ReLU()\n        \n    def forward(self, x, t, ):\n        # First Conv\n        h = self.bnorm1(self.relu(self.conv1(x)))\n        # Time embedding\n        time_emb = self.relu(self.time_mlp(t))\n        # Extend last 2 dimensions\n        time_emb = time_emb[(..., ) + (None, ) * 2]\n        # Add time channel\n        h = h + time_emb\n        # Second Conv\n        h = self.bnorm2(self.relu(self.conv2(h)))\n        # Down or Upsample\n        return self.transform(h)\n\n# Positional Embedding formule : retourne un vecteur décrivant la position d'un index dans une liste\nclass SinusoidalPositionEmbeddings(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n    \n    def forward(self, time):\n        device = time.device\n        half_dim = self.dim // 2\n        embeddings = math.log(10000) / (half_dim - 1)\n        embeddings = torch.exp(torch.arange(half_dim, device=device) * -embeddings)\n        embeddings = time[:, None] * embeddings[None, :]\n        embeddings = torch.cat((embeddings.sin(), embeddings.cos()), dim=-1)\n        # TODO: Double check the ordering here\n        return embeddings\n\nclass SimpleUnet(nn.Module):\n    \"\"\"\n    A simplified variant of the Unet architecture.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        image_channels = 1 # Nombre de channel de l'image en entrée, 3 = (R, G, B)\n        down_channels = (64, 128, 256, 512, 1024) # Nombre de filtres dans les layers de convolution,\n        # plus il y en a, plus la profondeur du tensor est grande. Jusqu'à 1024 channels.\n        up_channels = (1024, 512, 256, 128, 64) # Puis on réduit la taille\n        out_dim = 1 \n        time_emb_dim = 32\n\n        # Time embedding / Le résultat sera un vecteur\n        self.time_mlp = nn.Sequential(\n                SinusoidalPositionEmbeddings(time_emb_dim),\n                nn.Linear(time_emb_dim, time_emb_dim),\n                nn.ReLU()\n            )\n        \n        # Initial projection\n        self.conv0 = nn.Conv2d(1, down_channels[0], 3, padding=1)\n\n        # Downsample\n        self.downs = nn.ModuleList([Block(down_channels[i], down_channels[i+1], \\\n                                    time_emb_dim) \\\n                    for i in range(len(down_channels)-1)])\n        # Upsample\n        self.ups = nn.ModuleList([Block(up_channels[i], up_channels[i+1], \\\n                                        time_emb_dim, up=True) \\\n                    for i in range(len(up_channels)-1)])\n\n        self.output = nn.Conv2d(up_channels[-1], 1, out_dim)\n    \n    # On itère sur les objets Block, il prend x : image en entrée et retourne une image\n    # à jour plus ou moins grande avec plus de channels \n    def forward(self, x, timestep):\n        # Embedd time\n        t = self.time_mlp(timestep)\n        # Initial conv\n        x = self.conv0(x)\n        # Unet\n        residual_inputs = []\n        for down in self.downs:\n            x = down(x, t)\n            residual_inputs.append(x)\n        for up in self.ups:\n            residual_x = residual_inputs.pop()\n            # Add residual x as additional channels\n            x = torch.cat((x, residual_x), dim=1)           \n            x = up(x, t)\n        return self.output(x)\n\nmodel = SimpleUnet()\nprint(\"Num params: \", sum(p.numel() for p in model.parameters()))\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:09.23386Z","iopub.execute_input":"2022-12-14T09:24:09.234264Z","iopub.status.idle":"2022-12-14T09:24:09.819874Z","shell.execute_reply.started":"2022-12-14T09:24:09.234226Z","shell.execute_reply":"2022-12-14T09:24:09.818821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Loss Function\n> L1 Loss entre le bruit, et le bruit prédit","metadata":{}},{"cell_type":"code","source":"#Prend une image, un instant donné t, un model et retourne la perte entre le bruit préduit et celui calculé\ndef get_loss(model, x_0, t):\n    x_noisy, noise = forward_diffusion_sample(x_0, t, device)\n    noise_pred = model(x_noisy, t)\n    return F.l1_loss(noise, noise_pred)","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:09.821512Z","iopub.execute_input":"2022-12-14T09:24:09.822253Z","iopub.status.idle":"2022-12-14T09:24:09.828065Z","shell.execute_reply.started":"2022-12-14T09:24:09.822213Z","shell.execute_reply":"2022-12-14T09:24:09.826688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Fonctions pour générer les images et les afficher","metadata":{}},{"cell_type":"markdown","source":"> Le décorateur @torch.no_grad() permet de ne pas prendre en compte les précédentes images pour calculer le gradient (pytorch le fait par défaut) : ce qui nous évite de saturer la mémoire vive.","metadata":{}},{"cell_type":"code","source":"@torch.no_grad() \ndef sample_timestep(x, t):\n    \"\"\"\n    Calls the model to predict the noise in the image and returns \n    the denoised image. \n    Applies noise to this image, if we are not in the last step yet.\n    \"\"\"\n    betas_t = get_index_from_list(betas, t, x.shape)\n    sqrt_one_minus_alphas_cumprod_t = get_index_from_list(\n        sqrt_one_minus_alphas_cumprod, t, x.shape\n    )\n    sqrt_recip_alphas_t = get_index_from_list(sqrt_recip_alphas, t, x.shape)\n    \n    # Call model (current image - noise prediction)\n    model_mean = sqrt_recip_alphas_t * (\n        x - betas_t * model(x, t) / sqrt_one_minus_alphas_cumprod_t\n    )\n    posterior_variance_t = get_index_from_list(posterior_variance, t, x.shape)\n    \n    if t == 0:\n        return model_mean\n    else:\n        noise = torch.randn_like(x)\n        return model_mean + torch.sqrt(posterior_variance_t) * noise \n\n@torch.no_grad()\n#On ne prend en compte lors de l'affichage que les 30 dernières images (ou 30 premières)\ndef sample_plot_image():\n    # Sample noise\n    img_size = size\n    img = torch.randn((1, 1, img_size, img_size), device=device)\n    plt.figure(figsize=(15,15))\n    plt.axis('off')\n    num_images = 10\n    #stepsize = int(T/num_images)\n    stepsize = 30/num_images #T = 30\n    \n    for i in range(0,30)[::-1]: # On ne prend que 30 images\n        t = torch.full((1,), i, device=device, dtype=torch.long)\n        img = sample_timestep(img, t)\n        if i % stepsize == 0:\n            plt.subplot(1, num_images, int(i/stepsize)+1)\n            show_tensor_image(img.detach().cpu())\n    plt.show()            ","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:24:11.389738Z","iopub.execute_input":"2022-12-14T09:24:11.390901Z","iopub.status.idle":"2022-12-14T09:24:11.402565Z","shell.execute_reply.started":"2022-12-14T09:24:11.39085Z","shell.execute_reply":"2022-12-14T09:24:11.401441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.optim import Adam\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel.to(device)\noptimizer = Adam(model.parameters(), lr=0.001)\nepochs = 100 # Try more!\n            \nfor epoch in range(epochs):\n    for step, batch in enumerate(dataloader):\n        optimizer.zero_grad()\n        t = torch.randint(0, 30, (BATCH_SIZE,), device=device).long() #T = 30\n        loss = get_loss(model, batch, t)\n        loss.backward()\n        optimizer.step()\n        if step%5 == 0:\n            print(f\"Epoch {epoch} | step {step:03d} Loss: {loss.item()} \")\n            sample_plot_image()","metadata":{"execution":{"iopub.status.busy":"2022-12-14T09:27:57.031084Z","iopub.execute_input":"2022-12-14T09:27:57.031679Z","iopub.status.idle":"2022-12-14T17:30:54.023318Z","shell.execute_reply.started":"2022-12-14T09:27:57.031629Z","shell.execute_reply":"2022-12-14T17:30:54.02107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Le temps de calcul est beaucoup trop long sur les serveurs de Kaggle (cpu)... Il y a 80 steps pour chaque epoch, avec 100 epochs au total, sachant qu'il faut 4 minutes pour faire 1 step environ. Je ne peux donc pas vérifier si mes choix de paramètre sont bons ou non. Toutefois on remarque que le coût semble diminuer ?","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}