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

使用CrossEncoderTrainer恢复LoRA微调时遇ValueError问题求助

基于LoRA微调CrossEncoder时恢复训练的问题与解决方案

问题背景

在Kaggle平台(12小时运行限制)使用sentence-transformers库,基于LoRA对CrossEncoder模型进行微调,尝试从训练器生成的检查点恢复训练时,调用trainer.train(resume_from_checkpoint=...)触发ValueError。已尝试用model.model.load_adapter(checkpoint)加载权重,但无法恢复优化器状态;检查点由同一训练器生成,文件权限正常,期望能像全量微调那样恢复后续训练轮次。

检查点包含文件

['adapter_model.safetensors', 'trainer_state.json', 'training_args.bin', 'adapter_config.json', 'README.md', 'tokenizer.json', 'tokenizer_config.json', 'scaler.pt', 'scheduler.pt', 'special_tokens_map.json', 'optimizer.pt', 'rng_state.pth']

错误堆栈

ValueError: Unrecognized model in /kaggle/input/stock-dataset-qibot/checkpoint/checkpoint8336_r8_alpha32_dropout0.1_1e3_batch64_mnrLoss. Should have a `model_type` key in its config.json, or contain one of the following strings in its name: albert, align, altclip, aria, aria_text, audio-spectrogram-transformer, autoformer, aya_vision, bamba, bark, bart, beit, bert, bert-generation, big_bird, bigbird_pegasus, biogpt, bit, bitnet, blenderbot, blenderbot-small, blip, blip-2, blip_2_qformer, bloom, bridgetower, bros, camembert, canine, chameleon, chinese_clip, chinese_clip_vision_model, clap, clip, clip_text_model, clip_vision_model, clipseg, clvp, code_llama, codegen, cohere, cohere2, colpali, conditional_detr, convbert, convnext, convnextv2, cpmant, csm, ctrl, cvt, d_fine, dab-detr, dac, data2vec-audio, data2vec-text, data2vec-vision, dbrx, deberta, deberta-v2, decision_transformer, deepseek_v3, deformable_detr, deit, depth_anything, depth_pro, deta, detr, diffllama, dinat, dinov2, dinov2_with_registers, distilbert, donut-swin, dpr, dpt, efficientformer, efficientnet, electra, emu3, encodec, encoder-decoder, ernie, ernie_m, esm, falcon, falcon_mamba, fastspeech2_conformer, flaubert, flava, fnet, focalnet, fsmt, funnel, fuyu, gemma, gemma2, gemma3, gemma3_text, git, glm, glm4, glpn, got_ocr2, gpt-sw3, gpt2, gpt_bigcode, gpt_neo, gpt_neox, gpt_neox_japanese, gptj, gptsan-japanese, granite, granite_speech, granitemoe, granitemoehybrid, granitemoeshared, granitevision, graphormer, grounding-dino, groupvit, helium, hgnet_v2, hiera, hubert, ibert, idefics, idefics2, idefics3, idefics3_vision, ijepa, imagegpt, informer, instructblip, instructblipvideo, internvl, internvl_vision, jamba, janus, jetmoe, jukebox, kosmos-2, layoutlm, layoutlmv2, layoutlmv3, led, levit, lilt, llama, llama4, llama4_text, llava, llava_next, llava_next_video, llava_onevision, longformer, longt5, luke, lxmert, m2m_100, mamba, mamba2, marian, markuplm, mask2former, maskformer, maskformer-swin, mbart, mctct, mega, megatron-bert, mgp-str, mimi, mistral, mistral3, mixtral, mlcd, mllama, mobilebert, mobilenet_v1, mobilenet_v2, mobilevit, mobilevitv2, modernbert, moonshine, moshi, mpnet, mpt, mra, mt5, musicgen, musicgen_melody, mvp, nat, nemotron, nezha, nllb-moe, nougat, nystromformer, olmo, olmo2, olmoe, omdet-turbo, oneformer, open-llama, openai-gpt, opt, owlv2, owlvit, paligemma, patchtsmixer, patchtst, pegasus, pegasus_x, perceiver, persimmon, phi, phi3, phi4_multimodal, phimoe, pix2struct, pixtral, plbart, poolformer, pop2piano, prompt_depth_anything, prophetnet, pvt, pvt_v2, qdqbert, qwen2, qwen2_5_omni, qwen2_5_vl, qwen2_5_vl_text, qwen2_audio, qwen2_audio_encoder, qwen2_moe, qwen2_vl, qwen2_vl_text, qwen3, qwen3_moe, rag, realm, recurrent_gemma, reformer, regnet, rembert, resnet, retribert, roberta, roberta-prelayernorm, roc_bert, roformer, rt_detr, rt_detr_resnet, rt_detr_v2, rwkv, sam, sam_hq, sam_hq_vision_model, sam_vision_model, seamless_m4t, seamless_m4t_v2, segformer, seggpt, sew, sew-d, shieldgemma2, siglip, siglip2, siglip_vision_model, smolvlm, smolvlm_vision, speech-encoder-decoder, speech_to_text, speech_to_text_2, speecht5, splinter, squeezebert, stablelm, starcoder2, superglue, superpoint, swiftformer, swin, swin2sr, swinv2, switch_transformers, t5, table-transformer, tapas, textnet, time_series_transformer, timesfm, timesformer, timm_backbone, timm_wrapper, trajectory_transformer, transfo-xl, trocr, tvlt, tvp, udop, umt5, unispeech, unispeech-sat, univnet, upernet, van, video_llava, videomae, vilt, vipllava, vision-encoder-decoder, vision-text-dual-encoder, visual_bert, vit, vit_hybrid, vit_mae, vit_msn, vitdet, vitmatte, vitpose, vitpose_backbone, vits, vivit, wav2vec2, wav2vec2-bert, wav2vec2-conformer, wavlm, whisper, xclip, xglm, xlm, xlm-prophetnet, xlm-roberta, xlm-roberta-xl, xlnet, xmod, yolos, yoso, zamba, zamba2, zoedepth

问题解答

1. 如何正确使用CrossEncoderTrainer恢复LoRA微调?

由于resume_from_checkpoint默认逻辑不兼容LoRA检查点,需手动初始化模型并加载训练状态,步骤如下:

from sentence_transformers import CrossEncoder
from sentence_transformers.cross_encoder import CrossEncoderTrainer, CrossEncoderTrainingArguments
from peft import LoraConfig, get_peft_model
import torch
import json

checkpoint_path = '/path/to/checkpoint8336_r8_alpha32_dropout0.1_1e3_batch64_mnrLoss'

# 1. 完全复刻训练时的模型初始化流程(参数必须完全一致)
model = CrossEncoder('Alibaba-NLP/gte-multilingual-reranker-base', max_length=512, trust_remote_code=True)
lora_config = LoraConfig(
    r=16, lora_alpha=32, target_modules=['qkv_proj'],
    lora_dropout=0.1, bias="none", task_type="SEQ_CLS"
)
model.model = get_peft_model(model.model, lora_config)

# 2. 加载LoRA适配器权重并保持可训练
model.model.load_adapter(checkpoint_path, is_trainable=True)

# 3. 初始化训练参数(需与之前训练完全一致)
args = CrossEncoderTrainingArguments(
    output_dir='./output',
    per_device_train_batch_size=64,
    # 其他所有训练参数保持和之前一致
)

# 4. 初始化训练器
trainer = CrossEncoderTrainer(
    model=model,
    args=args,
    train_dataset=train_dataset,  # 传入你的训练数据集
    eval_dataset=eval_dataset     # 传入你的验证数据集(可选)
)

# 5. 手动加载优化器、调度器、混合精度scaler状态
trainer.optimizer.load_state_dict(torch.load(f"{checkpoint_path}/optimizer.pt"))
trainer.lr_scheduler.load_state_dict(torch.load(f"{checkpoint_path}/scheduler.pt"))
if trainer.scaler is not None:
    trainer.scaler.load_state_dict(torch.load(f"{checkpoint_path}/scaler.pt"))

# 6. 恢复训练器状态(当前步数、epoch等)
with open(f"{checkpoint_path}/trainer_state.json", 'r') as f:
    trainer_state = json.load(f)
trainer.state = trainer.state.from_dict(trainer_state)

# 7. 启动训练(设resume_from_checkpoint=False,避免触发默认的错误逻辑)
trainer.train(resume_from_checkpoint=False)

2. 该问题是否与sentence-transformers处理LoRA检查点的方式有关?

是的,核心原因有两点:

  • CrossEncoderTrainer底层依赖Hugging Face Trainer,默认resume_from_checkpoint会尝试加载完整模型的配置和权重,但LoRA检查点仅保存适配器权重和训练状态(优化器、调度器等),缺少基础模型的config.json,导致训练器无法识别模型类型,触发报错。
  • sentence-transformers的CrossEncoder包装了基础模型,LoRA微调生成的检查点结构与全量微调不同:全量微调检查点包含完整模型权重文件(如pytorch_model.bin),而LoRA检查点只有adapter_model.safetensors,进一步导致训练器的模型识别逻辑失效。

3. 此场景下load_adapter()与resume_from_checkpoint有何区别?

  • load_adapter()(PEFT库方法):仅负责加载LoRA适配器的权重,将其注入已初始化的基础模型中,让模型具备微调后的适配器参数,但不会恢复训练的其他状态(优化器参数、学习率调度器进度、训练步数、随机数种子等)。适合推理阶段加载微调后的适配器,或需要基于适配器权重重新开始训练的场景。
  • resume_from_checkpoint(Trainer方法):针对全量微调设计,会自动加载模型权重、优化器状态、调度器状态、训练器状态(当前步数、epoch),直接恢复到上次中断的训练状态继续训练。但在LoRA场景下,由于检查点缺少完整模型的配置和权重,会触发模型类型识别错误,无法正常工作。

内容的提问来源于stack exchange,提问作者Tuan Anh Pham

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 10:18:10