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

更换自定义Tokenizer后Seq2SeqTrainer输出EvalPrediction异常求助

问题

基于中文BART模型(fnlp/bart-base-chinese)使用Seq2SeqTrainer训练时,替换为自定义Tokenizer后,compute_metrics中EvalPrediction的preds解码结果为无意义乱码,但使用原Tokenizer(BertTokenizer.from_pretrained(checkpoint))时输出正常。推测模型无法识别自定义Tokenizer生成的Token ID,目标是用自定义Tokenizer完成训练。

模型与训练配置代码

model = BartForConditionalGeneration.from_pretrained(checkpoint)
model.resize_token_embeddings(len(tokenizer))
model.config.vocab_size = len(tokenizer)

steps = 500 # small value for debug purpose
batch_size = 4
training_args = CustomSeq2SeqTrainingArguments(
    output_dir = "my_output_dir",
    evaluation_strategy = IntervalStrategy.STEPS,
    optim = "adamw_torch",
    eval_steps = steps,
    logging_steps = steps,
    save_steps = steps,
    learning_rate = 2e-5,
    per_device_train_batch_size = batch_size,
    per_device_eval_batch_size = batch_size,
    weight_decay = 0.01,
    save_total_limit = 1,
    num_train_epochs = 30,
    predict_with_generate = True,
    remove_unused_columns = False, 
    fp16 = True, # save memory
    metric_for_best_model = "bleu",
    load_best_model_at_end = True,
    report_to = "wandb",
    # HuggingFace Hub related
    hub_token = hf_token,
    push_to_hub = True,
    save_safetensors = True,
)

trainer = Seq2SeqTrainer(
    model = model,
    args = training_args,
    train_dataset = tokenized_train_dataset,
    eval_dataset = tokenized_eval_dataset,
    tokenizer = tokenizer,
    data_collator = data_collator,
    compute_metrics = compute_metrics,
    callbacks = [EarlyStoppingCallback(early_stopping_patience=3)],
)

计算指标代码

def postprocess_text(preds, labels):
    preds = [pred.strip() for pred in preds]
    labels = [[label.strip()] for label in labels]

    return preds, labels

def compute_metrics(eval_preds):
    preds, labels = eval_preds

    print("Preds and Labels:", preds[0], labels[0])
    
    if isinstance(preds, tuple):
        preds = preds[0]
    decoded_preds = tokenizer.batch_decode(preds, skip_special_tokens=True)

    labels = np.where(labels != -100, labels, tokenizer.pad_token_id)
    decoded_labels = tokenizer.batch_decode(labels, skip_special_tokens=True)

    print("Decoded Preds (before postprocess):", decoded_preds[0])
    print("Decoded Labels (before postprocess):", decoded_labels[0])

    decoded_preds, decoded_labels = postprocess_text(decoded_preds, decoded_labels)
    print("Decoded Preds:", decoded_preds[0])
    print("Decoded Labels:", decoded_labels[0])

    result_bleu = metric_bleu.compute(predictions=decoded_preds, references=decoded_labels, tokenize='zh')
    result_chrf = metric_chrf.compute(predictions=decoded_preds, references=decoded_labels, word_order=2)
    results = {"bleu": result_bleu["score"], "chrf": result_chrf["score"]}

    prediction_lens = [np.count_nonzero(pred != tokenizer.pad_token_id) for pred in preds]
    results["gen_len"] = np.mean(prediction_lens)
    results = {k: round(v, 4) for k, v in results.items()}
    return results

解决方案

1. 对齐自定义Tokenizer与原模型的特殊Token

BART模型依赖<s>、</s>、<pad>等特殊Token,自定义Tokenizer必须保证这些Token的ID与原fnlp/bart-base-chinese模型完全一致:

  • 检查自定义Tokenizer的pad_token_id、bos_token_id、eos_token_id,确保和原Tokenizer的对应ID相同。
  • 新增词汇时,将其ID放在原词汇表之后,避免覆盖原有Token的ID映射关系。

2. 同步模型配置与自定义Tokenizer

加载自定义Tokenizer后,手动将模型的特殊Token配置同步为Tokenizer的对应值:

# 同步特殊Token到模型配置
model.config.bos_token_id = tokenizer.bos_token_id
model.config.eos_token_id = tokenizer.eos_token_id
model.config.pad_token_id = tokenizer.pad_token_id
model.config.decoder_start_token_id = tokenizer.bos_token_id  # BART解码起始Token默认是bos_token

3. 正确执行token_embeddings扩容

调用resize_token_embeddings时,确保原有Token的权重被保留,新增Token的权重正确初始化:

# 扩容embedding矩阵,自动保留原权重,新增Token用随机初始化
model.resize_token_embeddings(len(tokenizer))
# 可选:若新增领域特定词汇,可手动初始化其embedding(比如用相近词的embedding均值)

4. 验证数据集Tokenization正确性

随机抽取训练/验证集样本,用自定义Tokenizer解码其input_ids和labels,确认解码结果与原文本一致;同时检查labels中的-100掩码仅应用于padding部分,其他位置为有效Token ID。

5. 确保compute_metrics使用正确的Tokenizer实例

避免全局变量引用错误,可将Tokenizer作为参数传入compute_metrics:

# 修改compute_metrics,接收tokenizer参数
def compute_metrics(eval_preds, tokenizer):
    preds, labels = eval_preds
    # ... 原有逻辑 ...

# 初始化trainer时,用partial传入tokenizer
from functools import partial
trainer = Seq2SeqTrainer(
    # ... 其他参数 ...
    compute_metrics=partial(compute_metrics, tokenizer=tokenizer),
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 16:30:07