如何用Hugging Face T5模型输出替换输入序列中的掩码Token
使用T5模型替换输入中的掩码Token
问题场景
我正在使用Hugging Face Transformers库的T5模型,输入序列包含<extra_id_0>、<extra_id_1>这类掩码Token,需要将模型生成输出中对应掩码的内容提取出来,替换输入里的掩码,得到最终通顺的文本。
原代码及输出:
from transformers import T5Tokenizer, T5ForConditionalGeneration tokenizer = T5Tokenizer.from_pretrained("t5-small") model = T5ForConditionalGeneration.from_pretrained("t5-small") input_data = "The <extra_id_0> walks in <extra_id_1> park" input_ids = tokenizer(input_data, return_tensors="pt").input_ids sequence_ids = model.generate(input_ids) output_sequences = tokenizer.batch_decode(sequence_ids) print(output_sequences)
输出结果:
['<pad><extra_id_0> park offers<extra_id_1> the<extra_id_2> park.</s>']
期望最终输出:
The park offers walks in the park.
实现代码
以下是实现掩码替换逻辑的完整代码:
from transformers import T5Tokenizer, T5ForConditionalGeneration tokenizer = T5Tokenizer.from_pretrained("t5-small") model = T5ForConditionalGeneration.from_pretrained("t5-small") input_data = "The <extra_id_0> walks in <extra_id_1> park" input_ids = tokenizer(input_data, return_tensors="pt").input_ids # 模型生成 sequence_ids = model.generate(input_ids) output_sequences = tokenizer.batch_decode(sequence_ids, skip_special_tokens=False)[0] # 处理输出序列,提取各掩码对应的内容 # 移除无关特殊Token clean_output = output_sequences.replace("<pad>", "").replace("</s>", "") # 按掩码标识分割文本 segments = clean_output.split("<extra_id_") # 构建掩码与对应内容的映射字典 mask_map = {} for seg in segments[1:]: parts = seg.split(">", 1) if len(parts) == 2: mask_num = parts[0] content = parts[1].strip() if content: mask_map[f"<extra_id_{mask_num}>"] = content # 替换输入中的掩码Token final_text = input_data for mask, content in mask_map.items(): if mask in final_text: final_text = final_text.replace(mask, content) print(final_text)
逻辑说明
- 清洗输出文本:先移除生成结果中的
<pad>和</s>特殊Token,避免干扰后续内容提取。 - 提取掩码对应内容:以
<extra_id_为分隔符拆分输出文本,遍历拆分后的片段,提取每个掩码编号和对应的填充内容,构建映射关系。 - 替换输入掩码:遍历映射字典,将输入文本中的每个掩码Token替换为对应的生成内容,得到最终通顺的文本。
内容的提问来源于stack exchange,提问作者littleworth
相关产品推荐
相关产品推荐

