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

如何让PyTorch自定义高斯热力图函数支持批处理并行计算

解决PyTorch批量生成视频高斯热力图的并行化问题

核心方案:用向量化操作替代for循环

直接利用PyTorch的广播机制,将所有计算在批量维度上并行展开,彻底消除逐样本遍历的开销。关键是确保输入张量包含批量维度,并且所有计算步骤兼容广播规则。

支持批量的高斯热力图生成函数

假设你的centers张量形状为(B, 16, 3, 2)(B是批量大小,16对应帧数,3对应通道数,最后一维是中心的(x,y)坐标),以下是修改后的并行化实现:

import torch

def generate_batch_heatmap(centers, sigma=1.0, img_size=(224, 224)):
    B, T, C, _ = centers.shape
    H, W = img_size

    # 生成单帧单通道的坐标网格:(H, W, 2)
    x = torch.linspace(0, W-1, W, device=centers.device)
    y = torch.linspace(0, H-1, H, device=centers.device)
    xx, yy = torch.meshgrid(x, y, indexing='xy')
    grid = torch.stack([xx, yy], dim=-1)

    # 扩展网格维度以匹配批量/帧/通道的广播要求
    grid = grid.unsqueeze(0).unsqueeze(0).unsqueeze(0)  # (1, 1, 1, H, W, 2)
    # 扩展中心坐标维度,使其能和网格广播计算
    centers = centers.unsqueeze(-2).unsqueeze(-2)  # (B, T, C, 1, 1, 2)

    # 向量化计算所有点到中心的距离平方
    dist_sq = torch.sum((grid - centers) ** 2, dim=-1)
    # 生成高斯热力图
    heatmap = torch.exp(-dist_sq / (2 * sigma ** 2))

    return heatmap  # 输出形状:(B, 16, 3, 224, 224)

使用示例

# 构造批量输入的中心坐标(示例:随机生成0-223的坐标)
batch_size = 10
centers = torch.rand(batch_size, 16, 3, 2) * 223

# 生成批量热力图
mask = generate_batch_heatmap(centers)
print(mask.shape)  # 输出: torch.Size([10, 16, 3, 224, 224])

关键细节说明

  1. 维度对齐:通过unsqueeze扩展网格和中心坐标的维度,让两者能通过广播完成逐批量、逐帧、逐通道的并行计算,无需手动循环。
  2. 设备一致性:确保centers和生成的网格在同一设备(CPU/GPU)上,避免跨设备计算的维度错误。
  3. 灵活适配:如果你的centers原本没有批量维度(比如所有样本用相同中心),可以用centers = centers.unsqueeze(0).repeat(batch_size, 1, 1, 1)快速扩展批量维度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 11:10:29