fix : 修复策略相关文件的类型注解与抽象方法体

- 修复 strategy.py 单元素 Union 与缺失的参数/返回类型注解
- 修复 train_context.py 8 个 default=None 字段缺 Optional 标记
- 修复 sample.py/packing.py/position_id.py 方法缺参数及返回类型注解
- 修复 factory.py _resolve_type/list_registered 缺类型注解
- 修复 train_config.py 裸 dict/list 缺泛型参数
- abstractmethod body 从 ... 改为 raise NotImplementedError
- feat : checkpoint meta.json 保存 TrainConfig 超参供人工查阅
This commit is contained in:
2026-06-14 16:20:10 +08:00
parent a2512f8a5a
commit fec376b0dd
8 changed files with 70 additions and 30 deletions
+4 -4
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field, fields
from typing import Callable, List, Optional
from typing import Any, Callable, Dict, List, Optional
import torch.nn as nn
from torch.optim import Optimizer
@@ -40,7 +40,7 @@ class TrainConfig(BaseConfig):
max_grad_norm: float = field(
default=1.0, metadata={"help": "Maximum gradient norm."}
)
gradient_checkpointing_modules: list = field(
gradient_checkpointing_modules: List[str] = field(
default_factory=list,
metadata={"help": "Module types to enable activation checkpointing for."},
)
@@ -133,11 +133,11 @@ class TrainConfig(BaseConfig):
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
)
executor_kwargs: dict = field(
executor_kwargs: Dict[str, Any] = field(
default_factory=dict,
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
)
extra_kwargs: dict = field(
extra_kwargs: Dict[str, Any] = field(
default_factory=dict, metadata={"help": "Other arguments."}
)