PyTorch中如何合并两个张量以包含全部元素(类似列表拼接)
PyTorch张量按行拼接实现方法
你需要的是将两个张量按行方向(维度0)拼接,PyTorch中+运算符是元素逐位相加,无法实现列表式的拼接,可通过torch.cat()函数完成需求。
代码示例
首先定义输入张量:
import torch a = torch.tensor([[101, 103], [101, 1045]]) b = torch.tensor([[101, 777], [101, 888]])
执行拼接操作:
c = torch.cat([a, b], dim=0)
得到的结果张量c即为目标:
tensor([[ 101, 103], [ 101, 1045], [ 101, 777], [ 101, 888]])
补充说明
torch.cat()的dim参数用于指定拼接维度:dim=0对应行方向拼接(扩展行数),dim=1对应列方向拼接(扩展列数)- 拼接的张量在非拼接维度上的形状必须一致,比如这里a和b都是2行2列,非拼接维度(列数)都是2,满足拼接条件
内容的提问来源于stack exchange,提问作者Brana
相关产品推荐
相关产品推荐

