PyTorch中按指定索引规则合并张量A与B生成目标张量
张量拼接操作实现方案
问题说明
现有形状为[4, 4096]的张量A和形状为[5, 4096]的张量B,需完成以下操作:
- 沿第0轴取出B的每个元素,复制该元素后分别堆叠到A的第0轴首尾,得到形状为
[6, 4096]的临时张量; - 对B的所有元素重复上述操作,最终将所有临时张量拼接为形状
[30, 4096]的张量T。
结构可视化
A = [A1, A2, A3, A4] # 每个A_i为[4096]维度张量 B = [B1, B2, B3, B4, B5] # 每个B_i为[4096]维度张量 最终张量T的结构: [ B1, A1, A2, A3, A4, B1, B2, A1, A2, A3, A4, B2, B3, A1, A2, A3, A4, B3, B4, A1, A2, A3, A4, B4, B5, A1, A2, A3, A4, B5 ] 维度对应: - B: [5, 4096] - A: [4, 4096] - T: [30, 4096](5组×6个元素)
实现方法
基础循环实现(PyTorch)
适合新手理解逻辑,逐元素处理:
import torch # 构造示例张量(实际使用时替换为你的张量) A = torch.randn(4, 4096) B = torch.randn(5, 4096) temp_tensors = [] for b_element in B: # 将单个B元素升维为[1, 4096],满足拼接维度要求 b_expanded = b_element.unsqueeze(0) # 拼接得到B_i + A + B_i的临时张量 temp = torch.cat([b_expanded, A, b_expanded], dim=0) temp_tensors.append(temp) # 拼接所有临时张量得到最终结果 T = torch.cat(temp_tensors, dim=0) # 验证形状 print(T.shape) # 输出: torch.Size([30, 4096])
高效向量化实现(PyTorch)
避免循环,利用广播机制提升运算效率:
import torch A = torch.randn(4, 4096) B = torch.randn(5, 4096) # 将A复制5次,形状变为[5, 4, 4096] A_repeated = A.unsqueeze(0).repeat(5, 1, 1) # 将B调整为[5, 1, 4096],方便和A_repeated拼接 B_expanded = B.unsqueeze(1) # 沿第1轴拼接B、重复后的A、B,再展平前两维得到[30, 4096] T = torch.cat([B_expanded, A_repeated, B_expanded], dim=1).flatten(0, 1) print(T.shape) # 输出: torch.Size([30, 4096])
TensorFlow版本实现
如果使用TensorFlow,逻辑类似:
import tensorflow as tf A = tf.random.normal((4, 4096)) B = tf.random.normal((5, 4096)) # 向量化实现 A_repeated = tf.expand_dims(A, 0) A_repeated = tf.tile(A_repeated, [5, 1, 1]) B_expanded = tf.expand_dims(B, 1) T = tf.concat([B_expanded, A_repeated, B_expanded], axis=1) T = tf.reshape(T, (-1, 4096)) print(tf.shape(T)) # 输出: Tensor([30, 4096])
内容的提问来源于stack exchange,提问作者SAUMYA BHANDARY
相关产品推荐
相关产品推荐

