From 5a942527b245489db5519b51540109252da9d757 Mon Sep 17 00:00:00 2001 From: ccx1324 <202431060119@qq.com> Date: Mon, 20 Jul 2026 17:00:24 +0800 Subject: [PATCH] fix: inject LoRA before loading checkpoint state_dict move inject_lora() before load_state_dict in _before_wrap so that LoRA adapter weights from a checkpoint are properly restored on training resume. Previously, inject happened after load, causing lora_A/lora_B keys to be silently ignored (strict=False). Co-Authored-By: ccx1324 <2424441089@qq.com> --- astrai/trainer/train_context.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index ac62be5..d67c43d 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -110,8 +110,6 @@ class TrainContextBuilder: def _before_wrap(m): m = m.to(device=device) - if preloaded_state_dict is not None: - m.load_state_dict(preloaded_state_dict, strict=False) if cfg.lora is not None: inject_lora( m, @@ -119,6 +117,8 @@ class TrainContextBuilder: alpha=cfg.lora.alpha, target_modules=set(cfg.lora.target_modules), ) + if preloaded_state_dict is not None: + m.load_state_dict(preloaded_state_dict, strict=False) return m context = TrainContext(