refactor: 修改参数传递方案

This commit is contained in:
2026-02-28 18:09:00 +08:00
parent a33d086883
commit b17cc6a6fb
3 changed files with 28 additions and 49 deletions
+6 -25
View File
@@ -22,15 +22,14 @@ class TrainConfig:
default=None,
metadata={"help": "Dataset for training."}
)
optimizer: Optimizer = field(
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
default=None,
metadata={"help": "Optimizer for training."}
metadata={"help": "Optimizer factory for training."}
)
scheduler: LRScheduler = field(
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
default=None,
metadata={"help": "Scheduler for training."}
metadata={"help": "Scheduler factory for training."}
)
n_epoch: int = field(
default=1,
metadata={"help": "Number of epochs for training."}
@@ -105,19 +104,10 @@ class TrainConfig:
default=None,
metadata={"help": "Parallel function for training."}
)
state_dict_wrapper: Optional[Callable] = field(
state_dict_fn: Optional[Callable] = field(
default=None,
metadata={"help": "Parallel function for state dict saving."}
)
optimizer_factory: Optional[Callable[[nn.Module], Optimizer]] = field(
default=None,
metadata={"help": "Optimizer factory for training."}
)
scheduler_factory: Optional[Callable[[Optimizer], LRScheduler]] = field(
default=None,
metadata={"help": "Scheduler factory for training."}
)
# others
device_ids: Optional[List[int]] = field(
@@ -137,19 +127,10 @@ class TrainConfig:
self.validate()
def validate(self):
required_fields = ["model", "strategy", "dataset"]
required_fields = ["model", "strategy", "dataset", "optimizer_fn", "scheduler_fn"]
for field_name in required_fields:
if getattr(self, field_name) is None:
raise ValueError(f"{field_name} is required.")
factory_case = all([self.optimizer_factory, self.scheduler_factory])
argument_case = all([self.optimizer, self.scheduler])
self.nprocs = max(self.nprocs, 1)
if self.nprocs > 1 and not factory_case:
raise ValueError("Distributed training requires optimizer and scheduler factories.")
elif self.nprocs == 1 and not argument_case:
raise ValueError("Single process training requires optimizer and scheduler arguments.")
+6 -12
View File
@@ -36,8 +36,6 @@ class TrainContextBuilder:
self.config = config
self._context = TrainContext(
model=config.model,
optimizer=config.optimizer,
scheduler=config.scheduler,
world_size=get_world_size(),
rank=get_rank(),
)
@@ -46,20 +44,17 @@ class TrainContextBuilder:
self._context.model = self._context.model.to(device=device)
if self.config.nprocs > 1:
fn = self.config.parallel_wrapper
optimizer_fn = self.config.optimizer_factory
scheduler_fn = self.config.scheduler_factory
self._context.model = fn(self._context.model)
self._context.optimizer = optimizer_fn(self._context.model.parameters())
self._context.scheduler = scheduler_fn(self._context.optimizer)
self._context.optimizer = self.config.optimizer_fn(self._context.model.parameters())
self._context.scheduler = self.config.scheduler_fn(self._context.optimizer)
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
if checkpoint is None:
checkpoint = Checkpoint(
optimizer_state_dict=self.config.optimizer.state_dict(),
scheduler_state_dict=self.config.scheduler.state_dict() if self.config.scheduler is not None else None,
optimizer_state_dict=self._context.optimizer.state_dict(),
scheduler_state_dict=self._context.scheduler.state_dict(),
)
else:
# resume from the assigned checkpoint or assigned iteration
@@ -102,6 +97,5 @@ class TrainContextBuilder:
)
return self
def build(self) -> TrainContext:
return self._context