feat: 新增LoRA微调模块

- LoRALinear基于register_parameter托管base weight,state_dict路径不变
- inject_lora/merge_lora/save_lora/load_lora完备封装
- 24个单元测试覆盖注入、合并、存取、边界场景
This commit is contained in:
2026-05-25 20:15:31 +08:00
parent 7df6eb9211
commit a4688021bf
5 changed files with 576 additions and 0 deletions
+9
View File
@@ -6,6 +6,7 @@ from torch.utils.data import DataLoader
from astrai.config.train_config import TrainConfig
from astrai.dataset import ResumableDistributedSampler
from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
@@ -77,6 +78,14 @@ class TrainContextBuilder:
state_dict=context.model.state_dict(),
)
if cfg.lora is not None:
inject_lora(
context.model,
r=cfg.lora.r,
alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules),
)
context.optimizer = cfg.optimizer_fn(context.model)
context.scheduler = cfg.scheduler_fn(context.optimizer)