refactor(trainer): 优化trainer 结构
This commit is contained in:
@@ -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",
|
||||
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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 {
|
||||
|
||||
@@ -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."}
|
||||
)
|
||||
Reference in New Issue
Block a user