如何在PyTorch中对张量行所有组合应用通用函数?
高效实现PyTorch中矩阵行对的全量函数应用(支持自动求导)
核心思路
通过PyTorch的广播机制或**向量化映射(vmap)**替代Python双重循环,利用CUDA/CPU的批量运算能力大幅提升速度,同时完整保留自动求导支持。
方案1:利用广播机制改造函数
如果可以将函数f适配为支持批量输入的版本,广播是最高效的选择。步骤如下:
扩展张量维度:将两个输入矩阵的行维度分别扩展,触发PyTorch的广播机制,使每个行对能并行计算。
- 形状为
(k, d1)的source1用unsqueeze(1)扩展为(k, 1, d1) - 形状为
(k, d2)的source2用unsqueeze(0)扩展为(1, k, d2)
广播后两者形状均变为(k, k, d)(d为对应维度大小),可直接批量运算。
- 形状为
改造函数为批量版本:将原本处理单个行向量的
f,修改为处理批量张量的f_batch。
示例:点积运算
import torch # 原单样本函数 def f(a, b): return torch.dot(a, b) # 批量版本函数 def f_batch(a_batch, b_batch): # a_batch: (k, k, d), b_batch: (k, k, d) return torch.sum(a_batch * b_batch, dim=-1) # 输入矩阵 k, d = 100, 50 source1 = torch.randn(k, d, requires_grad=True) source2 = torch.randn(k, d, requires_grad=True) # 广播运算 source1_expanded = source1.unsqueeze(1) source2_expanded = source2.unsqueeze(0) result = f_batch(source1_expanded, source2_expanded) # 验证自动求导 result.sum().backward() print(source1.grad.shape) # 输出 (100, 50),符合预期
方案2:使用torch.vmap(保留原函数接口)
如果不想修改f的单样本接口,可使用PyTorch 1.10+提供的torch.vmap工具,自动将单样本函数向量化,适配批量输入。
示例:保留原f的实现
import torch from torch import vmap # 原单样本函数(无需修改) def f(a, b): # a: (d,), b: (d,) return torch.nn.functional.kl_div(a.log_softmax(dim=-1), b.softmax(dim=-1), reduction='sum') k, d = 100, 50 source1 = torch.randn(k, d, requires_grad=True) source2 = torch.randn(k, d, requires_grad=True) # 定义行级映射:对source1的一行,计算与source2所有行的f结果 def row_apply(a_row): return vmap(lambda b_row: f(a_row, b_row))(source2) # 对source1所有行应用row_apply result = vmap(row_apply)(source1) # 验证自动求导 result.sum().backward() print(source2.grad.shape) # 输出 (100, 50),符合预期
关键注意事项
- 自动求导支持:两种方案均基于PyTorch原生张量操作,自动求导会被完整追踪,无需额外处理。
- 相同张量输入:无论传入
apply_very_slow(M, M)还是apply_very_slow(M, N),广播和vmap都会自动处理内存共享,不会产生额外开销。 - 返回张量的情况:如果
f返回(m,)形状的张量,最终结果会是(k, k, m),两种方案均能正确处理。 - 性能对比:双重循环是Python级别的O(k²)迭代,广播/vmap是底层硬件加速的批量运算,当k>100时速度提升可达100x以上。
内容的提问来源于stack exchange,提问作者senseiwa
相关产品推荐
相关产品推荐

