{ "cells": [ { "cell_type": "markdown", "id": "255f7fbf-617f-40da-a338-d3f8fc5a0e92", "metadata": {}, "source": [ "# 现代卷积神经网络\n", "\n", "目录\n", "\n", "- 深度卷积神经网络(AlexNet)\n", "- 使用块的网络(VGG)\n", "- 网络中的网络(NiN)\n", "- 含并行连结的网络(GoogLeNet)\n", "- 批量规范化\n", "- 残差网络(ResNet)" ] }, { "cell_type": "markdown", "id": "3955813c-8792-42ae-af95-d88448189353", "metadata": { "jp-MarkdownHeadingCollapsed": true, "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 深度卷积神经网络(AlexNet)\n", "\n" ] }, { "cell_type": "code", "execution_count": 1, "id": "250437f3-532a-4eb2-9ed5-fed19a106e15", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from d2l import torch as d2l\n", "\n", "net = nn.Sequential(\n", " # 这里,我们使用一个11*11的更大窗口来捕捉对象。\n", " # 同时,步幅为4,以减少输出的高度和宽度。\n", " # 另外,输出通道的数目远大于LeNet\n", " nn.Conv2d(1, 96, kernel_size=11, stride=4, padding=1), nn.ReLU(),\n", " nn.MaxPool2d(kernel_size=3, stride=2),\n", " # 减小卷积窗口,使用填充为2来使得输入与输出的高和宽一致,且增大输出通道数\n", " nn.Conv2d(96, 256, kernel_size=5, padding=2), nn.ReLU(),\n", " nn.MaxPool2d(kernel_size=3, stride=2),\n", " # 使用三个连续的卷积层和较小的卷积窗口。\n", " # 除了最后的卷积层,输出通道的数量进一步增加。\n", " # 在前两个卷积层之后,汇聚层不用于减少输入的高度和宽度\n", " nn.Conv2d(256, 384, kernel_size=3, padding=1), nn.ReLU(),\n", " nn.Conv2d(384, 384, kernel_size=3, padding=1), nn.ReLU(),\n", " nn.Conv2d(384, 256, kernel_size=3, padding=1), nn.ReLU(),\n", " nn.MaxPool2d(kernel_size=3, stride=2),\n", " nn.Flatten(),\n", " # 这里,全连接层的输出数量是LeNet中的好几倍。使用dropout层来减轻过拟合\n", " nn.Linear(6400, 4096), nn.ReLU(),\n", " nn.Dropout(p=0.5),\n", " nn.Linear(4096, 4096), nn.ReLU(),\n", " nn.Dropout(p=0.5),\n", " # 最后是输出层。由于这里使用Fashion-MNIST,所以用类别数为10,而非论文中的1000\n", " nn.Linear(4096, 10))" ] }, { "cell_type": "markdown", "id": "9a6f3671-6eb9-44c3-bfe4-fd01bcfa2805", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "我们构造一个\n", "单通道数据,来观察每一层输出的形状" ] }, { "cell_type": "code", "execution_count": 2, "id": "96608413-913b-4dfe-8f23-9ffce9ea0582", "metadata": { "origin_pos": 6, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Conv2d output shape:\t torch.Size([1, 96, 54, 54])\n", "ReLU output shape:\t torch.Size([1, 96, 54, 54])\n", "MaxPool2d output shape:\t torch.Size([1, 96, 26, 26])\n", "Conv2d output shape:\t torch.Size([1, 256, 26, 26])\n", "ReLU output shape:\t torch.Size([1, 256, 26, 26])\n", "MaxPool2d output shape:\t torch.Size([1, 256, 12, 12])\n", "Conv2d output shape:\t torch.Size([1, 384, 12, 12])\n", "ReLU output shape:\t torch.Size([1, 384, 12, 12])\n", "Conv2d output shape:\t torch.Size([1, 384, 12, 12])\n", "ReLU output shape:\t torch.Size([1, 384, 12, 12])\n", "Conv2d output shape:\t torch.Size([1, 256, 12, 12])\n", "ReLU output shape:\t torch.Size([1, 256, 12, 12])\n", "MaxPool2d output shape:\t torch.Size([1, 256, 5, 5])\n", "Flatten output shape:\t torch.Size([1, 6400])\n", "Linear output shape:\t torch.Size([1, 4096])\n", "ReLU output shape:\t torch.Size([1, 4096])\n", "Dropout output shape:\t torch.Size([1, 4096])\n", "Linear output shape:\t torch.Size([1, 4096])\n", "ReLU output shape:\t torch.Size([1, 4096])\n", "Dropout output shape:\t torch.Size([1, 4096])\n", "Linear output shape:\t torch.Size([1, 10])\n" ] } ], "source": [ "X = torch.randn(1, 1, 224, 224)\n", "for layer in net:\n", " X=layer(X)\n", " print(layer.__class__.__name__,'output shape:\\t',X.shape)" ] }, { "cell_type": "markdown", "id": "356e7ffe-330d-48f9-ba76-3a54acddbeaa", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "Fashion-MNIST图像的分辨率\n", "低于ImageNet图像。\n", "我们将它们增加到$224 \\times 224$" ] }, { "cell_type": "code", "execution_count": 3, "id": "df2cf34a-9cd9-4384-a453-00043a33f530", "metadata": { "origin_pos": 9, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "batch_size = 128\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=224)" ] }, { "cell_type": "markdown", "id": "0bc4d037-556d-4847-b154-19f265aa660c", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练AlexNet。注意:这里是欠拟合的标准曲线" ] }, { "cell_type": "code", "execution_count": 4, "id": "967372e1-5db6-4762-a0d4-fc990ff5f0f7", "metadata": { "origin_pos": 11, "scrolled": true, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 0.333, train acc 0.878, test acc 0.878\n", "1762.7 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-03-21T08:52:06.634916\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "lr, num_epochs = 0.01, 10\n", "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] }, { "cell_type": "markdown", "id": "fc1d1ce8-9b00-4b87-9ff1-177da5ee139d", "metadata": { "jp-MarkdownHeadingCollapsed": true, "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 使用块的网络(VGG)\n", "\n", "VGG块" ] }, { "cell_type": "code", "execution_count": 1, "id": "cf0cb392-48d9-451d-a7ea-e1892865657d", "metadata": { "origin_pos": 4, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from d2l import torch as d2l\n", "\n", "\n", "def vgg_block(num_convs, in_channels, out_channels):\n", " layers = []\n", " for _ in range(num_convs):\n", " layers.append(nn.Conv2d(in_channels, out_channels,\n", " kernel_size=3, padding=1))\n", " layers.append(nn.ReLU())\n", " in_channels = out_channels\n", " layers.append(nn.MaxPool2d(kernel_size=2,stride=2))\n", " return nn.Sequential(*layers)" ] }, { "cell_type": "markdown", "id": "e40a3ce9-5838-491c-b02a-d4f626c5f024", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "VGG网络\n", "\n", "原始VGG网络有5个卷积块,其中前两个块各有一个卷积层,后三个块各包含两个卷积层。 第一个模块有64个输出通道,每个后续模块将输出通道数量翻倍,直到该数字达到512。由于该网络使用8个卷积层和3个全连接层,因此它通常被称为VGG-11。" ] }, { "cell_type": "code", "execution_count": 2, "id": "aad71cf3-3c74-4162-bf78-b5b2e313cf65", "metadata": { "origin_pos": 10, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "conv_arch = ((1, 64), (1, 128), (2, 256), (2, 512), (2, 512))" ] }, { "cell_type": "code", "execution_count": 3, "id": "df9c3499-e926-4821-8136-e66baf433731", "metadata": { "origin_pos": 10, "slideshow": { "slide_type": "slide" }, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def vgg(conv_arch):\n", " conv_blks = []\n", " in_channels = 1\n", " for (num_convs, out_channels) in conv_arch:\n", " conv_blks.append(vgg_block(num_convs, in_channels, out_channels))\n", " in_channels = out_channels\n", "\n", " return nn.Sequential(\n", " *conv_blks, nn.Flatten(),\n", " nn.Linear(out_channels * 7 * 7, 4096), nn.ReLU(), nn.Dropout(0.5),\n", " nn.Linear(4096, 4096), nn.ReLU(), nn.Dropout(0.5),\n", " nn.Linear(4096, 10))\n", "\n", "net = vgg(conv_arch)" ] }, { "cell_type": "markdown", "id": "18ceef7a-71d8-4eab-b340-c0848b4675b6", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "观察每个层输出的形状" ] }, { "cell_type": "code", "execution_count": 4, "id": "20065462-d243-406d-919e-5af13f73c088", "metadata": { "origin_pos": 14, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Sequential output shape:\t torch.Size([1, 64, 112, 112])\n", "Sequential output shape:\t torch.Size([1, 128, 56, 56])\n", "Sequential output shape:\t torch.Size([1, 256, 28, 28])\n", "Sequential output shape:\t torch.Size([1, 512, 14, 14])\n", "Sequential output shape:\t torch.Size([1, 512, 7, 7])\n", "Flatten output shape:\t torch.Size([1, 25088])\n", "Linear output shape:\t torch.Size([1, 4096])\n", "ReLU output shape:\t torch.Size([1, 4096])\n", "Dropout output shape:\t torch.Size([1, 4096])\n", "Linear output shape:\t torch.Size([1, 4096])\n", "ReLU output shape:\t torch.Size([1, 4096])\n", "Dropout output shape:\t torch.Size([1, 4096])\n", "Linear output shape:\t torch.Size([1, 10])\n" ] } ], "source": [ "X = torch.randn(size=(1, 1, 224, 224))\n", "for blk in net:\n", " X = blk(X)\n", " print(blk.__class__.__name__,'output shape:\\t',X.shape)" ] }, { "cell_type": "markdown", "id": "49d25e85-676d-4b59-8148-58596e407411", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "由于VGG-11比AlexNet计算量更大,因此我们构建了一个通道数较少的网络" ] }, { "cell_type": "code", "execution_count": 5, "id": "286b6e79-ca57-4ac4-b018-e64a210dd88e", "metadata": { "origin_pos": 17, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "ratio = 4\n", "small_conv_arch = [(pair[0], pair[1] // ratio) for pair in conv_arch]\n", "net = vgg(small_conv_arch)" ] }, { "cell_type": "markdown", "id": "0e9ec14c-f318-4c0e-969c-4a17e9ec09a7", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "模型训练。仔细观察:这里出现了过拟合" ] }, { "cell_type": "code", "execution_count": 6, "id": "e986bb65-7c5f-4402-af8d-55ccc4760662", "metadata": { "origin_pos": 20, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 0.172, train acc 0.936, test acc 0.914\n", "1119.5 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-03-21T09:16:51.006260\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "lr, num_epochs, batch_size = 0.05, 10, 128\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=224)\n", "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] }, { "cell_type": "markdown", "id": "051bf6dd-6153-4d4a-8a18-4878031998b0", "metadata": { "jp-MarkdownHeadingCollapsed": true, "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 网络中的网络(NiN)\n", "\n", "NiN块" ] }, { "cell_type": "code", "execution_count": 2, "id": "1c18dda6-ece9-4c0a-9235-83e4b5bd6f34", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from d2l import torch as d2l\n", "\n", "\n", "def nin_block(in_channels, out_channels, kernel_size, strides, padding):\n", " return nn.Sequential(\n", " nn.Conv2d(in_channels, out_channels, kernel_size, strides, padding),\n", " nn.ReLU(),\n", " nn.Conv2d(out_channels, out_channels, kernel_size=1), nn.ReLU(),\n", " nn.Conv2d(out_channels, out_channels, kernel_size=1), nn.ReLU())" ] }, { "cell_type": "markdown", "id": "9a43fb36-d21c-48d8-afc9-3940fc4151c8", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "NiN模型" ] }, { "cell_type": "code", "execution_count": 3, "id": "4d406faa-4e96-49f8-b226-a76be1f1cc40", "metadata": { "origin_pos": 6, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "net = nn.Sequential(\n", " nin_block(1, 96, kernel_size=11, strides=4, padding=0),\n", " nn.MaxPool2d(3, stride=2),\n", " nin_block(96, 256, kernel_size=5, strides=1, padding=2),\n", " nn.MaxPool2d(3, stride=2),\n", " nin_block(256, 384, kernel_size=3, strides=1, padding=1),\n", " nn.MaxPool2d(3, stride=2),\n", " nn.Dropout(0.5),\n", " nin_block(384, 10, kernel_size=3, strides=1, padding=1), # 标签类别数是10\n", " nn.AdaptiveAvgPool2d((1, 1)),\n", " nn.Flatten()) # 将四维的输出转成二维的输出,其形状为(批量大小,10)" ] }, { "cell_type": "markdown", "id": "554152fc-1a8c-40dc-9892-996c7039d471", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "查看每个块的输出形状" ] }, { "cell_type": "code", "execution_count": 4, "id": "73bc25b6-f379-4013-8ee2-b7e80c50a23f", "metadata": { "origin_pos": 10, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Sequential output shape:\t torch.Size([1, 96, 54, 54])\n", "MaxPool2d output shape:\t torch.Size([1, 96, 26, 26])\n", "Sequential output shape:\t torch.Size([1, 256, 26, 26])\n", "MaxPool2d output shape:\t torch.Size([1, 256, 12, 12])\n", "Sequential output shape:\t torch.Size([1, 384, 12, 12])\n", "MaxPool2d output shape:\t torch.Size([1, 384, 5, 5])\n", "Dropout output shape:\t torch.Size([1, 384, 5, 5])\n", "Sequential output shape:\t torch.Size([1, 10, 5, 5])\n", "AdaptiveAvgPool2d output shape:\t torch.Size([1, 10, 1, 1])\n", "Flatten output shape:\t torch.Size([1, 10])\n" ] } ], "source": [ "X = torch.rand(size=(1, 1, 224, 224))\n", "for layer in net:\n", " X = layer(X)\n", " print(layer.__class__.__name__,'output shape:\\t', X.shape)" ] }, { "cell_type": "markdown", "id": "feb20603-ba72-4fb4-9763-ba195b16da54", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练模型" ] }, { "cell_type": "code", "execution_count": 5, "id": "5ef0873e-78e2-4f65-9ce0-4dba0c5f48c9", "metadata": { "origin_pos": 13, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 2.303, train acc 0.100, test acc 0.100\n", "1445.5 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-03-22T09:47:52.697012\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "lr, num_epochs, batch_size = 0.1, 10, 128\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=224)\n", "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] }, { "cell_type": "markdown", "id": "3b78340b-c3f8-4617-876b-79f7e89d7d2e", "metadata": {}, "source": [ "注意:这是典型的训练失败案例。\n", "\n", "- 思考:如何(调参以实现)提高准确性?" ] }, { "cell_type": "markdown", "id": "0d7743da-c831-4ad6-bb10-bd55e05a7178", "metadata": { "jp-MarkdownHeadingCollapsed": true, "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 含并行连结的网络(GoogLeNet)" ] }, { "cell_type": "code", "execution_count": 1, "id": "8892c059-5c03-4cb2-9a2f-5d3cfecbf832", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from torch.nn import functional as F\n", "from d2l import torch as d2l" ] }, { "cell_type": "markdown", "id": "f161d0dd-95c0-4f0b-97e4-8e33bec1c5ac", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "Inception块" ] }, { "cell_type": "code", "execution_count": 2, "id": "4c8873f6-a2c8-43ed-bfb6-702d37fa897b", "metadata": { "origin_pos": 2, "slideshow": { "slide_type": "-" }, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "class Inception(nn.Module):\n", " # c1--c4是每条路径的输出通道数\n", " def __init__(self, in_channels, c1, c2, c3, c4, **kwargs):\n", " super(Inception, self).__init__(**kwargs)\n", " # 线路1,单1x1卷积层\n", " self.p1_1 = nn.Conv2d(in_channels, c1, kernel_size=1)\n", " # 线路2,1x1卷积层后接3x3卷积层\n", " self.p2_1 = nn.Conv2d(in_channels, c2[0], kernel_size=1)\n", " self.p2_2 = nn.Conv2d(c2[0], c2[1], kernel_size=3, padding=1)\n", " # 线路3,1x1卷积层后接5x5卷积层\n", " self.p3_1 = nn.Conv2d(in_channels, c3[0], kernel_size=1)\n", " self.p3_2 = nn.Conv2d(c3[0], c3[1], kernel_size=5, padding=2)\n", " # 线路4,3x3最大汇聚层后接1x1卷积层\n", " self.p4_1 = nn.MaxPool2d(kernel_size=3, stride=1, padding=1)\n", " self.p4_2 = nn.Conv2d(in_channels, c4, kernel_size=1)\n", "\n", " def forward(self, x):\n", " p1 = F.relu(self.p1_1(x))\n", " p2 = F.relu(self.p2_2(F.relu(self.p2_1(x))))\n", " p3 = F.relu(self.p3_2(F.relu(self.p3_1(x))))\n", " p4 = F.relu(self.p4_2(self.p4_1(x)))\n", " # 在通道维度上连结输出\n", " return torch.cat((p1, p2, p3, p4), dim=1)" ] }, { "cell_type": "markdown", "id": "1f734be6-59ed-4dd2-b60f-5f7616942401", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "GoogLeNet模型" ] }, { "cell_type": "code", "execution_count": 6, "id": "2ae4258c-339b-4bb2-b78a-7ad180697576", "metadata": { "origin_pos": 22, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "b1 = nn.Sequential(nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3),\n", " nn.ReLU(),\n", " nn.MaxPool2d(kernel_size=3, stride=2, padding=1))\n", "\n", "b2 = nn.Sequential(nn.Conv2d(64, 64, kernel_size=1),\n", " nn.ReLU(),\n", " nn.Conv2d(64, 192, kernel_size=3, padding=1),\n", " nn.ReLU(),\n", " nn.MaxPool2d(kernel_size=3, stride=2, padding=1))\n", "\n", "b3 = nn.Sequential(Inception(192, 64, (96, 128), (16, 32), 32),\n", " Inception(256, 128, (128, 192), (32, 96), 64),\n", " nn.MaxPool2d(kernel_size=3, stride=2, padding=1))\n", "\n", "b4 = nn.Sequential(Inception(480, 192, (96, 208), (16, 48), 64),\n", " Inception(512, 160, (112, 224), (24, 64), 64),\n", " Inception(512, 128, (128, 256), (24, 64), 64),\n", " Inception(512, 112, (144, 288), (32, 64), 64),\n", " Inception(528, 256, (160, 320), (32, 128), 128),\n", " nn.MaxPool2d(kernel_size=3, stride=2, padding=1))\n", "\n", "b5 = nn.Sequential(Inception(832, 256, (160, 320), (32, 128), 128),\n", " Inception(832, 384, (192, 384), (48, 128), 128),\n", " nn.AdaptiveAvgPool2d((1,1)),\n", " nn.Flatten())\n", "\n", "net = nn.Sequential(b1, b2, b3, b4, b5, nn.Linear(1024, 10))" ] }, { "cell_type": "markdown", "id": "0d0f8c4d-c45f-4610-9d0d-eccc76839e4a", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "为了使Fashion-MNIST上的训练短小精悍,我们将输入的高和宽从224降到96" ] }, { "cell_type": "code", "execution_count": 7, "id": "23870be5-1556-405a-bcef-2c0fe26cb645", "metadata": { "origin_pos": 26, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Sequential output shape:\t torch.Size([1, 64, 24, 24])\n", "Sequential output shape:\t torch.Size([1, 192, 12, 12])\n", "Sequential output shape:\t torch.Size([1, 480, 6, 6])\n", "Sequential output shape:\t torch.Size([1, 832, 3, 3])\n", "Sequential output shape:\t torch.Size([1, 1024])\n", "Linear output shape:\t torch.Size([1, 10])\n" ] } ], "source": [ "X = torch.rand(size=(1, 1, 96, 96))\n", "for layer in net:\n", " X = layer(X)\n", " print(layer.__class__.__name__,'output shape:\\t', X.shape)" ] }, { "cell_type": "markdown", "id": "9b814549-6649-4548-9297-d3847abaced1", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练模型" ] }, { "cell_type": "code", "execution_count": 8, "id": "54f90511-0268-4d50-9316-a8ddbc953670", "metadata": { "origin_pos": 29, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 0.244, train acc 0.908, test acc 0.896\n", "3490.2 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-01-17T01:55:44.273318\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "lr, num_epochs, batch_size = 0.1, 10, 128\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=96)\n", "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] }, { "cell_type": "markdown", "id": "7edc7652-ed97-4851-aac8-ed51bacae82f", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "## 批量规范化\n", "\n", "从零实现" ] }, { "cell_type": "code", "execution_count": 1, "id": "c59dccc1-434e-4134-a988-b4f8c4ec63f0", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from d2l import torch as d2l" ] }, { "cell_type": "code", "execution_count": 2, "id": "d6629b56-83fb-47bf-93cf-98becf2a088f", "metadata": { "origin_pos": 2, "slideshow": { "slide_type": "slide" }, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def batch_norm(X, gamma, beta, moving_mean, moving_var, eps, momentum):\n", " if not torch.is_grad_enabled(): # 判断当前模式是训练模式还是预测模式\n", " # 如果是在预测模式下,直接使用传入的移动平均所得的均值和方差\n", " X_hat = (X - moving_mean) / torch.sqrt(moving_var + eps)\n", " else:\n", " assert len(X.shape) in (2, 4)\n", " if len(X.shape) == 2: # 使用全连接层的情况,计算特征维上的均值和方差\n", " mean = X.mean(dim=0)\n", " var = ((X - mean) ** 2).mean(dim=0)\n", " else: # 使用二维卷积层的情况,计算通道维上(axis=1)的均值和方差\n", " # 这里我们需要保持X的形状以便后面可以做广播运算\n", " mean = X.mean(dim=(0, 2, 3), keepdim=True)\n", " var = ((X - mean) ** 2).mean(dim=(0, 2, 3), keepdim=True)\n", " # 训练模式下,用当前的均值和方差做标准化\n", " X_hat = (X - mean) / torch.sqrt(var + eps)\n", " # 更新移动平均的均值和方差\n", " moving_mean = momentum * moving_mean + (1.0 - momentum) * mean\n", " moving_var = momentum * moving_var + (1.0 - momentum) * var\n", " Y = gamma * X_hat + beta # 缩放和移位\n", " return Y, moving_mean.data, moving_var.data" ] }, { "cell_type": "markdown", "id": "f4872957-57c1-40c0-be54-6d00862c9d2f", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "创建一个正确的`BatchNorm`层" ] }, { "cell_type": "code", "execution_count": 3, "id": "f4c0ef6f-7546-4949-bfd9-b1e10f3abade", "metadata": { "origin_pos": 6, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "class BatchNorm(nn.Module):\n", " def __init__(self, num_features, num_dims):\n", " # num_features:全连接层的输出数量或卷积层的输出通道数。\n", " # num_dims:2表示全连接层,4表示卷积层\n", " super().__init__()\n", " if num_dims == 2:\n", " shape = (1, num_features)\n", " else:\n", " shape = (1, num_features, 1, 1)\n", " # 参与求梯度和迭代的拉伸和偏移参数,分别初始化成1和0\n", " self.gamma = nn.Parameter(torch.ones(shape))\n", " self.beta = nn.Parameter(torch.zeros(shape))\n", " # 非模型参数:均值、方差\n", " self.moving_mean = torch.zeros(shape)\n", " self.moving_var = torch.ones(shape)\n", "\n", " def forward(self, X):\n", " if self.moving_mean.device != X.device:\n", " # 如果X不在内存上,将moving_mean和moving_var复制到X所在显存上\n", " self.moving_mean = self.moving_mean.to(X.device)\n", " self.moving_var = self.moving_var.to(X.device)\n", " # 保存更新过的moving_mean和moving_var\n", " Y, self.moving_mean, self.moving_var = batch_norm(\n", " X, self.gamma, self.beta, self.moving_mean,\n", " self.moving_var, eps=1e-5, momentum=0.9)\n", " return Y" ] }, { "cell_type": "markdown", "id": "cf94bbae-4049-443e-82f9-b483b2b25a8e", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "应用`BatchNorm`\n", "于LeNet模型" ] }, { "cell_type": "code", "execution_count": 4, "id": "cc24a7dc-08ac-4a95-bfbc-b8aada3842e7", "metadata": { "origin_pos": 10, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "net = nn.Sequential(\n", " nn.Conv2d(1, 6, kernel_size=5), BatchNorm(6, num_dims=4), nn.Sigmoid(),\n", " nn.AvgPool2d(kernel_size=2, stride=2),\n", " nn.Conv2d(6, 16, kernel_size=5), BatchNorm(16, num_dims=4), nn.Sigmoid(),\n", " nn.AvgPool2d(kernel_size=2, stride=2), nn.Flatten(),\n", " nn.Linear(16*4*4, 120), BatchNorm(120, num_dims=2), nn.Sigmoid(),\n", " nn.Linear(120, 84), BatchNorm(84, num_dims=2), nn.Sigmoid(),\n", " nn.Linear(84, 10))" ] }, { "cell_type": "markdown", "id": "fbcd9945-0abd-4443-962f-56e928685200", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "在Fashion-MNIST数据集上训练网络。注意:学习率大很多" ] }, { "cell_type": "code", "execution_count": 5, "id": "3ab1ecba-57e6-4cec-a902-e68dc49b64d1", "metadata": { "origin_pos": 13, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 0.263, train acc 0.902, test acc 0.857\n", "31378.5 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-03-22T19:59:51.169965\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "lr, num_epochs, batch_size = 1.0, 10, 256\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size)\n", "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] }, { "cell_type": "markdown", "id": "b8bc077b-1747-4de4-87c1-4da4e574df34", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "拉伸参数`gamma`和偏移参数`beta`" ] }, { "cell_type": "code", "execution_count": 6, "id": "ed6c4a56-a677-4c93-9d13-8a3b5f5c489d", "metadata": { "origin_pos": 17, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "(tensor([2.9182, 3.9908, 2.5709, 2.1960, 3.0324, 2.0170], device='cuda:0',\n", " grad_fn=),\n", " tensor([-0.4980, 2.6303, -2.3259, -1.4185, -3.2187, -2.1353], device='cuda:0',\n", " grad_fn=))" ] }, "execution_count": 6, "metadata": {}, "output_type": "execute_result" } ], "source": [ "net[1].gamma.reshape((-1,)), net[1].beta.reshape((-1,))" ] }, { "cell_type": "markdown", "id": "57715920-0def-4217-96b1-4425f07b02f6", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "简明实现" ] }, { "cell_type": "code", "execution_count": 7, "id": "ff5c985a-76c9-4a4a-bccc-feb3179e3c4b", "metadata": { "origin_pos": 21, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "net = nn.Sequential(\n", " nn.Conv2d(1, 6, kernel_size=5), nn.BatchNorm2d(6), nn.Sigmoid(),\n", " nn.AvgPool2d(kernel_size=2, stride=2),\n", " nn.Conv2d(6, 16, kernel_size=5), nn.BatchNorm2d(16), nn.Sigmoid(),\n", " nn.AvgPool2d(kernel_size=2, stride=2), nn.Flatten(),\n", " nn.Linear(256, 120), nn.BatchNorm1d(120), nn.Sigmoid(),\n", " nn.Linear(120, 84), nn.BatchNorm1d(84), nn.Sigmoid(),\n", " nn.Linear(84, 10))" ] }, { "cell_type": "markdown", "id": "c991fda9-53a2-4510-9ee5-2968b3f521a5", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "使用相同超参数来训练模型。注意:通常高级API运行速度快得多,因为它的代码已编译为C++或CUDA" ] }, { "cell_type": "code", "execution_count": 8, "id": "2e6a118b-c4db-4db0-a15f-ab361efb2051", "metadata": { "origin_pos": 24, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 0.281, train acc 0.896, test acc 0.843\n", "55970.4 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-03-22T20:00:16.979860\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] }, { "cell_type": "markdown", "id": "be2ce66c-0694-4c20-8f9c-e4c62b9e92b1", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "## 残差网络(ResNet)\n", "\n", "残差块" ] }, { "cell_type": "code", "execution_count": 1, "id": "b253049d-96a4-4871-a6a1-c9d2de3a9379", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from torch.nn import functional as F\n", "from d2l import torch as d2l\n", "\n", "\n", "class Residual(nn.Module): \n", " def __init__(self, input_channels, num_channels,\n", " use_1x1conv=False, strides=1):\n", " super().__init__()\n", " self.conv1 = nn.Conv2d(input_channels, num_channels,\n", " kernel_size=3, padding=1, stride=strides)\n", " self.conv2 = nn.Conv2d(num_channels, num_channels,\n", " kernel_size=3, padding=1)\n", " if use_1x1conv:\n", " self.conv3 = nn.Conv2d(input_channels, num_channels,\n", " kernel_size=1, stride=strides)\n", " else:\n", " self.conv3 = None\n", " self.bn1 = nn.BatchNorm2d(num_channels)\n", " self.bn2 = nn.BatchNorm2d(num_channels)\n", "\n", " def forward(self, X):\n", " Y = F.relu(self.bn1(self.conv1(X)))\n", " Y = self.bn2(self.conv2(Y))\n", " if self.conv3:\n", " X = self.conv3(X)\n", " Y += X\n", " return F.relu(Y)" ] }, { "cell_type": "markdown", "id": "72824a52-1ed9-4f12-b826-e60d2ce42568", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "输入和输出形状一致" ] }, { "cell_type": "code", "execution_count": 2, "id": "02cc50d9-5710-4e2f-bf0f-568f1046321e", "metadata": { "origin_pos": 6, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "torch.Size([4, 3, 6, 6])" ] }, "execution_count": 2, "metadata": {}, "output_type": "execute_result" } ], "source": [ "blk = Residual(3,3)\n", "X = torch.rand(4, 3, 6, 6)\n", "Y = blk(X)\n", "Y.shape" ] }, { "cell_type": "markdown", "id": "c09c8e89-dfdb-4e6f-a75b-b949f8e04d6c", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "增加输出通道数的同时,减半输出的高和宽" ] }, { "cell_type": "code", "execution_count": 3, "id": "57b5f05e-853c-460e-afb7-3a3483c7e776", "metadata": { "origin_pos": 10, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "torch.Size([4, 6, 3, 3])" ] }, "execution_count": 3, "metadata": {}, "output_type": "execute_result" } ], "source": [ "blk = Residual(3,6, use_1x1conv=True, strides=2)\n", "blk(X).shape" ] }, { "cell_type": "markdown", "id": "57bd6688-aa3a-4758-b0f6-67b786db8c7a", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "ResNet模型" ] }, { "cell_type": "code", "execution_count": 7, "id": "95fcb71a-7682-4efe-9434-688e4f1606e0", "metadata": { "origin_pos": 26, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "b1 = nn.Sequential(nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3),\n", " nn.BatchNorm2d(64), nn.ReLU(),\n", " nn.MaxPool2d(kernel_size=3, stride=2, padding=1))\n", "\n", "def resnet_block(input_channels, num_channels, num_residuals,\n", " first_block=False):\n", " blk = []\n", " for i in range(num_residuals):\n", " if i == 0 and not first_block:\n", " blk.append(Residual(input_channels, num_channels,\n", " use_1x1conv=True, strides=2))\n", " else:\n", " blk.append(Residual(num_channels, num_channels))\n", " return blk\n", "\n", "b2 = nn.Sequential(*resnet_block(64, 64, 2, first_block=True))\n", "b3 = nn.Sequential(*resnet_block(64, 128, 2))\n", "b4 = nn.Sequential(*resnet_block(128, 256, 2))\n", "b5 = nn.Sequential(*resnet_block(256, 512, 2))\n", "\n", "net = nn.Sequential(b1, b2, b3, b4, b5,\n", " nn.AdaptiveAvgPool2d((1,1)),\n", " nn.Flatten(), nn.Linear(512, 10))" ] }, { "cell_type": "markdown", "id": "ead75736-665d-48b8-9ba9-7634f5ed1395", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "观察一下ResNet中不同模块的输入形状是如何变化的" ] }, { "cell_type": "code", "execution_count": 8, "id": "caaf4445-c02e-4d4e-af70-ead31c246332", "metadata": { "origin_pos": 30, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Sequential output shape:\t torch.Size([1, 64, 56, 56])\n", "Sequential output shape:\t torch.Size([1, 64, 56, 56])\n", "Sequential output shape:\t torch.Size([1, 128, 28, 28])\n", "Sequential output shape:\t torch.Size([1, 256, 14, 14])\n", "Sequential output shape:\t torch.Size([1, 512, 7, 7])\n", "AdaptiveAvgPool2d output shape:\t torch.Size([1, 512, 1, 1])\n", "Flatten output shape:\t torch.Size([1, 512])\n", "Linear output shape:\t torch.Size([1, 10])\n" ] } ], "source": [ "X = torch.rand(size=(1, 1, 224, 224))\n", "for layer in net:\n", " X = layer(X)\n", " print(layer.__class__.__name__,'output shape:\\t', X.shape)" ] }, { "cell_type": "markdown", "id": "057bd90d-193f-4577-8d4b-e922d8e1ae18", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练模型" ] }, { "cell_type": "code", "execution_count": 9, "id": "bd0e57ed-ed08-4d18-8cdf-4062fad7ffe1", "metadata": { "origin_pos": 33, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "loss 0.009, train acc 0.998, test acc 0.922\n", "4702.7 examples/sec on cuda:0\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2022-01-17T01:46:46.812949\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.3.3, https://matplotlib.org/\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "\n" ], "text/plain": [ "
" ] }, "metadata": { "needs_background": "light" }, "output_type": "display_data" } ], "source": [ "lr, num_epochs, batch_size = 0.05, 10, 256\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=96)\n", "d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3 (ipykernel)", "language": "python", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.8.16" } }, "nbformat": 4, "nbformat_minor": 5 }