{"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":"from PIL import Image\nimport numpy as np \nimport pandas as pd \nimport torch\nimport torchvision\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom zipfile import ZipFile \nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensor\nimport tensorflow as tf\nfrom tensorflow import keras\nimport matplotlib.pyplot as plt\nfrom albumentations.pytorch import ToTensor\nimport matplotlib.pyplot as plt\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-10-04T15:14:48.848533Z","iopub.execute_input":"2021-10-04T15:14:48.848847Z","iopub.status.idle":"2021-10-04T15:14:48.877084Z","shell.execute_reply.started":"2021-10-04T15:14:48.848818Z","shell.execute_reply":"2021-10-04T15:14:48.875748Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Giải nén dataset","metadata":{}},{"cell_type":"code","source":"train_zip = \"/kaggle/input/carvana-image-masking-challenge/train.zip\"\nwith ZipFile(train_zip, 'r') as zip_: \n    zip_.extractall('/kaggle/working')","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:44:52.257069Z","iopub.execute_input":"2021-10-04T08:44:52.257381Z","iopub.status.idle":"2021-10-04T08:45:02.090785Z","shell.execute_reply.started":"2021-10-04T08:44:52.257351Z","shell.execute_reply":"2021-10-04T08:45:02.089736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_mask_zip = \"/kaggle/input/carvana-image-masking-challenge/train_masks.zip\"\nwith ZipFile(train_mask_zip, 'r') as zip_: \n    zip_.extractall('/kaggle/working')","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:45:42.485401Z","iopub.execute_input":"2021-10-04T08:45:42.485749Z","iopub.status.idle":"2021-10-04T08:45:44.050285Z","shell.execute_reply.started":"2021-10-04T08:45:42.485719Z","shell.execute_reply":"2021-10-04T08:45:44.049419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Train set:  \", len(os.listdir(\"/kaggle/working/train\")))\nprint(\"Train masks:\", len(os.listdir(\"/kaggle/working/train_masks\")))","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:45:45.940266Z","iopub.execute_input":"2021-10-04T08:45:45.94072Z","iopub.status.idle":"2021-10-04T08:45:45.956989Z","shell.execute_reply.started":"2021-10-04T08:45:45.940678Z","shell.execute_reply":"2021-10-04T08:45:45.956209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Show đường dẫn của ảnh input và output","metadata":{}},{"cell_type":"code","source":"car_ids = []\npaths = []\nfor dirname, _, filenames in os.walk('/kaggle/working/train'):\n    for filename in filenames:\n        path = os.path.join(dirname, filename)    \n        paths.append(path)\n        \n        car_id = filename.split(\".\")[0]\n        car_ids.append(car_id)\n\nd = {\"id\": car_ids, \"car_path\": paths}\ndf = pd.DataFrame(data = d)\ndf = df.set_index('id')\ndf","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:45:48.832272Z","iopub.execute_input":"2021-10-04T08:45:48.832588Z","iopub.status.idle":"2021-10-04T08:45:48.882469Z","shell.execute_reply.started":"2021-10-04T08:45:48.832559Z","shell.execute_reply":"2021-10-04T08:45:48.881512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"car_ids = []\nmask_path = []\nfor dirname, _, filenames in os.walk('/kaggle/working/train_masks'):\n    for filename in filenames:\n        path = os.path.join(dirname, filename)\n        mask_path.append(path)\n        \n        car_id = filename.split(\".\")[0]\n        car_id = car_id.split(\"_mask\")[0]\n        car_ids.append(car_id)\n\n        \nd = {\"id\": car_ids,\"mask_path\": mask_path}\nmask_df = pd.DataFrame(data = d)\nmask_df = mask_df.set_index('id')\nmask_df","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:45:50.393134Z","iopub.execute_input":"2021-10-04T08:45:50.393478Z","iopub.status.idle":"2021-10-04T08:45:50.432565Z","shell.execute_reply.started":"2021-10-04T08:45:50.393447Z","shell.execute_reply":"2021-10-04T08:45:50.431774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"mask_path\"] = mask_df[\"mask_path\"]\ndf","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:45:51.676827Z","iopub.execute_input":"2021-10-04T08:45:51.677194Z","iopub.status.idle":"2021-10-04T08:45:51.692751Z","shell.execute_reply.started":"2021-10-04T08:45:51.677163Z","shell.execute_reply":"2021-10-04T08:45:51.691781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Xử lí data trước khi train","metadata":{}},{"cell_type":"markdown","source":"Code 1\n","metadata":{}},{"cell_type":"code","source":"train_df, val_df = train_test_split(df, test_size=0.25, shuffle = True)\nval_df, set_df = train_test_split(val_df, test_size=0.25, shuffle = True)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:46:04.044923Z","iopub.execute_input":"2021-10-04T08:46:04.04536Z","iopub.status.idle":"2021-10-04T08:46:04.055736Z","shell.execute_reply.started":"2021-10-04T08:46:04.045321Z","shell.execute_reply":"2021-10-04T08:46:04.054476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = [512,512]\n\ndef data_augmentation(car_img, mask_img):\n\n    if tf.random.uniform(()) > 0.5:\n        car_img = tf.image.flip_left_right(car_img)\n        mask_img = tf.image.flip_left_right(mask_img)\n\n    return car_img, mask_img\n\ndef preprocessing(car_path, mask_path):\n    car_img = tf.io.read_file(car_path) \n    car_img = tf.image.decode_jpeg(car_img, channels=3)\n    car_img = tf.image.resize(car_img, img_size)\n    car_img = tf.cast(car_img, tf.float32) / 255.0\n    \n    mask_img = tf.io.read_file(mask_path)\n    mask_img = tf.image.decode_jpeg(mask_img, channels=3)\n    mask_img = tf.image.resize(mask_img, img_size)\n    mask_img = mask_img[:,:,:1]    \n    mask_img = tf.math.sign(mask_img)\n    \n    \n    return car_img, mask_img\n\ndef create_dataset(df, train = False):\n    if not train:\n        ds = tf.data.Dataset.from_tensor_slices((df[\"car_path\"].values, df[\"mask_path\"].values))\n        ds = ds.map(preprocessing, tf.data.AUTOTUNE)\n    else:\n        ds = tf.data.Dataset.from_tensor_slices((df[\"car_path\"].values, df[\"mask_path\"].values))\n        ds = ds.map(preprocessing, tf.data.AUTOTUNE)\n        ds = ds.map(data_augmentation, tf.data.AUTOTUNE)\n\n    return ds\n\ndef display(display_list):\n    plt.figure(figsize=(15, 15))\n\n    title = ['Input Image', 'True Mask', 'Predicted Mask']\n\n    for i in range(len(display_list)):\n        plt.subplot(1, len(display_list), i+1)\n        plt.title(title[i])\n        plt.imshow(display_list[i]) # tf.keras.preprocessing.image.array_to_img(display_list[i])\n        plt.axis('off')\n    plt.show()\n    \ntrain = create_dataset(train_df, train = True)\nvalid = create_dataset(val_df)\n\nfor i in range(2):\n   for image, mask in train.take(i):\n        sample_image, sample_mask = image, mask\n        display([sample_image, sample_mask])\n        \n\n","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:46:06.436769Z","iopub.execute_input":"2021-10-04T08:46:06.437089Z","iopub.status.idle":"2021-10-04T08:46:13.122372Z","shell.execute_reply.started":"2021-10-04T08:46:06.437058Z","shell.execute_reply":"2021-10-04T08:46:13.118046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Code 2","metadata":{}},{"cell_type":"code","source":"# Dataset\nclass CustomDataset(Dataset):\n    def __init__(self, df, transform = None):\n        super(CustomDataset, self).__init__()\n        self.df = df\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        car_path = self.df.iloc[idx, :]['car_path']\n        mask_path = self.df.iloc[idx, :]['mask_path']\n        car = np.array(Image.open(car_path).convert('RGB'))\n        mask = np.array(Image.open(mask_path).convert('L'))\n        if self.transform:\n            transformed = self.transform(image=np.array(car), mask=np.array(mask))\n            car = transformed['image']\n            mask = transformed['mask']\n        car = transforms.ToTensor()(car)\n        mask = transforms.ToTensor()(mask)\n        return car, mask","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:46:25.922248Z","iopub.execute_input":"2021-10-04T08:46:25.922566Z","iopub.status.idle":"2021-10-04T08:46:25.930562Z","shell.execute_reply.started":"2021-10-04T08:46:25.922537Z","shell.execute_reply":"2021-10-04T08:46:25.929426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# increase idensity for dataset and so , decrease overfiting \ntrain_transform = A.Compose(\n    [\n    A.Resize(height=512, width= 512),\n    A.Rotate(limit=35, p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.1),\n    #It takes an array in as an input and normalizes its values between 0 and 1. It then returns an output array with the same dimensions as the input.\n    A.Normalize(\n        mean=[0.0, 0.0, 0.0],\n        std=[1.0, 1.0, 1.0],\n        max_pixel_value=255.0,\n    ),\n    ],\n)\n\nval_transform = A.Compose(\n    [\n    A.Resize(height=512, width= 512),\n    A.Normalize(\n        mean=[0.0, 0.0, 0.0],\n        std=[1.0, 1.0, 1.0],\n        max_pixel_value=255.0,\n    ),\n    ], \n)\n","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:46:29.769609Z","iopub.execute_input":"2021-10-04T08:46:29.76996Z","iopub.status.idle":"2021-10-04T08:46:29.776333Z","shell.execute_reply.started":"2021-10-04T08:46:29.769927Z","shell.execute_reply":"2021-10-04T08:46:29.775266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CustomDataset(train_df, train_transform)\nval_dataset = CustomDataset(val_df, val_transform)\nset_dataset = CustomDataset(set_df, val_transform)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:46:32.427919Z","iopub.execute_input":"2021-10-04T08:46:32.428299Z","iopub.status.idle":"2021-10-04T08:46:32.433577Z","shell.execute_reply.started":"2021-10-04T08:46:32.428269Z","shell.execute_reply":"2021-10-04T08:46:32.432474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=4, pin_memory=True )\nval_dataloader = DataLoader(val_dataset, batch_size=4, shuffle=False, num_workers=4, pin_memory=True )\nset_dataloader = DataLoader(set_dataset, pin_memory=True )","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:46:35.272766Z","iopub.execute_input":"2021-10-04T08:46:35.27312Z","iopub.status.idle":"2021-10-04T08:46:35.277819Z","shell.execute_reply.started":"2021-10-04T08:46:35.273085Z","shell.execute_reply":"2021-10-04T08:46:35.276834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"dataiter = iter(set_dataloader)\ncars, masks = next(dataiter)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:27:36.916639Z","iopub.execute_input":"2021-10-04T07:27:36.91701Z","iopub.status.idle":"2021-10-04T07:27:37.071397Z","shell.execute_reply.started":"2021-10-04T07:27:36.916976Z","shell.execute_reply":"2021-10-04T07:27:37.070485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nprint(cars.shape)\nprint(masks[0].shape)\n\n","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:27:38.687897Z","iopub.execute_input":"2021-10-04T07:27:38.688313Z","iopub.status.idle":"2021-10-04T07:27:38.698168Z","shell.execute_reply.started":"2021-10-04T07:27:38.688259Z","shell.execute_reply":"2021-10-04T07:27:38.697124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.unique(masks)\n# output = torch.unique(torch.tensor([1, 3, 2, 3], dtype=torch.long))\n# output\n# tensor([ 2,  3,  1])\n# => mean that maks is the tensor includes in 0. , 1. pixels","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:27:44.543306Z","iopub.execute_input":"2021-10-04T07:27:44.543652Z","iopub.status.idle":"2021-10-04T07:27:44.580435Z","shell.execute_reply.started":"2021-10-04T07:27:44.543622Z","shell.execute_reply":"2021-10-04T07:27:44.579284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cars","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:06.24999Z","iopub.execute_input":"2021-10-04T07:28:06.250346Z","iopub.status.idle":"2021-10-04T07:28:06.261926Z","shell.execute_reply.started":"2021-10-04T07:28:06.250311Z","shell.execute_reply":"2021-10-04T07:28:06.260798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"masks","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:08.863128Z","iopub.execute_input":"2021-10-04T07:28:08.863518Z","iopub.status.idle":"2021-10-04T07:28:08.870451Z","shell.execute_reply.started":"2021-10-04T07:28:08.863484Z","shell.execute_reply":"2021-10-04T07:28:08.869544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"car, mask = train_dataset[0]","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:33.434914Z","iopub.execute_input":"2021-10-04T07:28:33.435238Z","iopub.status.idle":"2021-10-04T07:28:33.505751Z","shell.execute_reply.started":"2021-10-04T07:28:33.435207Z","shell.execute_reply":"2021-10-04T07:28:33.504832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"car.shape","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:35.416462Z","iopub.execute_input":"2021-10-04T07:28:35.416779Z","iopub.status.idle":"2021-10-04T07:28:35.422125Z","shell.execute_reply.started":"2021-10-04T07:28:35.416748Z","shell.execute_reply":"2021-10-04T07:28:35.421072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.unique(mask)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:41.05601Z","iopub.execute_input":"2021-10-04T07:28:41.056363Z","iopub.status.idle":"2021-10-04T07:28:41.066734Z","shell.execute_reply.started":"2021-10-04T07:28:41.056332Z","shell.execute_reply":"2021-10-04T07:28:41.065629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask.dtype\ncar.dtype","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:50.103547Z","iopub.execute_input":"2021-10-04T07:28:50.103903Z","iopub.status.idle":"2021-10-04T07:28:50.112293Z","shell.execute_reply.started":"2021-10-04T07:28:50.103865Z","shell.execute_reply":"2021-10-04T07:28:50.111058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_df)\n","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:53.30297Z","iopub.execute_input":"2021-10-04T07:28:53.303334Z","iopub.status.idle":"2021-10-04T07:28:53.310224Z","shell.execute_reply.started":"2021-10-04T07:28:53.3033Z","shell.execute_reply":"2021-10-04T07:28:53.308977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(val_df)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:28:56.151052Z","iopub.execute_input":"2021-10-04T07:28:56.151394Z","iopub.status.idle":"2021-10-04T07:28:56.159485Z","shell.execute_reply.started":"2021-10-04T07:28:56.15136Z","shell.execute_reply":"2021-10-04T07:28:56.15623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***MODEL***","metadata":{}},{"cell_type":"code","source":"# Model\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms.functional as TF\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, 1, 1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, 1, 1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.conv(x)\n\nclass UNET(nn.Module):\n    def __init__(\n            self, in_channels=3, out_channels=1, features=[64, 128, 256, 512],\n    ):\n        super(UNET, self).__init__()\n        self.ups = nn.ModuleList() # Holds submodules in a list\n        self.downs = nn.ModuleList()\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n        #down part of Unet\n        for feature in features:\n            self.downs.append(DoubleConv(in_channels, feature)) # add Doubleconv in list downs \n            in_channels = feature # example with 1 -> 64 -> 64 => DoubleConv(1,64) and in_channels = 64 , features=[64, 128, 256, 512]\n\n        #Up part of Unet\n        for feature in reversed(features): # The reversed() method returns the reversed iterator of the given sequence. It is the same as the iter() method but in reverse order\n            self.ups.append(\n                # Transposed convolution that help us encode the previous feature maps into the more details feature \n                nn.ConvTranspose2d(\n                    feature*2, feature, kernel_size=2, stride=2,\n                )\n            )\n            self.ups.append(DoubleConv(feature*2, feature))\n\n        self.bottleneck = DoubleConv(features[-1], features[-1]*2) # features[-1] = 512 \n        self.final_conv = nn.Conv2d(features[0], out_channels, kernel_size=1) # features[0] = 64 \n\n    def forward(self, x):\n        skip_connections = [] # make to copy and crop in Unet\n\n        for down in self.downs:\n            x = down(x)\n            skip_connections.append(x)\n            x = self.pool(x)\n\n        x = self.bottleneck(x)\n        skip_connections = skip_connections[::-1] # [::-1] = list[<start>:<stop>:<step>]\n\n        for idx in range(0, len(self.ups),2):\n            x = self.ups[idx](x)\n            skip_connection=skip_connections[idx//2] # // chia lam tron\n\n            if x.shape != skip_connection.shape:\n                x = TF.resize(x, size=skip_connection.shape[2:]) # [2:] = reshape hight and width\n\n            concat_skip = torch.cat((skip_connection, x), dim=1)\n            x = self.ups[idx+1](concat_skip)\n\n        return  self.final_conv(x)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:54:38.616728Z","iopub.execute_input":"2021-10-04T08:54:38.617069Z","iopub.status.idle":"2021-10-04T08:54:38.632729Z","shell.execute_reply.started":"2021-10-04T08:54:38.617031Z","shell.execute_reply":"2021-10-04T08:54:38.631739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNET(in_channels=3, out_channels=1)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:54:41.757481Z","iopub.execute_input":"2021-10-04T08:54:41.757815Z","iopub.status.idle":"2021-10-04T08:54:42.007691Z","shell.execute_reply.started":"2021-10-04T08:54:41.757783Z","shell.execute_reply":"2021-10-04T08:54:42.006818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:54:43.390586Z","iopub.execute_input":"2021-10-04T08:54:43.390955Z","iopub.status.idle":"2021-10-04T08:54:43.400044Z","shell.execute_reply.started":"2021-10-04T08:54:43.390924Z","shell.execute_reply":"2021-10-04T08:54:43.398887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2021-10-04T06:09:01.137244Z","iopub.execute_input":"2021-10-04T06:09:01.137625Z","iopub.status.idle":"2021-10-04T06:09:10.496408Z","shell.execute_reply.started":"2021-10-04T06:09:01.137557Z","shell.execute_reply":"2021-10-04T06:09:10.49508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchsummary import summary","metadata":{"execution":{"iopub.status.busy":"2021-10-04T06:09:14.325127Z","iopub.execute_input":"2021-10-04T06:09:14.32558Z","iopub.status.idle":"2021-10-04T06:09:14.338669Z","shell.execute_reply.started":"2021-10-04T06:09:14.325545Z","shell.execute_reply":"2021-10-04T06:09:14.337333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(model, (3, 512, 512), 1,'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-10-04T06:09:17.387656Z","iopub.execute_input":"2021-10-04T06:09:17.388056Z","iopub.status.idle":"2021-10-04T06:09:33.981496Z","shell.execute_reply.started":"2021-10-04T06:09:17.388024Z","shell.execute_reply":"2021-10-04T06:09:33.980411Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test():\n    x = torch.randn((3, 1, 160, 16))\n    model = UNET(in_channels=1, out_channels=1)\n    preds = model(x)\n    print(preds.shape)\n    print(x.shape)\n    #assert  preds.shape == x.shape\n\nif __name__ == \"__main__\":\n    test()","metadata":{"execution":{"iopub.status.busy":"2021-10-04T07:29:21.365715Z","iopub.execute_input":"2021-10-04T07:29:21.366058Z","iopub.status.idle":"2021-10-04T07:29:21.865572Z","shell.execute_reply.started":"2021-10-04T07:29:21.366027Z","shell.execute_reply":"2021-10-04T07:29:21.864629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Train the Model***","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:54:15.896203Z","iopub.execute_input":"2021-10-04T08:54:15.896559Z","iopub.status.idle":"2021-10-04T08:54:15.905121Z","shell.execute_reply.started":"2021-10-04T08:54:15.89652Z","shell.execute_reply":"2021-10-04T08:54:15.904217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:55:22.497275Z","iopub.execute_input":"2021-10-04T08:55:22.497644Z","iopub.status.idle":"2021-10-04T08:55:22.546862Z","shell.execute_reply.started":"2021-10-04T08:55:22.497603Z","shell.execute_reply":"2021-10-04T08:55:22.545867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Once the model is created, you can config the model with losses and metrics with model.compile(), \ntrain the model with model.fit(), \nor use the model to do prediction with model.predict().","metadata":{}},{"cell_type":"markdown","source":"Training","metadata":{}},{"cell_type":"code","source":"criterion= nn.BCEWithLogitsLoss()\noptimizer= torch.optim.Adam(model.parameters(),lr=1e-3)\n# Train model cách 1\nfor i, batch in enumerate(train_dataloader):\n    cars, masks = batch\n    cars = cars.to(device)\n    masks = masks.to(device)\n    preds = model(cars)\n    loss = criterion(preds, masks)\n    \n    optimizer.zero_grad()\n    loss.backward()\n    optimizer.step()\n    print(f\"loss: {loss.item()}\")\n    ","metadata":{"execution":{"iopub.status.busy":"2021-10-04T09:04:39.057477Z","iopub.execute_input":"2021-10-04T09:04:39.057857Z","iopub.status.idle":"2021-10-04T09:18:29.798111Z","shell.execute_reply.started":"2021-10-04T09:04:39.057815Z","shell.execute_reply":"2021-10-04T09:18:29.797205Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Show ảnh dự đoán sau khi train xong","metadata":{}},{"cell_type":"code","source":"to_pil_image = transforms.ToPILImage()","metadata":{"execution":{"iopub.status.busy":"2021-10-04T09:20:47.322891Z","iopub.execute_input":"2021-10-04T09:20:47.323384Z","iopub.status.idle":"2021-10-04T09:20:47.332159Z","shell.execute_reply.started":"2021-10-04T09:20:47.323331Z","shell.execute_reply":"2021-10-04T09:20:47.331016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_imgs(\n    loader, model, device=\"cuda\"\n):\n  model.eval()\n  for idx, (x, y) in enumerate(loader):\n    x = x.to(device=device)\n    with torch.no_grad():\n      preds = torch.sigmoid(model(x))\n      preds = (preds > 0.5).float()\n      img = np.array(to_pil_image(preds[0]))\n      fig = plt.imshow(img)\n\n    #model.train()","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:21:47.667243Z","iopub.execute_input":"2021-10-04T08:21:47.667606Z","iopub.status.idle":"2021-10-04T08:21:47.674253Z","shell.execute_reply.started":"2021-10-04T08:21:47.667571Z","shell.execute_reply":"2021-10-04T08:21:47.673175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_imgs(val_dataloader, model, device='cuda')","metadata":{"execution":{"iopub.status.busy":"2021-10-04T08:21:51.60524Z","iopub.execute_input":"2021-10-04T08:21:51.605598Z","iopub.status.idle":"2021-10-04T08:23:14.369567Z","shell.execute_reply.started":"2021-10-04T08:21:51.605566Z","shell.execute_reply":"2021-10-04T08:23:14.368717Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataiter = iter(val_dataloader)\ncars, masks = next(dataiter)\ncars = cars.to(device)\nmasks = masks.to(device)\npreds = model(cars)\npreds = (preds > 0.5).float()\n\nfor j in range(4): \n    display_list = cars[j], masks[j], preds[j] \n    title = ['Input Image', 'True Mask', 'Predicted Mask']\n    for i in range(len(display_list)):\n        plt.subplot(1, len(display_list), i+1)\n        plt.title(title[i])\n        img = np.array(to_pil_image(display_list[i]))\n        fig1 = plt.imshow(img)\n        plt.axis('off')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-10-04T09:20:53.823676Z","iopub.execute_input":"2021-10-04T09:20:53.823992Z","iopub.status.idle":"2021-10-04T09:20:56.608663Z","shell.execute_reply.started":"2021-10-04T09:20:53.823959Z","shell.execute_reply":"2021-10-04T09:20:56.607753Z"},"trusted":true},"execution_count":null,"outputs":[]}]}