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

torch.vmap内基于输入形状创建张量的问题及优化问询

针对torch.vmap中创建匹配批量维度张量的优化方案与官方支持说明

更优雅的实现方式

要解决函数内创建的张量与输入BatchedTensor批量维度不兼容的问题,核心是显式将输入的批量维度注入到新创建的张量形状中,无需依赖torch.zeros_like。具体实现如下:

假设输入polynomial的形状为[B1, B2, ..., Bk, N](前k个是批量维度,最后一个是特征维度),要创建的companion张量形状为[B1, B2, ..., Bk, N-1, N-1],可以按以下方式编写函数:

import torch

def process_polynomial(polynomial):
    # 获取特征维度大小
    feat_dim = polynomial.size(-1)
    # 提取所有批量维度的形状
    batch_dims = polynomial.shape[:-1]
    # 构造目标张量形状:批量维度 + 目标矩阵维度
    companion_shape = batch_dims + (feat_dim - 1, feat_dim - 1)
    # 创建带批量维度的零张量,自动对齐输入的设备与数据类型
    companion = torch.zeros(companion_shape, device=polynomial.device, dtype=polynomial.dtype)
    
    # 此处添加你的后续处理逻辑,比如填充companion矩阵的特定位置
    # 示例:填充次对角线为1
    if feat_dim > 1:
        companion[..., 1:, :-1] = torch.eye(feat_dim - 2, device=polynomial.device)
    
    return companion

# 用vmap包装得到批量处理函数
batched_process = torch.vmap(process_polynomial)

# 测试:创建批量输入
batch_poly = torch.randn(32, 5)  # 32个批量,每个多项式5个系数
result = batched_process(batch_poly)
print(result.shape)  # 输出: torch.Size([32, 4, 4])

这种方式的优势是:

  • 完全基于输入的动态形状构造目标张量,适配任意数量的批量维度
  • 自动对齐输入的设备(CPU/GPU)和数据类型,避免兼容性问题
  • 逻辑清晰,比临时解决方案更易维护

官方支持情况

目前PyTorch的vmap机制在处理函数内全新创建张量时,不会自动推断并注入批量维度——这是因为vmap的核心逻辑是追踪输入张量的批量维度,并将其传播到后续操作中,但全新创建的张量没有关联的输入批量信息,因此需要用户显式处理。

关于未来是否会支持自动注入批量维度的场景:PyTorch团队在vmap的迭代规划中,确实考虑过增强对这类动态张量创建的支持(比如通过上下文感知批量维度),但目前没有明确的发布时间表。

内容的提问来源于stack exchange,提问作者harsanyidani

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 04:01:00