{ "cells": [ { "cell_type": "markdown", "id": "6b83a647-6801-4667-a822-0656b8e9556a", "metadata": {}, "source": [ "# 模型选择\n", "\n", "目录\n", "\n", "- 欠拟合和过拟合\n", "- 权重衰减\n", "- 暂退法(Dropout)" ] }, { "cell_type": "markdown", "id": "30394f0b-c857-48e6-b3f5-b7ceba0411c8", "metadata": { "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 欠拟合和过拟合\n", "\n", "通过多项式拟合来探索这些概念" ] }, { "cell_type": "code", "execution_count": 1, "id": "70d264d4-0dae-48f1-9288-375f9d84aa10", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import math\n", "import numpy as np\n", "import torch\n", "from torch import nn\n", "from d2l import torch as d2l" ] }, { "cell_type": "markdown", "id": "68ca9f7a-b179-4fd0-a183-65ce654526d7", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "使用以下三阶多项式来生成训练和测试数据的标签:\n", "$$y = 5 + 1.2x - 3.4\\frac{x^2}{2!} + 5.6 \\frac{x^3}{3!} + \\epsilon \\text{ where }\n", "\\epsilon \\sim \\mathcal{N}(0, 0.1^2)$$" ] }, { "cell_type": "code", "execution_count": 2, "id": "513f2378-8df6-4eaa-b98a-59c2b890fbae", "metadata": { "origin_pos": 5, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "max_degree = 20 # 多项式的最大阶数\n", "n_train, n_test = 100, 100 # 训练和测试数据集大小\n", "true_w = np.zeros(max_degree) # 分配大量的空间\n", "true_w[0:4] = np.array([5, 1.2, -3.4, 5.6])\n", "\n", "features = np.random.normal(size=(n_train + n_test, 1))\n", "np.random.shuffle(features)\n", "poly_features = np.power(features, np.arange(max_degree).reshape(1, -1))\n", "for i in range(max_degree):\n", " poly_features[:, i] /= math.gamma(i + 1) # gamma(n)=(n-1)!\n", "# labels的维度:(n_train+n_test,)\n", "labels = np.dot(poly_features, true_w)\n", "labels += np.random.normal(scale=0.1, size=labels.shape)" ] }, { "cell_type": "markdown", "id": "eda1d879-8d72-40bb-9d01-8aaa5ca28f4e", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "查看一下前2个样本" ] }, { "cell_type": "code", "execution_count": 3, "id": "a1988415-3bc9-4e2b-bb2c-8f10d26b97e1", "metadata": { "origin_pos": 8, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "(tensor([[0.7135],\n", " [0.4976]]),\n", " tensor([[1.0000e+00, 7.1346e-01, 2.5451e-01, 6.0527e-02, 1.0796e-02, 1.5405e-03,\n", " 1.8318e-04, 1.8670e-05, 1.6650e-06, 1.3199e-07, 9.4170e-09, 6.1079e-10,\n", " 3.6314e-11, 1.9930e-12, 1.0156e-13, 4.8308e-15, 2.1541e-16, 9.0403e-18,\n", " 3.5833e-19, 1.3455e-20],\n", " [1.0000e+00, 4.9764e-01, 1.2382e-01, 2.0540e-02, 2.5554e-03, 2.5434e-04,\n", " 2.1095e-05, 1.4997e-06, 9.3288e-08, 5.1582e-09, 2.5670e-10, 1.1613e-11,\n", " 4.8160e-13, 1.8436e-14, 6.5531e-16, 2.1741e-17, 6.7620e-19, 1.9794e-20,\n", " 5.4725e-22, 1.4334e-23]]),\n", " tensor([5.3958, 5.5251]))" ] }, "execution_count": 3, "metadata": {}, "output_type": "execute_result" } ], "source": [ "# NumPy ndarray转换为tensor\n", "true_w, features, poly_features, labels = [torch.tensor(x, dtype=\n", " torch.float32) for x in [true_w, features, poly_features, labels]]\n", "\n", "features[:2], poly_features[:2, :], labels[:2]" ] }, { "cell_type": "markdown", "id": "f39a6b66-4877-415f-88e7-fee01a15283f", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "实现一个函数来评估模型在给定数据集上的损失" ] }, { "cell_type": "code", "execution_count": 4, "id": "a3a49a16-7428-46f9-89e9-44ae554a182d", "metadata": { "origin_pos": 11, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def evaluate_loss(net, data_iter, loss): \n", " \"\"\"评估给定数据集上模型的损失\"\"\"\n", " metric = d2l.Accumulator(2) # 损失的总和,样本数量\n", " for X, y in data_iter:\n", " out = net(X)\n", " y = y.reshape(out.shape)\n", " l = loss(out, y)\n", " metric.add(l.sum(), l.numel())\n", " return metric[0] / metric[1]" ] }, { "cell_type": "markdown", "id": "1d3ab43b-025e-47fa-a787-b8e051703282", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "定义训练函数" ] }, { "cell_type": "code", "execution_count": 5, "id": "3b3a4d31-223a-43ba-9e05-751078977a7c", "metadata": { "origin_pos": 14, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def train(train_features, test_features, train_labels, test_labels,\n", " num_epochs=400):\n", " loss = nn.MSELoss(reduction='none')\n", " input_shape = train_features.shape[-1]\n", " # 不设置偏置,因为我们已经在多项式中实现了它\n", " net = nn.Sequential(nn.Linear(input_shape, 1, bias=False))\n", " batch_size = min(10, train_labels.shape[0])\n", " train_iter = d2l.load_array((train_features, train_labels.reshape(-1,1)),\n", " batch_size)\n", " test_iter = d2l.load_array((test_features, test_labels.reshape(-1,1)),\n", " batch_size, is_train=False)\n", " trainer = torch.optim.SGD(net.parameters(), lr=0.01)\n", " animator = d2l.Animator(xlabel='epoch', ylabel='loss', yscale='log',\n", " xlim=[1, num_epochs], ylim=[1e-3, 1e2],\n", " legend=['train', 'test'])\n", " for epoch in range(num_epochs):\n", " d2l.train_epoch_ch3(net, train_iter, loss, trainer)\n", " if epoch == 0 or (epoch + 1) % 20 == 0:\n", " animator.add(epoch + 1, (evaluate_loss(net, train_iter, loss),\n", " evaluate_loss(net, test_iter, loss)))\n", " print('weight:', net[0].weight.data.numpy())" ] }, { "cell_type": "markdown", "id": "12e3f5b7-701b-48d7-b962-cd307dd3f93e", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "三阶多项式函数拟合(正常)" ] }, { "cell_type": "code", "execution_count": 6, "id": "4dbd5c85-6b5e-40d7-a21e-c486758926e1", "metadata": { "origin_pos": 17, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "weight: [[ 4.9961076 1.1977049 -3.383854 5.6232586]]\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:16:47.839828\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \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": {}, "output_type": "display_data" } ], "source": [ "# 从多项式特征中选择前4个维度,即1,x,x^2/2!,x^3/3!\n", "train(poly_features[:n_train, :4], poly_features[n_train:, :4],\n", " labels[:n_train], labels[n_train:])" ] }, { "cell_type": "markdown", "id": "56cd7a76-7d08-44bf-a0eb-f5249642c893", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "线性函数拟合(欠拟合)" ] }, { "cell_type": "code", "execution_count": 7, "id": "255bb596-d3fc-479f-8296-5da42c6a333a", "metadata": { "origin_pos": 19, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "weight: [[3.1164684 4.298125 ]]\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:17:04.064205\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \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": {}, "output_type": "display_data" } ], "source": [ "# 从多项式特征中选择前2个维度,即1和x\n", "train(poly_features[:n_train, :2], poly_features[n_train:, :2],\n", " labels[:n_train], labels[n_train:])" ] }, { "cell_type": "markdown", "id": "c001f342-1be2-4b1a-a76a-41198de00ea4", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "高阶多项式函数拟合(过拟合)" ] }, { "cell_type": "code", "execution_count": 8, "id": "37034ea3-9a8f-41bd-8af4-a36612b2c82a", "metadata": { "origin_pos": 21, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "weight: [[ 4.99766 1.2872416 -3.3798084 5.19445 -0.01223703 1.1709547\n", " 0.18639933 -0.03800117 0.22357842 -0.19907166 -0.04570822 -0.11779588\n", " 0.01551174 0.08947276 0.10764467 0.2204279 -0.15191121 -0.09354579\n", " -0.12271424 -0.10525914]]\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:18:09.170658\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \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": {}, "output_type": "display_data" } ], "source": [ "# 从多项式特征中选取所有维度\n", "train(poly_features[:n_train, :], poly_features[n_train:, :],\n", " labels[:n_train], labels[n_train:], num_epochs=1500)" ] }, { "cell_type": "markdown", "id": "8389c4a8-fa32-4615-9bc0-e39c4d158736", "metadata": { "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 权重衰减\n", "\n", "权重衰减是最广泛使用的正则化的技术之一" ] }, { "cell_type": "code", "execution_count": 9, "id": "c76286b2-e78b-4a9e-8f94-f714086da60e", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "%matplotlib inline\n", "import torch\n", "from torch import nn\n", "from d2l import torch as d2l" ] }, { "cell_type": "markdown", "id": "659e2837-48e9-4674-b70f-356a7673c99f", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "像以前一样生成一些数据\n", "$$y = 0.05 + \\sum_{i = 1}^d 0.01 x_i + \\epsilon \\text{ where }\n", "\\epsilon \\sim \\mathcal{N}(0, 0.01^2)$$" ] }, { "cell_type": "code", "execution_count": 10, "id": "ba4a796e-7965-4c8a-b0cb-11522498e8e6", "metadata": { "origin_pos": 5, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "n_train, n_test, num_inputs, batch_size = 20, 100, 200, 5\n", "true_w, true_b = torch.ones((num_inputs, 1)) * 0.01, 0.05\n", "train_data = d2l.synthetic_data(true_w, true_b, n_train)\n", "train_iter = d2l.load_array(train_data, batch_size)\n", "test_data = d2l.synthetic_data(true_w, true_b, n_test)\n", "test_iter = d2l.load_array(test_data, batch_size, is_train=False)" ] }, { "cell_type": "markdown", "id": "832601c4-cad0-4e90-9efa-8cb03e1da391", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "初始化模型参数。注意:$w$的标准差异常高,容易导致优化问题" ] }, { "cell_type": "code", "execution_count": 11, "id": "f38cd2b7-2537-4c6c-9157-804cdde17ff9", "metadata": { "origin_pos": 8, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def init_params():\n", " w = torch.normal(0, 1, size=(num_inputs, 1), requires_grad=True)\n", " b = torch.zeros(1, requires_grad=True)\n", " return [w, b]" ] }, { "cell_type": "markdown", "id": "e5875ed8-6ae1-46ca-82fc-7b0eaf870aeb", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "定义$L_2$范数惩罚" ] }, { "cell_type": "code", "execution_count": 12, "id": "9e70a101-c049-481f-888c-76c492432cc1", "metadata": { "origin_pos": 12, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def l2_penalty(w):\n", " return torch.sum(w.pow(2)) / 2" ] }, { "cell_type": "markdown", "id": "1dcb1126-4b76-416c-ac79-4813f759f890", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "定义训练代码实现" ] }, { "cell_type": "code", "execution_count": 13, "id": "970f04ea-eaab-436c-8044-6cc01390e530", "metadata": { "origin_pos": 16, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def train(lambd):\n", " w, b = init_params()\n", " net, loss = lambda X: d2l.linreg(X, w, b), d2l.squared_loss\n", " num_epochs, lr = 100, 0.003\n", " animator = d2l.Animator(xlabel='epochs', ylabel='loss', yscale='log',\n", " xlim=[5, num_epochs], legend=['train', 'test'])\n", " for epoch in range(num_epochs):\n", " for X, y in train_iter:\n", " # 广播机制使l2_penalty(w)成为一个长度为batch_size的向量\n", " l = loss(net(X), y) + lambd * l2_penalty(w) # 增加了L2范数惩罚项\n", " l.sum().backward()\n", " d2l.sgd([w, b], lr, batch_size)\n", " if (epoch + 1) % 5 == 0:\n", " animator.add(epoch + 1, (d2l.evaluate_loss(net, train_iter, loss),\n", " d2l.evaluate_loss(net, test_iter, loss)))\n", " print('w的L2范数是:', torch.norm(w).item())" ] }, { "cell_type": "markdown", "id": "14e879f7-bee3-49bf-acc8-0b33510de08f", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "忽略正则化直接训练。非常典型的欠拟合:训练误差还在下降阶段" ] }, { "cell_type": "code", "execution_count": 14, "id": "7670e9cc-28ea-4dc5-ad88-53189f3ce1b3", "metadata": { "origin_pos": 19, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "w的L2范数是: 13.431254386901855\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:18:19.916991\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "train(lambd=0)" ] }, { "cell_type": "markdown", "id": "ced15153-7f8c-4b6e-8114-0ed74d803dba", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "使用权重衰减。测试误差仍在下降阶段。" ] }, { "cell_type": "code", "execution_count": 15, "id": "20f01d1a-e18f-4af1-86c0-63d762d64cb1", "metadata": { "origin_pos": 21, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "w的L2范数是: 0.3891151249408722\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:18:27.934603\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "train(lambd=3)" ] }, { "cell_type": "markdown", "id": "75f01a5f-fe45-44d1-a7b8-1f57ebd4543c", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "简洁实现" ] }, { "cell_type": "code", "execution_count": 16, "id": "bf65bdb2-46ee-4618-b40b-6408fc54637b", "metadata": { "origin_pos": 27, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def train_concise(wd):\n", " net = nn.Sequential(nn.Linear(num_inputs, 1))\n", " for param in net.parameters():\n", " param.data.normal_() # 默认值恰好是(mean=0, std=1)\n", " loss = nn.MSELoss(reduction='none')\n", " num_epochs, lr = 100, 0.003\n", " # 偏置参数没有衰减\n", " trainer = torch.optim.SGD([\n", " {\"params\":net[0].weight,'weight_decay': wd},\n", " {\"params\":net[0].bias}], lr=lr)\n", " animator = d2l.Animator(xlabel='epochs', ylabel='loss', yscale='log',\n", " xlim=[5, num_epochs], legend=['train', 'test'])\n", " for epoch in range(num_epochs):\n", " for X, y in train_iter:\n", " trainer.zero_grad()\n", " l = loss(net(X), y)\n", " l.mean().backward()\n", " trainer.step()\n", " if (epoch + 1) % 5 == 0:\n", " animator.add(epoch + 1,\n", " (d2l.evaluate_loss(net, train_iter, loss),\n", " d2l.evaluate_loss(net, test_iter, loss)))\n", " print('w的L2范数:', net[0].weight.norm().item())" ] }, { "cell_type": "markdown", "id": "081cf988-6c1b-46fe-8af7-986b387581dc", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "这些图看起来与之前类似。但训练误差更快趋向水平:库函数有内部优化" ] }, { "cell_type": "code", "execution_count": 17, "id": "809320c9-7ef9-414c-9df2-a10b36f6e8ce", "metadata": { "origin_pos": 30, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "w的L2范数: 14.588664054870605\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:18:42.182739\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "train_concise(0)" ] }, { "cell_type": "code", "execution_count": 18, "id": "d256f3af-1647-4398-84d8-357e8c31629f", "metadata": { "origin_pos": 31, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "w的L2范数: 0.373628705739975\n" ] }, { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:18:54.170364\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "train_concise(3)" ] }, { "cell_type": "markdown", "id": "5dea2d34-f4a7-4ae1-87d3-09b1fdd7dadb", "metadata": { "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 暂退法(Dropout)" ] }, { "cell_type": "code", "execution_count": 19, "id": "4ba0141a-e3b9-43d8-9b38-78b2d788d707", "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 dropout_layer(X, dropout):\n", " \"\"\"以`dropout`的概率丢弃张量输入`X`中的元素\"\"\"\n", " assert 0 <= dropout <= 1\n", " if dropout == 1: # 在本情况中,所有元素都被丢弃\n", " return torch.zeros_like(X)\n", " if dropout == 0: # 在本情况中,所有元素都被保留\n", " return X\n", " # 为什么要掩码?并行计算两个分支\n", " mask = (torch.rand(X.shape) > dropout).float()\n", " return mask * X / (1.0 - dropout)" ] }, { "cell_type": "markdown", "id": "88bd8f63-4716-4314-bcaf-9fe4062b3fb1", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "测试`dropout_layer`函数" ] }, { "cell_type": "code", "execution_count": 20, "id": "1186ef28-60fb-4bf3-8c25-3db7e19d4aee", "metadata": { "origin_pos": 6, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "tensor([[ 0., 1., 2., 3., 4., 5., 6., 7.],\n", " [ 8., 9., 10., 11., 12., 13., 14., 15.]])\n", "tensor([[ 0., 1., 2., 3., 4., 5., 6., 7.],\n", " [ 8., 9., 10., 11., 12., 13., 14., 15.]])\n", "tensor([[ 0., 2., 0., 6., 8., 10., 0., 14.],\n", " [ 0., 0., 0., 0., 24., 0., 0., 0.]])\n", "tensor([[0., 0., 0., 0., 0., 0., 0., 0.],\n", " [0., 0., 0., 0., 0., 0., 0., 0.]])\n" ] } ], "source": [ "X= torch.arange(16, dtype = torch.float32).reshape((2, 8))\n", "print(X)\n", "print(dropout_layer(X, 0.))\n", "print(dropout_layer(X, 0.5))\n", "print(dropout_layer(X, 1.))" ] }, { "cell_type": "markdown", "id": "4f18f009-335f-429a-b794-9e28e653fa39", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "定义具有两个隐藏层的多层感知机,每个隐藏层包含256个单元" ] }, { "cell_type": "code", "execution_count": 21, "id": "26c4ba3a-48e1-4eac-86da-c4846e0ac04e", "metadata": { "origin_pos": 14, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "num_inputs, num_outputs, num_hiddens1, num_hiddens2 = 784, 10, 256, 256\n", "\n", "dropout1, dropout2 = 0.2, 0.5\n", "\n", "class Net(nn.Module):\n", " def __init__(self, num_inputs, num_outputs, num_hiddens1, num_hiddens2,\n", " is_training = True):\n", " super(Net, self).__init__()\n", " self.num_inputs = num_inputs\n", " self.training = is_training\n", " self.lin1 = nn.Linear(num_inputs, num_hiddens1)\n", " self.lin2 = nn.Linear(num_hiddens1, num_hiddens2)\n", " self.lin3 = nn.Linear(num_hiddens2, num_outputs)\n", " self.relu = nn.ReLU()\n", "\n", " def forward(self, X):\n", " H1 = self.relu(self.lin1(X.reshape((-1, self.num_inputs))))\n", " if self.training == True: # 只有在训练模型时才使用dropout\n", " H1 = dropout_layer(H1, dropout1) # 在第一个全连接层之后添加一个dropout层\n", " H2 = self.relu(self.lin2(H1))\n", " if self.training == True:\n", " H2 = dropout_layer(H2, dropout2) # 在第一个全连接层之后添加一个dropout层\n", " out = self.lin3(H2)\n", " return out\n", "\n", "\n", "net = Net(num_inputs, num_outputs, num_hiddens1, num_hiddens2)" ] }, { "cell_type": "markdown", "id": "48ea088d-77a0-4373-af8c-2545e7a6cccf", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练和测试" ] }, { "cell_type": "code", "execution_count": 22, "id": "c08ed34e-a50e-44dc-832c-6bc1341cec58", "metadata": { "origin_pos": 18, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:22:37.456317\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "num_epochs, lr, batch_size = 10, 0.5, 256\n", "loss = nn.CrossEntropyLoss(reduction='none')\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size)\n", "trainer = torch.optim.SGD(net.parameters(), lr=lr)\n", "d2l.train_ch3(net, train_iter, test_iter, loss, num_epochs, trainer)" ] }, { "cell_type": "markdown", "id": "539c3a27-31ef-4dd6-9655-3d161aaafb79", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "简洁实现" ] }, { "cell_type": "code", "execution_count": 23, "id": "3d46073f-8828-4c26-b1e5-98c0fcddae43", "metadata": { "origin_pos": 22, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "net = nn.Sequential(nn.Flatten(),\n", " nn.Linear(784, 256),\n", " nn.ReLU(),\n", " nn.Dropout(dropout1), # 在第一个全连接层之后添加一个dropout层\n", " nn.Linear(256, 256),\n", " nn.ReLU(),\n", " nn.Dropout(dropout2), # 在第二个全连接层之后添加一个dropout层\n", " nn.Linear(256, 10))\n", "\n", "def init_weights(m):\n", " if type(m) == nn.Linear:\n", " nn.init.normal_(m.weight, std=0.01)\n", "\n", "net.apply(init_weights);" ] }, { "cell_type": "markdown", "id": "49b9f82e-14c3-433e-8ee8-3637ffafb2b1", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "对模型进行训练和测试" ] }, { "cell_type": "code", "execution_count": 24, "id": "66683b72-aa3c-4cb8-97ef-ced571cca688", "metadata": { "origin_pos": 26, "tab": [ "pytorch" ], "tags": [] }, "outputs": [ { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-28T09:26:30.099732\n", " image/svg+xml\n", " \n", " \n", " Matplotlib v3.5.1, 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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "trainer = torch.optim.SGD(net.parameters(), lr=lr)\n", "d2l.train_ch3(net, train_iter, test_iter, loss, num_epochs, trainer)" ] } ], "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 }