refactor: deduplicate low-risk code paths

This commit is contained in:
2026-08-19 16:17:40 +08:00
parent f252af495c
commit 398e8a3ea3
12 changed files with 64 additions and 114 deletions
+12 -12
View File
@@ -11,10 +11,10 @@ from torch.utils.data import Dataset
from astrai.config.base import BaseConfig from astrai.config.base import BaseConfig
from astrai.model.components.lora import LoRAConfig from astrai.model.components.lora import LoRAConfig
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"}) TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"}) PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
_BACKENDS = frozenset({"nccl", "gloo"}) BACKENDS = frozenset({"nccl", "gloo"})
_START_METHODS = frozenset({"spawn", "fork", "forkserver"}) START_METHODS = frozenset({"spawn", "fork", "forkserver"})
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"}) _COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
@@ -129,31 +129,31 @@ class TrainConfig(BaseConfig):
@field_validator("strategy") @field_validator("strategy")
def _validate_strategy(cls, v: str) -> str: def _validate_strategy(cls, v: str) -> str:
if v not in _TRAIN_TYPES: if v not in TRAIN_TYPES:
raise ValueError( raise ValueError(
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}" f"strategy must be one of {sorted(TRAIN_TYPES)}, got {v!r}"
) )
return v return v
@field_validator("parallel_mode") @field_validator("parallel_mode")
def _validate_parallel_mode(cls, v: str) -> str: def _validate_parallel_mode(cls, v: str) -> str:
if v not in _PARALLEL_MODES: if v not in PARALLEL_MODES:
raise ValueError( raise ValueError(
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}" f"parallel_mode must be one of {sorted(PARALLEL_MODES)}, got {v!r}"
) )
return v return v
@field_validator("backend") @field_validator("backend")
def _validate_backend(cls, v: str) -> str: def _validate_backend(cls, v: str) -> str:
if v not in _BACKENDS: if v not in BACKENDS:
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}") raise ValueError(f"backend must be one of {sorted(BACKENDS)}, got {v!r}")
return v return v
@field_validator("start_method") @field_validator("start_method")
def _validate_start_method(cls, v: str) -> str: def _validate_start_method(cls, v: str) -> str:
if v not in _START_METHODS: if v not in START_METHODS:
raise ValueError( raise ValueError(
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}" f"start_method must be one of {sorted(START_METHODS)}, got {v!r}"
) )
return v return v
+4 -4
View File
@@ -383,10 +383,10 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
transform = _build_jsonl_transform(load_path, tokenizer_path) transform = _build_jsonl_transform(load_path, tokenizer_path)
if transform is None: if transform is None:
raise FileNotFoundError( raise FileNotFoundError(
f"JSONL dataset config not found. Expected " "JSONL dataset config not found. Expected "
f"dataset_config.json alongside *.jsonl files, pass " "dataset_config.json alongside *.jsonl files, pass "
f"tokenizer_path= for the built-in messages config, or " "tokenizer_path= for the built-in messages config, or "
f"use processor= for lazy on-the-fly tokenisation." "use processor= for lazy on-the-fly tokenisation."
) )
store.load(load_path, transform=transform, **kwargs) store.load(load_path, transform=transform, **kwargs)
else: else:
+1 -1
View File
@@ -217,7 +217,7 @@ class Store(ABC):
""" """
if self._window_size <= 0: if self._window_size <= 0:
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)") raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
if self._window_size <= 0 or self._length <= self._window_size: if self._length <= self._window_size:
raise IndexError( raise IndexError(
f"Data too short for window: token_count={self._length}, " f"Data too short for window: token_count={self._length}, "
f"window_size={self._window_size}" f"window_size={self._window_size}"
-2
View File
@@ -16,8 +16,6 @@ from abc import ABC, abstractmethod
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Callable, Dict, List, Optional, OrderedDict from typing import Callable, Dict, List, Optional, OrderedDict
import torch
from astrai.inference.cache.buffer import ReqToTokenPool from astrai.inference.cache.buffer import ReqToTokenPool
# ---- data contract: per-task slot state ---- # ---- data contract: per-task slot state ----
+1 -7
View File
@@ -29,12 +29,6 @@ class FFNOutput(TypedDict):
router_stats: Optional[RouterStats] router_stats: Optional[RouterStats]
class RoutedOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
router_stats: Optional[RouterStats]
@FFNFactory.register("mlp") @FFNFactory.register("mlp")
class MLP(nn.Module): class MLP(nn.Module):
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02): def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
@@ -122,7 +116,7 @@ class DeepSeekMoE(nn.Module):
/ self.n_shared_experts / self.n_shared_experts
) )
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput: def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> FFNOutput:
N, D = x.shape N, D = x.shape
K = self.n_activated_experts K = self.n_activated_experts
E = self.n_routed_experts E = self.n_routed_experts
+7 -11
View File
@@ -247,18 +247,15 @@ class LocalStrategy(LaunchStrategy):
ctx.join() ctx.join()
def _detect_launcher() -> str: def _is_external_launcher() -> bool:
"""Detect the distributed launcher from environment. """Whether an external launcher (torchrun/elastic/manual env) started us."""
Returns one of: "torchelastic", "torchrun", "external", "local".
"""
if dist.is_torchelastic_launched(): if dist.is_torchelastic_launched():
return "torchelastic" return True
if "LOCAL_WORLD_SIZE" in os.environ: if "LOCAL_WORLD_SIZE" in os.environ:
return "torchrun" return True
if "RANK" in os.environ and "WORLD_SIZE" in os.environ: if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
return "external" return True
return "local" return False
def spawn_parallel_fn( def spawn_parallel_fn(
@@ -273,8 +270,7 @@ def spawn_parallel_fn(
): ):
if master_port is None: if master_port is None:
master_port = find_free_port() master_port = find_free_port()
launcher = _detect_launcher() if _is_external_launcher():
if launcher in ("torchelastic", "torchrun", "external"):
strategy = TorchrunStrategy( strategy = TorchrunStrategy(
world_size, backend, master_addr, master_port, device_type, start_method world_size, backend, master_addr, master_port, device_type, start_method
) )
+12 -20
View File
@@ -5,7 +5,7 @@ import json
import time import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Union from typing import Any, Callable, Dict, Optional, Union
import safetensors.torch as st import safetensors.torch as st
import torch import torch
@@ -22,39 +22,31 @@ def save_safetensors(state_dict: dict, path: Union[str, Path]):
st.save_file(state_dict, str(path)) st.save_file(state_dict, str(path))
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict: def _broadcast_load(loader: Callable[[], dict], broadcast: bool) -> dict:
"""Load on rank 0 and broadcast the object to all ranks."""
if not broadcast or not dist.is_initialized(): if not broadcast or not dist.is_initialized():
return st.load_file(str(path)) return loader()
rank = get_rank() rank = get_rank()
if rank == 0: if rank == 0:
state_dict = st.load_file(str(path)) data = loader()
else: else:
state_dict = {} data = {}
tmp = [state_dict] tmp = [data]
dist.broadcast_object_list(tmp, src=0) dist.broadcast_object_list(tmp, src=0)
return tmp[0] return tmp[0]
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
return _broadcast_load(lambda: st.load_file(str(path)), broadcast)
def save_json(data: dict, path: Union[str, Path]): def save_json(data: dict, path: Union[str, Path]):
with open(str(path), "w") as f: with open(str(path), "w") as f:
json.dump(data, f, indent=2) json.dump(data, f, indent=2)
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict: def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
if not broadcast or not dist.is_initialized(): return _broadcast_load(lambda: json.loads(Path(path).read_text()), broadcast)
with open(str(path), "r") as f:
return json.load(f)
rank = get_rank()
if rank == 0:
with open(str(path), "r") as f:
data = json.load(f)
else:
data = {}
tmp = [data]
dist.broadcast_object_list(tmp, src=0)
return tmp[0]
def save_torch(obj: Any, path: Union[str, Path]): def save_torch(obj: Any, path: Union[str, Path]):
+2 -18
View File
@@ -94,21 +94,5 @@ def ctx_get_grad_snr(ctx):
return tracker.snr return tracker.snr
def ctx_get_moe_aux_loss(ctx): def ctx_get_moe_metric(ctx, key):
return ctx.strategy._moe_metrics.get("aux_loss") return ctx.strategy._moe_metrics.get(key)
def ctx_get_router_entropy(ctx):
return ctx.strategy._moe_metrics.get("router_entropy")
def ctx_get_dead_expert_fraction(ctx):
return ctx.strategy._moe_metrics.get("dead_expert_fraction")
def ctx_get_load_imbalance_mean(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_mean")
def ctx_get_load_imbalance_max(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_max")
+2 -15
View File
@@ -1,6 +1,6 @@
"""Training strategy implementations with factory pattern.""" """Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod from abc import ABC
from typing import Callable, Dict, List, Optional, TypedDict, Union from typing import Callable, Dict, List, Optional, TypedDict, Union
import torch import torch
@@ -187,7 +187,6 @@ class BaseStrategy(ABC):
self.extra_kwargs = kwargs self.extra_kwargs = kwargs
self._rollout_runner = None self._rollout_runner = None
@abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
"""Compute loss for the given batch. """Compute loss for the given batch.
@@ -197,7 +196,7 @@ class BaseStrategy(ABC):
Returns: Returns:
Computed loss tensor Computed loss tensor
""" """
raise NotImplementedError return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
return self._normalize_output(self.compute_loss(batch)) return self._normalize_output(self.compute_loss(batch))
@@ -328,9 +327,6 @@ class SEQStrategy(BaseStrategy):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
self.label_smoothing = label_smoothing self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
input_ids, target_ids = batch["input_ids"], batch["target_ids"] input_ids, target_ids = batch["input_ids"], batch["target_ids"]
@@ -369,9 +365,6 @@ class SFTStrategy(BaseStrategy):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
self.label_smoothing = label_smoothing self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
input_ids, target_ids, position_ids, loss_mask = ( input_ids, target_ids, position_ids, loss_mask = (
@@ -426,9 +419,6 @@ class DPOStrategy(BaseStrategy):
self.beta = beta self.beta = beta
self.reduction = reduction self.reduction = reduction
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"] chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
@@ -553,9 +543,6 @@ class GRPOStrategy(BaseStrategy):
if state_dict is not None: if state_dict is not None:
self.old_model.load_state_dict(state_dict) self.old_model.load_state_dict(state_dict)
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
prompts = batch["prompts"] prompts = batch["prompts"]
+11 -10
View File
@@ -3,6 +3,7 @@ import logging
import os import os
import sys import sys
import time import time
from functools import partial
from pathlib import Path from pathlib import Path
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
@@ -17,15 +18,11 @@ from astrai.parallel import only_on_rank
from astrai.parallel.setup import get_current_device from astrai.parallel.setup import get_current_device
from astrai.serialization import Checkpoint from astrai.serialization import Checkpoint
from astrai.trainer.metric_util import ( from astrai.trainer.metric_util import (
ctx_get_dead_expert_fraction,
ctx_get_grad_norm, ctx_get_grad_norm,
ctx_get_grad_snr, ctx_get_grad_snr,
ctx_get_load_imbalance_max,
ctx_get_load_imbalance_mean,
ctx_get_loss, ctx_get_loss,
ctx_get_lr, ctx_get_lr,
ctx_get_moe_aux_loss, ctx_get_moe_metric,
ctx_get_router_entropy,
ctx_get_val_loss, ctx_get_val_loss,
) )
from astrai.trainer.train_context import TrainContext from astrai.trainer.train_context import TrainContext
@@ -262,11 +259,15 @@ class MetricCallback(TrainCallback):
"val_loss": ctx_get_val_loss, "val_loss": ctx_get_val_loss,
"grad_norm": ctx_get_grad_norm, "grad_norm": ctx_get_grad_norm,
"grad_snr": ctx_get_grad_snr, "grad_snr": ctx_get_grad_snr,
"moe_aux_loss": ctx_get_moe_aux_loss, "moe_aux_loss": partial(ctx_get_moe_metric, key="aux_loss"),
"router_entropy": ctx_get_router_entropy, "router_entropy": partial(ctx_get_moe_metric, key="router_entropy"),
"dead_expert_fraction": ctx_get_dead_expert_fraction, "dead_expert_fraction": partial(
"load_imbalance_mean": ctx_get_load_imbalance_mean, ctx_get_moe_metric, key="dead_expert_fraction"
"load_imbalance_max": ctx_get_load_imbalance_max, ),
"load_imbalance_mean": partial(
ctx_get_moe_metric, key="load_imbalance_mean"
),
"load_imbalance_max": partial(ctx_get_moe_metric, key="load_imbalance_max"),
} }
def _metrics(self, context: TrainContext, names): def _metrics(self, context: TrainContext, names):
-1
View File
@@ -204,7 +204,6 @@ class TrainContextBuilder:
def _create_dataloaders( def _create_dataloaders(
self, context: TrainContext, train_dataset, val_dataset self, context: TrainContext, train_dataset, val_dataset
) -> None: ) -> None:
cfg = self.config
sampler_offset = context.consumed_samples // context.world_size sampler_offset = context.consumed_samples // context.world_size
if self._resume and sampler_offset > 0: if self._resume and sampler_offset > 0:
samples_per_replica = ( samples_per_replica = (
+12 -13
View File
@@ -11,6 +11,12 @@ from click.core import ParameterSource
from torch import optim from torch import optim
from astrai.config import AutoRegressiveLMConfig, TrainConfig from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.config.train_config import (
BACKENDS,
PARALLEL_MODES,
START_METHODS,
TRAIN_TYPES,
)
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM from astrai.model import AutoRegressiveLM
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
@@ -92,12 +98,12 @@ def _merge_yaml_into_kwargs(
return merged return merged
_TRAIN_TYPE = ["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"] _TRAIN_TYPE = sorted(TRAIN_TYPES)
_PARALLEL = ["none", "ddp", "fsdp"] _PARALLEL = sorted(PARALLEL_MODES)
_SCHEDULES = ["cosine", "sgdr", "wsd"] _SCHEDULES = ["cosine", "sgdr", "wsd"]
_OPTIMIZERS = OptimizerFactory.list_registered() _OPTIMIZERS = OptimizerFactory.list_registered()
_BACKENDS = ["nccl", "gloo"] _BACKENDS = sorted(BACKENDS)
_START_METHODS = ["spawn", "fork", "forkserver"] _START_METHODS = sorted(START_METHODS)
@click.command( @click.command(
@@ -651,17 +657,10 @@ def train(
decay_steps: int, decay_steps: int,
**kwargs, **kwargs,
): ):
if train_type not in [ if train_type not in _TRAIN_TYPE:
"seq",
"sft",
"dpo",
"grpo",
"online_grpo",
"online_dpo",
]:
raise ValueError( raise ValueError(
f"Invalid train_type '{train_type}'. " f"Invalid train_type '{train_type}'. "
f"Must be one of: seq, sft, dpo, grpo, online_grpo, online_dpo" f"Must be one of: {', '.join(_TRAIN_TYPE)}"
) )
if not os.path.exists(param_path): if not os.path.exists(param_path):
raise FileNotFoundError(f"Model directory not found: {param_path}") raise FileNotFoundError(f"Model directory not found: {param_path}")