1. 从一次张量维度对齐的“翻车”说起
在PyTorch里做张量运算,最常遇到的“坑”之一就是维度不匹配。我记得有一次,我需要将一个形状为[batch_size, 1, feature_dim]的中间特征张量,与另一个形状为[batch_size, num_heads, feature_dim]的注意力权重张量进行逐元素相乘。直觉上,我觉得[1, feature_dim]这个维度应该能通过广播机制自动扩展到[num_heads, feature_dim]。于是,我信心满满地写了intermediate_feat * attention_weight,结果直接抛出了一个RuntimeError: The size of tensor a (1) must match the size of tensor b (num_heads) at non-singleton dimension 1。
问题出在哪?广播机制确实存在,但它的规则是从后往前(从最右边的维度开始)逐维度比较。对于我的例子,比较过程是:
feature_dim与feature_dim:相等,没问题。1与num_heads:不相等,且1不等于1?等等,这里1是intermediate_feat的第二个维度,而num_heads是attention_weight的第二个维度。根据广播规则,当两个维度不相等时,其中一个必须为1,才能进行广播。这里1和num_heads都不为1(num_heads显然大于1),所以广播失败。
我的需求本质上是希望将intermediate_feat在“头”这个维度上复制num_heads次,使其形状变为[batch_size, num_heads, feature_dim],然后再进行运算。这时,torch.repeat()函数就是解决这类问题的“瑞士军刀”。它不像view()或reshape()那样改变数据的解读方式,也不像expand()那样只在逻辑上扩展而不实际复制数据(在某些情况下)。repeat()是实打实地在内存中复制数据,生成一个全新的张量,其行为非常直观和确定:告诉我每个维度你要复制几次,我就给你一个复制好的新张量。
2.torch.repeat()的核心机制与参数详解
torch.repeat()的函数签名非常简单:tensor.repeat(*sizes)。这里的*sizes表示一个可变参数,你传入多少个数字,就代表你希望结果张量在每个维度上的尺寸是原始张量对应维度的多少倍。
2.1 基本规则与底层逻辑
它的工作逻辑遵循一个清晰的两步过程:
- 维度对齐:如果传入的
sizes参数长度(记为len_sizes)大于原始张量的维度数(记为dim_tensor),repeat()会自动在原始张量的前面(即左侧)添加大小为1的维度,直到两者的维度数相等。这个行为与许多其他PyTorch函数(如torch.squeeze())的语义一致。 - 逐维度复制:对齐后,对于结果张量的第
i个维度,其大小等于原始尺寸[i] * sizes[i]。复制是在物理内存层面进行的,原始张量中的数据块会被重复填充到新张量的对应位置。
我们通过一个一维张量的例子来建立直观感受:
import torch x = torch.tensor([1, 2, 3]) # shape: [3] print(x.repeat(4)) # 输出:tensor([1, 2, 3, 1, 2, 3, 1, 2, 3, 1, 2, 3]), shape: [12]这里,sizes是(4,),len_sizes=1,dim_tensor=1,维度相等。结果就是在第0维(也是唯一的一维)上,将[1,2,3]这个序列重复了4次,拼接成一个长度为12的一维张量。
2.2 不同维度张量的repeat示例与解析
理解高维张量的repeat最好的方式就是动手实验。我们构建一个基础张量A,其形状为(2, 3),内容清晰,便于追踪。
A = torch.tensor([[1, 2, 3], [4, 5, 6]]) # shape: (2, 3) print(‘A:\n‘, A)场景一:扩展行和列
result = A.repeat(2, 3) # 在维度0(行)复制2次,在维度1(列)复制3次 print(‘A.repeat(2, 3) shape:‘, result.shape) print(‘A.repeat(2, 3):\n‘, result)输出:
A.repeat(2, 3) shape: torch.Size([4, 9]) A.repeat(2, 3): tensor([[1, 2, 3, 1, 2, 3, 1, 2, 3], [4, 5, 6, 4, 5, 6, 4, 5, 6], [1, 2, 3, 1, 2, 3, 1, 2, 3], [4, 5, 6, 4, 5, 6, 4, 5, 6]])我们来拆解这个过程:
原始形状 (2,3),sizes=(2,3), 维度对齐,无需补维。- 第0维(行):原始有2行
[R1, R2]。复制2次,得到[R1, R2, R1, R2]。这就是结果中的4行。 - 第1维(列):对于结果中的每一行,其原始行有3个元素
[a,b,c]。复制3次,得到[a,b,c, a,b,c, a,b,c]。 所以,最终是一个4行9列的张量,你可以看到它是一个2x3的“瓷砖”模式铺满了整个空间。
场景二:增加新的批次维度这是非常常见的用法,特别是在处理单个样本数据,需要将其扩展为一个批次时。
result = A.repeat(3, 1, 1) # 注意这里传入了三个数字 print(‘A.repeat(3, 1, 1) shape:‘, result.shape) print(‘A.repeat(3, 1, 1):\n‘, result)输出:
A.repeat(3, 1, 1) shape: torch.Size([3, 2, 3]) A.repeat(3, 1, 1): tensor([[[1, 2, 3], [4, 5, 6]], [[1, 2, 3], [4, 5, 6]], [[1, 2, 3], [4, 5, 6]]])关键点在于维度对齐:
原始形状 (2,3),dim_tensor=2。sizes=(3,1,1),len_sizes=3。- 因为
len_sizes > dim_tensor,所以系统自动在A的前面添加(3-2)=1个维度,将其视为形状为(1, 2, 3)的张量。 - 然后对这个
(1,2,3)的张量执行repeat(3,1,1):- 新第0维:
1 * 3 = 3 - 新第1维:
2 * 1 = 2 - 新第2维:
3 * 1 = 3最终我们得到了一个形状为(3,2,3)的张量,可以理解为有3个完全相同的“样本”,每个样本就是原来的矩阵A。这在数据预处理或模型推理单样本时非常有用。
- 新第0维:
注意:
repeat()的参数顺序始终对应结果张量从最左(最高)维度到最右(最低)维度的复制倍数。对于A.repeat(3,1,1),第一个参数3对应的是新增加的批次维,而不是原来的行维。
2.3 与view()/reshape()和expand()的关键区别
很多初学者容易混淆这几个函数,这里彻底厘清。
view()/reshape():改变形状,不改变数据内容与顺序。它们要求新形状的总元素数必须与原张量一致。你可以把它们理解为给同一块内存数据“换一种解读方式”。例如,一个(6,)的张量,可以被view(2,3)解读为一个2行3列的矩阵。它绝对无法实现(2,3)到(4,9)的转换,因为元素数量从6个变成了36个。expand():逻辑扩展,通常不复制数据。它可以将大小为1的维度扩展到任意大小,且原始张量在该维度上的唯一元素会被“广播”到新尺寸。它返回的是一个原张量的“视图”,在某些情况下(如后续进行写操作)可能会触发隐式复制。它的限制是只能将维度从1扩展到大,不能将非1的维度(如大小为3)改变。B = torch.tensor([[1, 2, 3]]) # shape: (1, 3) expanded = B.expand(4, 3) # shape: (4, 3), 内存中可能仍然只有 [1,2,3] 这一行数据 # 尝试 A.expand(4,3) 会报错,因为A的第0维是2,不是1,无法扩展。repeat():物理复制,创建新张量。它是最“暴力”也是最直接的方式,无视原始维度是否为1,直接按照指定的倍数在各个维度上进行数据复制。它总是会分配新的内存。当你需要确切的、独立的数据副本,或者需要增加非1维度的尺寸时,就必须使用repeat()。
用一个表格总结:
| 特性 | view()/reshape() | expand() | repeat() |
|---|---|---|---|
| 核心作用 | 改变张量的形状视图 | 将大小为1的维度逻辑扩展 | 在所有维度物理复制数据 |
| 内存 | 共享底层数据 | 通常共享,条件写时可能复制 | 总是创建新内存副本 |
| 维度变化 | 元素总数必须不变 | 只能将1维扩展为N维 | 可将任何维度 M 变为 M*N |
| 数据变化 | 无,数据顺序不变 | 无,数据通过广播填充 | 有,数据被重复复制 |
| 典型用途 | 调整网络层间数据形状 | 广播机制的高效实现 | 创建重复模式的数据、扩展批次 |
3. 实战场景:repeat()在深度学习任务中的应用
理解了原理,我们来看看repeat()在真实项目中如何大显身手。它绝不仅仅是一个简单的复制工具。
3.1 场景一:构造空间位置编码(Spatial Position Encoding)
在视觉Transformer(ViT)或目标检测模型中,我们经常需要为特征图的每个位置生成一个唯一的编码。假设我们有一个基于正弦余弦的、针对一维序列的位置编码矩阵pos_1d,形状为[max_len, d_model]。现在我们要将其应用到二维图像特征[batch, height, width, d_model]上。
一种常见方法是分别生成行编码和列编码,然后相加。repeat()在这里扮演了关键角色。
import torch import math def create_2d_sincos_position_embedding(height, width, dim): """ 创建二维正弦余弦位置编码 Args: height: 特征图高度 width: 特征图宽度 dim: 编码维度,需为偶数 Returns: pos_embed: 形状为 [height, width, dim] """ assert dim % 2 == 0, “维度必须是偶数” pos_embed = torch.zeros(height, width, dim) # 1. 分别创建行和列的位置索引 rows = torch.arange(height).float() cols = torch.arange(width).float() # 2. 计算频率因子 div_term = torch.exp(torch.arange(0, dim, 2).float() * -(math.log(10000.0) / dim)) # 3. 计算行编码 (shape: [height, 1, dim//2]) rows = rows.unsqueeze(1) # [height, 1] rows_sin = torch.sin(rows * div_term) # [height, dim//2] rows_cos = torch.cos(rows * div_term) # [height, dim//2] # 交错合并sin和cos,并扩展维度 rows_encoding = torch.stack([rows_sin, rows_cos], dim=2).view(height, 1, dim) # [height, 1, dim] # 4. 计算列编码 (shape: [1, width, dim//2]) cols = cols.unsqueeze(0) # [1, width] cols_sin = torch.sin(cols * div_term) # [1, width, dim//2] cols_cos = torch.cos(cols * div_term) # [1, width, dim//2] cols_encoding = torch.stack([cols_sin, cols_cos], dim=2).view(1, width, dim) # [1, width, dim] # 5. 使用 repeat 将行编码扩展到所有列,列编码扩展到所有行,然后相加 # rows_encoding: [height, 1, dim] -> repeat(1, width, 1) -> [height, width, dim] # cols_encoding: [1, width, dim] -> repeat(height, 1, 1) -> [height, width, dim] pos_embed = rows_encoding.repeat(1, width, 1) + cols_encoding.repeat(height, 1, 1) return pos_embed # 使用示例 H, W, D = 4, 6, 8 pos_2d = create_2d_sincos_position_embedding(H, W, D) print(f“二维位置编码形状: {pos_2d.shape}“) # torch.Size([4, 6, 8])在这个例子中,rows_encoding.repeat(1, width, 1)将每一行的编码复制到所有列上,cols_encoding.repeat(height, 1, 1)将每一列的编码复制到所有行上,两者相加就得到了每个(row, col)位置的唯一编码。这种“复制+相加”的模式在构建多维参数时非常高效。
3.2 场景二:数据增强中的样本复制与权重分配
在训练不平衡数据集时,我们可能会对少数类样本进行过采样。假设我们有一个批次的数据X和标签y,我们想将其中标签为class_idx的样本复制repeat_times份,并追加到原批次后面。
def oversample_minority_class(X, y, class_idx, repeat_times): """ 对指定类别的样本进行过采样 Args: X: 输入特征,形状 [batch, ...] y: 标签,形状 [batch] class_idx: 需要过采样的类别索引 repeat_times: 复制次数(包含原始样本,如2表示再复制1份) Returns: X_aug: 增强后的特征 y_aug: 增强后的标签 """ # 1. 找出少数类样本的掩码 minority_mask = (y == class_idx) X_minority = X[minority_mask] # 形状 [minority_count, ...] y_minority = y[minority_mask] # 形状 [minority_count] # 2. 使用 repeat 复制样本。注意:X_minority 可能有多维,我们需要在批次维度(第0维)复制 # 构建 repeat 参数:第0维复制 repeat_times 次,其他所有维度复制1次。 repeat_dims = [repeat_times] + [1] * (X_minority.dim() - 1) X_minority_repeated = X_minority.repeat(*repeat_dims) y_minority_repeated = y_minority.repeat(repeat_times) # 3. 拼接原始数据和复制数据 X_aug = torch.cat([X, X_minority_repeated], dim=0) y_aug = torch.cat([y, y_minority_repeated], dim=0) return X_aug, y_aug # 模拟数据 batch_size = 10 feat_dim = 5 X = torch.randn(batch_size, feat_dim) y = torch.randint(0, 3, (batch_size,)) # 3个类别 print(“原始标签分布:“, torch.bincount(y)) # 对类别1过采样,复制3次(即额外增加2份) X_aug, y_aug = oversample_minority_class(X, y, class_idx=1, repeat_times=3) print(“增强后标签分布:“, torch.bincount(y_aug)) print(f“X shape: {X.shape} -> {X_aug.shape}“)这里的关键技巧是动态构建repeat_dims列表:[repeat_times] + [1] * (X_minority.dim() - 1)。这确保了无论特征张量X_minority有多少个维度(比如对于图像是[N, C, H, W]),我们都只在第0维(批次维)进行复制,其他维度保持不变。这是一种非常通用和安全的写法。
3.3 场景三:为注意力机制准备键值对缓存(KV Cache)
在自回归模型(如GPT)的推理优化中,KV Cache是加速解码的核心技术。为了避免在生成每个新token时重新计算所有历史token的Key和Value,我们会缓存它们。当批次中有多个序列,且长度不一致时(需要padding),我们需要将当前步计算的KV(形状为[batch, 1, num_heads, head_dim])正确地存入一个形状为[batch, max_seq_len, num_heads, head_dim]的缓存中。这里,repeat()可以帮助我们处理某些特殊的注意力模式。
例如,在分组查询注意力(Grouped-Query Attention, GQA)中,Key和Value的头数(num_kv_heads)可能少于查询的头数(num_heads)。为了与标准的多头注意力计算兼容,我们需要将KV在“头”维度上进行复制。
def prepare_kv_cache_for_gqa(kv_state, num_heads, num_kv_heads): """ 为GQA准备KV缓存:将KV状态在头维度上复制,以匹配查询的头数。 Args: kv_state: 当前步计算的Key或Value,形状 [batch, 1, num_kv_heads, head_dim] num_heads: 查询的头数 num_kv_heads: Key/Value的头数 (num_kv_heads <= num_heads, 且 num_heads % num_kv_heads == 0) Returns: kv_state_expanded: 扩展后的Key或Value,形状 [batch, 1, num_heads, head_dim] """ batch, _, kv_heads, head_dim = kv_state.shape assert num_heads % num_kv_heads == 0, “num_heads must be divisible by num_kv_heads“ repeat_ratio = num_heads // num_kv_heads # 在头维度(第2维)上复制 repeat_ratio 次 # 我们希望形状从 [batch, 1, num_kv_heads, head_dim] 变为 [batch, 1, num_heads, head_dim] # 因此 repeat 参数为:批次维1倍,序列维1倍,头维 repeat_ratio 倍,特征维1倍。 kv_state_expanded = kv_state.repeat(1, 1, repeat_ratio, 1) # 或者更清晰地:kv_state_expanded = kv_state.repeat_interleave(repeat_ratio, dim=2) # repeat_interleave 是另一种复制方式,语义更清晰。 return kv_state_expanded # 模拟GQA场景 batch = 2 num_heads = 8 num_kv_heads = 2 # 分组查询,KV头数较少 head_dim = 64 current_k = torch.randn(batch, 1, num_kv_heads, head_dim) k_for_attention = prepare_kv_cache_for_gqa(current_k, num_heads, num_kv_heads) print(f“原始K形状: {current_k.shape}“) print(f“扩展后K形状: {k_for_attention.shape}“) # torch.Size([2, 1, 8, 64])在这个场景中,repeat(1, 1, repeat_ratio, 1)精确地控制了只在第三个维度(头维度)进行复制。这使得计算注意力分数时,每个查询头都能找到对应的(复制的)键头。虽然这里也可以用repeat_interleave,但repeat()在需要同时处理多个维度复制时,其参数化方式更加统一和灵活。
4. 高级技巧、性能陷阱与替代方案
torch.repeat()虽然强大,但盲目使用也会带来问题。下面是一些实战中积累的经验和需要避开的“坑”。
4.1 性能陷阱:无谓的大张量复制与内存爆炸
这是使用repeat()时最容易犯的错误。因为它进行的是物理复制,所以复制的倍数会以乘积方式放大内存占用。
# 危险示例:一个不小心的操作可能导致OOM(内存溢出) large_tensor = torch.randn(256, 256, 3) # 一张256x256的RGB图,约0.5MB # 假设你想把它变成一个4张图的“批次” batch_tensor = large_tensor.repeat(4, 1, 1) # 形状 [4, 256, 256, 3], 约2MB, 可以接受 # 但如果你手滑了... dangerous_tensor = large_tensor.repeat(4, 4, 4) # 形状 [1024, 1024, 12], 内存爆炸!教训:在使用repeat()前,一定要清楚每个维度的复制倍数,并估算结果张量的大致内存占用(元素数量 * 每个元素字节数)。对于非常大的张量,考虑是否真的需要物理复制,也许expand()或广播机制就能满足需求。
4.2 与expand()和广播的协同与选择
如何决定用repeat()还是expand()?遵循以下决策链:
- 目标是否只是为了让形状兼容以进行运算?如果是,并且原始张量在需要扩展的维度上大小恰好为1,优先使用
expand()。它更高效(可能零拷贝)。# 好例子:使用 expand 进行高效广播 stats = torch.tensor([[[0.5, 0.2]]]) # shape: [1, 1, 2] 均值和方差 batch_data = torch.randn(32, 10, 2) # 归一化:将 stats 广播到 batch_data 的形状 normalized = (batch_data - stats.expand_as(batch_data)) # 高效,逻辑扩展 - 需要扩展的维度大小不为1,或者你需要一份独立的数据副本以避免后续的梯度传播问题?使用
repeat()。# 需要物理副本的例子 template = torch.tensor([1, 0, 1, 0]) # shape: [4] # 创建一个 3x4 的掩码,每行都是 [1,0,1,0] mask = template.repeat(3, 1) # 形状 [3, 4] # 后续对 mask 的修改不会影响 template mask[0, 0] = 99 print(template) # 仍然是 tensor([1, 0, 1, 0]) - 利用PyTorch的自动广播。很多时候,我们甚至不需要显式调用
expand()或repeat()。PyTorch的运算符(+,-,*,/,@等)会自动应用广播规则。文章开头我踩的坑,其实可以用.unsqueeze()解决:intermediate_feat = torch.randn(8, 1, 64) # [batch, 1, feat] attention_weight = torch.randn(8, 4, 64) # [batch, heads, feat] # 错误: result = intermediate_feat * attention_weight # 正确: 利用广播,但需要对齐维度 # 将 intermediate_feat 的“头”维度显式补1并扩展 result = intermediate_feat.expand_as(attention_weight) * attention_weight # 或者更简洁地,利用广播自动完成 expand result = intermediate_feat * attention_weight.unsqueeze(1) # 这不行,维度不对 # 正确做法是: result = intermediate_feat.expand(-1, 4, -1) * attention_weight # 使用expand # 或者,如果你确定需要物理复制: result = intermediate_feat.repeat(1, 4, 1) * attention_weight # 使用repeat
4.3repeat_interleave():更精细的复制控制
torch.repeat_interleave()是repeat()的一个更灵活的变体。两者的核心区别在于复制的模式:
tensor.repeat(a, b, c...):在整个张量层面进行区块复制。它先复制整个张量a次,然后在次维度上复制b次,以此类推。torch.repeat_interleave(tensor, repeats, dim):在指定维度dim上,对该维度的每个元素进行复制。repeats可以是一个整数(所有元素复制相同次数),也可以是一个列表(指定每个元素复制的次数)。
x = torch.tensor([[1, 2], [3, 4]]) # repeat 模式:整体复制 print(‘x.repeat(2, 3):\n‘, x.repeat(2, 3)) # 输出: # tensor([[1, 2, 1, 2, 1, 2], # [3, 4, 3, 4, 3, 4], # [1, 2, 1, 2, 1, 2], # [3, 4, 3, 4, 3, 4]]) # 可以看作把 [[1,2],[3,4]] 这个2x2的块,先向下复制2次,再向右复制3次。 # repeat_interleave 模式:元素级复制 print(‘torch.repeat_interleave(x, 2, dim=0):\n‘, torch.repeat_interleave(x, 2, dim=0)) # 输出: # tensor([[1, 2], # [1, 2], # [3, 4], # [3, 4]]) # 在第0维(行),将第0行[1,2]复制2次,再将第1行[3,4]复制2次。 print(‘torch.repeat_interleave(x, [1, 3], dim=1):\n‘, torch.repeat_interleave(x, [1, 3], dim=1)) # 输出: # tensor([[1, 2, 2, 2], # [3, 4, 4, 4]]) # 在第1维(列),对于第一行[1,2]:第0列元素‘1‘复制1次,第1列元素‘2‘复制3次。如何选择:如果你需要的是“平铺”或“区块复制”效果,用repeat()。如果你需要的是“按元素重复”或“交错复制”,用repeat_interleave()。例如,将序列[a, b, c]的每个元素重复两次得到[a, a, b, b, c, c],这就是repeat_interleave的典型用例。
4.4 梯度传播问题
由于repeat()创建了新的物理存储,其梯度传播行为是符合直觉的:最终结果张量的梯度会平均分配到原始张量的每一个复制源元素上。
x = torch.tensor([1.0, 2.0], requires_grad=True) y = x.repeat(3) # y = [1., 2., 1., 2., 1., 2.] z = y.sum() # z = 9.0 z.backward() print(x.grad) # tensor([3., 3.])z对x的梯度计算:x[0]在y中出现了3次,每次的梯度贡献是1,所以总梯度是3。x[1]同理。这符合自动微分的链式法则。
最后,分享一个我调试repeat()相关bug时的小技巧:当结果形状不符合预期时,我通常会先打印出tensor.shape和我要传入的sizes元组,然后在脑子里或草稿纸上执行前面提到的“维度对齐”和“逐维度相乘”两步,几乎能立刻定位问题所在。对于复杂操作,先用一个小规模的、数据有规律的张量(比如像本文示例一样用torch.arange生成)进行测试,验证repeat的效果,再应用到真实数据上,能节省大量排查时间。