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

BERT模型训练:高效截断数据集并添加<PRE>令牌的方法

解决方案:直接在令牌层面操作,无需解码重编码

你完全不需要走「编码→截断→解码→加令牌→重编码」的弯路,直接在tokenizer输出的令牌数组上做插入操作就行——既然你已经把<PRE>添加为特殊令牌,就能直接拿到它的ID,全程不用碰文本解码。

先指出你现有代码的问题

  • tokenized["input_ids"][0] = tokenizer("<PRE>")是错误写法:tokenizer("<PRE>")返回的是包含input_ids的字典,不是单个令牌ID,而且这行是替换原序列的第一个令牌,不是在开头插入。
  • 没显式设置max_length,默认截断长度可能不是511,导致加了<PRE>后总长度超过512。

修正后的完整代码

additional_special_tokens = ["<PRE>", "<SUF>", "<MID>"]

model_name = "your-model-name-here"  # 替换为你的目标模型名称
tokenizer = AutoTokenizer.from_pretrained(model_name, truncation_side="left")
# 正确添加特殊令牌的方式:用add_special_tokens方法,而非直接赋值属性
tokenizer.add_special_tokens({"additional_special_tokens": additional_special_tokens})

small_eval_dataset = full_dataset["validation"].shuffle(42).select(range(1))

# 获取<PRE>对应的令牌ID
pre_token_id = tokenizer.convert_tokens_to_ids("<PRE>")

def build_training_data(examples):
    to_tokenized = examples["context"] + "<SUF><MID>" + examples["gt"]
    # 编码时直接截断到511个令牌,留1个位置给<PRE>
    tokenized = tokenizer(
        to_tokenized,
        truncation=True,
        max_length=511,
        padding=False  # 无需padding,我们要精确控制长度
    )
    
    # 在input_ids开头插入<PRE>的ID
    tokenized["input_ids"] = [pre_token_id] + tokenized["input_ids"]
    # 对应的attention_mask开头插入1(表示该令牌需要被模型关注)
    tokenized["attention_mask"] = [1] + tokenized["attention_mask"]
    
    # 如果tokenizer返回token_type_ids(BERT默认会返回),也在开头插入0
    if "token_type_ids" in tokenized:
        tokenized["token_type_ids"] = [0] + tokenized["token_type_ids"]
    
    return tokenized

small_eval_dataset = small_eval_dataset.map(build_training_data)

关键说明

  1. 添加特殊令牌时必须用tokenizer.add_special_tokens()方法,而非直接赋值属性——这样tokenizer会正确处理这些令牌的编码、解码逻辑,避免潜在的令牌映射错误。
  2. 编码时设置max_length=511,配合左侧截断,刚好保留最后511个令牌,插入<PRE>后总长度为512,完美匹配BERT类模型的输入限制。
  3. 全程只在令牌数组层面操作,完全跳过文本解码步骤,效率大幅提升。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 19:07:04