一字一诗
LESSON 05 / 26

看懂模型里的数字形状

这一课完成什么

看懂模型数据的尺寸。PyTorch 把多层排列的数字叫做“张量”,可以先理解为带有明确形状的数字表。一个数是零维,一列数是一维,行列组成二维,再往外可以加“第几首诗”这样的维度。

数字的三条轴,各自代表什么

可以先把一批输入想成几本小册子:一本对应一首诗,一行对应一个字,一行里的格子对应这个字的特征。

数字的三条轴,各自代表什么
用文字逐步读这张图
  1. 批数量 B:同时处理几首;例如 2 首
  2. 字数 T:每首取几个位置;例如 3 个字
  3. 特征宽度 C:每个字用几个数;例如 4 个数
  4. 输入形状:[B, T, C];[2, 3, 4]
  5. 变换矩阵:[C, D];例如 [4, 6]
  6. 输出形状:[B, T, D];得到 [2, 3, 6]
改变特征宽度,会自动多生成几个字吗?

不会。特征宽度是每个位置内部的表示大小,和序列字数是不同的轴。

给每一层数字命名

本课输入尺寸为 [2,3,2]:一次放入 2 首诗,每首暂看 3 个位置,每个位置用 2 个数字表示。注意:这已经是查表后的数字表示。查表前的字编号只有 [2,3]

名称本课数值表示什么
B / batch2一次处理几条练习题
T / time3每条有几个字的位置
C / channels2每个位置有几项数字

字母是代码中的简写,意义由所在运算决定。本课会一直写出它们的含义,不要求先背缩写。

一次矩阵乘法

python
x = torch.arange(12, dtype=torch.float32).view(2, 3, 2)
w = torch.tensor([[1.,0.,2.,1.], [0.,1.,1.,2.]])
result = x @ w

arange(12) 产生 0 到 11;view 在保持数字数量不变的情况下重新排列。右侧矩阵尺寸是 [2,4],它将每个位置的 2 项信息变成 4 项。结果尺寸为 [2,3,4],诗的数量和位置数量都没有变化。

第一首第一个位置为 [0,1],分别与矩阵四列相乘并相加,得到 [0,1,1,2]。第一首第二个位置 [2,3] 则得到 [2,3,7,8]。同一份矩阵被所有位置共用,这就是后面线性层的一部分。

动手与检查

bash
./poet lesson 5
.venv/bin/python lessons/05/experiment.py

修改实验,把 12 个数变成 18 个,把形状改为 [3,3,2],其余矩阵不变。先预测结果应为 [3,3,4],再运行。接着故意把右侧矩阵变成 3 行,观察错误信息如何指出不能相乘。

三种容易混淆的操作

view 只重新解释数字排列;transpose 交换两个维度;矩阵乘法会真正计算出新数值。后面多头注意力会连续使用这三种操作。交换维度后,内存排列可能不连续,代码会先 .contiguous().view(...),避免把不兼容的排列强行重新解释。

模型参数是需要学习的数字;中间输出也是数字,但不等于新增加了参数。处理更长的诗会增加中间结果数量,并不会自动新增一份独立的线性矩阵。

验收与排错

小练习与答案

如果一次处理 8 首诗,每首 32 个位置,每个位置 128 项数字,经过 [128,256] 的变换后尺寸是什么?

查看答案

结果是 [8,32,256]。只有最后的表示宽度从 128 变成 256;同一个变换作用于所有批次和位置。变换矩阵有 128×256 项参数,若带偏置再加 256 项,而不是乘上 8 或 32。

动手补全一小段

先运行上面的完整实验,再复制本课起始文件为自己的练习。starter 有意留空 solve 函数;补全后运行文件,底部检查会告诉你是否符合本课要求。卡住时打开参考答案,比较每一步。

bash
cp lessons/05/starter.py lessons/05/my_exercise.py
# 编辑 my_exercise.py 中的 solve 函数,然后运行:
.venv/bin/python lessons/05/my_exercise.py
# 对照完整答案:
.venv/bin/python lessons/05/solution.py

下载起始代码 · 下载参考答案

展开本课补全练习的完整参考答案
python
"""第 05 课补全练习:变换每个位置的表示。修改 solve,保持下方检查不变。"""
import math, io, argparse
import torch
from torch import nn
from torch.nn import functional as F
torch.set_num_threads(2)
torch.manual_seed(26)

def solve(x, weights):
    return x @ weights

torch.testing.assert_close(solve(torch.tensor([[2.,3.]]),torch.tensor([[1.,0.],[0.,2.]])),torch.tensor([[2.,6.]]))
print("本课补全练习通过。")
展开本课完整、可独立运行的实验代码
python
"""本课独立实验;在仓库根目录执行 .venv/bin/python lessons/05/experiment.py。"""
import sys, json, math
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
import torch
from torch import nn
from torch.nn import functional as F
from poetry_gpt.common import ROOT, DEFAULT_DATA, read_json, read_jsonl
from poetry_gpt.model import ModelConfig, PoetryGPT, CausalAttention, Block
from poetry_gpt.data import Tokenizer, SPECIAL, clean_record, keywords_for
from poetry_gpt.labs import show
torch.set_num_threads(2)
torch.manual_seed(26)

x = torch.arange(12, dtype=torch.float32).view(2, 3, 2)
w = torch.tensor([[1.0, 0.0, 2.0, 1.0], [0.0, 1.0, 1.0, 2.0]])
show('输入 [批,字,特征]', list(x.shape))
show('变换 [入,出]', list(w.shape))
show('输出尺寸', list((x @ w).shape))
show('第一首第一个字的输出', (x @ w)[0, 0].tolist())
下载本课实验