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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 17:41:31