PyTorch中合并两个张量的标准方法是什么?(对应维度元素相邻)
合并两个2D张量(沿dim=1对应部分相邻排列)的标准方法
针对你描述的「将两个2D张量沿dim=1的对应列彼此相邻排列」的需求,主流深度学习框架(PyTorch/TensorFlow)都有简洁的实现方式,核心思路是先将对应列配对到新维度,再合并目标维度与新维度,以下是具体操作:
PyTorch 实现
假设两个输入张量形状均为 (N, C)(N为样本数,C为列数),合并后目标形状为 (N, 2*C):
方法1:torch.stack + flatten
通过stack在新维度(dim=2)将两个张量的对应列配对,再展平dim=1与dim=2维度:
import torch # 示例张量 A = torch.randn(3, 2) # 形状 (3, 2) B = torch.randn(3, 2) # 形状 (3, 2) # 沿dim=2堆叠,得到形状 (3, 2, 2) stacked = torch.stack([A, B], dim=2) # 展平dim=1及之后的维度,得到目标形状 (3, 4) merged = stacked.flatten(start_dim=1)
方法2:torch.cat + 维度扩展
先给两个张量添加一个新维度,再在该维度拼接,最后展平:
merged = torch.cat([A.unsqueeze(2), B.unsqueeze(2)], dim=2).flatten(start_dim=1)
TensorFlow 实现
逻辑与PyTorch一致,对应API略有差异:
import tensorflow as tf # 示例张量 A = tf.random.normal((3, 2)) B = tf.random.normal((3, 2)) # 方法1:stack后reshape stacked = tf.stack([A, B], axis=2) merged = tf.reshape(stacked, (tf.shape(A)[0], -1)) # 方法2:concat后reshape merged = tf.reshape( tf.concat([tf.expand_dims(A, 2), tf.expand_dims(B, 2)], axis=2), (tf.shape(A)[0], -1) )
通用说明
如果处理更高维度的张量,只需调整stack/concat的目标维度即可——核心是先将需要相邻的对应元素配对到新维度,再合并该新维度与原目标维度。
内容的提问来源于stack exchange,提问作者user3435407
相关产品推荐
相关产品推荐

