AI 工程基础体系 · 第 57/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。
深度学习张量与形状:维度、广播、布局、批次和调试方法
在机器学习、深度学习和生成式 AI 系统中,模型通常接收的不是单个数字,而是带有多个轴的张量。图像可能表示为 [批次, 通道, 高度, 宽度],文本通常表示为 [批次, 序列长度] 或 [批次, 序列长度, 隐藏维度],多头注意力还会引入头数这一轴。
许多看似不同的错误——mat1 and mat2 shapes cannot be multiplied、The size of tensor a must match...、显存突然增长、模型输出错位——本质上都与以下问题有关:
- 每个轴分别表示什么;
- 运算要求哪些轴相等、哪些轴可以广播;
- 张量的物理布局是否允许某种视图操作;
- 批次轴是否被正确保留;
- 出错时如何定位是数据、模型还是评测流程的问题。
本文使用 PyTorch 语义说明这些机制。具体 API 的行为应以对应版本的 PyTorch Documentation 和 PyTorch Tutorials 为准。
一、张量、维度、轴、形状和秩
1.1 张量是什么
张量可以先理解为带有规则索引结构的数值容器:
- 标量只有一个值;
- 向量有一个索引;
- 矩阵有两个索引;
- 更高阶张量有更多索引。
例如一个形状为 [2, 3, 4] 的张量 x,可以写成:
其中:
它总共包含:
个元素。
在工程语境中需要区分四个词:
- 轴(axis):一个索引方向,例如第 0 轴、第 1 轴;
- 维度(dimension):中文中既可能指轴,也可能指某个轴的长度;
- 形状(shape):各轴长度组成的元组,例如
(2, 3, 4); - 秩(rank):轴的数量,例如
(2, 3, 4)的秩是 3。
因此,x.shape == (2, 3, 4) 表示:
x是秩为 3 的张量;- 有 3 个轴;
- 三个轴的长度分别是 2、3、4;
- 但它没有说明每个轴的语义。
[2, 3, 4] 可以表示:
- 2 个样本、每个样本 3 个时间步、每步 4 个特征;
- 2 个通道、图像高度 3、图像宽度 4;
- 2 个头、序列长度 3、序列长度 4。
形状相同不代表语义相同。模型能否正确运行,首先取决于轴的语义是否正确。
1.2 PyTorch 中的基本检查
import torch
x = torch.randn(2, 3, 4)
print(x.shape) # torch.Size([2, 3, 4])
print(x.ndim) # 3
print(x.numel()) # 24
print(x.dtype) # 通常是 torch.float32
print(x.device) # 通常是 cpu
这里:
shape描述每个轴的长度;ndim是秩;numel()是元素总数;dtype决定每个元素如何解释;device决定数据位于 CPU、CUDA 等设备。
形状正确并不保证运算正确。例如,float32 模型和 int64 输入可能产生类型错误;CPU 张量和 CUDA 张量混用会产生设备错误。形状是核心问题,但不是唯一的张量契约。
1.3 从语义命名轴
对于复杂模型,单独记住整数位置容易出错。可以在代码和文档中明确写出约定:
x: [B, S, D]
B = batch size,批次大小
S = sequence length,序列长度
D = feature size 或 hidden size,特征/隐藏维度
例如:
x = torch.randn(8, 128, 768)
只有结合约定,才能知道它表示:
- 8 个样本;
- 每个样本 128 个 token;
- 每个 token 用 768 维向量表示。
二、形状运算的基本规则
张量运算并非都遵循同一种形状规则。逐元素运算、矩阵乘法、拼接、归约和索引分别有不同的契约。
2.1 逐元素运算要求对应元素存在
如果两个张量逐元素相加:
那么每个输出位置都需要找到对应的 x 元素和 y 元素。若两个张量形状完全相同,例如:
x = torch.tensor([[1., 2., 3.],
[4., 5., 6.]])
y = torch.tensor([[10., 20., 30.],
[40., 50., 60.]])
z = x + y
则:
z =
[[11, 22, 33],
[44, 55, 66]]
x 和 y 的形状都是 [2, 3],位置一一对应。
但是 PyTorch 还支持广播,使某些形状不同的张量能够参与逐元素运算。
三、广播:隐式扩展而不是复制数据
3.1 广播的形式化条件
设两个张量从右向左对齐。对齐后,每一对对应轴必须满足以下条件之一:
- 两个轴长度相等;
- 其中一个轴长度为 1;
- 某个张量在该位置没有轴,视为长度为 1。
如果所有轴都满足条件,则可以广播。
例如:
A: [2, 3, 4]
B: [3, 4]
从右侧对齐后相当于:
A: [2, 3, 4]
B: [1, 3, 4]
第 0 轴上,1 可以广播成 2,因此结果形状为 [2, 3, 4]。
相反:
A: [2, 3, 4]
B: [2, 4]
对齐后:
A: [2, 3, 4]
B: [1, 2, 4]
第 1 轴比较 3 和 2,既不相等,也没有一个是 1,因此不能广播。
3.2 完整算例
import torch
x = torch.tensor([
[[1., 2., 3., 4.],
[5., 6., 7., 8.],
[9., 10., 11., 12.]],
[[13., 14., 15., 16.],
[17., 18., 19., 20.],
[21., 22., 23., 24.]]
]) # [2, 3, 4]
bias = torch.tensor([100., 200., 300., 400.]) # [4]
y = x + bias # [2, 3, 4]
print(y.shape)
print(y[0, 0])
bias 的形状 [4] 会从左侧补成 [1, 1, 4]。它并不是只加到某一个位置,而是对 [2, 3] 所有组合都使用同一个长度为 4 的偏置:
x[batch, position, feature] + bias[feature]
因此第一行输出为:
tensor([101., 202., 303., 404.])
这正是线性层偏置的典型形式。若隐藏状态是 [B, S, D],偏置通常是 [D],可以广播到所有批次和位置。
3.3 广播不是任意的复制
广播在语义上可以理解为把长度为 1 的轴扩展到目标长度,但实际实现通常不会立刻物理复制全部数据。expand 可以构造一个共享存储的视图:
x = torch.tensor([[1., 2., 3.]]) # [1, 3]
y = x.expand(4, 3) # [4, 3]
print(y)
print(y.stride())
y 的四行通常引用同一行数据。这样可以节省内存,但也意味着:
y不是独立副本;- 对广播视图进行原地修改可能不安全;
- 多个逻辑位置可能映射到同一个物理元素。
因此,不应把广播结果当作普通可独立写入的矩阵。若确实需要独立存储,应显式使用 clone(),但这会产生内存成本:
y_copy = x.expand(4, 3).clone()
3.4 原地操作的额外限制
非原地运算可以生成广播后的输出:
x = torch.zeros(2, 3)
b = torch.ones(3)
x = x + b
但原地运算要求左侧张量的形状不能因为广播而改变:
x = torch.zeros(2, 3)
b = torch.ones(3)
x += b # 可以,结果仍然是 [2, 3]
以下模式可能失败:
x = torch.zeros(1, 3)
b = torch.ones(2, 3)
x += b # 通常报错
原因是原地操作必须把结果写回 x 的存储,而 x 的形状 [1, 3] 无法容纳 [2, 3] 的结果。非原地写法:
x = x + b # 可以,生成新的 [2, 3]
3.5 广播的反例:形状相近但不能广播
a = torch.randn(2, 3, 4)
b = torch.randn(2, 5, 4)
a + b
错误发生在中间轴:
a: [2, 3, 4]
b: [2, 5, 4]
^ ^
3 != 5
正确修复不是盲目调用 reshape,而是先判断语义:
- 如果
3和5本来代表不同序列长度,可能需要 padding 或 mask; - 如果它们本来代表同一个特征轴,可能是上游层配置错误;
- 如果希望做矩阵乘法,则逐元素加法本身就不是正确操作。
四、矩阵乘法与 matmul 的形状契约
逐元素乘法 * 和矩阵乘法 @ 完全不同:
a * b # 逐元素乘法,需要相同或可广播的形状
a @ b # 矩阵乘法,需要满足乘法维度契约
二维矩阵:
相乘后:
中间维度 k 必须相等。
a = torch.randn(2, 3)
b = torch.randn(3, 4)
c = a @ b
print(c.shape) # torch.Size([2, 4])
计算过程是:
若 b 的形状为 [2, 4],则无法完成,因为 a 的最后一维为 3,而 b 的倒数第二维为 2:
a = torch.randn(2, 3)
b = torch.randn(2, 4)
a @ b # 形状错误
4.1 批量矩阵乘法
对更高阶输入,torch.matmul 通常把最后两个轴当作矩阵轴,前面的轴当作批次轴,并尝试对批次轴广播。
a = torch.randn(8, 16, 32) # 8 个批次,每个是 [16, 32]
b = torch.randn(8, 32, 10) # 8 个批次,每个是 [32, 10]
c = a @ b
print(c.shape) # [8, 16, 10]
对每个批次 i,执行:
如果 b 的形状是 [32, 10],则它可以被所有批次共享:
b = torch.randn(32, 10)
c = a @ b
print(c.shape) # [8, 16, 10]
但不能仅凭“维度数量相同”推断可乘。必须明确哪两个轴是矩阵的行、列,哪几个轴是批次轴。
五、reshape、view、permute 与内存布局
5.1 逻辑形状和物理布局
张量的 shape 只描述逻辑索引范围,不描述元素在内存中如何排列。布局通常由:
- 底层存储;
stride;- 是否连续(contiguous);
- dtype 和设备
共同决定。
对形状 [2, 3] 的连续张量:
[[a, b, c],
[d, e, f]]
如果每个元素按行连续存储,那么相邻列元素相距 1 个位置,相邻行元素相距 3 个位置,stride 通常为:
(3, 1)
索引地址可以抽象为:
5.2 permute 只改变轴解释
x = torch.arange(24).reshape(2, 3, 4)
y = x.permute(0, 2, 1)
print(x.shape) # [2, 3, 4]
print(y.shape) # [2, 4, 3]
print(x.stride())
print(y.stride())
permute(0, 2, 1) 将原来的第 1 轴和第 2 轴交换。它通常只改变:
- shape;
- stride;
- 轴的解释顺序。
它不一定复制数据。因此,y 常常是非连续布局。
这也是为什么下面的代码可能失败:
x = torch.arange(24).reshape(2, 3, 4)
y = x.permute(0, 2, 1)
z = y.view(2, 12) # 可能报错
view 要求新形状能够按照当前 stride 直接解释同一块存储。经过 permute 后,逻辑顺序与物理顺序不再兼容时,view 无法完成。
可以使用:
z = y.contiguous().view(2, 12)
或者:
z = y.reshape(2, 12)
5.3 view 和 reshape 的差异
view 的核心语义是“不复制存储地改变视图”。如果布局不满足条件,它会报错。
reshape 更灵活:
- 如果当前布局允许,它可能返回视图;
- 如果不允许,它可能创建一个新的连续副本。
因此:
z = y.reshape(2, 12)
代码更方便,但不能假设它永远零拷贝。对于大张量,这个隐式副本可能带来:
- 更高的峰值显存;
- 额外的数据搬运;
- 训练或推理吞吐下降。
如果布局是否连续对性能或内存很重要,应显式检查:
print(y.is_contiguous())
print(y.stride())
并在需要时明确写出:
z = y.contiguous().view(2, 12)
这表示“我知道这里需要一次连续化复制”。
5.4 典型的图像布局转换
很多卷积接口默认使用:
NCHW = [batch, channel, height, width]
而某些数据源或算子使用:
NHWC = [batch, height, width, channel]
转换不能只改变 shape 字符串,必须改变轴顺序:
x_nchw = torch.randn(8, 3, 224, 224)
x_nhwc = x_nchw.permute(0, 2, 3, 1)
print(x_nhwc.shape) # [8, 224, 224, 3]
如果随后某个操作要求连续内存,可以:
x_nhwc = x_nhwc.contiguous()
错误地使用 reshape(8, 224, 224, 3) 代替 permute 不会完成轴交换;它只会按照原有线性存储重新解释元素,导致像素和通道数据错位。代码可能运行,但结果语义已经错误,这是比直接抛异常更危险的情况。
六、批次:一个轴背后的数据契约
6.1 批次轴是什么
批次(batch)是同时处理的一组独立样本。若单个样本形状为 [D],批次输入通常是 [B, D];若单个图像形状为 [C, H, W],批次输入通常是 [B, C, H, W]。
批次轴的加入通常不改变“单个样本的语义”,只是把同一种计算并行应用于多个样本:
但是,批次不是任意添加的维度。每个模型接口都规定了哪些轴是批次、哪些轴是特征或空间。
6.2 Dataset、DataLoader 和最后一个批次
一个常见流程是:
from torch.utils.data import DataLoader, TensorDataset
features = torch.randn(10, 4)
labels = torch.randint(0, 2, (10,))
dataset = TensorDataset(features, labels)
loader = DataLoader(dataset, batch_size=4, shuffle=True)
for batch_x, batch_y in loader:
print(batch_x.shape, batch_y.shape)
预期批次形状为:
torch.Size([4, 4]) torch.Size([4])
torch.Size([4, 4]) torch.Size([4])
torch.Size([2, 4]) torch.Size([2])
最后一个批次只有 2 个样本。不能默认每个批次都等于 batch_size。如果模型或并行策略确实要求固定批次,可以使用 drop_last=True,但这会丢弃最后不足一个完整批次的数据:
loader = DataLoader(
dataset,
batch_size=4,
drop_last=True
)
这不是纯粹的形状修复,而是数据策略改变。训练、评测和生产推理应分别评估这种取舍。
6.3 不要用固定批次大小破坏通用代码
错误模式:
for x, y in loader:
x = x.reshape(32, 4)
当最后一个批次不是 32 时,这段代码会失败;即使元素数恰好能重排,也可能掩盖批次语义错误。
更安全的写法是保留动态批次:
for x, y in loader:
batch_size = x.shape[0]
assert x.shape == (batch_size, 4)
如果需要展平除批次轴之外的所有轴,可使用:
x = x.flatten(start_dim=1)
例如输入 [B, C, H, W] 会变成 [B, C*H*W],不会错误地把不同样本混在一起。
6.4 squeeze 的隐患
squeeze() 会删除所有长度为 1 的轴:
x = torch.randn(1, 1, 8)
print(x.squeeze().shape) # [8]
这可能同时删除批次轴和其他语义轴。如果批次大小从 8 变为 1,代码行为就会改变。
更明确的写法是指定轴:
x = x.squeeze(dim=1)
或者只在确认该轴必须为 1 时删除:
assert x.shape[1] == 1
x = x.squeeze(1)
同理,unsqueeze(0) 是明确地插入批次轴:
sample = torch.randn(3, 224, 224)
batch = sample.unsqueeze(0)
print(batch.shape) # [1, 3, 224, 224]
七、序列与 Transformer 中的形状
7.1 常见的序列表示
Transformer 中经常使用:
X: [B, S, D]
B = 批次大小
S = 序列长度
D = 隐藏维度
如果有 H 个注意力头,通常要求:
其中 D_h 是每个头的维度。
查询、键、值经过线性变换后可以重排为:
Q, K, V: [B, H, S, Dh]
自注意力分数为:
对 K 的最后两个轴交换:
K: [B, H, S, Dh]
K.transpose(-2, -1): [B, H, Dh, S]
于是:
Q @ K.transpose(-2, -1): [B, H, S, S]
这里最后两个轴满足:
Q: [B, H, S, Dh]
Kᵀ: [B, H, Dh, S]
输出: [B, H, S, S]
中间的 Dh 被求和,两个 S 分别表示查询位置和键位置。
7.2 注意力 mask 的广播
一个 padding mask 可能形状为:
[B, S]
而注意力分数是:
[B, H, S, S]
为了让 mask 表示“每个批次中哪些键位置可见”,通常需要显式插入轴:
padding_mask = torch.tensor([
[True, True, True, False],
[True, True, False, False],
]) # [B, S]
key_mask = padding_mask[:, None, None, :] # [B, 1, 1, S]
它可以广播到:
[B, H, S, S]
因为:
[B, 1, 1, S]
[B, H, S, S]
对应轴分别为:
B == B
1 可以广播到 H
1 可以广播到 S
S == S
若把 [B, S] 直接与 [B, H, S, S] 混合,右对齐后可能会把 B 错误地对齐到最后几个轴,造成报错或更隐蔽的语义错误。因此,mask 的 unsqueeze 位置必须由“它要约束哪个轴”决定,而不是凭经验添加。
7.3 变长序列不是简单的形状问题
一个批次中的序列长度可能不同。常见做法是 padding 到批次内最大长度:
token_ids: [B, Smax]
attention_mask: [B, Smax]
其中 Smax 是该批次中的最大序列长度。padding token 参与张量计算,但应该通过 mask 排除。
这会产生两个成本:
- 计算量通常随
Smax增长; - 自注意力分数
[B, H, Smax, Smax]的元素数量随序列长度平方增长。
因此,增大序列长度不仅改变形状,也改变显存和计算成本。生成式模型中的 KV cache 还会按层、按头、按历史序列长度保存键和值,其形状设计必须与注意力实现保持一致;形状错位可能表现为注意力结果错误,也可能直接造成显存异常。
八、归约、拼接和索引也有形状契约
8.1 归约会删除或保留轴
sum、mean、max 等归约操作会沿指定轴聚合:
x = torch.randn(2, 3, 4) # [B, S, D]
a = x.mean(dim=1)
print(a.shape) # [2, 4]
b = x.mean(dim=1, keepdim=True)
print(b.shape) # [2, 1, 4]
dim=1 表示对序列位置聚合:
keepdim=True 保留长度为 1 的轴,便于后续广播。例如 [B, 1, D] 可以与 [B, S, D] 相加,而 [B, D] 需要额外插入轴才能表达相同语义。
8.2 cat 和 stack 不是一回事
torch.cat 沿已有轴拼接:
a = torch.randn(2, 3)
b = torch.randn(2, 5)
c = torch.cat([a, b], dim=1)
print(c.shape) # [2, 8]
除拼接轴外,其他轴必须相等。
torch.stack 则新增一个轴:
a = torch.randn(3)
b = torch.randn(3)
c = torch.stack([a, b], dim=0)
print(c.shape) # [2, 3]
如果有 4 个 [3] 向量,stack 得到 [4, 3],这通常表示新增一个批次或样本轴;cat 得到 [12],表示沿已有轴连接。误用二者会让代码仍然运行,但下游语义发生变化。
九、一个可运行的形状调试示例
下面的程序演示一个常见的数据流:批次输入、线性层、广播偏置、轴交换、注意力分数,以及错误定位。
import torch
from torch import nn
torch.manual_seed(0)
B, S, D, H = 2, 4, 8, 2
Dh = D // H
x = torch.randn(B, S, D) # [B, S, D]
projection = nn.Linear(D, 3 * D)
qkv = projection(x) # [B, S, 3D]
q, k, v = qkv.chunk(3, dim=-1) # 各为 [B, S, D]
def split_heads(t):
# [B, S, H*Dh] -> [B, H, S, Dh]
return t.reshape(B, S, H, Dh).transpose(1, 2)
q = split_heads(q)
k = split_heads(k)
v = split_heads(v)
print("x:", x.shape)
print("q:", q.shape)
print("k:", k.shape)
print("v:", v.shape)
scores = q @ k.transpose(-2, -1)
scores = scores / (Dh ** 0.5)
print("scores:", scores.shape) # [B, H, S, S]
valid_tokens = torch.tensor([
[True, True, True, False],
[True, True, False, False],
]) # [B, S]
key_mask = valid_tokens[:, None, None, :] # [B, 1, 1, S]
scores = scores.masked_fill(~key_mask, torch.finfo(scores.dtype).min)
weights = torch.softmax(scores, dim=-1)
output = weights @ v
print("weights:", weights.shape) # [B, H, S, S]
print("output:", output.shape) # [B, H, S, Dh]
逐步检查如下:
x的最后一维是D=8,因此nn.Linear(D, 3*D)合法;- 线性层只作用于最后一维,输入
[B,S,D]输出[B,S,3D]; chunk(3, dim=-1)把最后一维拆成三个[B,S,D];reshape(B,S,H,Dh)使用D=H*Dh的条件;transpose(1,2)将[B,S,H,Dh]变成[B,H,S,Dh];q @ k.transpose(-2,-1)的中间维度Dh相等,输出[B,H,S,S];key_mask的形状[B,1,1,S]可以广播到分数形状;weights @ v的输出为[B,H,S,Dh]。
如果将:
key_mask = valid_tokens[:, None, None, :]
错误地改成:
key_mask = valid_tokens[:, :, None, None]
形状会变成 [B,S,1,1]。它并不表示最后一个注意力轴上的 key 位置,而是把序列轴放到了第二维,可能导致广播失败,或者在某些尺寸碰巧兼容时产生错误 mask。
十、系统化调试形状错误
10.1 先记录“进入模块前”和“离开模块后”
不要只看最终异常。应该在数据管道和关键模块边界记录:
def describe(name, x):
print(
f"{name}: shape={tuple(x.shape)}, "
f"dtype={x.dtype}, device={x.device}, "
f"stride={x.stride()}, contiguous={x.is_contiguous()}"
)
使用:
describe("input", x)
describe("q", q)
describe("scores", scores)
这能同时发现:
- 轴数量错误;
- 轴顺序错误;
- dtype 不匹配;
- device 不匹配;
- 非连续布局;
- 某一步意外产生了副本或视图。
生产环境中不应无条件打印全部张量值,因为这会泄露数据并造成日志开销。通常记录 shape、dtype、device 和有限的统计量更合适。
10.2 用断言表达模型契约
assert x.ndim == 3, f"expected [B,S,D], got {x.shape}"
assert x.shape[-1] == D, f"expected hidden size {D}, got {x.shape[-1]}"
assert q.shape == (B, H, S, Dh)
断言应尽量靠近问题产生的位置,而不是等到损失函数或 CUDA kernel 才报错。
对批次数据还可以检查:
assert batch_x.shape[0] == batch_y.shape[0]
assert batch_x.dtype == torch.float32
assert batch_y.dtype == torch.int64
最后两个检查与分类损失常见契约有关,但具体要求取决于所使用的损失函数。形状调试应同时检查数据类型和语义。
10.3 把每个 reshape 写成可解释的变换
不推荐:
x = x.reshape(-1, 64)
这里 -1 可能掩盖批次、序列或空间轴的错误合并。
更清晰的写法:
B, S, D = x.shape
assert D == 64
x = x.reshape(B * S, D)
这样明确表示:只合并批次轴和序列轴,不改变特征轴。若要恢复:
x = x.reshape(B, S, D)
但前提是中间运算没有改变元素数量或样本顺序。
10.4 区分三种失败
形状相关的问题至少有三类:
第一类:立即报错。
例如矩阵乘法中间轴不相等。这类问题通常容易定位,应读取错误中给出的实际尺寸,并回溯最近一次改变 shape 的操作。
第二类:运行成功但语义错误。
例如把 NHWC 当成 NCHW,或把 [B,S] 的 mask 广播到错误轴。这类问题不会由 Python 类型系统自动阻止,必须依靠轴命名、断言、可视化样本和小规模人工验证。
第三类:运行成功但资源异常。
例如误把 [B,S,D] 展平为 [B*S,D] 后扩大了后续输出,或生成了 [B,H,S,S] 的巨大注意力矩阵。此时应检查每一步元素数量:
def count(name, x):
print(name, tuple(x.shape), x.numel())
count("x", x)
count("scores", scores)
如果一个操作使元素数量出现非预期数量级增长,应继续检查广播、重复、repeat、拼接和序列长度。
10.5 用最小输入复现
调试大型模型时,可以固定:
B, S, D = 2, 4, 8
先验证:
- 单个算子;
- 一个批次;
- 一个前向步骤;
- 一次反向传播;
- 一个边界批次,例如
B=1、S=1、最后不足整批。
尤其要测试 B=1。很多不安全的 squeeze()、错误索引和隐含广播只会在批次为 1 时暴露。
十一、形状设计中的常见误区
11.1 把“维度”理解成“轴数量”
“这个张量是三维的”可能表示:
- 它有三个轴;
- 它的某个空间维度是 3;
- 它使用三维坐标。
在技术文档和代码中最好使用“秩为 3”“第 1 轴长度为 3”“隐藏维度为 768”等精确表述。
11.2 用 reshape 修复未知语义
reshape 只能改变索引解释或在必要时复制数据,不能自动推断轴语义。它不会把通道轴移动到正确位置,也不会生成缺失的 batch 轴。
如果问题是轴顺序,应使用 permute 或 transpose;如果问题是缺少轴,应使用 unsqueeze;如果问题是元素总数不匹配,则应回到数据或模型契约检查。
11.3 用 repeat 代替广播
bias = torch.randn(8)
expanded = bias.repeat(2, 4, 1)
这会实际创建更多数据。若只是逐元素计算,通常可以直接使用:
x + bias
或使用不复制存储的 expand。不过,是否允许使用视图取决于后续操作是否需要可写的独立存储。
11.4 只在训练集验证形状
训练数据可能总是固定尺寸,而生产输入会出现:
- 单条请求;
- 空文本或极短文本;
- 超长文本;
- 不同图像分辨率;
- 最后一个不完整批次;
- 不同 dtype 或设备。
因此形状测试应覆盖代表性边界,而不是只覆盖一个“正常批次”。在生成式 AI 服务中,还要单独验证预填充(prefill)、增量解码(decode)和 KV cache 的形状,因为这几个阶段的序列长度和缓存状态不同。
十二、形状与生产系统的关系
形状不是只存在于模型代码中的局部细节。它贯穿数据、模型、评测、权限和成本:
- 数据层需要定义样本轴、特征轴和可变长度处理方式;
- 模型层需要保证每个模块的输入输出契约;
- 评测层需要保证预测和标签在样本轴上严格对齐;
- 权限层需要避免在调试日志中输出原始输入或完整中间张量;
- 成本层需要关注批次大小、序列长度、广播副本和注意力矩阵规模。
例如,评测时如果预测形状为 [B, C],标签却是 [B, 1],某些逐元素比较可能触发广播。代码可能不报错,但比较结果不一定符合预期。应先明确任务契约:
assert logits.shape == (batch_size, num_classes)
assert labels.shape == (batch_size,)
再根据损失函数或评测指标进行必要的变换,而不是依赖隐式广播“碰巧运行”。
一个可靠的张量接口至少应明确:
名称:input_ids
形状:[B, S]
dtype:torch.int64
语义:每个批次中的 token ID
padding:使用何种 token
mask:mask 的形状及其约束的轴
device:与模型参数一致
当这些契约被写入模块边界、断言和测试后,形状错误就不再只是某条异常信息,而会变成可以定位、验证和恢复的工程状态。
张量形状描述的是“有多少个轴以及每个轴多长”,广播描述的是“不等长轴如何参与逐元素计算”,布局描述的是“这些元素如何映射到底层存储”,批次描述的是“哪个轴承载独立样本”,调试则是把这些隐含假设显式化。掌握这四层关系,才能从“让代码跑起来”进一步判断结果是否仍然保持了正确的数学和业务语义。
系列导航与关联阅读
- 系列入口:AI 工程完整学习路线:从机器学习与 Transformer 到 RAG、Agent 和生产治理
- 上一篇:因果推断基础:相关与因果、DAG、混杂、实验和反事实
- 下一篇:自动微分与反向传播:计算图、梯度累积、截断和数值检查
官方资料
本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。

评论
0 条讨论