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

如何优化非均匀批量下基于权重矩阵的张量高效计算?

问题

我有一个分属不同批次的1D Token张量,批次大小不均匀。每个批次需要和对应的权重矩阵相乘,当前用batch pointer向量、对应唯一指针的不同权重矩阵加for循环实现。想要高效得到形状为[num_tokens, output_dim]的结果(每个权重矩阵形状是[input_dim, output_dim]),同时为了利用NVIDIA Tensor Cores,会把输入填充到8的整数倍。

当前实现的示例代码(修正笔误后):

# shape [num_tokens,]
input_dim, output_dim = 4, 8
ptr = torch.tensor([0, 1, 1, 2, 2, 2, 3, 3, 3, -1, -1, -1]) # -1 means padding
features = torch.randn(ptr.shape[0], input_dim)
weights = [torch.randn(input_dim, output_dim) for _ in range(4)]

unique = torch.unique(ptr, sorted=False, return_inverse=False, return_counts=False)
unique = unique[unique != -1] # ignore padding 

results = []

for i in unique:
    split = features[ptr == i, :]
    # pad each split to multiple of 8 for NVIDIA A100
    pad = (
        torch.empty((-split.size(0)) % 8, split.size(-1))
        .uniform_()
        .to(split.device)
    )
    padded_split = torch.cat((split, pad), dim=0)
    attn_mask = torch.cat((torch.ones(split.size(0)), torch.zeros(pad.size(0)))).to(
        torch.bool
    )

    # forward pass
    result = padded_split @ weights[i]
    # strip padding so I can create a 2D result tensor of correct dimension again
    results.append(result[attn_mask, :])

results = torch.cat(results, dim=0)

当前方案在推理前向传播时性能下降明显,怀疑是padding操作导致的。考虑过用ptr作为索引的scatter操作,但现有方法只支持求和、均值、最大值等基础归约操作,请问该怎么优化?

优化方案

1. 全局批量处理,避免逐批次padding

核心思路是把所有非padding的特征按组预处理,一次性构建填充后的大张量,搭配对应权重矩阵的批量映射,减少CUDA kernel调用次数:

具体实现:

import torch

input_dim, output_dim = 4, 8
ptr = torch.tensor([0, 1, 1, 2, 2, 2, 3, 3, 3, -1, -1, -1]) # -1 means padding
features = torch.randn(ptr.shape[0], input_dim)
weights = [torch.randn(input_dim, output_dim) for _ in range(4)]
weights_tensor = torch.stack(weights)  # 转为张量:[4, input_dim, output_dim]

# 过滤padding部分
valid_mask = ptr != -1
valid_features = features[valid_mask]
valid_ptr = ptr[valid_mask]

# 计算每个组的填充长度和目标长度
unique_ptr, counts = torch.unique(valid_ptr, sorted=True, return_counts=True)
pad_lengths = (-counts) % 8
target_lengths = counts + pad_lengths

# 构建填充后的特征张量和权重索引
padded_features = []
weight_indices = []
for idx in range(len(unique_ptr)):
    group = valid_features[valid_ptr == unique_ptr[idx]]
    pad = torch.empty((pad_lengths[idx], input_dim)).uniform_().to(group.device)
    padded_group = torch.cat([group, pad], dim=0)
    padded_features.append(padded_group)
    # 记录填充后每个位置对应的权重索引
    weight_indices.extend([unique_ptr[idx]] * target_lengths[idx])

padded_features = torch.cat(padded_features, dim=0)
weight_indices = torch.tensor(weight_indices, device=padded_features.device)

# 批量矩阵乘法,一次性完成所有计算
selected_weights = weights_tensor[weight_indices]  # [total_padded, input_dim, output_dim]
padded_results = torch.bmm(padded_features.unsqueeze(1), selected_weights).squeeze(1)

# 提取有效结果,忽略填充部分
valid_result = padded_results[:len(valid_features)]

# 若需要保留原张量的padding位置,将结果填充回去
final_results = torch.zeros((features.shape[0], output_dim), device=features.device)
final_results[valid_mask] = valid_result

2. 跳过padding,直接映射权重计算

完全避免padding操作,利用索引直接匹配每个token对应的权重矩阵,通过广播或 einsum 完成批量计算:

具体实现:

import torch

input_dim, output_dim = 4, 8
ptr = torch.tensor([0, 1, 1, 2, 2, 2, 3, 3, 3, -1, -1, -1]) # -1 means padding
features = torch.randn(ptr.shape[0], input_dim)
weights = [torch.randn(input_dim, output_dim) for _ in range(4)]
weights_tensor = torch.stack(weights)

# 过滤padding
valid_mask = ptr != -1
valid_features = features[valid_mask]
valid_ptr = ptr[valid_mask]

# 直接获取每个有效token对应的权重矩阵
selected_weights = weights_tensor[valid_ptr]  # [num_valid, input_dim, output_dim]

# 用einsum完成特征与对应权重的乘法
valid_result = torch.einsum('bi,bio->bo', valid_features, selected_weights)

# 填充回原张量(若需要保留padding位置)
final_results = torch.zeros((features.shape[0], output_dim), device=features.device)
final_results[valid_mask] = valid_result

这种方法无需padding,若要利用Tensor Cores,只需确保input_dim和output_dim是8的倍数(若不是,可提前对特征和权重做全局padding到最近的8的倍数)。

3. Tensor Cores适配优化

  • 使用torch.float16或torch.bfloat16数据类型,Tensor Cores对低精度计算的加速效果更显著
  • 确保input_dim和output_dim为8的整数倍,若原始维度不满足,可对特征和权重矩阵做全局padding
  • 尽量合并运算为大张量操作,减少CUDA kernel的调用开销

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 04:54:55