[{"data":1,"prerenderedAt":1863},["ShallowReactive",2],{"content-\u002Fcontents\u002Fpytorch-advanced":3,"surroundPost-\u002Fcontents\u002Fpytorch-advanced":1854},{"id":4,"title":5,"body":6,"createdAt":1840,"description":1841,"draft":1842,"extension":1843,"meta":1844,"navigation":155,"path":1845,"seo":1846,"stem":1847,"tags":1848,"thumbnail":1852,"updatedAt":1840,"__hash__":1853},"contents\u002Fcontents\u002Fpytorch-advanced.md","【PyTorch】モデルの可視化・保存方法について学ぶ",{"type":7,"value":8,"toc":1833},"minimark",[9,13,25,28,39,44,53,56,62,69,93,96,428,431,588,591,594,601,608,611,619,622,630,633,636,639,647,650,653,656,1126,1129,1154,1157,1164,1167,1170,1173,1180,1207,1210,1217,1232,1235,1285,1288,1291,1367,1370,1414,1421,1475,1478,1482,1488,1491,1499,1505,1508,1511,1514,1517,1605,1611,1614,1619,1622,1661,1664,1670,1753,1759,1766,1770,1775,1778,1781,1800,1803,1806,1809,1823,1826,1829],[10,11,12],"p",{},"本記事では、PyTorch でよく使うモデルの可視化や保存方法を紹介します。",[10,14,15,16,20,21,24],{},"また、たまに使うけどよくわからない",[17,18,19],"code",{},"register_buffer","や",[17,22,23],{},"torch.lerp","についても調べてみました。",[10,26,27],{},"本記事では、前回使用した MLP モデルを使っていきます。",[29,30,31],"ul",{},[32,33,34],"li",{},[35,36,38],"a",{"href":37},"pytorch-beginer","【学び直し】Pytorch の基本と MLP で MNIST の分類・可視化の実装まで",[40,41,43],"h2",{"id":42},"torchsummay-でモデルを可視化","torchsummay でモデルを可視化",[10,45,46,52],{},[35,47,51],{"href":48,"rel":49},"https:\u002F\u002Fgithub.com\u002Fsksq96\u002Fpytorch-summary",[50],"nofollow","torchsummary","というモジュールを利用することで、モデルを可視化することができます。",[10,54,55],{},"複雑なモデルを定義していると入力や出力の shape がわからなくなったり、「これメモリに乗るのかな」ということがあります。",[10,57,58,59,61],{},"そういう時にこの",[17,60,51],{},"を利用します。",[10,63,64,65,68],{},"インストールは",[17,66,67],{},"pip","でできます。",[70,71,76],"pre",{"className":72,"code":73,"language":74,"meta":75,"style":75},"language-bash shiki shiki-themes github-dark","pip install torchsummary\n","bash","",[17,77,78],{"__ignoreMap":75},[79,80,83,86,90],"span",{"class":81,"line":82},"line",1,[79,84,67],{"class":85},"svObZ",[79,87,89],{"class":88},"sU2Wk"," install",[79,91,92],{"class":88}," torchsummary\n",[10,94,95],{},"使い方はこんな感じです。前回のコードを流用します。",[70,97,101],{"className":98,"code":99,"language":100,"meta":75,"style":75},"language-py shiki shiki-themes github-dark","import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.tensorboard import SummaryWriter\nfrom torch.utils.data import DataLoader\nimport torchvision.transforms as transforms\nfrom torchvision import datasets, transforms\n\n# 追加============================\nimport os\nfrom torchsummary import summary\n# ===============================\n\nfrom datetime import datetime\n\nprint(torch.__version__) # 1.5.0\n\n# colabでgoogle driveをマウントしてない場合のパス\nroot=\"content\u002F\"\n\n# dataの変換方法を定義\ntrans = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])\n\n# dataをダウンロード\ntrain_set = datasets.MNIST(root=root, train=True, transform=trans, download=True)\ntest_set = datasets.MNIST(root=root, train=False, transform=trans, download=True)\n\n# cpuかgpuか\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n# dataloaderを定義\ntrain_loader = DataLoader(train_set, batch_size=100, shuffle=True)\ntest_loader = DataLoader(test_set, batch_size=100, shuffle=False)\n\n# Networkを定義\nclass MLPNet (nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.fc1 = nn.Linear(1 * 28 * 28, 512)\n        self.fc2 =nn.Linear(512, 512)\n        self.fc3 = nn.Linear(512, 10)\n        self.dropout1=nn.Dropout2d(0.2)\n        self.dropout2=nn.Dropout2d(0.2)\n\n    def forward(self, x):\n        x = F.relu(self.fc1(x))\n        x = self.dropout1(x)\n        x = F.relu(self.fc2(x))\n        x = self.dropout2(x)\n        return F.relu(self.fc3(x))\n\nnet = MLPNet().to(device)\n\n# torchsummaryを使った可視化\nsummary(net, input_size=(1,1 * 28 * 28))\n","py",[17,102,103,108,114,120,126,132,138,144,150,157,163,169,175,181,186,192,197,203,208,214,220,225,231,237,242,248,254,260,265,271,277,282,288,294,300,305,311,317,323,329,335,341,347,353,359,364,370,376,382,388,394,400,405,411,416,422],{"__ignoreMap":75},[79,104,105],{"class":81,"line":82},[79,106,107],{},"import torch\n",[79,109,111],{"class":81,"line":110},2,[79,112,113],{},"import torch.nn as nn\n",[79,115,117],{"class":81,"line":116},3,[79,118,119],{},"import torch.nn.functional as F\n",[79,121,123],{"class":81,"line":122},4,[79,124,125],{},"import torch.optim as optim\n",[79,127,129],{"class":81,"line":128},5,[79,130,131],{},"from torch.utils.tensorboard import SummaryWriter\n",[79,133,135],{"class":81,"line":134},6,[79,136,137],{},"from torch.utils.data import DataLoader\n",[79,139,141],{"class":81,"line":140},7,[79,142,143],{},"import torchvision.transforms as transforms\n",[79,145,147],{"class":81,"line":146},8,[79,148,149],{},"from torchvision import datasets, transforms\n",[79,151,153],{"class":81,"line":152},9,[79,154,156],{"emptyLinePlaceholder":155},true,"\n",[79,158,160],{"class":81,"line":159},10,[79,161,162],{},"# 追加============================\n",[79,164,166],{"class":81,"line":165},11,[79,167,168],{},"import os\n",[79,170,172],{"class":81,"line":171},12,[79,173,174],{},"from torchsummary import summary\n",[79,176,178],{"class":81,"line":177},13,[79,179,180],{},"# ===============================\n",[79,182,184],{"class":81,"line":183},14,[79,185,156],{"emptyLinePlaceholder":155},[79,187,189],{"class":81,"line":188},15,[79,190,191],{},"from datetime import datetime\n",[79,193,195],{"class":81,"line":194},16,[79,196,156],{"emptyLinePlaceholder":155},[79,198,200],{"class":81,"line":199},17,[79,201,202],{},"print(torch.__version__) # 1.5.0\n",[79,204,206],{"class":81,"line":205},18,[79,207,156],{"emptyLinePlaceholder":155},[79,209,211],{"class":81,"line":210},19,[79,212,213],{},"# colabでgoogle driveをマウントしてない場合のパス\n",[79,215,217],{"class":81,"line":216},20,[79,218,219],{},"root=\"content\u002F\"\n",[79,221,223],{"class":81,"line":222},21,[79,224,156],{"emptyLinePlaceholder":155},[79,226,228],{"class":81,"line":227},22,[79,229,230],{},"# dataの変換方法を定義\n",[79,232,234],{"class":81,"line":233},23,[79,235,236],{},"trans = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])\n",[79,238,240],{"class":81,"line":239},24,[79,241,156],{"emptyLinePlaceholder":155},[79,243,245],{"class":81,"line":244},25,[79,246,247],{},"# dataをダウンロード\n",[79,249,251],{"class":81,"line":250},26,[79,252,253],{},"train_set = datasets.MNIST(root=root, train=True, transform=trans, download=True)\n",[79,255,257],{"class":81,"line":256},27,[79,258,259],{},"test_set = datasets.MNIST(root=root, train=False, transform=trans, download=True)\n",[79,261,263],{"class":81,"line":262},28,[79,264,156],{"emptyLinePlaceholder":155},[79,266,268],{"class":81,"line":267},29,[79,269,270],{},"# cpuかgpuか\n",[79,272,274],{"class":81,"line":273},30,[79,275,276],{},"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n",[79,278,280],{"class":81,"line":279},31,[79,281,156],{"emptyLinePlaceholder":155},[79,283,285],{"class":81,"line":284},32,[79,286,287],{},"# dataloaderを定義\n",[79,289,291],{"class":81,"line":290},33,[79,292,293],{},"train_loader = DataLoader(train_set, batch_size=100, shuffle=True)\n",[79,295,297],{"class":81,"line":296},34,[79,298,299],{},"test_loader = DataLoader(test_set, batch_size=100, shuffle=False)\n",[79,301,303],{"class":81,"line":302},35,[79,304,156],{"emptyLinePlaceholder":155},[79,306,308],{"class":81,"line":307},36,[79,309,310],{},"# Networkを定義\n",[79,312,314],{"class":81,"line":313},37,[79,315,316],{},"class MLPNet (nn.Module):\n",[79,318,320],{"class":81,"line":319},38,[79,321,322],{},"    def __init__(self):\n",[79,324,326],{"class":81,"line":325},39,[79,327,328],{},"        super().__init__()\n",[79,330,332],{"class":81,"line":331},40,[79,333,334],{},"        self.fc1 = nn.Linear(1 * 28 * 28, 512)\n",[79,336,338],{"class":81,"line":337},41,[79,339,340],{},"        self.fc2 =nn.Linear(512, 512)\n",[79,342,344],{"class":81,"line":343},42,[79,345,346],{},"        self.fc3 = nn.Linear(512, 10)\n",[79,348,350],{"class":81,"line":349},43,[79,351,352],{},"        self.dropout1=nn.Dropout2d(0.2)\n",[79,354,356],{"class":81,"line":355},44,[79,357,358],{},"        self.dropout2=nn.Dropout2d(0.2)\n",[79,360,362],{"class":81,"line":361},45,[79,363,156],{"emptyLinePlaceholder":155},[79,365,367],{"class":81,"line":366},46,[79,368,369],{},"    def forward(self, x):\n",[79,371,373],{"class":81,"line":372},47,[79,374,375],{},"        x = F.relu(self.fc1(x))\n",[79,377,379],{"class":81,"line":378},48,[79,380,381],{},"        x = self.dropout1(x)\n",[79,383,385],{"class":81,"line":384},49,[79,386,387],{},"        x = F.relu(self.fc2(x))\n",[79,389,391],{"class":81,"line":390},50,[79,392,393],{},"        x = self.dropout2(x)\n",[79,395,397],{"class":81,"line":396},51,[79,398,399],{},"        return F.relu(self.fc3(x))\n",[79,401,403],{"class":81,"line":402},52,[79,404,156],{"emptyLinePlaceholder":155},[79,406,408],{"class":81,"line":407},53,[79,409,410],{},"net = MLPNet().to(device)\n",[79,412,414],{"class":81,"line":413},54,[79,415,156],{"emptyLinePlaceholder":155},[79,417,419],{"class":81,"line":418},55,[79,420,421],{},"# torchsummaryを使った可視化\n",[79,423,425],{"class":81,"line":424},56,[79,426,427],{},"summary(net, input_size=(1,1 * 28 * 28))\n",[10,429,430],{},"出力は以下のようになります。",[70,432,434],{"className":72,"code":433,"language":74,"meta":75,"style":75},"----------------------------------------------------------------\ntitle: 【PyTorch】モデルの可視化・保存方法について学ぶ\ncreatedAt: '2020-06-06'\nupdatedAt: '2020-06-06'\ntags: ['PyTorch', 'Python', '機械学習']\ndraft: false\ndescription:  'PyTorchを使った少々実践的な内容をまとめました。モデルの可視化や保存方法について説明します。また、たまに見かけるtorch.lerpやregister_bufferについてもコード付きで紹介します。'\nthumbnail: '\u002Fimg\u002Ftwitter-card.png'\n---\n# 【PyTorch】モデルの可視化・保存方法について学ぶ\n-------------------------------------------------------------\ntitle: 【PyTorch】モデルの可視化・保存方法について学ぶ\ncreatedAt: '2020-06-06'\nupdatedAt: '2020-06-06'\ntags: ['PyTorch', 'Python', '機械学習']\ndraft: false\ndescription:  'PyTorchを使った少々実践的な内容をまとめました。モデルの可視化や保存方法について説明します。また、たまに見かけるtorch.lerpやregister_bufferについてもコード付きで紹介します。'\nthumbnail: '\u002Fimg\u002Ftwitter-card.png'\n---\n# 【PyTorch】モデルの可視化・保存方法について学ぶ\n-------------------------------------------------------------\n",[17,435,436,441,449,457,464,485,494,502,510,515,521,526,532,538,544,558,564,570,576,580,584],{"__ignoreMap":75},[79,437,438],{"class":81,"line":82},[79,439,440],{"class":85},"----------------------------------------------------------------\n",[79,442,443,446],{"class":81,"line":110},[79,444,445],{"class":85},"title:",[79,447,448],{"class":88}," 【PyTorch】モデルの可視化・保存方法について学ぶ\n",[79,450,451,454],{"class":81,"line":116},[79,452,453],{"class":85},"createdAt:",[79,455,456],{"class":88}," '2020-06-06'\n",[79,458,459,462],{"class":81,"line":122},[79,460,461],{"class":85},"updatedAt:",[79,463,456],{"class":88},[79,465,466,469,473,476,479,482],{"class":81,"line":128},[79,467,468],{"class":85},"tags:",[79,470,472],{"class":471},"s95oV"," [",[79,474,475],{"class":88},"'PyTorch'",[79,477,478],{"class":471},", ",[79,480,481],{"class":88},"'Python',",[79,483,484],{"class":88}," '機械学習']\n",[79,486,487,490],{"class":81,"line":134},[79,488,489],{"class":85},"draft:",[79,491,493],{"class":492},"sDLfK"," false\n",[79,495,496,499],{"class":81,"line":140},[79,497,498],{"class":85},"description:",[79,500,501],{"class":88},"  'PyTorchを使った少々実践的な内容をまとめました。モデルの可視化や保存方法について説明します。また、たまに見かけるtorch.lerpやregister_bufferについてもコード付きで紹介します。'\n",[79,503,504,507],{"class":81,"line":146},[79,505,506],{"class":85},"thumbnail:",[79,508,509],{"class":88}," '\u002Fimg\u002Ftwitter-card.png'\n",[79,511,512],{"class":81,"line":152},[79,513,514],{"class":85},"---\n",[79,516,517],{"class":81,"line":159},[79,518,520],{"class":519},"sAwPA","# 【PyTorch】モデルの可視化・保存方法について学ぶ\n",[79,522,523],{"class":81,"line":165},[79,524,525],{"class":85},"-------------------------------------------------------------\n",[79,527,528,530],{"class":81,"line":171},[79,529,445],{"class":85},[79,531,448],{"class":88},[79,533,534,536],{"class":81,"line":177},[79,535,453],{"class":85},[79,537,456],{"class":88},[79,539,540,542],{"class":81,"line":183},[79,541,461],{"class":85},[79,543,456],{"class":88},[79,545,546,548,550,552,554,556],{"class":81,"line":188},[79,547,468],{"class":85},[79,549,472],{"class":471},[79,551,475],{"class":88},[79,553,478],{"class":471},[79,555,481],{"class":88},[79,557,484],{"class":88},[79,559,560,562],{"class":81,"line":194},[79,561,489],{"class":85},[79,563,493],{"class":492},[79,565,566,568],{"class":81,"line":199},[79,567,498],{"class":85},[79,569,501],{"class":88},[79,571,572,574],{"class":81,"line":205},[79,573,506],{"class":85},[79,575,509],{"class":88},[79,577,578],{"class":81,"line":210},[79,579,514],{"class":85},[79,581,582],{"class":81,"line":216},[79,583,520],{"class":519},[79,585,586],{"class":81,"line":222},[79,587,525],{"class":85},[10,589,590],{},"非常にわかりやすいです。",[10,592,593],{},"特に他人にモデルの説明をするときにあると重宝します。",[10,595,596,597,600],{},"注意としては、今回のモデルのように入力が 1 次元の場合はそのまま入力サイズに",[17,598,599],{},"input_size=(1*28*28)","とするとエラーになります。",[10,602,603,604,607],{},"なので、チャネルの次元を加えて",[17,605,606],{},"input_size=(1, 1*28*28)","とします。",[40,609,610],{"id":610},"学習済みモデルの保存",[10,612,613,618],{},[35,614,617],{"href":615,"rel":616},"https:\u002F\u002Fpytorch.org\u002Ftutorials\u002Frecipes\u002Frecipes\u002Fsaving_and_loading_models_for_inference.html",[50],"公式","に詳しく書いてありますが念のため。",[10,620,621],{},"まずモデルの保存を行う目的は２つあります。",[29,623,624,627],{},[32,625,626],{},"学習済みモデルを使って推論を行う",[32,628,629],{},"保存済みモデルの学習を再開する",[10,631,632],{},"目的によって、保存しておくべき内容が違います。",[10,634,635],{},"次に PyTorch でモデルを保存する方法について確認してきます。",[10,637,638],{},"PyTorch ではモデルを保存する方法が 2 通りあります。",[29,640,641,644],{},[32,642,643],{},"モデル全体を保存する",[32,645,646],{},"モデルのパラメータを保存する",[10,648,649],{},"さらに保存する際には GPU か CPU なのかを注意する必要があります。",[10,651,652],{},"ややこしいですが、認識しておく必要があります。",[10,654,655],{},"まずは普通にモデルを保存してみます。",[70,657,659],{"className":98,"code":658,"language":100,"meta":75,"style":75},"net.apply(init_weights)　# 追加\n\n# loss関数\ncriterion = nn.CrossEntropyLoss()\n\n# 最適化方法\noptimizer = optim.SGD(net.parameters(), lr=0.01, momentum=0.9)\n\n# log用フォルダを毎回生成\n# tensorboardの可視化用\nnow = datetime.now()\nlog_path = \".\u002Fruns\u002F\" + now.strftime(\"%Y%m%d-%H%M%S\") + \"\u002F\"\nprint(log_path)\n\n# tensorboard用のwriter\nwriter = SummaryWriter(log_path)\n\nepochs = 30\n\nfor epoch in range(epochs):\n    train_loss = 0\n    train_acc = 0\n    val_loss = 0\n    val_acc = 0\n\n    # train dataで訓練\n    net.train()\n    for i, (images, labels) in enumerate(train_loader):\n\n        images, labels = images.view(-1, 28*28*1).to(device), labels.to(device)\n\n        # 勾配を０にリセット\n        optimizer.zero_grad()\n\n        # 順伝搬\n        out = net(images)\n\n        # loss計算\n        loss = criterion(out, labels)\n\n        # 計算したlossとaccの値を入れる\n        train_loss += loss.item()\n        train_acc += (out.max(1)[1] == labels).sum().item()\n\n        # 誤差逆伝搬\n        loss.backward()\n\n        # 重みの更新\n        optimizer.step()\n\n        # 平均のlossとacc計算\n        avg_train_loss = train_loss \u002F len(train_loader.dataset)\n        avg_train_acc = train_acc \u002F len(train_loader.dataset)\n\n    # validation dataで評価\n    net.eval()\n\n    with torch.no_grad():\n        for (images, labels) in test_loader:\n            images, labels = images.view(-1, 28*28*1).to(device), labels.to(device)\n            out = net(images)\n            loss = criterion(out, labels)\n            val_loss += loss.item()\n            acc = (out.max(1)[1] == labels).sum()\n            val_acc += acc.item()\n    avg_val_loss = val_loss \u002F len(test_loader.dataset)\n    avg_val_acc = val_acc \u002F len(test_loader.dataset)\n\n    # print log\n    print ('Epoch [{}\u002F{}], Loss: {loss:.4f}, val_loss: {val_loss:.4f}, val_acc: {val_acc:.4f}'\n                   .format(epoch+1, epochs, loss=avg_train_loss, val_loss=avg_val_loss, val_acc=avg_val_acc))\n\n    # tensorboard用\n    writer.add_scalars('loss', {'train_loss':avg_train_loss, 'val_loss':avg_val_loss},epoch+1)\n    writer.add_scalars('accuracy', {'train_acc':avg_train_acc, 'val_acc':avg_val_acc}, epoch+1)\n\nwriter.close()\n\n# 追加部分\ndir_name = 'output'\n\nif not os.path.exists(dir_name):\n    os.mkdir(dir_name)\n\nmodel_save_path = os.path.join(dir_name, \"model_full.pt\")\n\n# モデル保存\ntorch.save(net, model_save_path)\n\n# モデルロード\nmodel_full = torch.load(model_save_path)\n",[17,660,661,666,670,675,680,684,689,694,698,703,708,713,718,723,727,732,737,741,746,750,755,760,765,770,775,779,784,789,794,798,803,807,812,817,821,826,831,835,840,845,849,854,859,864,868,873,878,882,887,892,896,901,906,911,915,920,925,930,936,942,948,954,960,966,972,978,984,990,995,1001,1007,1013,1018,1024,1030,1036,1041,1047,1052,1058,1064,1069,1075,1081,1086,1092,1097,1103,1109,1114,1120],{"__ignoreMap":75},[79,662,663],{"class":81,"line":82},[79,664,665],{},"net.apply(init_weights)　# 追加\n",[79,667,668],{"class":81,"line":110},[79,669,156],{"emptyLinePlaceholder":155},[79,671,672],{"class":81,"line":116},[79,673,674],{},"# loss関数\n",[79,676,677],{"class":81,"line":122},[79,678,679],{},"criterion = nn.CrossEntropyLoss()\n",[79,681,682],{"class":81,"line":128},[79,683,156],{"emptyLinePlaceholder":155},[79,685,686],{"class":81,"line":134},[79,687,688],{},"# 最適化方法\n",[79,690,691],{"class":81,"line":140},[79,692,693],{},"optimizer = optim.SGD(net.parameters(), lr=0.01, momentum=0.9)\n",[79,695,696],{"class":81,"line":146},[79,697,156],{"emptyLinePlaceholder":155},[79,699,700],{"class":81,"line":152},[79,701,702],{},"# log用フォルダを毎回生成\n",[79,704,705],{"class":81,"line":159},[79,706,707],{},"# tensorboardの可視化用\n",[79,709,710],{"class":81,"line":165},[79,711,712],{},"now = datetime.now()\n",[79,714,715],{"class":81,"line":171},[79,716,717],{},"log_path = \".\u002Fruns\u002F\" + now.strftime(\"%Y%m%d-%H%M%S\") + \"\u002F\"\n",[79,719,720],{"class":81,"line":177},[79,721,722],{},"print(log_path)\n",[79,724,725],{"class":81,"line":183},[79,726,156],{"emptyLinePlaceholder":155},[79,728,729],{"class":81,"line":188},[79,730,731],{},"# tensorboard用のwriter\n",[79,733,734],{"class":81,"line":194},[79,735,736],{},"writer = SummaryWriter(log_path)\n",[79,738,739],{"class":81,"line":199},[79,740,156],{"emptyLinePlaceholder":155},[79,742,743],{"class":81,"line":205},[79,744,745],{},"epochs = 30\n",[79,747,748],{"class":81,"line":210},[79,749,156],{"emptyLinePlaceholder":155},[79,751,752],{"class":81,"line":216},[79,753,754],{},"for epoch in range(epochs):\n",[79,756,757],{"class":81,"line":222},[79,758,759],{},"    train_loss = 0\n",[79,761,762],{"class":81,"line":227},[79,763,764],{},"    train_acc = 0\n",[79,766,767],{"class":81,"line":233},[79,768,769],{},"    val_loss = 0\n",[79,771,772],{"class":81,"line":239},[79,773,774],{},"    val_acc = 0\n",[79,776,777],{"class":81,"line":244},[79,778,156],{"emptyLinePlaceholder":155},[79,780,781],{"class":81,"line":250},[79,782,783],{},"    # train dataで訓練\n",[79,785,786],{"class":81,"line":256},[79,787,788],{},"    net.train()\n",[79,790,791],{"class":81,"line":262},[79,792,793],{},"    for i, (images, labels) in enumerate(train_loader):\n",[79,795,796],{"class":81,"line":267},[79,797,156],{"emptyLinePlaceholder":155},[79,799,800],{"class":81,"line":273},[79,801,802],{},"        images, labels = images.view(-1, 28*28*1).to(device), labels.to(device)\n",[79,804,805],{"class":81,"line":279},[79,806,156],{"emptyLinePlaceholder":155},[79,808,809],{"class":81,"line":284},[79,810,811],{},"        # 勾配を０にリセット\n",[79,813,814],{"class":81,"line":290},[79,815,816],{},"        optimizer.zero_grad()\n",[79,818,819],{"class":81,"line":296},[79,820,156],{"emptyLinePlaceholder":155},[79,822,823],{"class":81,"line":302},[79,824,825],{},"        # 順伝搬\n",[79,827,828],{"class":81,"line":307},[79,829,830],{},"        out = net(images)\n",[79,832,833],{"class":81,"line":313},[79,834,156],{"emptyLinePlaceholder":155},[79,836,837],{"class":81,"line":319},[79,838,839],{},"        # loss計算\n",[79,841,842],{"class":81,"line":325},[79,843,844],{},"        loss = criterion(out, labels)\n",[79,846,847],{"class":81,"line":331},[79,848,156],{"emptyLinePlaceholder":155},[79,850,851],{"class":81,"line":337},[79,852,853],{},"        # 計算したlossとaccの値を入れる\n",[79,855,856],{"class":81,"line":343},[79,857,858],{},"        train_loss += loss.item()\n",[79,860,861],{"class":81,"line":349},[79,862,863],{},"        train_acc += (out.max(1)[1] == labels).sum().item()\n",[79,865,866],{"class":81,"line":355},[79,867,156],{"emptyLinePlaceholder":155},[79,869,870],{"class":81,"line":361},[79,871,872],{},"        # 誤差逆伝搬\n",[79,874,875],{"class":81,"line":366},[79,876,877],{},"        loss.backward()\n",[79,879,880],{"class":81,"line":372},[79,881,156],{"emptyLinePlaceholder":155},[79,883,884],{"class":81,"line":378},[79,885,886],{},"        # 重みの更新\n",[79,888,889],{"class":81,"line":384},[79,890,891],{},"        optimizer.step()\n",[79,893,894],{"class":81,"line":390},[79,895,156],{"emptyLinePlaceholder":155},[79,897,898],{"class":81,"line":396},[79,899,900],{},"        # 平均のlossとacc計算\n",[79,902,903],{"class":81,"line":402},[79,904,905],{},"        avg_train_loss = train_loss \u002F len(train_loader.dataset)\n",[79,907,908],{"class":81,"line":407},[79,909,910],{},"        avg_train_acc = train_acc \u002F len(train_loader.dataset)\n",[79,912,913],{"class":81,"line":413},[79,914,156],{"emptyLinePlaceholder":155},[79,916,917],{"class":81,"line":418},[79,918,919],{},"    # validation dataで評価\n",[79,921,922],{"class":81,"line":424},[79,923,924],{},"    net.eval()\n",[79,926,928],{"class":81,"line":927},57,[79,929,156],{"emptyLinePlaceholder":155},[79,931,933],{"class":81,"line":932},58,[79,934,935],{},"    with torch.no_grad():\n",[79,937,939],{"class":81,"line":938},59,[79,940,941],{},"        for (images, labels) in test_loader:\n",[79,943,945],{"class":81,"line":944},60,[79,946,947],{},"            images, labels = images.view(-1, 28*28*1).to(device), labels.to(device)\n",[79,949,951],{"class":81,"line":950},61,[79,952,953],{},"            out = net(images)\n",[79,955,957],{"class":81,"line":956},62,[79,958,959],{},"            loss = criterion(out, labels)\n",[79,961,963],{"class":81,"line":962},63,[79,964,965],{},"            val_loss += loss.item()\n",[79,967,969],{"class":81,"line":968},64,[79,970,971],{},"            acc = (out.max(1)[1] == labels).sum()\n",[79,973,975],{"class":81,"line":974},65,[79,976,977],{},"            val_acc += acc.item()\n",[79,979,981],{"class":81,"line":980},66,[79,982,983],{},"    avg_val_loss = val_loss \u002F len(test_loader.dataset)\n",[79,985,987],{"class":81,"line":986},67,[79,988,989],{},"    avg_val_acc = val_acc \u002F len(test_loader.dataset)\n",[79,991,993],{"class":81,"line":992},68,[79,994,156],{"emptyLinePlaceholder":155},[79,996,998],{"class":81,"line":997},69,[79,999,1000],{},"    # print log\n",[79,1002,1004],{"class":81,"line":1003},70,[79,1005,1006],{},"    print ('Epoch [{}\u002F{}], Loss: {loss:.4f}, val_loss: {val_loss:.4f}, val_acc: {val_acc:.4f}'\n",[79,1008,1010],{"class":81,"line":1009},71,[79,1011,1012],{},"                   .format(epoch+1, epochs, loss=avg_train_loss, val_loss=avg_val_loss, val_acc=avg_val_acc))\n",[79,1014,1016],{"class":81,"line":1015},72,[79,1017,156],{"emptyLinePlaceholder":155},[79,1019,1021],{"class":81,"line":1020},73,[79,1022,1023],{},"    # tensorboard用\n",[79,1025,1027],{"class":81,"line":1026},74,[79,1028,1029],{},"    writer.add_scalars('loss', {'train_loss':avg_train_loss, 'val_loss':avg_val_loss},epoch+1)\n",[79,1031,1033],{"class":81,"line":1032},75,[79,1034,1035],{},"    writer.add_scalars('accuracy', {'train_acc':avg_train_acc, 'val_acc':avg_val_acc}, epoch+1)\n",[79,1037,1039],{"class":81,"line":1038},76,[79,1040,156],{"emptyLinePlaceholder":155},[79,1042,1044],{"class":81,"line":1043},77,[79,1045,1046],{},"writer.close()\n",[79,1048,1050],{"class":81,"line":1049},78,[79,1051,156],{"emptyLinePlaceholder":155},[79,1053,1055],{"class":81,"line":1054},79,[79,1056,1057],{},"# 追加部分\n",[79,1059,1061],{"class":81,"line":1060},80,[79,1062,1063],{},"dir_name = 'output'\n",[79,1065,1067],{"class":81,"line":1066},81,[79,1068,156],{"emptyLinePlaceholder":155},[79,1070,1072],{"class":81,"line":1071},82,[79,1073,1074],{},"if not os.path.exists(dir_name):\n",[79,1076,1078],{"class":81,"line":1077},83,[79,1079,1080],{},"    os.mkdir(dir_name)\n",[79,1082,1084],{"class":81,"line":1083},84,[79,1085,156],{"emptyLinePlaceholder":155},[79,1087,1089],{"class":81,"line":1088},85,[79,1090,1091],{},"model_save_path = os.path.join(dir_name, \"model_full.pt\")\n",[79,1093,1095],{"class":81,"line":1094},86,[79,1096,156],{"emptyLinePlaceholder":155},[79,1098,1100],{"class":81,"line":1099},87,[79,1101,1102],{},"# モデル保存\n",[79,1104,1106],{"class":81,"line":1105},88,[79,1107,1108],{},"torch.save(net, model_save_path)\n",[79,1110,1112],{"class":81,"line":1111},89,[79,1113,156],{"emptyLinePlaceholder":155},[79,1115,1117],{"class":81,"line":1116},90,[79,1118,1119],{},"# モデルロード\n",[79,1121,1123],{"class":81,"line":1122},91,[79,1124,1125],{},"model_full = torch.load(model_save_path)\n",[10,1127,1128],{},"以下コードの部分が保存用のコードです。",[70,1130,1132],{"className":98,"code":1131,"language":100,"meta":75,"style":75},"# モデル保存\ntorch.save(net, model_save_path)\n\n# モデルロード\nmodel_full = torch.load(model_save_path)\n",[17,1133,1134,1138,1142,1146,1150],{"__ignoreMap":75},[79,1135,1136],{"class":81,"line":82},[79,1137,1102],{},[79,1139,1140],{"class":81,"line":110},[79,1141,1108],{},[79,1143,1144],{"class":81,"line":116},[79,1145,156],{"emptyLinePlaceholder":155},[79,1147,1148],{"class":81,"line":122},[79,1149,1119],{},[79,1151,1152],{"class":81,"line":128},[79,1153,1125],{},[10,1155,1156],{},"これが一番単純な方法です。",[10,1158,1159,1160],{},"しかし、",[1161,1162,1163],"strong",{},"この保存方法は公式で推奨されてません。",[10,1165,1166],{},"非推奨の理由をいろいろ調べてみると、どうもこの方法でやると保存時の GPU にロード時も読み込まれてしまうらしい。",[10,1168,1169],{},"つまり GPU がない場合は詰んでしまう可能性がある。",[10,1171,1172],{},"あとはもう一つの方法に比べてサイズが大きい。",[10,1174,1175,1176,1179],{},"なので保存時は公式推奨の",[17,1177,1178],{},"state_dict()","の方法で行う。",[70,1181,1183],{"className":98,"code":1182,"language":100,"meta":75,"style":75},"# モデル保存\ntorch.save(net.state_dict(), model_save_path)\n\n# モデルロード\nmodel.load_state_dict(torch.load(model_save_path))\n",[17,1184,1185,1189,1194,1198,1202],{"__ignoreMap":75},[79,1186,1187],{"class":81,"line":82},[79,1188,1102],{},[79,1190,1191],{"class":81,"line":110},[79,1192,1193],{},"torch.save(net.state_dict(), model_save_path)\n",[79,1195,1196],{"class":81,"line":116},[79,1197,156],{"emptyLinePlaceholder":155},[79,1199,1200],{"class":81,"line":122},[79,1201,1119],{},[79,1203,1204],{"class":81,"line":128},[79,1205,1206],{},"model.load_state_dict(torch.load(model_save_path))\n",[10,1208,1209],{},"一応、GPU で保存してしまっても CPU で読み出す方法はあるらしいが、失敗するのが怖いので CPU で保存しておくのが無難。",[10,1211,1212,1213,1216],{},"やり方は以下のように保存時に",[17,1214,1215],{},"to('cpu)","をつける。",[70,1218,1220],{"className":98,"code":1219,"language":100,"meta":75,"style":75},"torch.save(net.to('cpu').state_dict(), model_save_path)\nmodel_cpu.load_state_dict(torch.load(model_save_path))\n",[17,1221,1222,1227],{"__ignoreMap":75},[79,1223,1224],{"class":81,"line":82},[79,1225,1226],{},"torch.save(net.to('cpu').state_dict(), model_save_path)\n",[79,1228,1229],{"class":81,"line":110},[79,1230,1231],{},"model_cpu.load_state_dict(torch.load(model_save_path))\n",[10,1233,1234],{},"学習を再開するための checkpoint を作りたい場合は以下のようにします。",[70,1236,1238],{"className":98,"code":1237,"language":100,"meta":75,"style":75},"if epoch % 3 == 0:\n        file_name = 'epoch_{}.pt'.format(epoch)\n        path = os.path.join(checkPoint_dir, file_name)\n        torch.save({\n            'epoch' : epoch,\n            'model_state_dict' : net.state_dict(),\n            'optimaizer_state_dict': optimizer.state_dict(),\n            'loss': avg_train_loss\n        }, path)\n",[17,1239,1240,1245,1250,1255,1260,1265,1270,1275,1280],{"__ignoreMap":75},[79,1241,1242],{"class":81,"line":82},[79,1243,1244],{},"if epoch % 3 == 0:\n",[79,1246,1247],{"class":81,"line":110},[79,1248,1249],{},"        file_name = 'epoch_{}.pt'.format(epoch)\n",[79,1251,1252],{"class":81,"line":116},[79,1253,1254],{},"        path = os.path.join(checkPoint_dir, file_name)\n",[79,1256,1257],{"class":81,"line":122},[79,1258,1259],{},"        torch.save({\n",[79,1261,1262],{"class":81,"line":128},[79,1263,1264],{},"            'epoch' : epoch,\n",[79,1266,1267],{"class":81,"line":134},[79,1268,1269],{},"            'model_state_dict' : net.state_dict(),\n",[79,1271,1272],{"class":81,"line":140},[79,1273,1274],{},"            'optimaizer_state_dict': optimizer.state_dict(),\n",[79,1276,1277],{"class":81,"line":146},[79,1278,1279],{},"            'loss': avg_train_loss\n",[79,1281,1282],{"class":81,"line":152},[79,1283,1284],{},"        }, path)\n",[10,1286,1287],{},"保存するタイミングは適当に決めます。",[10,1289,1290],{},"こちらが素直に学習した場合の出力です。",[70,1292,1294],{"className":72,"code":1293,"language":74,"meta":75,"style":75},"Epoch [0\u002F10], Loss: 0.0060, val_loss: 0.0019, val_acc: 0.9452\nEpoch [1\u002F10], Loss: 0.0020, val_loss: 0.0014, val_acc: 0.9600\nEpoch [2\u002F10], Loss: 0.0015, val_loss: 0.0012, val_acc: 0.9645\nEpoch [3\u002F10], Loss: 0.0012, val_loss: 0.0009, val_acc: 0.9715\nEpoch [4\u002F10], Loss: 0.0010, val_loss: 0.0009, val_acc: 0.9738\nEpoch [5\u002F10], Loss: 0.0009, val_loss: 0.0008, val_acc: 0.9743\nEpoch [6\u002F10], Loss: 0.0008, val_loss: 0.0008, val_acc: 0.9754\nEpoch [7\u002F10], Loss: 0.0007, val_loss: 0.0008, val_acc: 0.9762\nEpoch [8\u002F10], Loss: 0.0007, val_loss: 0.0007, val_acc: 0.9773\nEpoch [9\u002F10], Loss: 0.0006, val_loss: 0.0007, val_acc: 0.9775\n",[17,1295,1296,1304,1311,1318,1325,1332,1339,1346,1353,1360],{"__ignoreMap":75},[79,1297,1298,1301],{"class":81,"line":82},[79,1299,1300],{"class":85},"Epoch",[79,1302,1303],{"class":471}," [0\u002F10], Loss: 0.0060, val_loss: 0.0019, val_acc: 0.9452\n",[79,1305,1306,1308],{"class":81,"line":110},[79,1307,1300],{"class":85},[79,1309,1310],{"class":471}," [1\u002F10], Loss: 0.0020, val_loss: 0.0014, val_acc: 0.9600\n",[79,1312,1313,1315],{"class":81,"line":116},[79,1314,1300],{"class":85},[79,1316,1317],{"class":471}," [2\u002F10], Loss: 0.0015, val_loss: 0.0012, val_acc: 0.9645\n",[79,1319,1320,1322],{"class":81,"line":122},[79,1321,1300],{"class":85},[79,1323,1324],{"class":471}," [3\u002F10], Loss: 0.0012, val_loss: 0.0009, val_acc: 0.9715\n",[79,1326,1327,1329],{"class":81,"line":128},[79,1328,1300],{"class":85},[79,1330,1331],{"class":471}," [4\u002F10], Loss: 0.0010, val_loss: 0.0009, val_acc: 0.9738\n",[79,1333,1334,1336],{"class":81,"line":134},[79,1335,1300],{"class":85},[79,1337,1338],{"class":471}," [5\u002F10], Loss: 0.0009, val_loss: 0.0008, val_acc: 0.9743\n",[79,1340,1341,1343],{"class":81,"line":140},[79,1342,1300],{"class":85},[79,1344,1345],{"class":471}," [6\u002F10], Loss: 0.0008, val_loss: 0.0008, val_acc: 0.9754\n",[79,1347,1348,1350],{"class":81,"line":146},[79,1349,1300],{"class":85},[79,1351,1352],{"class":471}," [7\u002F10], Loss: 0.0007, val_loss: 0.0008, val_acc: 0.9762\n",[79,1354,1355,1357],{"class":81,"line":152},[79,1356,1300],{"class":85},[79,1358,1359],{"class":471}," [8\u002F10], Loss: 0.0007, val_loss: 0.0007, val_acc: 0.9773\n",[79,1361,1362,1364],{"class":81,"line":159},[79,1363,1300],{"class":85},[79,1365,1366],{"class":471}," [9\u002F10], Loss: 0.0006, val_loss: 0.0007, val_acc: 0.9775\n",[10,1368,1369],{},"ロードは以下のようにします。",[70,1371,1373],{"className":98,"code":1372,"language":100,"meta":75,"style":75},"tmp_path = 'checkPoint\u002Fepoch_3.pt'\n\nif os.path.exists(tmp_path):\n    checkpoint = torch.load(tmp_path)\n    net.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimaizer_state_dict'])\n    epoch_num = checkpoint['epoch']\n    loss = checkpoint['loss']\n",[17,1374,1375,1380,1384,1389,1394,1399,1404,1409],{"__ignoreMap":75},[79,1376,1377],{"class":81,"line":82},[79,1378,1379],{},"tmp_path = 'checkPoint\u002Fepoch_3.pt'\n",[79,1381,1382],{"class":81,"line":110},[79,1383,156],{"emptyLinePlaceholder":155},[79,1385,1386],{"class":81,"line":116},[79,1387,1388],{},"if os.path.exists(tmp_path):\n",[79,1390,1391],{"class":81,"line":122},[79,1392,1393],{},"    checkpoint = torch.load(tmp_path)\n",[79,1395,1396],{"class":81,"line":128},[79,1397,1398],{},"    net.load_state_dict(checkpoint['model_state_dict'])\n",[79,1400,1401],{"class":81,"line":134},[79,1402,1403],{},"    optimizer.load_state_dict(checkpoint['optimaizer_state_dict'])\n",[79,1405,1406],{"class":81,"line":140},[79,1407,1408],{},"    epoch_num = checkpoint['epoch']\n",[79,1410,1411],{"class":81,"line":146},[79,1412,1413],{},"    loss = checkpoint['loss']\n",[10,1415,1416,1417,1420],{},"そしてこちらが",[17,1418,1419],{},"epoch=3","の時の checkpoint をロードした時の結果です。",[70,1422,1424],{"className":72,"code":1423,"language":74,"meta":75,"style":75},"Epoch [3\u002F10], Loss: 0.0010, val_loss: 0.0010, val_acc: 0.9699\nEpoch [4\u002F10], Loss: 0.0009, val_loss: 0.0009, val_acc: 0.9728\nEpoch [5\u002F10], Loss: 0.0008, val_loss: 0.0008, val_acc: 0.9750\nEpoch [6\u002F10], Loss: 0.0007, val_loss: 0.0007, val_acc: 0.9771\nEpoch [7\u002F10], Loss: 0.0006, val_loss: 0.0007, val_acc: 0.9783\nEpoch [8\u002F10], Loss: 0.0006, val_loss: 0.0007, val_acc: 0.9791\nEpoch [9\u002F10], Loss: 0.0005, val_loss: 0.0006, val_acc: 0.9798\n",[17,1425,1426,1433,1440,1447,1454,1461,1468],{"__ignoreMap":75},[79,1427,1428,1430],{"class":81,"line":82},[79,1429,1300],{"class":85},[79,1431,1432],{"class":471}," [3\u002F10], Loss: 0.0010, val_loss: 0.0010, val_acc: 0.9699\n",[79,1434,1435,1437],{"class":81,"line":110},[79,1436,1300],{"class":85},[79,1438,1439],{"class":471}," [4\u002F10], Loss: 0.0009, val_loss: 0.0009, val_acc: 0.9728\n",[79,1441,1442,1444],{"class":81,"line":116},[79,1443,1300],{"class":85},[79,1445,1446],{"class":471}," [5\u002F10], Loss: 0.0008, val_loss: 0.0008, val_acc: 0.9750\n",[79,1448,1449,1451],{"class":81,"line":122},[79,1450,1300],{"class":85},[79,1452,1453],{"class":471}," [6\u002F10], Loss: 0.0007, val_loss: 0.0007, val_acc: 0.9771\n",[79,1455,1456,1458],{"class":81,"line":128},[79,1457,1300],{"class":85},[79,1459,1460],{"class":471}," [7\u002F10], Loss: 0.0006, val_loss: 0.0007, val_acc: 0.9783\n",[79,1462,1463,1465],{"class":81,"line":134},[79,1464,1300],{"class":85},[79,1466,1467],{"class":471}," [8\u002F10], Loss: 0.0006, val_loss: 0.0007, val_acc: 0.9791\n",[79,1469,1470,1472],{"class":81,"line":140},[79,1471,1300],{"class":85},[79,1473,1474],{"class":471}," [9\u002F10], Loss: 0.0005, val_loss: 0.0006, val_acc: 0.9798\n",[10,1476,1477],{},"乱数固定しなかったので微妙にずれてますが、学習が再開されていることがわかります。",[40,1479,1481],{"id":1480},"register_buffer-とは","register_buffer とは",[10,1483,1484,1485,1487],{},"次は",[17,1486,19],{},"です。",[10,1489,1490],{},"たまに論文実装のコードをみるとモデルに書いてあります。",[10,1492,1493,1498],{},[35,1494,1497],{"href":1495,"rel":1496},"https:\u002F\u002Fpytorch.org\u002Fdocs\u002Fstable\u002Fnn.html",[50],"公式の説明","によると",[1500,1501,1502],"blockquote",{},[10,1503,1504],{},"This is typically used to register a buffer that should not to be considered a model parameter.",[10,1506,1507],{},"とあります。",[10,1509,1510],{},"model のパラメーター ではないけどモデルに持っておきたい値を保存する際に使うようです。",[10,1512,1513],{},"利用シーンとしては Batchnormalization の計算のためのバッチごとの計算結果を保持するのに使われます。",[10,1515,1516],{},"ネットワークに少し追加をして実験しました。",[70,1518,1520],{"className":98,"code":1519,"language":100,"meta":75,"style":75},"# Networkを定義\nclass MLPNet (nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.fc1 = nn.Linear(1 * 28 * 28, 512)\n        self.fc2 =nn.Linear(512, 512)\n        self.fc3 = nn.Linear(512, 10)\n        self.dropout1=nn.Dropout2d(0.2)\n        self.dropout2=nn.Dropout2d(0.2)\n\n        # 追加部分\n        self.mean_val = 0 # 比較用\n        self.register_buffer('count', torch.ones(2,2))\n\n    def forward(self, x):\n        x = F.relu(self.fc1(x))\n        x = self.dropout1(x)\n        x = F.relu(self.fc2(x))\n        x = self.dropout2(x)\n        return F.relu(self.fc3(x))\n",[17,1521,1522,1526,1530,1534,1538,1542,1546,1550,1554,1558,1562,1567,1572,1577,1581,1585,1589,1593,1597,1601],{"__ignoreMap":75},[79,1523,1524],{"class":81,"line":82},[79,1525,310],{},[79,1527,1528],{"class":81,"line":110},[79,1529,316],{},[79,1531,1532],{"class":81,"line":116},[79,1533,322],{},[79,1535,1536],{"class":81,"line":122},[79,1537,328],{},[79,1539,1540],{"class":81,"line":128},[79,1541,334],{},[79,1543,1544],{"class":81,"line":134},[79,1545,340],{},[79,1547,1548],{"class":81,"line":140},[79,1549,346],{},[79,1551,1552],{"class":81,"line":146},[79,1553,352],{},[79,1555,1556],{"class":81,"line":152},[79,1557,358],{},[79,1559,1560],{"class":81,"line":159},[79,1561,156],{"emptyLinePlaceholder":155},[79,1563,1564],{"class":81,"line":165},[79,1565,1566],{},"        # 追加部分\n",[79,1568,1569],{"class":81,"line":171},[79,1570,1571],{},"        self.mean_val = 0 # 比較用\n",[79,1573,1574],{"class":81,"line":177},[79,1575,1576],{},"        self.register_buffer('count', torch.ones(2,2))\n",[79,1578,1579],{"class":81,"line":183},[79,1580,156],{"emptyLinePlaceholder":155},[79,1582,1583],{"class":81,"line":188},[79,1584,369],{},[79,1586,1587],{"class":81,"line":194},[79,1588,375],{},[79,1590,1591],{"class":81,"line":199},[79,1592,381],{},[79,1594,1595],{"class":81,"line":205},[79,1596,387],{},[79,1598,1599],{"class":81,"line":210},[79,1600,393],{},[79,1602,1603],{"class":81,"line":216},[79,1604,399],{},[10,1606,1607,1608,1610],{},"普通にクラス変数を定義した場合と",[17,1609,19],{},"の場合を書いてみました。",[10,1612,1613],{},"学習中にこの 2 つをインクリメントして、保存後に値をみるという検証です。",[10,1615,1616,1618],{},[17,1617,19],{},"を使うとパラメータ同様に保存されることを確かめます。",[10,1620,1621],{},"学習後にモデルの中身をみると",[70,1623,1625],{"className":98,"code":1624,"language":100,"meta":75,"style":75},"print(net.mean_val)\nprint(net.count)\n\n# >>>\n# 6000\n# tensor([[6001., 6001.],\n#       [6001., 6001.]], device='cuda:0')\n",[17,1626,1627,1632,1637,1641,1646,1651,1656],{"__ignoreMap":75},[79,1628,1629],{"class":81,"line":82},[79,1630,1631],{},"print(net.mean_val)\n",[79,1633,1634],{"class":81,"line":110},[79,1635,1636],{},"print(net.count)\n",[79,1638,1639],{"class":81,"line":116},[79,1640,156],{"emptyLinePlaceholder":155},[79,1642,1643],{"class":81,"line":122},[79,1644,1645],{},"# >>>\n",[79,1647,1648],{"class":81,"line":128},[79,1649,1650],{},"# 6000\n",[79,1652,1653],{"class":81,"line":134},[79,1654,1655],{},"# tensor([[6001., 6001.],\n",[79,1657,1658],{"class":81,"line":140},[79,1659,1660],{},"#       [6001., 6001.]], device='cuda:0')\n",[10,1662,1663],{},"ちゃんとインクリメントされた値が入ってます。",[10,1665,1666,1667,1669],{},"これをいったん",[17,1668,1178],{},"で保存し、保存したモデルを再度呼び出します。",[70,1671,1673],{"className":98,"code":1672,"language":100,"meta":75,"style":75},"dir_name = 'output'\n\nif not os.path.exists(dir_name):\n    os.mkdir(dir_name)\n\nmodel_save_path = os.path.join(dir_name, \"model.pt\")\ntorch.save(net.state_dict(), model_save_path)\n\nmodel = MLPNet()\nmodel.load_state_dict(torch.load(model_save_path))\n\nprint(model.mean_val)\nprint(model.count)\n\n# >>>\n# 0\n# tensor([[6001., 6001.],\n#        [6001., 6001.]])\n\n",[17,1674,1675,1679,1683,1687,1691,1695,1700,1704,1708,1713,1717,1721,1726,1731,1735,1739,1744,1748],{"__ignoreMap":75},[79,1676,1677],{"class":81,"line":82},[79,1678,1063],{},[79,1680,1681],{"class":81,"line":110},[79,1682,156],{"emptyLinePlaceholder":155},[79,1684,1685],{"class":81,"line":116},[79,1686,1074],{},[79,1688,1689],{"class":81,"line":122},[79,1690,1080],{},[79,1692,1693],{"class":81,"line":128},[79,1694,156],{"emptyLinePlaceholder":155},[79,1696,1697],{"class":81,"line":134},[79,1698,1699],{},"model_save_path = os.path.join(dir_name, \"model.pt\")\n",[79,1701,1702],{"class":81,"line":140},[79,1703,1193],{},[79,1705,1706],{"class":81,"line":146},[79,1707,156],{"emptyLinePlaceholder":155},[79,1709,1710],{"class":81,"line":152},[79,1711,1712],{},"model = MLPNet()\n",[79,1714,1715],{"class":81,"line":159},[79,1716,1206],{},[79,1718,1719],{"class":81,"line":165},[79,1720,156],{"emptyLinePlaceholder":155},[79,1722,1723],{"class":81,"line":171},[79,1724,1725],{},"print(model.mean_val)\n",[79,1727,1728],{"class":81,"line":177},[79,1729,1730],{},"print(model.count)\n",[79,1732,1733],{"class":81,"line":183},[79,1734,156],{"emptyLinePlaceholder":155},[79,1736,1737],{"class":81,"line":188},[79,1738,1645],{},[79,1740,1741],{"class":81,"line":194},[79,1742,1743],{},"# 0\n",[79,1745,1746],{"class":81,"line":199},[79,1747,1655],{},[79,1749,1750],{"class":81,"line":205},[79,1751,1752],{},"#        [6001., 6001.]])\n",[10,1754,1755,1756,1758],{},"結果をみると、普通にモデル内に定義した変数の値は保持されていません。一方で、",[17,1757,19],{},"の方は保存した時の値がちゃんと残ってます。",[10,1760,1761,1762,1765],{},"したがって、",[17,1763,1764],{},"state_dict","などで後からモデルを呼び出す際に、パラメータじゃないけど必要な値をモデルに入れておきたいを値を使う際に役立ちます。",[40,1767,1769],{"id":1768},"torchlerp-とは","torch.lerp とは",[10,1771,1772,1774],{},[17,1773,23],{},"は線形補完を行う関数です。",[10,1776,1777],{},"線形補完は式で表すと以下のようになります。",[10,1779,1780],{},"$$\nout_i = v_1 + w (v_2 - v_1)\n$$",[70,1782,1784],{"className":98,"code":1783,"language":100,"meta":75,"style":75},"torch.lerp(torch.tensor([1,1],dtype=float), torch.tensor([4,4],dtype=float), 0.5)\n\n# >>> tensor([2.5000, 2.5000], dtype=torch.float64)\n",[17,1785,1786,1791,1795],{"__ignoreMap":75},[79,1787,1788],{"class":81,"line":82},[79,1789,1790],{},"torch.lerp(torch.tensor([1,1],dtype=float), torch.tensor([4,4],dtype=float), 0.5)\n",[79,1792,1793],{"class":81,"line":110},[79,1794,156],{"emptyLinePlaceholder":155},[79,1796,1797],{"class":81,"line":116},[79,1798,1799],{},"# >>> tensor([2.5000, 2.5000], dtype=torch.float64)\n",[10,1801,1802],{},"上の式に代入すると同じ結果になります。\n２つのベクトルの間を$w$で動かす感じです。",[40,1804,1805],{"id":1805},"まとめ",[10,1807,1808],{},"今回は少し発展的な PyTorch の内容を説明しました。",[29,1810,1811,1814,1817,1820],{},[32,1812,1813],{},"torchsummary の使い方",[32,1815,1816],{},"モデルの保存方法",[32,1818,1819],{},"register_buffer の使い方",[32,1821,1822],{},"torch.lerp について",[10,1824,1825],{},"深層学習は数学的な難しさもありますし、層が増えるとその分コードでモデルを構築するのも難しいです。",[10,1827,1828],{},"フレームワークやライブラリをうまく使って本質的な問題解決に時間を当てたいですね。",[1830,1831,1832],"style",{},"html pre.shiki code .svObZ, html code.shiki .svObZ{--shiki-default:#B392F0}html pre.shiki code .sU2Wk, html code.shiki .sU2Wk{--shiki-default:#9ECBFF}html .default .shiki span {color: var(--shiki-default);background: var(--shiki-default-bg);font-style: var(--shiki-default-font-style);font-weight: var(--shiki-default-font-weight);text-decoration: var(--shiki-default-text-decoration);}html .shiki span {color: var(--shiki-default);background: var(--shiki-default-bg);font-style: var(--shiki-default-font-style);font-weight: var(--shiki-default-font-weight);text-decoration: var(--shiki-default-text-decoration);}html pre.shiki code .s95oV, html code.shiki .s95oV{--shiki-default:#E1E4E8}html pre.shiki code .sDLfK, html code.shiki .sDLfK{--shiki-default:#79B8FF}html pre.shiki code .sAwPA, html code.shiki .sAwPA{--shiki-default:#6A737D}",{"title":75,"searchDepth":110,"depth":110,"links":1834},[1835,1836,1837,1838,1839],{"id":42,"depth":110,"text":43},{"id":610,"depth":110,"text":610},{"id":1480,"depth":110,"text":1481},{"id":1768,"depth":110,"text":1769},{"id":1805,"depth":110,"text":1805},"2020-06-06","PyTorchを使った少々実践的な内容をまとめました。モデルの可視化や保存方法について説明します。また、たまに見かけるtorch.lerpやregister_bufferについてもコード付きで紹介します。",false,"md",{},"\u002Fcontents\u002Fpytorch-advanced",{"title":5,"description":1841},"contents\u002Fpytorch-advanced",[1849,1850,1851],"PyTorch","Python","機械学習","\u002Fimg\u002Ftwitter-card.png","H1NAyXMPJ7HdsX6QZqCM21LMGddS_DxRMeD2TspNpW0",[1855,1859],{"title":1856,"path":1857,"stem":1858,"children":-1},"Pythonでログ出力する方法について","\u002Fcontents\u002Fpythonlogger","contents\u002Fpythonlogger",{"title":1860,"path":1861,"stem":1862,"children":-1},"【学び直し】Pytorchの基本とMLPでMNISTの分類・可視化の実装まで","\u002Fcontents\u002Fpytorch-beginer","contents\u002Fpytorch-beginer",1784936718506]