ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch张量复制:torch.repeat()机制详解与实战应用

PyTorch张量复制:torch.repeat()机制详解与实战应用 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都不为1num_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_tensorrepeat()会自动在原始张量的前面即左侧添加大小为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_sizes1dim_tensor1维度相等。结果就是在第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_tensor2。sizes(3,1,1)len_sizes3。因为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。这在数据预处理或模型推理单样本时非常有用。注意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在视觉TransformerViT或目标检测模型中我们经常需要为特征图的每个位置生成一个唯一的编码。假设我们有一个基于正弦余弦的、针对一维序列的位置编码矩阵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], dim2).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], dim2).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], dim0) y_aug torch.cat([y, y_minority_repeated], dim0) 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_idx1, repeat_times3) 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, dim2) # 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 # 使用repeat4.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, dim0):\n‘, torch.repeat_interleave(x, 2, dim0)) # 输出 # 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], dim1):\n‘, torch.repeat_interleave(x, [1, 3], dim1)) # 输出 # 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_gradTrue) 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的效果再应用到真实数据上能节省大量排查时间。
返回列表