如何拼接两个指定尺寸的PyTorch张量得到尺寸为[16,121]的张量
PyTorch张量按指定维度拼接实现方案
你需要使用PyTorch内置的torch.cat()接口,指定沿第1维度(维度索引从0开始计数)拼接两个张量即可,两个张量除拼接维度外其余维度尺寸均为16,完全满足拼接要求,最终得到的张量维度正好符合你要的torch.Size([16, 121])。
具体实现代码
- 基础用法(张量a在前,张量b在后):
import torch # 假设a、b是你已定义好的两个张量 result = torch.cat([a, b], dim=1)
- 调整拼接顺序(张量b在前,张量a在后):
result = torch.cat([b, a], dim=1)
结果验证
拼接完成后可以直接打印结果张量的尺寸确认:print(result.shape)
输出将为torch.Size([16, 121])
注意:
torch.cat()仅会在指定的维度对张量进行拼接,不会改变张量的总维度数,如果你误用torch.stack()接口会额外新增一个维度,无法得到你需要的尺寸。
内容的提问来源于stack exchange,提问作者LamaMo
相关产品推荐
相关产品推荐

