{
"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"
],
"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"
],
"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"
],
"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"
],
"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"
],
"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"
],
"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"
],
"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"
],
"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"
],
"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
}