如何以最快方式将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
相关产品推荐
相关产品推荐

