{"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":"markdown","source":"# 1.5 「ファインチューニング」で精度向上を実現する方法  \n- 本ファイルでは、学習済みのVGGモデルを使用し、ファインチューニングでアリとハチの画像を分類するモデルを学習します","metadata":{}},{"cell_type":"markdown","source":"# 学習目標  \n1.PyTorchでGPUを使用する実装コードを書けるようになる  \n2.最適化手法の設定において、層ごとに異なる学習率を設定したファインチューニングを実装できるようになる  \n3.学習したネットワークを保存・ロードできるようになる ","metadata":{}},{"cell_type":"code","source":"# パッケージのimport\nimport os\nimport numpy as np\nimport json\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n%matplotlib inline\nimport random\n\n\nfrom tqdm import tqdm\n\nimport torch\nimport torchvision\nimport torch.utils.data as data\nimport torch.nn as nn\nimport torch.optim as optim\nimport pandas as pd\nfrom torchvision import models, transforms\nfrom sklearn.model_selection import train_test_split ","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:40.830794Z","iopub.execute_input":"2021-11-13T08:53:40.831258Z","iopub.status.idle":"2021-11-13T08:53:40.844527Z","shell.execute_reply.started":"2021-11-13T08:53:40.831191Z","shell.execute_reply":"2021-11-13T08:53:40.843447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# PyTorchのバージョン確認\nprint(\"PyTorch Version: \",torch.__version__)\nprint(\"Torchvision Version: \",torchvision.__version__)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:40.847172Z","iopub.execute_input":"2021-11-13T08:53:40.847761Z","iopub.status.idle":"2021-11-13T08:53:40.860337Z","shell.execute_reply.started":"2021-11-13T08:53:40.847719Z","shell.execute_reply":"2021-11-13T08:53:40.859294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 乱数のシードを設定\ntorch.manual_seed(1234)\nnp.random.seed(1234)\nrandom.seed(1234)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:40.861708Z","iopub.execute_input":"2021-11-13T08:53:40.862471Z","iopub.status.idle":"2021-11-13T08:53:40.868035Z","shell.execute_reply.started":"2021-11-13T08:53:40.862405Z","shell.execute_reply":"2021-11-13T08:53:40.867061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 入力画像の前処理クラスを作成","metadata":{}},{"cell_type":"code","source":"\nresize = 224\nmean = (0.485, 0.456, 0.406)\nstd = (0.229, 0.224, 0.225)\n\ntrain_transforms = transforms.Compose([\n                transforms.RandomResizedCrop(\n                    resize, scale=(0.5, 1.0)),  # データオーギュメンテーション\n                transforms.RandomHorizontalFlip(),  # データオーギュメンテーション\n                transforms.ToTensor(),  # テンソルに変換\n                transforms.Normalize(mean, std)  # 標準化\n    ])\nval_transforms = transforms.Compose([\n                transforms.Resize(resize),  # リサイズ\n                transforms.CenterCrop(resize),  # 画像中央をresize×resizeで切り取り\n                transforms.ToTensor(),  # テンソルに変換\n                transforms.Normalize(mean, std)  # 標準化\n    ])","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:40.869483Z","iopub.execute_input":"2021-11-13T08:53:40.870070Z","iopub.status.idle":"2021-11-13T08:53:40.880468Z","shell.execute_reply.started":"2021-11-13T08:53:40.870030Z","shell.execute_reply":"2021-11-13T08:53:40.879508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# テストデータを表示","metadata":{}},{"cell_type":"code","source":"# 画像前処理の動作を確認\n\n# 1. 画像読み込み\nimage_file_path = '../input/cassava-leaf-disease-classification/test_images/2216849948.jpg'\nimg = Image.open(image_file_path)  # [高さ][幅][色RGB]\n\n# 2. 元の画像の表示\nplt.imshow(img)\nplt.show()\n\n# 3. 画像の前処理と処理済み画像の表示\nimg_transformed = train_transforms(img)  # torch.Size([3, 224, 224])\n\n# (色、高さ、幅)を (高さ、幅、色)に変換し、0-1に値を制限して表示\nimg_transformed = img_transformed.numpy().transpose((1, 2, 0))\nimg_transformed = np.clip(img_transformed, 0, 1)\nplt.imshow(img_transformed)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:40.883035Z","iopub.execute_input":"2021-11-13T08:53:40.883757Z","iopub.status.idle":"2021-11-13T08:53:41.261052Z","shell.execute_reply.started":"2021-11-13T08:53:40.883715Z","shell.execute_reply":"2021-11-13T08:53:41.260163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataSetを作成","metadata":{}},{"cell_type":"code","source":"PATH = \"../input/cassava-leaf-disease-classification/train_images/\"\nIMG_SIZE = 224\n\nclass CassavaDataset(data.Dataset):\n    def __init__(self,path,image_ids,labels,image_size, mode='val'):\n        self.image_ids = image_ids\n        self.labels = labels\n        self.path = path\n        self.image_size = image_size\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self,item):\n      image_ids = str(self.image_ids[item])\n      labels = self.labels[item]\n      img = Image.open(self.path+image_ids)\n      \n      if self.mode==\"train\":\n        return train_transforms(img),torch.tensor(labels,dtype=torch.long)\n      else:\n        return val_transforms(img),torch.tensor(labels,dtype=torch.long)\n        #return torch.tensor(img,dtype=torch.float),torch.tensor(labels,dtype=torch.long)\n    \ndfx = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\nxtrain, xval, ytrain, yval = train_test_split(dfx[\"image_id\"].values,\n                                              dfx.label.values,\n                                              test_size = 0.1, random_state=0)\nIMG_SIZE = 224\n\n# 実行\ntrain_dataset = CassavaDataset(PATH,xtrain,ytrain,IMG_SIZE, mode=\"train\")\nval_dataset   = CassavaDataset(PATH,xval,yval,IMG_SIZE)\nprint(dfx)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:41.262866Z","iopub.execute_input":"2021-11-13T08:53:41.263391Z","iopub.status.idle":"2021-11-13T08:53:41.377225Z","shell.execute_reply.started":"2021-11-13T08:53:41.263350Z","shell.execute_reply":"2021-11-13T08:53:41.376312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataLoaderを作成","metadata":{}},{"cell_type":"code","source":"IMG_SIZE = 224\n\n# DataLoaderを作成する\nbatch_size = 32\n\ntrain_dataloader = torch.utils.data.DataLoader(\n    train_dataset, batch_size=batch_size, shuffle=True)\n\nval_dataloader = torch.utils.data.DataLoader(\n    val_dataset, batch_size=batch_size, shuffle=False)\n\n# 辞書オブジェクトにまとめる\ndataloaders_dict = {\"train\": train_dataloader, \"val\": val_dataloader}","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:41.378751Z","iopub.execute_input":"2021-11-13T08:53:41.379346Z","iopub.status.idle":"2021-11-13T08:53:41.387784Z","shell.execute_reply.started":"2021-11-13T08:53:41.379305Z","shell.execute_reply":"2021-11-13T08:53:41.386631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ネットワークモデルの作成","metadata":{}},{"cell_type":"code","source":"# 学習済みのVGG-16モデルをロード\n\n# VGG-16モデルのインスタンスを生成\nuse_pretrained = True  # 学習済みのパラメータを使用\n#use_pretrained = False  # 学習済みのパラメータを使用\nnet = models.vgg16(pretrained=use_pretrained)\n\n# VGG16の最後の出力層の出力ユニットを病気4種と健康の5つに付け替える\nnet.classifier[6] = nn.Linear(in_features=4096, out_features=5)\n\n# 訓練モードに設定\nnet.train()\n\nprint('ネットワーク設定完了：学習済みの重みをロードし、訓練モードに設定しました')","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:41.389174Z","iopub.execute_input":"2021-11-13T08:53:41.389618Z","iopub.status.idle":"2021-11-13T08:53:42.822474Z","shell.execute_reply.started":"2021-11-13T08:53:41.389551Z","shell.execute_reply":"2021-11-13T08:53:42.821474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 損失関数を定義\n* 損失関数は、モデルがいかに問題を学習・推論できるかを図る指標”ロス(損失)”の計算方法を指定する","metadata":{}},{"cell_type":"code","source":"# 損失関数の設定\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:42.823831Z","iopub.execute_input":"2021-11-13T08:53:42.824363Z","iopub.status.idle":"2021-11-13T08:53:42.829251Z","shell.execute_reply.started":"2021-11-13T08:53:42.824322Z","shell.execute_reply":"2021-11-13T08:53:42.828248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 最適化手法を設定\n* OptimizerがVGGモデルのどこのパラメータを更新(＝学習・最適化)するかを指定する\n* Optimizerは”誤差逆伝搬法”を用いて、損失関数を最小化するようにVGGモデルのパラメータを更新(＝学習・最適化)役割を担う","metadata":{}},{"cell_type":"code","source":"# ファインチューニングで学習させるパラメータを、変数params_to_updateの1～3に格納する\n\nparams_to_update_1 = []\nparams_to_update_2 = []\nparams_to_update_3 = []\n\n# 学習させる層のパラメータ名を指定\nupdate_param_names_1 = [\"features\"]\nupdate_param_names_2 = [\"classifier.0.weight\",\n                        \"classifier.0.bias\", \"classifier.3.weight\", \"classifier.3.bias\"]\nupdate_param_names_3 = [\"classifier.6.weight\", \"classifier.6.bias\"]\n\n# パラメータごとに各リストに格納する\nfor name, param in net.named_parameters():\n    if update_param_names_1[0] in name:\n        param.requires_grad = True\n        params_to_update_1.append(param)\n        print(\"params_to_update_1に格納：\", name)\n\n    elif name in update_param_names_2:\n        param.requires_grad = True\n        params_to_update_2.append(param)\n        print(\"params_to_update_2に格納：\", name)\n\n    elif name in update_param_names_3:\n        param.requires_grad = True\n        params_to_update_3.append(param)\n        print(\"params_to_update_3に格納：\", name)\n\n    else:\n        param.requires_grad = False\n        print(\"勾配計算なし。学習しない：\", name)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:42.831318Z","iopub.execute_input":"2021-11-13T08:53:42.832094Z","iopub.status.idle":"2021-11-13T08:53:42.853803Z","shell.execute_reply.started":"2021-11-13T08:53:42.832054Z","shell.execute_reply":"2021-11-13T08:53:42.853060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 最適化手法の設定\noptimizer = optim.SGD([\n    {'params': params_to_update_1, 'lr': 1e-4},\n    {'params': params_to_update_2, 'lr': 5e-4},\n    {'params': params_to_update_3, 'lr': 1e-3}\n], momentum=0.9)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:42.859273Z","iopub.execute_input":"2021-11-13T08:53:42.859509Z","iopub.status.idle":"2021-11-13T08:53:42.866858Z","shell.execute_reply.started":"2021-11-13T08:53:42.859486Z","shell.execute_reply":"2021-11-13T08:53:42.865670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 学習・検証を実施","metadata":{}},{"cell_type":"code","source":"# モデルを学習させる関数を作成\n\n\ndef train_model(net, dataloaders_dict, criterion, optimizer, num_epochs):\n\n    # 初期設定\n    # GPUが使えるかを確認\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    print(\"使用デバイス：\", device)\n\n    # ネットワークをGPUへ\n    net.to(device)\n\n    # ネットワークがある程度固定であれば、高速化させる\n    torch.backends.cudnn.benchmark = True\n\n    # epochのループ\n    for epoch in range(num_epochs):\n        print('Epoch {}/{}'.format(epoch+1, num_epochs))\n        print('-------------')\n\n        # epochごとの訓練と検証のループ\n        for phase in ['train', 'val']:\n            if phase == 'train':\n                net.train()  # モデルを訓練モードに\n            else:\n                net.eval()   # モデルを検証モードに\n\n            epoch_loss = 0.0  # epochの損失和\n            epoch_corrects = 0  # epochの正解数\n\n            # 未学習時の検証性能を確かめるため、epoch=0の訓練は省略\n            if (epoch == 0) and (phase == 'train'):\n                continue\n\n            # データローダーからミニバッチを取り出すループ\n            for inputs, labels in tqdm(dataloaders_dict[phase]):\n\n                # GPUが使えるならGPUにデータを送る\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n                                \n                # optimizerを初期化\n                optimizer.zero_grad()\n\n                # 順伝搬（forward）計算\n                with torch.set_grad_enabled(phase == 'train'):\n                    outputs = net(inputs)\n                    loss = criterion(outputs, labels)  # 損失を計算\n                    _, preds = torch.max(outputs, 1)  # ラベルを予測\n\n                    # 訓練時はバックプロパゲーション\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n\n                    # 結果の計算\n                    epoch_loss += loss.item() * inputs.size(0)  # lossの合計を更新\n                    # 正解数の合計を更新\n                    epoch_corrects += torch.sum(preds == labels.data)\n\n            # epochごとのlossと正解率を表示\n            epoch_loss = epoch_loss / len(dataloaders_dict[phase].dataset)\n            epoch_acc = epoch_corrects.double(\n            ) / len(dataloaders_dict[phase].dataset)\n\n            print('{} Loss: {:.4f} Acc: {:.4f}'.format(\n                phase, epoch_loss, epoch_acc))\n            ","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:42.870642Z","iopub.execute_input":"2021-11-13T08:53:42.871282Z","iopub.status.idle":"2021-11-13T08:53:42.886810Z","shell.execute_reply.started":"2021-11-13T08:53:42.871241Z","shell.execute_reply":"2021-11-13T08:53:42.885937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 学習・検証を実行する\nnum_epochs=1\ntrain_model(net, dataloaders_dict, criterion, optimizer, num_epochs=num_epochs)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:53:42.889427Z","iopub.execute_input":"2021-11-13T08:53:42.890058Z","iopub.status.idle":"2021-11-13T08:54:13.347868Z","shell.execute_reply.started":"2021-11-13T08:53:42.890017Z","shell.execute_reply":"2021-11-13T08:54:13.343613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 学習したネットワークを保存・ロード","metadata":{}},{"cell_type":"code","source":"# PyTorchのネットワークパラメータの保存\nsave_path = './weights_fine_tuning.pth'\ntorch.save(net.state_dict(), save_path)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:54:13.349693Z","iopub.execute_input":"2021-11-13T08:54:13.350549Z","iopub.status.idle":"2021-11-13T08:54:15.835892Z","shell.execute_reply.started":"2021-11-13T08:54:13.350506Z","shell.execute_reply":"2021-11-13T08:54:15.834708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# PyTorchのネットワークパラメータのロード\nload_path = './weights_fine_tuning.pth'\nload_weights = torch.load(load_path)\nnet.load_state_dict(load_weights)\n\n# GPU上で保存された重みをCPU上でロードする場合\nload_weights = torch.load(load_path, map_location={'cuda:0': 'cpu'})\nnet.load_state_dict(load_weights)","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:54:15.837736Z","iopub.execute_input":"2021-11-13T08:54:15.838531Z","iopub.status.idle":"2021-11-13T08:54:16.687129Z","shell.execute_reply.started":"2021-11-13T08:54:15.838484Z","shell.execute_reply":"2021-11-13T08:54:16.686288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset ,DataLoader\n\nIMG_SIZE=224\nTEST_FILE_PATH = \"../input/cassava-leaf-disease-classification/test_images/\"\n    \nclass CassavaTestDataset(data.Dataset):\n    def __init__(self,path,image_ids,image_size, mode='val'):\n        print(path)\n        print(image_ids)\n        print(image_size)\n        self.image_ids = image_ids\n        self.path = path\n        self.image_size = image_size\n\n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self,item):\n      image_ids = str(self.image_ids[item])\n      img = Image.open(self.path+image_ids)\n    \n      return val_transforms(img)\n        \nsample = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\ntest_dataset = CassavaTestDataset(TEST_FILE_PATH,sample.image_id,sample.label,IMG_SIZE)\ntest_loader = DataLoader(test_dataset,\n                      batch_size=1,\n                      shuffle=False)\n\n# 初期設定\n# GPUが使えるかを確認\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(\"使用デバイス：\", device)\n\n# ネットワークをGPUへ\nnet.to(device)\n    \nfin_outputs = []\n\nfor inputs in test_loader:\n\n    # GPUが使えるならGPUにデータを送る\n    inputs = inputs.to(device)\n    \n    outputs = net(inputs)\n    outputs = nn.Softmax(dim=-1)(outputs)\n    outputs = torch.argmax(outputs,dim=1)\n    fin_outputs.append(outputs.cpu().detach().numpy())\n                \nsample[\"label\"] = np.array(fin_outputs).reshape(-1)\nsample[[\"image_id\",\"label\"]].to_csv(\"submission.csv\",index=False)\nsample.head()","metadata":{"execution":{"iopub.status.busy":"2021-11-13T08:54:16.688636Z","iopub.execute_input":"2021-11-13T08:54:16.689230Z","iopub.status.idle":"2021-11-13T08:54:16.741814Z","shell.execute_reply.started":"2021-11-13T08:54:16.689176Z","shell.execute_reply":"2021-11-13T08:54:16.740913Z"},"trusted":true},"execution_count":null,"outputs":[]}]}