You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 10:35:20