如何优化非均匀批量下基于权重矩阵的张量高效计算?
问题
我有一个分属不同批次的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
相关产品推荐
相关产品推荐

