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

如何以最快方式将PyTorch生成的补全结果追加至原始种子张量?

PyTorch向量化实现种子文本扩展优化

我有一个由n个已分词种子组成的2D张量,生成预测logits后用argsort()排序取前m个候选。想找到最优方法,把每个预测追加到生成它的种子上,得到包含n*m个新种子的2D张量(每个种子长度增加一个token)。刚接触PyTorch,想知道有没有内置的向量化方法替代原代码里的嵌套循环。原逻辑的伪代码如下:

def predict(seeds, coherence_threshold, batch_size=16):
    """Takes in a tensor of tokenized seeds, outputs a tensor of tokenized seeds with completions"""

    dataloader = torch.utils.data.DataLoader(seeds, batch_size=batch_size, shuffle=False)

    new_seeds = torch.tensor([], dtype=int)
    with torch.no_grad():
        for batch in dataloader:
            batch_tensors = reference_gpt2(batch)
            batch_preds = batch_tensors.argsort(descending=True)
            batch_preds_pruned = batch_preds[:,-1,:coherence_threshold]
            # TODO come up with a more efficient way to do this
            for i in range(len(batch)):
                for j in range(len(batch_preds[i])):
                    new_seed = torch.concat((batch[i], batch_preds[i,j]))
                    new_seeds = torch.concat((new_seeds, [new_seed]))
                    
    return(new_seeds)

优化方案:用向量化操作替代嵌套循环

原代码里的嵌套循环效率极低,每次循环拼接张量都会触发内存重新分配。下面用PyTorch的广播、维度扩展等向量化操作实现,全程无循环,效率拉满:

def predict(seeds, coherence_threshold, batch_size=16):
    """输入分词后的种子张量,输出追加了候选token的新种子张量"""
    dataloader = torch.utils.data.DataLoader(seeds, batch_size=batch_size, shuffle=False)
    # 初始化结果列表,用来收集每个batch的输出
    new_seeds_list = []
    
    with torch.no_grad():
        for batch in dataloader:
            # 获取模型预测的logits,假设shape是[batch_size, seq_len, vocab_size]
            batch_tensors = reference_gpt2(batch)
            # 对最后一维(词汇表维度)降序排序,取最后一个位置(当前要预测的token)的前m个候选
            # batch_preds_pruned shape: [batch_size, coherence_threshold]
            batch_preds_pruned = batch_tensors.argsort(descending=True)[:, -1, :coherence_threshold]
            
            # 核心向量化操作开始
            batch_size_current = batch.shape[0]
            m = coherence_threshold
            
            # 1. 扩展种子张量:每个种子复制m份,shape从[batch_size, seq_len]变成[batch_size*m, seq_len]
            seeds_expanded = batch.repeat_interleave(m, dim=0)
            
            # 2. 调整候选张量形状:从[batch_size, m]变成[batch_size*m, 1],方便拼接
            preds_reshaped = batch_preds_pruned.flatten().unsqueeze(1)
            
            # 3. 拼接种子和候选token,得到每个种子追加对应候选的结果
            batch_new_seeds = torch.cat([seeds_expanded, preds_reshaped], dim=1)
            
            # 4. 将当前batch的结果加入列表
            new_seeds_list.append(batch_new_seeds)
    
    # 把所有batch的结果拼接成最终张量
    return torch.cat(new_seeds_list, dim=0)

关键步骤解释

  • repeat_interleave:把每个种子重复m次,让每个种子和它的m个候选一一对应,避免循环复制
  • flatten() + unsqueeze(1):把候选张量从二维展平后再增加一个维度,确保和扩展后的种子张量在维度上匹配,能直接拼接
  • 用列表收集结果再一次性拼接:避免原代码里每次循环都拼接new_seeds,减少内存分配次数,提升效率

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 05:05:17