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

如何在PyTorch中对张量行所有组合应用通用函数?

高效实现PyTorch中矩阵行对的全量函数应用(支持自动求导)

核心思路

通过PyTorch的广播机制或**向量化映射(vmap)**替代Python双重循环,利用CUDA/CPU的批量运算能力大幅提升速度,同时完整保留自动求导支持。


方案1:利用广播机制改造函数

如果可以将函数f适配为支持批量输入的版本,广播是最高效的选择。步骤如下:

  1. 扩展张量维度:将两个输入矩阵的行维度分别扩展,触发PyTorch的广播机制,使每个行对能并行计算。

    • 形状为(k, d1)的source1用unsqueeze(1)扩展为(k, 1, d1)
    • 形状为(k, d2)的source2用unsqueeze(0)扩展为(1, k, d2)
      广播后两者形状均变为(k, k, d)(d为对应维度大小),可直接批量运算。
  2. 改造函数为批量版本:将原本处理单个行向量的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 14:22:40