refactor(trainer): 优化trainer 结构

This commit is contained in:
2025-12-07 21:23:05 +08:00
parent 82e65ccc21
commit c98b175cd5
18 changed files with 314 additions and 424 deletions
+1 -3
View File
@@ -1,5 +1,5 @@
from khaosz.config.model_config import ModelConfig
from khaosz.config.param_config import BaseModelIO, ModelParameter, Checkpoint, ParameterLoader
from khaosz.config.param_config import BaseModelIO, ModelParameter
from khaosz.config.schedule_config import ScheduleConfig, CosineScheduleConfig, SGDRScheduleConfig
from khaosz.config.train_config import TrainConfig
@@ -7,8 +7,6 @@ from khaosz.config.train_config import TrainConfig
__all__ = [
"BaseModelIO",
"ModelParameter",
"Checkpoint",
"ParameterLoader",
"ModelConfig",
"TrainConfig",
+4 -143
View File
@@ -1,11 +1,8 @@
import pickle as pkl
import matplotlib.pyplot as plt
import safetensors.torch as st
import torch.nn as nn
import torch.optim as optim
import safetensors.torch as st
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Self, Union
from typing import Optional, Self, Union
from pathlib import Path
from khaosz.data.tokenizer import BpeTokenizer
@@ -63,7 +60,7 @@ class BaseModelIO:
return self
def to(self, *args, **kwargs) -> Self:
def to(self, *args, **kwargs) -> "BaseModelIO":
"""Move model to device."""
if self.model is not None:
self.model.to(*args, **kwargs)
@@ -77,142 +74,6 @@ class ModelParameter(BaseModelIO):
def save(self, save_dir: Union[str, Path]):
self.save_components(save_dir)
def load(self, load_dir: Union[str, Path]) -> Self:
def load(self, load_dir: Union[str, Path]) -> "ModelParameter":
return self.load_components(load_dir)
@dataclass
class Checkpoint(BaseModelIO):
"""Extended model parameters with training state."""
optimizer_state: Dict[str, Any] = field(
default=None,
metadata={"help": "Optimizer state."}
)
scheduler_state: Dict[str, Any] = field(
default=None,
metadata={"help": "Sampler state."}
)
loss_list: List[float] = field(
default_factory=list,
metadata={"help": "List of training losses."}
)
epoch: int = field(
default=0,
metadata={"help": "Current epoch."}
)
batch_iter: int = field(
default=0,
metadata={"help": "Current iteration."}
)
def _get_training_paths(self, directory: Union[str, Path]) -> dict[str, Path]:
dir_path = Path(directory)
return {
"loss_plot": dir_path / "loss_plot.png",
"training_state": dir_path / "training_state.pkl"
}
def to_dict(self) -> Dict[str, Any]:
return {
"optimizer_state": self.optimizer_state,
"scheduler_state": self.scheduler_state,
"epoch": self.epoch,
"batch_iter": self.batch_iter,
"loss_list": self.loss_list,
}
def from_dict(self, data: Dict[str, Any]) -> Self:
self.optimizer_state = data["optimizer_state"]
self.scheduler_state = data["scheduler_state"]
self.epoch = data["epoch"]
self.batch_iter = data["batch_iter"]
self.loss_list = data["loss_list"]
def save_training_state(self, save_dir: Union[str, Path]):
paths = self._get_training_paths(save_dir)
# Save loss plot
self._plot_loss(str(paths["loss_plot"]))
# Save training state
with open(str(paths["training_state"]), "wb") as f:
pkl.dump(self.to_dict(), f)
def load_training_state(self, load_dir: Union[str, Path]) -> Self:
paths = self._get_training_paths(load_dir)
# Load training state
with open(str(paths["training_state"]), "rb") as f:
train_state = pkl.load(f)
self.from_dict(train_state)
return self
def _plot_loss(self, save_path: str):
"""Plot and save loss curve."""
if not self.loss_list:
return
batch_iter = len(self.loss_list)
plt.figure(figsize=(10, 6))
plt.plot(self.loss_list)
plt.title(f"Training Loss - Iteration {batch_iter}")
plt.xlabel("Batch")
plt.ylabel("Loss")
plt.grid(True)
plt.savefig(save_path, dpi=30, bbox_inches="tight")
plt.close()
def save(self, save_dir: Union[str, Path]):
"""Save complete checkpoint."""
self.save_components(save_dir)
self.save_training_state(save_dir)
def load(self, load_dir: Union[str, Path]) -> Self:
"""Load complete checkpoint."""
self.load_components(load_dir)
self.load_training_state(load_dir)
return self
class ParameterLoader:
"""Factory class for loading model parameters or checkpoints."""
@staticmethod
def load(load_dir: Union[str, Path]) -> Union[ModelParameter, Checkpoint]:
"""Load either ModelParameter or Checkpoint based on directory contents."""
load_dir = Path(load_dir)
# Check for training-specific files
loss_file = load_dir / "loss.pkl"
has_training_data = loss_file.exists()
# Create appropriate instance
if has_training_data:
checkpoint = Checkpoint()
checkpoint.load(str(load_dir))
return checkpoint
else:
params = ModelParameter()
params.load(str(load_dir))
return params
@staticmethod
def create_checkpoint(
model: nn.Module,
tokenizer: BpeTokenizer,
config: ModelConfig,
loss_list: Optional[list[float]] = None,
optimizer: Optional[optim.Optimizer] = None,
) -> Checkpoint:
"""Convenience method to create a training checkpoint."""
return Checkpoint(
model=model,
tokenizer=tokenizer,
config=config,
loss_list=loss_list or [],
optimizer_state=optimizer
)
+9 -3
View File
@@ -1,4 +1,4 @@
from typing import Any, Literal, Dict
from typing import Any, Dict
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
@@ -39,7 +39,10 @@ class CosineScheduleConfig(ScheduleConfig):
default=None,
metadata={"help": "Total training steps for cosine schedule."}
)
schedule_type: Literal["cosine"] = "cosine"
def __post_init__(self) -> None:
self.schedule_type = "cosine"
self.validate()
def get_kwargs(self) -> Dict[str, Any]:
if self.total_steps is None:
@@ -68,7 +71,10 @@ class SGDRScheduleConfig(ScheduleConfig):
default=2,
metadata={"help": "Multiplier for cycle length growth."}
)
schedule_type: Literal["sgdr"] = "sgdr"
def __post_init__(self) -> None:
self.schedule_type = "sgdr"
self.validate()
def get_kwargs(self) -> Dict[str, Any]:
return {
+33 -16
View File
@@ -1,15 +1,20 @@
from dataclasses import dataclass, field
from typing import Optional, TYPE_CHECKING
from torch import nn
from torch.utils.data import Dataset
from torch.optim import Optimizer
from torch.optim.lr_scheduler import LRScheduler
if TYPE_CHECKING:
from khaosz.trainer.strategy import BaseStrategy
from dataclasses import dataclass, field
from typing import Optional
@dataclass
class TrainConfig:
strategy: "BaseStrategy" = field(
# basic setting
model: nn.Module = field(
default=None,
metadata={"help": "Model for training."}
)
strategy: str = field(
default=None,
metadata={"help": "Training strategy."}
)
@@ -21,9 +26,9 @@ class TrainConfig:
default=None,
metadata={"help": "Optimizer for training."}
)
checkpoint_dir: str = field(
default="./checkpoint",
metadata={"help": "Checkpoint directory."}
scheduler: LRScheduler = field(
default=None,
metadata={"help": "Scheduler for training."}
)
n_epoch: int = field(
default=1,
@@ -33,6 +38,16 @@ class TrainConfig:
default=4,
metadata={"help": "Batch size for training."}
)
accumulation_steps: int = field(
default=1,
metadata={"help": "Number of iterations between steps."}
)
max_grad_norm: float = field(
default=1.0,
metadata={"help": "Maximum gradient norm."}
)
# checkpoint setting
start_epoch: int = field(
default=0,
metadata={"help": "Start epoch for training."}
@@ -41,18 +56,14 @@ class TrainConfig:
default=0,
metadata={"help": "Start batch iteration for training."}
)
checkpoint_dir: str = field(
default="./checkpoint",
metadata={"help": "Checkpoint directory."}
)
checkpoint_interval: int = field(
default=5000,
metadata={"help": "Number of iterations between checkpoints."}
)
accumulation_steps: int = field(
default=1,
metadata={"help": "Number of iterations between steps."}
)
max_grad_norm: float = field(
default=1.0,
metadata={"help": "Maximum gradient norm."}
)
# dataloader setting
random_seed: int = field(
@@ -76,4 +87,10 @@ class TrainConfig:
nprocs: int = field(
default=1,
metadata={"help": "Number of processes for distributed training."}
)
# others
kwargs: dict = field(
default_factory=dict,
metadata={"help": "Other arguments."}
)