feat: 新增LoRA微调模块
- LoRALinear基于register_parameter托管base weight,state_dict路径不变 - inject_lora/merge_lora/save_lora/load_lora完备封装 - 24个单元测试覆盖注入、合并、存取、边界场景
This commit is contained in:
@@ -7,6 +7,7 @@ from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.model.components.lora import LoRAConfig
|
||||
|
||||
|
||||
def required(**kw):
|
||||
@@ -56,6 +57,12 @@ class TrainConfig(BaseConfig):
|
||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||
)
|
||||
|
||||
# lora setting
|
||||
lora: Optional[LoRAConfig] = field(
|
||||
default=None,
|
||||
metadata={"help": "LoRA config. None means full fine-tuning."},
|
||||
)
|
||||
|
||||
# metric setting
|
||||
log_dir: str = field(
|
||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||
|
||||
Reference in New Issue
Block a user