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)
关键说明
- 添加特殊令牌时必须用
tokenizer.add_special_tokens()方法,而非直接赋值属性——这样tokenizer会正确处理这些令牌的编码、解码逻辑,避免潜在的令牌映射错误。 - 编码时设置
max_length=511,配合左侧截断,刚好保留最后511个令牌,插入<PRE>后总长度为512,完美匹配BERT类模型的输入限制。 - 全程只在令牌数组层面操作,完全跳过文本解码步骤,效率大幅提升。
内容的提问来源于stack exchange,提问作者Shafiq Jetha
相关产品推荐
相关产品推荐

