{"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":"import torch\nimport torch.nn as nn  # All neural network modules, nn.Linear, nn.Conv2d, BatchNorm, Loss functions\nimport torch.optim as optim  # For all Optimization algorithms, SGD, Adam, etc.\nimport torchvision.transforms as transforms  # Transformations we can perform on our dataset\nimport torchvision\nfrom PIL import Image, ImageEnhance\nimport math\nimport os\nimport pandas as pd\nfrom skimage import io\n\nfrom torch.utils.data import (\n    Dataset,\n    DataLoader,\n)  # Gives easier dataset managment and creates mini batches","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-21T16:32:35.417264Z","iopub.execute_input":"2022-07-21T16:32:35.418111Z","iopub.status.idle":"2022-07-21T16:32:35.426211Z","shell.execute_reply.started":"2022-07-21T16:32:35.418077Z","shell.execute_reply":"2022-07-21T16:32:35.424713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)), \n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # to tensor before normalize\n])\ndef load_image(file):\n    return Image.open(file)\ndef image_path(root, basename, extension):\n    return os.path.join(root, f'{basename}{extension}')\nos.listdir('../input/hsgs-hackathon2022/train_data/Train_labels')\nclass CustomImageDataset(Dataset):\n    def __init__(self, img_dir, label_dir, transform=transform):\n        video_files = os.listdir(img_dir)\n        self.labels = list()\n        self.imgs = list()\n        for file_name in video_files: \n            cur_csv = pd.read_csv(os.path.join(label_dir, file_name+'.csv'))\n            cur_len = len(cur_csv.index)\n            for itr in range(cur_len):\n                img_file_dir = os.path.join(os.path.join(img_dir, file_name), cur_csv.iloc[itr, 0] + '.PNG')\n                self.imgs.append(img_file_dir) # image directory\n                self.labels.append([cur_csv.iloc[itr, 0], cur_csv.iloc[itr, 1]]) # label\n        self.img_dir = img_dir\n        self.label_dir = label_dir\n        self.transform = transform\n\n            \n    def __getitem__(self, index):\n        with open(self.imgs[index], 'rb') as f:\n            image = load_image(f).convert('RGB')\n        label = self.labels[index][1]\n        filename = self.labels[index][0]\n        \n#         image, label = self.spm_transform(image, label)\n        \n        if self.transform is not None:\n            image = self.transform(image)\n\n        return image, label\n\n    def __len__(self):\n        return len(self.labels)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-21T16:32:35.432013Z","iopub.execute_input":"2022-07-21T16:32:35.432656Z","iopub.status.idle":"2022-07-21T16:32:35.458181Z","shell.execute_reply.started":"2022-07-21T16:32:35.432614Z","shell.execute_reply":"2022-07-21T16:32:35.456813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)\n\n# Hyperparameters\nin_channel = 3\nnum_classes = 2\nlearning_rate = 1e-3\nbatch_size = 64\nnum_epochs = 10\n\n# Load Data\n\ndataset = CustomImageDataset(\n    img_dir='../input/hsgs-hackathon2022/train_data/Train',\n    label_dir='../input/hsgs-hackathon2022/train_data/Train_labels'\n)\nn_train = math.floor(0.9*len(dataset)) \nn_val = len(dataset) - math.floor(0.9*len(dataset)) \ntrain_dataset, validation_dataset = torch.utils.data.random_split(dataset, [n_train, n_val])\ntrain=train_dataset\nval=validation_dataset\n\ntrain_loader = DataLoader(train_dataset, batch_size = batch_size, shuffle=True, drop_last = True)\nvalidation_loader = DataLoader(validation_dataset, batch_size = batch_size, shuffle=False, drop_last = True)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-21T16:32:35.461416Z","iopub.execute_input":"2022-07-21T16:32:35.46228Z","iopub.status.idle":"2022-07-21T16:32:38.93112Z","shell.execute_reply.started":"2022-07-21T16:32:35.462234Z","shell.execute_reply":"2022-07-21T16:32:38.929775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nmodel=torchvision.models.googlenet(pretrained=True)\nmodel=model.to(device)\n\n# Loss and optimizer\ncriterion = nn.CrossEntropyLoss()\ncriterion = criterion.to(device)\noptimizer = optim.Adam(model.parameters(), lr=learning_rate)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-21T16:32:38.933234Z","iopub.execute_input":"2022-07-21T16:32:38.9337Z","iopub.status.idle":"2022-07-21T16:32:39.211867Z","shell.execute_reply.started":"2022-07-21T16:32:38.933654Z","shell.execute_reply":"2022-07-21T16:32:39.210595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train Network\nfrom tqdm import tqdm\nmin_val_loss=100000000000\nfor epoch in range(num_epochs):\n  total_loss_train=0\n  total_acc_train=0\n  for x, y in tqdm(train_loader):\n\n    x = x.to(device)\n    y = y.to(device)\n    output = model(x.float())\n\n    batch_loss = criterion(output, y)\n    total_loss_train += batch_loss.item()\n\n    acc = (output.argmax(dim=1)==y).sum().item()\n    total_acc_train += acc\n\n    optimizer.zero_grad()\n    batch_loss.backward()\n    optimizer.step()\n  \n  total_loss_val=0\n  total_acc_val=0\n\n  with torch.no_grad():\n    for x, y in tqdm(validation_loader):\n\n      x = x.to(device)\n      y = y.to(device)\n      output = model(x)\n      batch_loss = criterion(output, y)\n      total_loss_val += batch_loss.item()\n\n      acc = (output.argmax(dim=1)==y).sum().item()\n      total_acc_val += acc\n  \n  print(\n      f'Epochs: {epoch+1} | Train Loss: {total_loss_train / len(train):.3f}\\\n      | Train Accuracy: {total_acc_train/len(train):.3f}\\\n      | Val Loss: {total_loss_val/len(val):.3f}\\\n      | Val Accuracy:{total_acc_val/len(val):.3f}'\n  )\n\n  if min_val_loss>total_loss_val/len(val):\n    min_val_loss = total_loss_val/len(val)\n    torch.save(model.state_dict(), \"ModelA.pt\")\n    print(f\"Save model because val loss improve loss {min_val_loss:.3f}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-21T16:32:39.216406Z","iopub.execute_input":"2022-07-21T16:32:39.216753Z","iopub.status.idle":"2022-07-21T16:32:40.810748Z","shell.execute_reply.started":"2022-07-21T16:32:39.216723Z","shell.execute_reply":"2022-07-21T16:32:40.808892Z"},"trusted":true},"execution_count":null,"outputs":[]}]}