如何修改Python文本截断代码,按Token数而非字符数截取?
基于Token截断文本的代码修改方案
问题背景
原有Python代码通过截取文本最后n个字符实现历史文本截断,现需改为截取最后n个Token(使用HuggingFace提供的EleutherAI/gpt-j-6B分词器),核心需求是修改text_end = text[-n:]部分,使text_end对应最后n个Token的文本内容。
原代码:
def truncate_history(text, n): """ :param text: The text to truncate :param n: The length of the history (from the end) :return: The truncated text """ text_end = text[-n:] alpha_index = text_end.find("person alpha:\n") beta_index = text_end.find("person beta:\n") if alpha_index == -1 and beta_index == -1: alpha_index = text.rfind("person alpha:\n") beta_index = text.rfind("person beta:\n") i = max(alpha_index, beta_index) elif alpha_index == -1: i = beta_index - n # alpha NOT found elif beta_index == -1: i = alpha_index - n # beta NOT found else: i = min(alpha_index, beta_index) - n # both FOUND return text[i:]
已初始化的分词器:
import transformers tokenizer = transformers.AutoTokenizer.from_pretrained("EleutherAI/gpt-j-6B", pad_token='<|endoftext|>', eos_token='<|endoftext|>')
修改方案
要获取文本最后n个Token对应的内容,需通过编码文本为Token序列→截取最后n个Token→解码回文本的流程实现,具体修改如下:
1. 替换text_end的生成逻辑
将原代码中的text_end = text[-n:]替换为以下代码:
# 编码文本为Token,返回包含input_ids的字典 encoded = tokenizer.encode_plus(text, return_tensors="pt") # 截取最后n个Token的input_ids last_n_token_ids = encoded["input_ids"][0][-n:] # 解码回文本,skip_special_tokens=True可忽略分词器的特殊标记 text_end = tokenizer.decode(last_n_token_ids, skip_special_tokens=True)
2. 调整后续索引计算逻辑(关键补充)
原代码中通过beta_index - n、alpha_index - n计算原文本的截取位置,但现在text_end是最后n个Token对应的文本,其字符长度不再等于n,因此需要重新计算原文本中text_end的起始位置,才能正确推导i的值。
修改后的完整代码:
def truncate_history(text, n, tokenizer): """ :param text: The text to truncate :param n: The length of the history (from the end, in tokens) :param tokenizer: HuggingFace tokenizer instance :return: The truncated text """ # 步骤1:获取最后n个Token对应的文本 encoded = tokenizer.encode_plus(text, return_tensors="pt") input_ids = encoded["input_ids"][0] # 处理n大于总Token数的边界情况 if len(input_ids) <= n: text_end = text start_pos_in_full_text = 0 else: last_n_token_ids = input_ids[-n:] text_end = tokenizer.decode(last_n_token_ids, skip_special_tokens=True) # 取前100个字符匹配,避免重复内容干扰定位 start_pos_in_full_text = text.rfind(text_end[:100]) # 步骤2:查找目标字符串位置 alpha_index = text_end.find("person alpha:\n") beta_index = text_end.find("person beta:\n") if alpha_index == -1 and beta_index == -1: alpha_index = text.rfind("person alpha:\n") beta_index = text.rfind("person beta:\n") i = max(alpha_index, beta_index) elif alpha_index == -1: # 转换为原文本中的绝对位置 i = start_pos_in_full_text + beta_index elif beta_index == -1: i = start_pos_in_full_text + alpha_index else: # 取两个位置中更早的那个 i = start_pos_in_full_text + min(alpha_index, beta_index) # 确保截取位置不小于0 return text[max(i, 0):]
关键说明
- 使用
encode_plus获取Token序列,若不需要PyTorch张量,可将return_tensors="pt"改为return_tensors=None,直接得到Token列表。 - 处理
n大于文本总Token数的边界情况,此时直接返回原文本。 - 通过
text.rfind(text_end[:100])定位text_end在原文本中的起始位置,避免文本中存在重复片段导致匹配错误。 - 最终计算
i时,将text_end内的相对索引加上start_pos_in_full_text,转换为原文本的绝对索引,确保截取位置正确。
内容的提问来源于stack exchange,提问作者BlackHawk
相关产品推荐
相关产品推荐

