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

如何修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 11:01:30