基于PyTorch张量变长索引的词概率高效计算需求
高效计算分词后词的概率(替代循环实现)
输入说明
token_list:维度为n_words × max_tokenization_length,每行代表一个词的分词结果,用pad token补全到最大长度pxhs:维度为n_words × (max_tokenization_length + 1) × |vocabulary|,每行对应一个词的多步预测概率(第i步对应生成第i个token的概率)next_word_token_ids:构成新词起始的token集合(比如所有空格开头的token)
示例输入
import torch pxhs = torch.rand((3,4,1000)) # 3个词,4步预测,词汇表大小1000 pad_token_id = 0 # 假设pad token id为0 word_token_list = [ [120, pad_token_id, pad_token_id], [131, 132, pad_token_id], [140, 141, 142], ] new_word_token_ids = [0,1,2,3,5]
期望计算逻辑
每个词的概率 = 词本身各有效token的概率乘积 × 后续生成新词起始token的概率和
- 词1:
pxhs[0,0,120] × pxhs[0,1,new_word_token_ids].sum() - 词2:
pxhs[1,0,131] × pxhs[1,1,132] × pxhs[1,2,new_word_token_ids].sum() - 词3:
pxhs[2,0,140] × pxhs[2,1,141] × pxhs[2,2,142] × pxhs[2,3,new_word_token_ids].sum()
高效向量化实现
用PyTorch张量批量操作替代循环,充分利用硬件加速:
import torch # 1. 把token列表转为张量,适配批量操作 token_tensor = torch.tensor(word_token_list) # shape: (3, 3) n_words, max_len = token_tensor.shape # 2. 计算每个词的有效token概率乘积 # 生成mask:标记非pad token的位置 mask = token_tensor != pad_token_id # shape: (3, 3) # 批量提取每个token对应的概率 token_probs = torch.gather( pxhs[:, 0:max_len, :], dim=2, index=token_tensor.unsqueeze(-1) ).squeeze(-1) # shape: (3, 3) # pad位置概率设为1(不影响乘积结果),再按行求乘积 token_probs = torch.where(mask, token_probs, torch.tensor(1.0, device=token_probs.device)) word_token_product = token_probs.prod(dim=1) # shape: (3,) # 3. 计算每个词对应的新词起始token概率和 # 提取当前词生成完成后,下一步的预测概率 final_step_logits = pxhs[:, max_len, :] # shape: (3, 1000) # 对指定token集合求和 new_word_prob_sum = final_step_logits[:, new_word_token_ids].sum(dim=1) # shape: (3,) # 4. 得到最终每个词的概率 word_probs = word_token_product * new_word_prob_sum
结果验证
可以用循环实现的结果做对比,确认逻辑一致:
# 循环实现(用于验证) loop_probs = [] for i in range(n_words): tokens = word_token_list[i] prod = 1.0 step = 0 for token in tokens: if token == pad_token_id: break prod *= pxhs[i, step, token] step += 1 prod *= pxhs[i, step, new_word_token_ids].sum() loop_probs.append(prod) # 验证结果是否一致 print(torch.allclose(word_probs, torch.tensor(loop_probs))) # 输出True表示一致
核心优化点
- 用
torch.gather批量提取所有token的概率,避免逐索引遍历 - 用mask处理pad token,跳过无效位置的计算
- 全程张量级批量操作,适配GPU加速,相比循环效率提升几个数量级
内容的提问来源于stack exchange,提问作者clemboclem
相关产品推荐
相关产品推荐

