{ "cells": [ { "cell_type": "markdown", "id": "8467bf5b-8c23-4084-adb3-a0e0c3b3e7a6", "metadata": {}, "source": [ "## 目录\n", "\n", "- 感知机分类器\n", "- 逻辑回归分类器\n", "- 图像分类数据集\n", "- softmax回归的从零开始实现\n", "- softmax回归的简洁实现" ] }, { "cell_type": "markdown", "id": "412d121a-f1f9-4150-ad97-39c5b0edc498", "metadata": { "tags": [] }, "source": [ "## 感知机分类器" ] }, { "cell_type": "markdown", "id": "13474e54-4e85-4abf-abc0-bead50177068", "metadata": {}, "source": [ "Perceptron是一个线性分类器。它的工作方式与没有隐藏层的神经网络相同(只有输入和输出)。\n", "\n", "![perceptron](perceptron.png)\n", "\n", "首先,它对给定的数据集进行权重训练,然后通过网络对一个新项目进行分类。\n", "\n", "> 注意,在分类问题中,每个节点代表一个类别。最终的分类是具有最大输出值的类/节点。" ] }, { "cell_type": "markdown", "id": "d05ceed2-8d2e-4bc6-ac1c-da2bb67ecb89", "metadata": { "tags": [] }, "source": [ "### 实现\n", "\n", "`PerceptionLinearLearner` 用于训练(计算)给定数据集的权重。\n", "\n", "函数`predict`用于对一个新的项目进行分类:该函数计算项目与外层每个节点的计算权重的(代数)点积。然后,它选择一个最大的值,将项目归入相应的类别。" ] }, { "cell_type": "code", "execution_count": 2, "id": "c3423455-3a8c-44e5-a9aa-592b7c6e285d", "metadata": { "tags": [] }, "outputs": [ { "data": { "text/html": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", "\n", "\n", "

\n", "\n", "
class PerceptionLinearLearner(LinearClassifier):\n",
       "    """\n",
       "    Perception linear classifier: hard threshold\n",
       "    """\n",
       "    def __init__(self, dataset, learning_rate=0.01, epochs=100):\n",
       "        self.idx_i = dataset.inputs\n",
       "        self.idx_t = dataset.target\n",
       "        self.examples = dataset.examples\n",
       "        self.num_examples = len(self.examples)\n",
       "        # initialize random weights\n",
       "        self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)\n",
       "        # learning loop\n",
       "        self.learn(learning_rate, epochs)\n",
       "        \n",
       "    def learn(self, learning_rate, epochs):\n",
       "        """ learning loop """\n",
       "        def loss(example, w, idx_i, idx_t):\n",
       "            """ error: difference between estimation and true value """\n",
       "            raise NotImplementedError\n",
       "        \n",
       "        def update(w, learning_rate, err, X_col, num_examples):\n",
       "            """ update weights """\n",
       "            raise NotImplementedError\n",
       "\n",
       "        def homogeneous(num_examples):\n",
       "            """ build homogeneous coordinates """\n",
       "            raise NotImplementedError\n",
       "\n",
       "        raise NotImplementedError\n",
       "\n",
       "    def predict(self, x):\n",
       "        """ make prediction """\n",
       "        return int(np.dot(self.w, [1] + x))\n",
       "
\n", "\n", "\n" ], "text/plain": [ "" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "import sys\n", "sys.path.insert(1, '../')\n", "from utils.dataset4learners import *\n", "from LinearClassifier_3 import *\n", "\n", "psource(PerceptionLinearLearner)" ] }, { "cell_type": "markdown", "id": "8c4fac32-dfa4-46d4-a2c4-8e5724f1def3", "metadata": {}, "source": [ "#### 实现要点\n", "\n", "```py\n", "def learn(self, learning_rate, epochs):\n", " \"\"\" learning loop \"\"\"\n", " def loss(example, w, idx_i, idx_t):\n", " \"\"\" error: difference between estimation and true value \"\"\"\n", " x = example的齐次坐标表示\n", " y = w * x,即预测值\n", " t = 真实值\n", " return y - t\n", "\n", " def update(w, learning_rate, err, X_col, num_examples):\n", " \"\"\" update weights \"\"\"\n", " 遍历w[i]:\n", " w[i] = w[i] - 学习率 * 误差值 * 自变量X[i] / 样本数量)\n", "\n", " def homogeneous(num_examples):\n", " \"\"\" build homogeneous coordinates \"\"\"\n", " return 齐次坐标表示\n", "\n", " X_col = 齐次坐标表示\n", " for epoch in range(epochs):\n", " err = []\n", " # pass over all examples\n", " for example in self.examples:\n", " err.append(loss(example, w, idx_i, idx_t))\n", "\n", " # update weights\n", " update(w, 学习率, err, X_col, 样本数量)\n", "```" ] }, { "cell_type": "markdown", "id": "d8f78b41-5136-4563-8359-67c1dae52642", "metadata": {}, "source": [ "> 注意,Perceptron是一个单层的神经网络,在后面课程中讲授。" ] }, { "cell_type": "markdown", "id": "ca1d7196-b054-4718-9402-c0a377fc53a6", "metadata": {}, "source": [ "### 例子\n", "\n", "我们将在`iris`数据集上训练感知机。\n", "\n", "尽管`BackPropagationLearner`使用的是整数索引而不是字符串,我们需要将类名转换成整数。" ] }, { "cell_type": "markdown", "id": "c37147ed-40e7-4b05-83e3-5c4db4c0a039", "metadata": { "tags": [] }, "source": [ "```py\n", "iris = DataSet(name=\"iris\")\n", "iris.classes_to_numbers()\n", "\n", "perceptron = PerceptionLinearLearner(iris)\n", "print(perceptron.predict([5, 3, 1, 0.1]))\n", "```" ] }, { "cell_type": "markdown", "id": "debca21b-5262-49ac-8502-8e692e9d26af", "metadata": {}, "source": [ "正确的输出是0,这意味着该物品属于第一类,\"setosa\"。注意Perceptron算法并不完美,可能会产生错误的分类。" ] }, { "cell_type": "markdown", "id": "c6ab0a90-bf04-497b-9795-31e62155054d", "metadata": { "tags": [] }, "source": [ "## 逻辑回归分类器" ] }, { "cell_type": "markdown", "id": "55e14e99-4e16-4ef2-a652-43a51364c144", "metadata": {}, "source": [ "逻辑回归将感知机采用的硬性阈值函数改为柔性阈值,可以更好地对分界区域的不确定性建模。\n", "\n", "### 实现\n" ] }, { "cell_type": "code", "execution_count": 3, "id": "360a50d6-647e-4e48-8b66-5feb22db7dd0", "metadata": { "tags": [] }, "outputs": [ { "data": { "text/html": [ "\n", "\n", "\n", "\n", " \n", " \n", " \n", "\n", "\n", "

\n", "\n", "
class LogisticLinearLeaner(LinearClassifier):\n",
       "    def __init__(self, dataset, learning_rate=0.01, epochs=100):\n",
       "        self.idx_i = dataset.inputs\n",
       "        self.idx_t = dataset.target\n",
       "        self.examples = dataset.examples\n",
       "        self.num_examples = len(self.examples)\n",
       "        # initialize random weights\n",
       "        self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)\n",
       "        # learning loop\n",
       "        self.learn(learning_rate, epochs)\n",
       "        \n",
       "    def learn(self, learning_rate, epochs):\n",
       "        """ learning loop """\n",
       "        def loss(example, w, idx_i, idx_t, h):\n",
       "            """ error: difference between estimation and true value """\n",
       "            raise NotImplementedError\n",
       "        \n",
       "        def update(w, learning_rate, err, h, X_col, num_examples):\n",
       "            """ update weights """\n",
       "            raise NotImplementedError\n",
       "\n",
       "        def homogeneous(num_examples):\n",
       "            """ build homogeneous coordinates """\n",
       "            raise NotImplementedError\n",
       "\n",
       "        raise NotImplementedError\n",
       "\n",
       "    def predict(self, x):\n",
       "        """ make prediction """\n",
       "        return int(np.dot(self.w, [1] + x))\n",
       "
\n", "\n", "\n" ], "text/plain": [ "" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "psource(LogisticLinearLeaner)" ] }, { "cell_type": "markdown", "id": "64d23b34-7203-4abe-bb3b-b3e65f192d62", "metadata": {}, "source": [ "#### 实现要点\n", "\n", "```py\n", "def learn(self, learning_rate, epochs):\n", " \"\"\" learning loop \"\"\"\n", " def loss(example, w, idx_i, idx_t, h):\n", " \"\"\" error: difference between estimation and true value \"\"\"\n", " x = example的齐次坐标表示\n", " y = Sigmoid(w * x),即预测值\n", " h.append(Sigmoid().derivative(y)):收集梯度值\n", " t = 真实值\n", " return y - t\n", "\n", " def update(w, learning_rate, err, h, X_col, num_examples):\n", " \"\"\" update weights \"\"\"\n", " 遍历w[i]:\n", " w[i] = w[i] - 学习率 * 误差值err[i] * 梯度值h[i] / 样本数量)\n", "\n", " def homogeneous(num_examples):\n", " \"\"\" build homogeneous coordinates \"\"\"\n", " 同感知机\n", "\n", " 同感知机\n", "```" ] }, { "cell_type": "markdown", "id": "34dde5da-d093-4021-ae8b-7a38cd5b11dc", "metadata": {}, "source": [ "该算法首先为输入变量分配一些随机权重,然后根据计算的误差更新每个变量的权重。最后用更新后的权重进行预测。" ] }, { "cell_type": "markdown", "id": "2b3f1358-fe11-4499-8fc1-f8a7fb6ef826", "metadata": { "tags": [] }, "source": [ "```py\n", "iris = DataSet(name=\"iris\")\n", "iris.classes_to_numbers()\n", "\n", "logisticer = LogisticLinearLeaner(iris)\n", "print(logisticer.predict([5, 3, 1, 0.1]))\n", "```" ] }, { "cell_type": "markdown", "id": "b0e066c5-554b-4286-ae5f-1f71dd7b95a8", "metadata": { "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## 图像分类数据集" ] }, { "cell_type": "markdown", "id": "19b612e3-6ede-4ffa-a9dd-14a0480a2c8f", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "MNIST数据集是图像分类中广泛使用的数据集之一,但作为基准数据集过于简单。\n", "\n", "我们将使用类似但更复杂的Fashion-MNIST数据集" ] }, { "cell_type": "code", "execution_count": 6, "id": "451140a0-6c93-45e0-8e19-8ffb51099e77", "metadata": { "origin_pos": 2, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "%matplotlib inline\n", "import torch\n", "import torchvision\n", "from torch.utils import data\n", "from torchvision import transforms\n", "from d2l import torch as d2l\n", "\n", "d2l.use_svg_display()" ] }, { "cell_type": "markdown", "id": "b5386450-4cdd-4c05-9524-25505b67f22c", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "通过框架中的内置函数将Fashion-MNIST数据集下载并读取到内存中" ] }, { "cell_type": "code", "execution_count": 8, "id": "af7b66a6-a2ca-4ec3-981e-8acaf5880cbe", "metadata": { "origin_pos": 9, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/train-images-idx3-ubyte.gz\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/train-images-idx3-ubyte.gz to ../data/FashionMNIST/raw/train-images-idx3-ubyte.gz\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "100.0%\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "Extracting ../data/FashionMNIST/raw/train-images-idx3-ubyte.gz to ../data/FashionMNIST/raw\n", "\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/train-labels-idx1-ubyte.gz\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/train-labels-idx1-ubyte.gz to ../data/FashionMNIST/raw/train-labels-idx1-ubyte.gz\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "100.0%\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "Extracting ../data/FashionMNIST/raw/train-labels-idx1-ubyte.gz to ../data/FashionMNIST/raw\n", "\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/t10k-images-idx3-ubyte.gz\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/t10k-images-idx3-ubyte.gz to ../data/FashionMNIST/raw/t10k-images-idx3-ubyte.gz\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "100.0%\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "Extracting ../data/FashionMNIST/raw/t10k-images-idx3-ubyte.gz to ../data/FashionMNIST/raw\n", "\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/t10k-labels-idx1-ubyte.gz\n", "Downloading http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/t10k-labels-idx1-ubyte.gz to ../data/FashionMNIST/raw/t10k-labels-idx1-ubyte.gz\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "100.0%\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "Extracting ../data/FashionMNIST/raw/t10k-labels-idx1-ubyte.gz to ../data/FashionMNIST/raw\n", "\n" ] }, { "data": { "text/plain": [ "(60000, 10000)" ] }, "execution_count": 8, "metadata": {}, "output_type": "execute_result" } ], "source": [ "# 通过ToTensor实例将图像数据从PIL类型变换成32位浮点数格式,\n", "# 并除以255使得所有像素的数值均在0到1之间\n", "trans = transforms.ToTensor()\n", "mnist_train = torchvision.datasets.FashionMNIST(\n", " root=\"../data\", train=True, transform=trans, download=True)\n", "mnist_test = torchvision.datasets.FashionMNIST(\n", " root=\"../data\", train=False, transform=trans, download=True)\n", "\n", "len(mnist_train), len(mnist_test)" ] }, { "cell_type": "code", "execution_count": 9, "id": "e5de50d6-1fe0-407d-87a9-8eec3c6c475d", "metadata": { "origin_pos": 12, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "torch.Size([1, 28, 28])" ] }, "execution_count": 9, "metadata": {}, "output_type": "execute_result" } ], "source": [ "mnist_train[0][0].shape" ] }, { "cell_type": "markdown", "id": "87d0348b-ec5d-4299-9d8c-08987de460f1", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "两个可视化数据集的函数" ] }, { "cell_type": "code", "execution_count": 10, "id": "0557638e-71b8-40ce-95e6-5d61038f1ffc", "metadata": { "origin_pos": 17, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def get_fashion_mnist_labels(labels): \n", " \"\"\"返回Fashion-MNIST数据集的文本标签\"\"\"\n", " text_labels = ['t-shirt', 'trouser', 'pullover', 'dress', 'coat',\n", " 'sandal', 'shirt', 'sneaker', 'bag', 'ankle boot']\n", " return [text_labels[int(i)] for i in labels]\n", "\n", "def show_images(imgs, num_rows, num_cols, titles=None, scale=1.5): \n", " \"\"\"绘制图像列表\"\"\"\n", " figsize = (num_cols * scale, num_rows * scale)\n", " _, axes = d2l.plt.subplots(num_rows, num_cols, figsize=figsize)\n", " axes = axes.flatten()\n", " for i, (ax, img) in enumerate(zip(axes, imgs)):\n", " if torch.is_tensor(img):\n", " ax.imshow(img.numpy()) # 图片张量\n", " else:\n", " ax.imshow(img) # PIL图片\n", " ax.axes.get_xaxis().set_visible(False)\n", " ax.axes.get_yaxis().set_visible(False)\n", " if titles:\n", " ax.set_title(titles[i])\n", " return axes" ] }, { "cell_type": "markdown", "id": "0c9255c0-c447-4232-a1e8-9cbdd63f2a62", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "几个样本的图像及其相应的标签" ] }, { "cell_type": "code", "execution_count": 11, "id": "6bd5a31f-a289-4ab2-9ee5-bbe8322ac216", "metadata": { "origin_pos": 20, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-26T14:13:22.029419\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", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \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": [ "X, y = next(iter(data.DataLoader(mnist_train, batch_size=18)))\n", "show_images(X.reshape(18, 28, 28), 2, 9, titles=get_fashion_mnist_labels(y));" ] }, { "cell_type": "markdown", "id": "9ce47a48-d8e3-46a6-8e2b-009c4e299341", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "读取一小批量数据,大小为`batch_size`" ] }, { "cell_type": "code", "execution_count": 12, "id": "b15b88d6-d561-4971-9c20-b30d2618da04", "metadata": { "origin_pos": 27, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "'20.21 sec'" ] }, "execution_count": 12, "metadata": {}, "output_type": "execute_result" } ], "source": [ "batch_size = 256\n", "\n", "def get_dataloader_workers(): \n", " \"\"\"使用4个进程来读取数据\"\"\"\n", " return 4\n", "\n", "train_iter = data.DataLoader(mnist_train, batch_size, shuffle=True,\n", " num_workers=get_dataloader_workers())\n", "\n", "timer = d2l.Timer()\n", "for X, y in train_iter:\n", " continue\n", "f'{timer.stop():.2f} sec'" ] }, { "cell_type": "markdown", "id": "1e0a3486-8788-4a87-b9b9-cb327ba7d2ca", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "定义`load_data_fashion_mnist`函数" ] }, { "cell_type": "code", "execution_count": 13, "id": "5f88e2a1-88b6-49c3-a8b8-2ac9602a71a8", "metadata": { "origin_pos": 33, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def load_data_fashion_mnist(batch_size, resize=None): \n", " \"\"\"下载Fashion-MNIST数据集,然后将其加载到内存中\"\"\"\n", " trans = [transforms.ToTensor()]\n", " if resize:\n", " trans.insert(0, transforms.Resize(resize))\n", " trans = transforms.Compose(trans)\n", " mnist_train = torchvision.datasets.FashionMNIST(\n", " root=\"../data\", train=True, transform=trans, download=True)\n", " mnist_test = torchvision.datasets.FashionMNIST(\n", " root=\"../data\", train=False, transform=trans, download=True)\n", " return (data.DataLoader(mnist_train, batch_size, shuffle=True,\n", " num_workers=get_dataloader_workers()),\n", " data.DataLoader(mnist_test, batch_size, shuffle=False,\n", " num_workers=get_dataloader_workers()))" ] }, { "cell_type": "code", "execution_count": 14, "id": "b60ce5c6-c1e3-4d51-9d2a-d35e114773ed", "metadata": { "origin_pos": 33, "slideshow": { "slide_type": "slide" }, "tab": [ "pytorch" ] }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "torch.Size([32, 1, 64, 64]) torch.float32 torch.Size([32]) torch.int64\n" ] } ], "source": [ "train_iter, test_iter = load_data_fashion_mnist(32, resize=64)\n", "for X, y in train_iter:\n", " print(X.shape, X.dtype, y.shape, y.dtype)\n", " break" ] }, { "cell_type": "markdown", "id": "2b3bc703-261c-4f28-8b13-10599cdfb648", "metadata": { "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## softmax回归的从零开始实现" ] }, { "cell_type": "markdown", "id": "461c158e-d492-4138-a1dc-a6dee27856c3", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "就像我们从零开始实现线性回归一样,你应该知道实现softmax回归的细节" ] }, { "cell_type": "code", "execution_count": 15, "id": "7c6ea6a5-00d4-4032-8b69-a23295a1c3b3", "metadata": { "origin_pos": 4, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from IPython import display\n", "from d2l import torch as d2l\n", "\n", "batch_size = 256\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size)" ] }, { "cell_type": "markdown", "id": "79517ad9-cddb-4ee2-9f61-43e0b6bb3341", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "将展平每个图像,把它们看作长度为784的向量。\n", "因为我们的数据集有10个类别,所以网络输出维度为10" ] }, { "cell_type": "code", "execution_count": 16, "id": "bf6ec527-910e-42ce-b24c-c72da6fa9c74", "metadata": { "origin_pos": 7, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "num_inputs = 784\n", "num_outputs = 10\n", "\n", "W = torch.normal(0, 0.01, size=(num_inputs, num_outputs), requires_grad=True)\n", "b = torch.zeros(num_outputs, requires_grad=True)" ] }, { "cell_type": "markdown", "id": "2499dc12-833a-4851-b5de-859d86f428c5", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "给定一个矩阵`X`,我们可以对所有元素求和" ] }, { "cell_type": "code", "execution_count": 17, "id": "e7fbccdf-9d88-4244-963a-33ebe05ba5ae", "metadata": { "origin_pos": 10, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "(tensor([[5., 7., 9.]]),\n", " tensor([[ 6.],\n", " [15.]]))" ] }, "execution_count": 17, "metadata": {}, "output_type": "execute_result" } ], "source": [ "X = torch.tensor([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])\n", "X.sum(0, keepdim=True), X.sum(1, keepdim=True)" ] }, { "cell_type": "markdown", "id": "7338e047-94ed-438d-a0b0-2477692b0c45", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "实现softmax\n", "\n", "$$\n", "\\mathrm{softmax}(\\mathbf{X})_{ij} = \\frac{\\exp(\\mathbf{X}_{ij})}{\\sum_k \\exp(\\mathbf{X}_{ik})}.\n", "$$" ] }, { "cell_type": "code", "execution_count": 18, "id": "073aadb5-9a24-44e3-ac4b-3bae84a3a9e2", "metadata": { "origin_pos": 14, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def softmax(X):\n", " X_exp = torch.exp(X)\n", " partition = X_exp.sum(1, keepdim=True)\n", " return X_exp / partition # 广播机制" ] }, { "cell_type": "markdown", "id": "17706f99-24f6-446f-9f4a-a17368170e3f", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "我们将每个元素变成一个非负数。\n", "此外,依据概率原理,每行总和为1" ] }, { "cell_type": "code", "execution_count": 19, "id": "2ff371fa-6e82-4540-8c7d-0bba75225e1f", "metadata": { "origin_pos": 16, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "(tensor([[0.3532, 0.0347, 0.0805, 0.0395, 0.4921],\n", " [0.6108, 0.0814, 0.1468, 0.0484, 0.1126]]),\n", " tensor([1., 1.]))" ] }, "execution_count": 19, "metadata": {}, "output_type": "execute_result" } ], "source": [ "X = torch.normal(0, 1, (2, 5))\n", "X_prob = softmax(X)\n", "X_prob, X_prob.sum(1)" ] }, { "cell_type": "markdown", "id": "02ba4bf7-0785-4958-afbf-9519f33ed8af", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "实现softmax回归模型" ] }, { "cell_type": "code", "execution_count": 20, "id": "ec1c7adf-b253-48cb-8abc-395098c3635e", "metadata": { "origin_pos": 19, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def net(X):\n", " return softmax(torch.matmul(X.reshape((-1, W.shape[0])), W) + b)" ] }, { "cell_type": "markdown", "id": "e1ee288d-3e6e-41de-b4f5-c020e83bd14e", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "创建一个数据样本`y_hat`,其中包含2个样本在3个类别的预测概率,\n", "以及它们对应的标签`y`。\n", "使用`y`作为`y_hat`中概率的索引" ] }, { "cell_type": "code", "execution_count": 21, "id": "83632ade-00af-4692-887f-cd13041c814d", "metadata": { "origin_pos": 21, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "tensor([0.1000, 0.5000])" ] }, "execution_count": 21, "metadata": {}, "output_type": "execute_result" } ], "source": [ "y = torch.tensor([0, 2]) # 在第二个样本中,第三类是正确的预测\n", "y_hat = torch.tensor([[0.1, 0.3, 0.6], [0.3, 0.2, 0.5]])\n", "y_hat[[0, 1], y] # 使用y作为y_hat中概率的索引" ] }, { "cell_type": "markdown", "id": "c40efbc6-376d-4557-8b46-35c6aff4f979", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "实现交叉熵损失函数" ] }, { "cell_type": "code", "execution_count": 22, "id": "f070bea1-2b36-402e-828f-29d2446b0c56", "metadata": { "origin_pos": 24, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "tensor([2.3026, 0.6931])" ] }, "execution_count": 22, "metadata": {}, "output_type": "execute_result" } ], "source": [ "def cross_entropy(y_hat, y):\n", " \"\"\" 避免低效的for循环 \"\"\"\n", " return - torch.log(y_hat[range(len(y_hat)), y])\n", "\n", "cross_entropy(y_hat, y) # 损失函数,越低越好" ] }, { "cell_type": "markdown", "id": "9374a06e-27ed-42fa-a1c9-0c6d90dd4a48", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "将预测类别与真实`y`元素进行比较" ] }, { "cell_type": "code", "execution_count": 23, "id": "97a56aa3-f780-4f32-98ca-086d6e1024f0", "metadata": { "origin_pos": 29, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "0.5" ] }, "execution_count": 23, "metadata": {}, "output_type": "execute_result" } ], "source": [ "def accuracy(y_hat, y): \n", " \"\"\"计算预测正确的数量\"\"\"\n", " if len(y_hat.shape) > 1 and y_hat.shape[1] > 1:\n", " y_hat = y_hat.argmax(axis=1)\n", " cmp = y_hat.type(y.dtype) == y # 数据类型转成一致\n", " return float(cmp.type(y.dtype).sum())\n", "\n", "accuracy(y_hat, y) / len(y)" ] }, { "cell_type": "markdown", "id": "585af5a4-d317-42ee-a5f1-255a594856fa", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "我们可以评估在任意模型`net`的精度" ] }, { "cell_type": "code", "execution_count": 24, "id": "540c6f10-7592-4fd3-8221-0ef56f88a0bf", "metadata": { "origin_pos": 32, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def evaluate_accuracy(net, data_iter): \n", " \"\"\"计算在指定数据集上模型的精度\"\"\"\n", " if isinstance(net, torch.nn.Module):\n", " net.eval() # 将模型设置为评估模式\n", " metric = Accumulator(2) # 正确预测数、预测总数\n", " with torch.no_grad():\n", " for X, y in data_iter:\n", " metric.add(accuracy(net(X), y), y.numel())\n", " return metric[0] / metric[1]" ] }, { "cell_type": "markdown", "id": "d7593436-0ca0-4fc3-ae2c-7d7311ddb9f6", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "`Accumulator`实例中创建了2个变量,\n", "分别用于存储正确预测的数量和预测的总数量" ] }, { "cell_type": "code", "execution_count": 25, "id": "bb38e91c-b8aa-4912-862e-9080a0bbda66", "metadata": { "origin_pos": 36, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "text/plain": [ "0.1004" ] }, "execution_count": 25, "metadata": {}, "output_type": "execute_result" } ], "source": [ "class Accumulator: \n", " \"\"\"在n个变量上累加\"\"\"\n", " def __init__(self, n):\n", " self.data = [0.0] * n\n", "\n", " def add(self, *args):\n", " self.data = [a + float(b) for a, b in zip(self.data, args)]\n", "\n", " def reset(self):\n", " self.data = [0.0] * len(self.data)\n", "\n", " def __getitem__(self, idx):\n", " return self.data[idx]\n", "\n", "evaluate_accuracy(net, test_iter) # 随机权重初始化,随机猜测应该接近0.1" ] }, { "cell_type": "markdown", "id": "9ca7e9be-2270-481f-b0c9-882b79efed78", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "Softmax回归的训练" ] }, { "cell_type": "code", "execution_count": 26, "id": "c4c903ea-72eb-49a9-9888-8e8d5fd6f8a6", "metadata": { "origin_pos": 39, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def train_epoch_ch3(net, train_iter, loss, updater): \n", " \"\"\"训练模型一个迭代周期(定义见第3章)\"\"\"\n", " if isinstance(net, torch.nn.Module):\n", " net.train() # 将模型设置为训练模式\n", " metric = Accumulator(3) # 训练损失总和、训练准确度总和、样本数\n", " for X, y in train_iter:\n", " # 计算梯度并更新参数\n", " y_hat = net(X)\n", " l = loss(y_hat, y)\n", " if isinstance(updater, torch.optim.Optimizer):\n", " # 使用PyTorch内置的优化器和损失函数\n", " updater.zero_grad()\n", " l.mean().backward()\n", " updater.step()\n", " else: # 使用定制的优化器和损失函数\n", " l.sum().backward()\n", " updater(X.shape[0])\n", " metric.add(float(l.sum()), accuracy(y_hat, y), y.numel())\n", " # 返回训练损失和训练精度\n", " return metric[0] / metric[2], metric[1] / metric[2]" ] }, { "cell_type": "markdown", "id": "dd06ef05-3914-4ecd-ba5e-59c4fea740c0", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "定义一个在动画中绘制数据的实用程序类" ] }, { "cell_type": "code", "execution_count": 27, "id": "3ecf8ca6-1797-497d-abc3-867a591bf464", "metadata": { "origin_pos": 42, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "class Animator: \n", " \"\"\"在动画中绘制数据\"\"\"\n", " def __init__(self, xlabel=None, ylabel=None, legend=None, xlim=None,\n", " ylim=None, xscale='linear', yscale='linear',\n", " fmts=('-', 'm--', 'g-.', 'r:'), nrows=1, ncols=1,\n", " figsize=(3.5, 2.5)):\n", " if legend is None:\n", " legend = [] # 增量地绘制多条线\n", " d2l.use_svg_display()\n", " self.fig, self.axes = d2l.plt.subplots(nrows, ncols, figsize=figsize)\n", " if nrows * ncols == 1:\n", " self.axes = [self.axes, ]\n", " # 使用lambda函数捕获参数\n", " self.config_axes = lambda: d2l.set_axes(\n", " self.axes[0], xlabel, ylabel, xlim, ylim, xscale, yscale, legend)\n", " self.X, self.Y, self.fmts = None, None, fmts\n", "\n", " def add(self, x, y):\n", " # 向图表中添加多个数据点\n", " if not hasattr(y, \"__len__\"):\n", " y = [y]\n", " n = len(y)\n", " if not hasattr(x, \"__len__\"):\n", " x = [x] * n\n", " if not self.X:\n", " self.X = [[] for _ in range(n)]\n", " if not self.Y:\n", " self.Y = [[] for _ in range(n)]\n", " for i, (a, b) in enumerate(zip(x, y)):\n", " if a is not None and b is not None:\n", " self.X[i].append(a)\n", " self.Y[i].append(b)\n", " self.axes[0].cla()\n", " for x, y, fmt in zip(self.X, self.Y, self.fmts):\n", " self.axes[0].plot(x, y, fmt)\n", " self.config_axes()\n", " display.display(self.fig)\n", " display.clear_output(wait=True)" ] }, { "cell_type": "markdown", "id": "91290f6c-081c-430f-8563-73166defe322", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练函数" ] }, { "cell_type": "code", "execution_count": 28, "id": "500cf101-f3cb-44c5-abc6-bbfb770e7953", "metadata": { "origin_pos": 44, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "def train_ch3(net, train_iter, test_iter, loss, num_epochs, updater): \n", " \"\"\"训练模型(定义见第3章)\"\"\"\n", " animator = Animator(xlabel='epoch', xlim=[1, num_epochs], ylim=[0.3, 0.9],\n", " legend=['train loss', 'train acc', 'test acc'])\n", " for epoch in range(num_epochs):\n", " train_metrics = train_epoch_ch3(net, train_iter, loss, updater)\n", " test_acc = evaluate_accuracy(net, test_iter)\n", " animator.add(epoch + 1, train_metrics + (test_acc,))\n", " train_loss, train_acc = train_metrics\n", " assert train_loss < 0.5, train_loss\n", " assert train_acc <= 1 and train_acc > 0.7, train_acc\n", " assert test_acc <= 1 and test_acc > 0.7, test_acc" ] }, { "cell_type": "markdown", "id": "09e83970-88ed-4d77-a10a-92fad6a89449", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "小批量随机梯度下降来优化模型的损失函数" ] }, { "cell_type": "code", "execution_count": 29, "id": "ed05114c-e815-44c4-bcce-10880f11286a", "metadata": { "origin_pos": 46, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "lr = 0.1\n", "\n", "def updater(batch_size):\n", " return d2l.sgd([W, b], lr, batch_size)" ] }, { "cell_type": "markdown", "id": "6e153647-c2d1-4497-a261-42dfc29ccfd9", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "训练模型10个迭代周期" ] }, { "cell_type": "code", "execution_count": 30, "id": "cc4149fd-a083-4a50-a164-d71513ea5edb", "metadata": { "origin_pos": 49, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-26T14:17:07.123256\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 = 10\n", "train_ch3(net, train_iter, test_iter, cross_entropy, num_epochs, updater)" ] }, { "cell_type": "markdown", "id": "2ec84ab6-8482-46a9-ab68-522f2faba21c", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "对图像进行分类预测" ] }, { "cell_type": "code", "execution_count": 31, "id": "f5909439-a31c-4ade-aab8-5bb3edb9fa23", "metadata": { "origin_pos": 51, "tab": [ "pytorch" ] }, "outputs": [ { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-26T14:17:16.183922\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" ], "text/plain": [ "
" ] }, "metadata": {}, "output_type": "display_data" } ], "source": [ "def predict_ch3(net, test_iter, n=8): \n", " \"\"\"预测标签(定义见第3章)\"\"\"\n", " for X, y in test_iter:\n", " break\n", " trues = d2l.get_fashion_mnist_labels(y)\n", " preds = d2l.get_fashion_mnist_labels(net(X).argmax(axis=1))\n", " titles = [true +'\\n' + pred for true, pred in zip(trues, preds)]\n", " d2l.show_images(\n", " X[0:n].reshape((n, 28, 28)), 1, n, titles=titles[0:n])\n", "\n", "predict_ch3(net, test_iter)" ] }, { "cell_type": "markdown", "id": "fb05ae93-ca43-4978-b77b-cca182c45c8f", "metadata": { "slideshow": { "slide_type": "-" }, "tags": [] }, "source": [ "## softmax回归的简洁实现" ] }, { "cell_type": "markdown", "id": "a7da7eb8-6bad-4d17-83a7-24ec09779ef0", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "通过深度学习框架的高级API能够使实现softmax回归变得更加容易" ] }, { "cell_type": "code", "execution_count": 32, "id": "b598a5ff-1618-4bc3-a8aa-c627a3d3caa7", "metadata": { "origin_pos": 4, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "import torch\n", "from torch import nn\n", "from d2l import torch as d2l\n", "\n", "batch_size = 256\n", "train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size)" ] }, { "cell_type": "markdown", "id": "dd341874-7652-4afe-aa0a-41207311318f", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "Softmax回归的输出层是一个全连接层" ] }, { "cell_type": "code", "execution_count": 33, "id": "ee566c3d-2669-4149-bdcb-d9832ed4c4d5", "metadata": { "origin_pos": 7, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "# PyTorch不会隐式地调整输入的形状。因此,\n", "# 我们在线性层前定义了展平层(flatten),来调整网络输入的形状\n", "net = nn.Sequential(nn.Flatten(), nn.Linear(784, 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": "1330ec1d-3c66-46a0-8422-948553a33c41", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "在交叉熵损失函数中传递未规范化的预测,并同时计算softmax及其对数" ] }, { "cell_type": "code", "execution_count": 34, "id": "fcac843e-7436-4e88-bdf3-480bdfa12f26", "metadata": { "origin_pos": 11, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "loss = nn.CrossEntropyLoss(reduction='none')" ] }, { "cell_type": "markdown", "id": "1c7344bf-3aa5-4a19-8b25-b0296f9e238d", "metadata": { "slideshow": { "slide_type": "-" } }, "source": [ "使用学习率为0.1的小批量随机梯度下降作为优化算法" ] }, { "cell_type": "code", "execution_count": 35, "id": "3c5df7de-6a41-4069-bab3-3b518da1f30f", "metadata": { "origin_pos": 15, "tab": [ "pytorch" ] }, "outputs": [], "source": [ "trainer = torch.optim.SGD(net.parameters(), lr=0.1)" ] }, { "cell_type": "markdown", "id": "eef015d8-1fbc-4dd1-b1c0-5bbb46ec6d34", "metadata": { "slideshow": { "slide_type": "slide" } }, "source": [ "调用\n", "之前\n", "定义的训练函数来训练模型" ] }, { "cell_type": "code", "execution_count": 36, "id": "2b28c827-05cd-4c4b-9dd7-db27d8067a27", "metadata": { "origin_pos": 18, "tab": [ "pytorch" ], "tags": [] }, "outputs": [ { "data": { "image/svg+xml": [ "\n", "\n", "\n", " \n", " \n", " \n", " \n", " 2023-02-26T14:19:54.955113\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 = 10\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 }