PyTorch张量第一维度扩展至n并填充零的更优实现方法
优化方案
你的需求可以通过更简洁高效的方式实现,避免循环拼接带来的冗余操作,以下是两种推荐方案:
方案一:直接初始化全零张量并赋值(最优)
这种方法直接创建目标形状的全零张量,再将原张量放入第一个位置,内存仅分配一次,效率最高:
import torch n = 5 a = torch.randn(1, 3, 8) # 形状(1,3,8)的随机张量 b = torch.zeros(n, 3, 8, device=a.device, dtype=a.dtype) # 创建目标形状的全零张量,保持设备和数据类型一致 b[0] = a # 将原张量赋值到第一个维度 print(b.shape) # 输出: torch.Size([5, 3, 8])
方案二:一次性拼接(更简洁)
通过列表生成式一次性生成需要拼接的张量列表,再调用一次torch.cat完成拼接,省去循环:
import torch n = 5 a = torch.randn(1, 3, 8) b = torch.cat([a] + [torch.zeros_like(a)] * (n - 1), dim=0) print(b.shape) # 输出: torch.Size([5, 3, 8])
对比原方案的优势
原方案通过循环多次调用torch.cat,每次拼接都会重新分配内存并复制数据,当n较大时会产生明显的性能损耗;而上述两种方案要么仅初始化一次内存,要么仅执行一次拼接操作,代码更简洁的同时效率也更高。
内容的提问来源于stack exchange,提问作者darth baba
相关产品推荐
相关产品推荐

