{"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":"!pip install torch_snippets > /dev/null\nfrom torch_snippets import Report\nimport os\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch import nn\nfrom torch import optim\nfrom torch_snippets import *\n\n# Chemin du dossier contenant les fichiers .npy\nmask_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/1000660467359258186/human_pixel_masks.npy\"\n\n\n# Charger l'image depuis le fichier .npy\nimage = np.load(mask_path)\nprint(image.shape)\n\nfor channel in range(0,image.shape[2]):\n\n    # Afficher l'image\n    new_shape = image.shape[0],image.shape[1]\n    plt.imshow(image[:,:,channel], cmap='gray')\n    plt.title(\"mask\")  # Utilise le nom du fichier comme titre\n    plt.show()\n    print(\"ok\")\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-28T13:05:06.030071Z","iopub.execute_input":"2023-05-28T13:05:06.030791Z","iopub.status.idle":"2023-05-28T13:05:18.473423Z","shell.execute_reply.started":"2023-05-28T13:05:06.030747Z","shell.execute_reply":"2023-05-28T13:05:18.472341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visuliser les masques","metadata":{}},{"cell_type":"markdown","source":"# Load dataset","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nimport json\n\nfolder = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/\"\n\nclass ContrailDataset(Dataset):\n    def __init__(self, ids, base_dir, bands=None, transforms:list=[], test_mode:bool=False):\n        self.ids = ids\n        self.base_dir = base_dir\n        self.transforms = transforms\n        self.bands = bands\n        self.permute = (2,0,1)\n        self.test_mode = test_mode\n        \n    def __getitem__(self, index):\n        record_id = self.ids[index]\n        \n        if self.bands is None:\n            band_list = [f'band_{band:02d}.npy' for band in range(8,17)]\n        else :\n            band_list = [f'band_{int(band):02d}.npy' for band in self.bands]\n        \n        x = list()\n        for band in band_list:\n            with open(os.path.join(self.base_dir, record_id, band), 'rb') as f:\n                x.append(np.load(f).transpose(self.permute))\n        x = np.stack(x,axis=0) ## X.shape = (Band,Time_frame,H,W)\n        \n        for transformation in self.transforms:\n            x = transformation(x)\n        x = torch.from_numpy(x.astype(np.float32))\n        \n        if self.test_mode:\n            return x\n        else:\n            with open(os.path.join(self.base_dir, record_id,'human_pixel_masks.npy'), 'rb') as f:\n                y = torch.from_numpy(np.load(f).squeeze().astype(np.float32))\n\n            return x[5,:,:,:], y\n\n    def __len__(self):\n        return len(self.ids)\n    \n    def ratio_0(self):\n        total_mean = 0\n        for i in range(self.__len__()):\n            _,y = self.__getitem__(i)\n            total_mean = y.mean().item()\n        return total_mean/self.__len__()\n        \n#Load \ntrain_ids = []\nwith open(os.path.join(folder,\"train_metadata.json\")) as f:\n    data = json.load(f)\n    train_ids = [record[\"record_id\"] for record in data]\n    \nval_ids = []\nwith open(os.path.join(folder,\"validation_metadata.json\")) as f:\n    data = json.load(f)\n    val_ids = [record[\"record_id\"] for record in data]\n    \n\n# Utiliser les données\nprint(\"Number of train exemple :\",len(train_ids))\nprint(\"Number of validation exemple :\",len(val_ids))\n\n# Datasets \ndataset_params = {}\ntrain_dataset = ContrailDataset(train_ids, os.path.join(folder,\"train\"), **dataset_params)\n\n\nval_dataset = ContrailDataset(val_ids,  os.path.join(folder,\"validation\"), **dataset_params)\n\n# DalaLoaders\ndataloader_params = {\n    \"batch_size\" : 32, \n    \"shuffle\" : True,\n    \"num_workers\": 0\n}\ntrain_loader = DataLoader(train_dataset, **dataloader_params)\nval_loader = DataLoader(val_dataset, **dataloader_params)\nprint(\"batch number  of train exemple :\",len(train_loader))\nprint(\"batch number batch of validation exemple :\",len(val_loader))\n","metadata":{"execution":{"iopub.status.busy":"2023-05-28T13:05:18.475965Z","iopub.execute_input":"2023-05-28T13:05:18.476336Z","iopub.status.idle":"2023-05-28T13:05:18.687433Z","shell.execute_reply.started":"2023-05-28T13:05:18.476299Z","shell.execute_reply":"2023-05-28T13:05:18.686343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions\n","metadata":{}},{"cell_type":"code","source":"import torch\n\ndevice = \"cuda\"\n\n\ndef UnetLoss(preds, targets):\n    ce = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([10],device=device))\n    ce_loss = ce(preds, targets)\n    \n    preds = torch.round(torch.sigmoid(preds))\n\n    correct = (preds == targets).float()\n    acc = correct.sum() / targets.numel() \n    \n    intersection = torch.sum(preds * targets) \n    union = torch.sum(preds) + torch.sum(targets)  \n    dice = (2 * intersection) / (union + intersection)\n    \n    ratio = intersection / torch.sum(torch.eq(targets, 1))\n    \n    return ce_loss, acc,ratio, dice\n\ndef train_batch(model, data, optimizer, criterion):\n    model.train()\n    ims, ce_masks = data\n    ims = ims.to(device)\n    ce_masks = ce_masks.to(device)\n    _masks = model(ims)\n    optimizer.zero_grad()\n    loss, acc,ratio, dice = criterion(_masks, ce_masks.view(ce_masks.shape[0],-1))\n    loss.backward()\n    optimizer.step()\n    return loss.item(), acc.item(),ratio.item(), dice.item()\n\n@torch.no_grad()\ndef validate_batch(model, data, criterion):\n    model.eval()\n    ims, masks = data\n    ims = ims.to(device)\n    masks = masks.to(device)\n    _masks = model(ims)\n    loss, acc,ratio,dice = criterion(_masks, masks.view(masks.shape[0],-1))\n    return loss.item(), acc.item(),ratio.item(),dice.item()\n\n                                        \n        \n        ","metadata":{"execution":{"iopub.status.busy":"2023-05-28T13:05:18.689174Z","iopub.execute_input":"2023-05-28T13:05:18.689555Z","iopub.status.idle":"2023-05-28T13:05:18.701459Z","shell.execute_reply.started":"2023-05-28T13:05:18.689521Z","shell.execute_reply":"2023-05-28T13:05:18.700557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\n\n","metadata":{}},{"cell_type":"code","source":"\nimport torchvision.models as models\n\n# Charger le modèle pré-entraîné ResNet\nmodel = models.resnet50(pretrained=True)\n\n# Figer les poids du modèle\nfor param in model.parameters():\n    param.requires_grad = False\n    \n# Modifier la couche de sortie pour la tâche de masquage\nmodel.fc = nn.Identity()\nmodel.avgpool = nn.Sequential(\n    nn.Conv2d(2048,256*256,kernel_size = (1,1),stride=(1, 1), bias=False),\n    nn.AdaptiveAvgPool2d((1,1))\n)\n\n    \nmodel.conv1 = nn.Sequential(\n    nn.Conv2d(8, 3, kernel_size=(1, 1), stride=(1, 1), padding=(1, 1), bias=False),\n    model.conv1\n)\n\n\n\n\nprint(model)\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-05-28T13:05:18.703885Z","iopub.execute_input":"2023-05-28T13:05:18.704657Z","iopub.status.idle":"2023-05-28T13:05:21.214516Z","shell.execute_reply.started":"2023-05-28T13:05:18.704622Z","shell.execute_reply":"2023-05-28T13:05:21.213566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train ","metadata":{}},{"cell_type":"code","source":"# Créer un objet Checkpoint en spécifiant le répertoire de sauvegarde\ncriterion = UnetLoss\noptimizer = optim.Adam(model.parameters(), lr=1e-5)\nn_epochs = 10\n\n\n\n\n\nlog = Report(n_epochs)\n\nfor ex in range(n_epochs):\n    print(\"----\")\n    N = len(train_loader)\n    for bx, data in enumerate(train_loader):\n        loss,acc,ratio,dice = train_batch(model, data, optimizer, criterion)\n        log.record(ex+(bx+1)/N, trn_loss=loss, trn_acc=acc,ratio=ratio,dice=dice, end='\\r')\n\n    N = len(val_loader)\n    for bx, data in enumerate(val_loader):\n        loss,acc,ratio,dice = validate_batch(model, data, criterion)\n        log.record(ex+(bx+1)/N, val_loss=loss, val_acc=acc,ratio=ratio,dice=dice, end='\\r')\n    \n    \n    log.report_avgs(ex+1)","metadata":{"execution":{"iopub.status.busy":"2023-05-28T13:05:21.216029Z","iopub.execute_input":"2023-05-28T13:05:21.216368Z","iopub.status.idle":"2023-05-28T13:05:32.553965Z","shell.execute_reply.started":"2023-05-28T13:05:21.216333Z","shell.execute_reply":"2023-05-28T13:05:32.55211Z"},"trusted":true},"execution_count":null,"outputs":[]}]}