From 164be9708b3e2643d1d2ca036bf2cd4416e560dd Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Thu, 30 Jul 2026 08:21:17 +0800 Subject: [PATCH] refactor: migrate config system to Pydantic dataclasses - Replace hand-rolled BaseConfig (from_dict/to_dict/_coerce/_unwrap_optional) with pydantic.dataclasses - from_dict now uses cls(**d), to_dict uses dataclasses.asdict + json.dumps filter - TrainConfig: required fields are now truly required (no default=None), delete manual validate()/__post_init__ - Remove dead required() helper and metadata={'help': ...} annotations - Fix gradient_checkpointing_modules type from List[str] to List[type] - Add pydantic>=2.0 as direct dependency in pyproject.toml - Add numpy-style Parameters docstrings to all config classes - Enable use_attribute_docstrings in BaseConfig for schema generation - LoRAConfig also migrated to pydantic dataclass --- astrai/config/base.py | 100 +++-------- astrai/config/model_config.py | 69 ++++++-- astrai/config/preprocess_config.py | 80 ++++----- astrai/config/train_config.py | 254 ++++++++++------------------ astrai/inference/api/tool_parser.py | 25 ++- astrai/model/components/lora.py | 3 +- pyproject.toml | 1 + 7 files changed, 218 insertions(+), 314 deletions(-) diff --git a/astrai/config/base.py b/astrai/config/base.py index 7758d54..ffdf880 100644 --- a/astrai/config/base.py +++ b/astrai/config/base.py @@ -1,92 +1,32 @@ import json -from dataclasses import MISSING, dataclass, fields +from dataclasses import asdict from pathlib import Path -from typing import Any, Dict, Optional, Self, Union, get_type_hints +from typing import Any, Dict, Self, Union + +from pydantic import ConfigDict +from pydantic.dataclasses import dataclass -@dataclass +@dataclass(config=ConfigDict(use_attribute_docstrings=True)) class BaseConfig: def to_dict(self) -> Dict[str, Any]: - d = {} - for fld in fields(self): - v = getattr(self, fld.name) - if isinstance(v, (str, int, float, bool)): - d[fld.name] = v - elif v is None: - d[fld.name] = None - elif isinstance(v, (dict, list, tuple)): - try: - val = list(v) if isinstance(v, tuple) else v - json.dumps(val) - d[fld.name] = val - except (TypeError, ValueError): - pass - elif isinstance(v, BaseConfig): - d[fld.name] = v.to_dict() - elif hasattr(v, "__dataclass_fields__"): - sub = {} - for f in fields(v): - a = getattr(v, f.name) - sub[f.name] = list(a) if isinstance(a, tuple) else a - d[fld.name] = sub - return d + result = {} + for k, v in asdict(self).items(): + if isinstance(v, tuple): + v = list(v) + try: + json.dumps(v) + result[k] = v + except (TypeError, ValueError): + # Skip non-serializable runtime objects (e.g. model_fn, dataset). + # TrainConfig mixes hyperparams with callables/datasets; only the + # JSON-serializable subset is written to checkpoint meta. + pass + return result @classmethod def from_dict(cls, d: Dict[str, Any]) -> Self: - hints = get_type_hints(cls) - inst = cls.__new__(cls) - for fld in fields(cls): - if fld.name in d: - v = d[fld.name] - target = cls._unwrap_optional(hints.get(fld.name)) - if target is not None: - try: - v = cls._coerce(v, target) - except (TypeError, ValueError): - pass - object.__setattr__(inst, fld.name, v) - elif fld.default is not MISSING: - object.__setattr__(inst, fld.name, fld.default) - elif fld.default_factory is not MISSING: - object.__setattr__(inst, fld.name, fld.default_factory()) - else: - object.__setattr__(inst, fld.name, None) - return inst - - @staticmethod - def _unwrap_optional(tp) -> Optional[type]: - if tp is None: - return None - origin = getattr(tp, "__origin__", None) - if origin is not None: - args = getattr(tp, "__args__", ()) - non_none = [a for a in args if a is not type(None)] - return non_none[0] if non_none else None - return tp - - @staticmethod - def _coerce(value: Any, target_type: type) -> Any: - if target_type is bool and isinstance(value, bool): - return value - if ( - target_type is int - and isinstance(value, (int, float)) - and not isinstance(value, bool) - ): - return int(value) - if ( - target_type is float - and isinstance(value, (int, float)) - and not isinstance(value, bool) - ): - return float(value) - if target_type is str and isinstance(value, str): - return value - if isinstance(value, target_type): - return value - if isinstance(value, dict) and issubclass(target_type, BaseConfig): - return target_type.from_dict(value) - raise TypeError + return cls(**d) @classmethod def from_file(cls, path: Union[str, Path]) -> Self: diff --git a/astrai/config/model_config.py b/astrai/config/model_config.py index a74842a..cc6f414 100644 --- a/astrai/config/model_config.py +++ b/astrai/config/model_config.py @@ -1,6 +1,7 @@ -from dataclasses import dataclass from typing import Any, Dict, Optional +from pydantic.dataclasses import dataclass + from astrai.config.base import BaseConfig from astrai.factory import BaseFactory @@ -17,7 +18,12 @@ class ConfigFactory(BaseFactory[BaseConfig]): @dataclass class BaseModelConfig(BaseConfig): - """Base config with ``model_type`` dispatch and file I/O.""" + """Base config with ``model_type`` dispatch and file I/O. + + Args: + model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None. + neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0. + """ model_type: Optional[str] = None neftune_alpha: float = 0.0 @@ -26,7 +32,34 @@ class BaseModelConfig(BaseConfig): @dataclass @ConfigFactory.register("autoregressive_lm") class AutoRegressiveLMConfig(BaseModelConfig): - """Configuration for autoregressive language model.""" + """Configuration for autoregressive language model. + + Args: + model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None. + neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0. + vocab_size (Optional[int]): Vocabulary size. Defaults to None. + hidden_size (Optional[int]): Hidden dimension size. Defaults to None. + num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None. + rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None. + intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None. + tie_word_embeddings (Optional[bool]): Whether to tie embedding and lm_head weights. Defaults to None. + max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None. + rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None. + rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None. + attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa". + num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None. + num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None. + use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None. + use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None. + kv_lora_rank (Optional[int]): KV compression rank, MLA only. Defaults to None. + qk_nope_head_dim (Optional[int]): Non-RoPE head dimension, MLA only. Defaults to None. + qk_rope_head_dim (Optional[int]): RoPE head dimension, MLA only. Defaults to None. + ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp". + n_routed_experts (Optional[int]): Number of routed experts, MoE only. Defaults to None. + n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None. + n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None. + topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None. + """ vocab_size: Optional[int] = None hidden_size: Optional[int] = None @@ -34,21 +67,17 @@ class AutoRegressiveLMConfig(BaseModelConfig): rms_norm_eps: Optional[float] = None intermediate_size: Optional[int] = None tie_word_embeddings: Optional[bool] = None - max_position_embeddings: Optional[int] = None rope_theta: Optional[float] = None rope_scaling: Optional[dict] = None - attn_type: str = "gqa" num_attention_heads: Optional[int] = None num_key_value_heads: Optional[int] = None use_qk_norm: Optional[bool] = None use_gated_attention: Optional[bool] = None - kv_lora_rank: Optional[int] = None qk_nope_head_dim: Optional[int] = None qk_rope_head_dim: Optional[int] = None - ffn_type: str = "mlp" n_routed_experts: Optional[int] = None n_shared_experts: Optional[int] = None @@ -59,24 +88,42 @@ class AutoRegressiveLMConfig(BaseModelConfig): @dataclass @ConfigFactory.register("embedding") class EncoderConfig(BaseModelConfig): - """Configuration for embedding encoder model.""" + """Configuration for embedding encoder model. + + Args: + model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None. + neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0. + vocab_size (Optional[int]): Vocabulary size. Defaults to None. + hidden_size (Optional[int]): Hidden dimension size. Defaults to None. + num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None. + rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None. + intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None. + max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None. + rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None. + rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None. + attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa". + num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None. + num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None. + use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None. + use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None. + ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp". + pooling_type (Optional[str]): Pooling strategy for embedding, e.g. 'mean', 'cls'. Defaults to None. + normalize_embeddings (Optional[bool]): Whether to L2-normalize output embeddings. Defaults to None. + """ vocab_size: Optional[int] = None hidden_size: Optional[int] = None num_hidden_layers: Optional[int] = None rms_norm_eps: Optional[float] = None intermediate_size: Optional[int] = None - max_position_embeddings: Optional[int] = None rope_theta: Optional[float] = None rope_scaling: Optional[dict] = None - attn_type: str = "gqa" num_attention_heads: Optional[int] = None num_key_value_heads: Optional[int] = None use_qk_norm: Optional[bool] = None use_gated_attention: Optional[bool] = None - ffn_type: str = "mlp" pooling_type: Optional[str] = None normalize_embeddings: Optional[bool] = None diff --git a/astrai/config/preprocess_config.py b/astrai/config/preprocess_config.py index 91219d4..28abcbe 100644 --- a/astrai/config/preprocess_config.py +++ b/astrai/config/preprocess_config.py @@ -5,9 +5,11 @@ modes, both driven declaratively through ``input.sections`` or ``input.sources``. """ -from dataclasses import dataclass, field +from dataclasses import field from typing import Dict, List, Optional +from pydantic.dataclasses import dataclass + from astrai.config.base import BaseConfig @@ -25,6 +27,10 @@ class InputConfig(BaseConfig): "chosen": {"sections": [{"field": "chosen", ...}]}, "rejected": {"sections": [{"field": "rejected", ...}]}, }}} + + Args: + sections (Optional[List[Dict]]): Section list for single-output mode. Defaults to None. + sources (Optional[Dict[str, Dict]]): Source map for multi-output mode, DPO/GRPO. Defaults to None. """ sections: Optional[List[Dict]] = None @@ -33,34 +39,17 @@ class InputConfig(BaseConfig): @dataclass class ProcessingConfig(BaseConfig): - """Processing configuration. + """Processing configuration for tokenization and packing. - Parameters - ---------- - max_seq_len : int - Maximum sequence length (default: 2048). - min_chars : int - Minimum number of characters to keep (default: 50). - max_chars : int - Maximum number of characters to keep (default: 2_000_000). - max_items : Optional[int] - Maximum number of items to process (default: None, unlimited). - batch_size : int - Number of records tokenized together (default: 256). - packing_strategy : str - How to pack sequences into a contiguous stream. - - - ``"simple"``: sequential concatenation (default, backward compatible). - - ``"bfd"``: best-fit decreasing bin packing, minimises wasted tokens. - - ``"bfd_split"``: BFD with over-length sequences split into chunks. - max_packed_len : int - Maximum length of a packed bin. Sequences longer than this are - truncated or split depending on ``packing_strategy`` (default: 8192). - truncation_mode : str - How to truncate sequences longer than ``max_packed_len``. - - - ``"keep_start"``: keep the first ``max_packed_len`` tokens (default). - - ``"keep_end"``: keep the last ``max_packed_len`` tokens. + Args: + max_seq_len (int): Maximum sequence length. Defaults to 2048. + min_chars (int): Minimum number of characters to keep. Defaults to 50. + max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000. + max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None. + batch_size (int): Number of records tokenized together. Defaults to 256. + packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple". + max_packed_len (int): Maximum length of a packed bin. Defaults to 8192. + truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start". """ max_seq_len: int = 2048 @@ -75,24 +64,14 @@ class ProcessingConfig(BaseConfig): @dataclass class OutputConfig(BaseConfig): - """Output configuration. + """Output configuration for storage. - Parameters - ---------- - domain_key : Optional[str] - Domain key for the output store (default: None). - storage_format : str - Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``). - max_tokens_per_shard : int - Maximum tokens per shard before splitting (default: 100_000_000). - dtype : Dict[str, str] - Per-key dtype overrides, e.g. ``{"input_ids": "int32"}`` (default: {}). - position_ids_mode : Optional[str] - How to compute position_ids in packed sequences. - - - ``"none"``: do not generate (default). - - ``"doc_reset"``: reset to 0 at each document boundary. - - ``"continuous"``: sequential 0, 1, 2, ... (pretrain, single doc). + Args: + domain_key (Optional[str]): Domain key for the output store. Defaults to None. + storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin". + max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000. + dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}. + position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset". """ domain_key: Optional[str] = None @@ -104,6 +83,17 @@ class OutputConfig(BaseConfig): @dataclass class PipelineConfig(BaseConfig): + """Top-level preprocessing pipeline config. + + Args: + version (int): Config schema version. Defaults to 1. + input (InputConfig): Input mapping config. + mask (Dict[str, str]): Per-field mask labels, e.g. {"system": "mask", "assistant": "train"}. Defaults to {}. + mask_default (str): Default mask label for unlisted fields. Defaults to "mask". + preprocessing (ProcessingConfig): Processing config. + output (OutputConfig): Output config. + """ + version: int = 1 input: InputConfig = field(default_factory=InputConfig) mask: Dict[str, str] = field(default_factory=dict) diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index 82d514d..d0e2e45 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -1,7 +1,9 @@ -from dataclasses import dataclass, field, fields +from dataclasses import field from typing import Any, Callable, Dict, List, Optional import torch.nn as nn +from pydantic import ConfigDict +from pydantic.dataclasses import dataclass from torch.optim import Optimizer from torch.optim.lr_scheduler import LRScheduler from torch.utils.data import Dataset @@ -10,175 +12,103 @@ from astrai.config.base import BaseConfig from astrai.model.components.lora import LoRAConfig -def required(**kw): - return {"required": True, **kw} - - -@dataclass +@dataclass(config=ConfigDict(arbitrary_types_allowed=True)) class TrainConfig(BaseConfig): - # basic setting - model_fn: Callable[[], nn.Module] = field( - default=None, metadata=required(help="Model factory for training.") - ) - strategy: str = field(default=None, metadata=required(help="Training strategy.")) - dataset: Dataset = field( - default=None, metadata=required(help="Dataset for training.") - ) - optimizer_fn: Callable[[nn.Module], Optimizer] = field( - default=None, metadata=required(help="Optimizer factory for training.") - ) - scheduler_fn: Callable[[Optimizer], LRScheduler] = field( - default=None, metadata=required(help="Scheduler factory for training.") - ) - n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."}) - batch_per_device: int = field( - default=4, metadata={"help": "Batch size per device."} - ) - grad_accum_steps: int = field( - default=1, metadata={"help": "Number of iterations between steps."} - ) - max_grad_norm: Optional[float] = field( - default=1.0, - metadata={"help": "Maximum gradient norm. None disables clipping."}, - ) - gradient_checkpointing_modules: List[str] = field( - default_factory=list, - metadata={"help": "Module types to enable activation checkpointing for."}, - ) - compile_mode: Optional[str] = field( - default=None, - metadata={ - "help": "torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None to disable." - }, - ) + """Training configuration. - # checkpoint setting - start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."}) - start_samples: int = field( - default=0, - metadata={ - "help": "Start samples count (per rank). Superseded by checkpoint consumed_samples." - }, - ) - ckpt_dir: str = field( - default="./checkpoint", metadata={"help": "Checkpoint directory."} - ) - ckpt_interval: int = field( - default=5000, - metadata={"help": "Number of optimizer steps between checkpoints."}, - ) + Combines hyperparameters with runtime objects (model_fn, dataset, etc.). + Only JSON-serializable fields are written to checkpoint meta via to_dict(). - # lora setting - lora: Optional[LoRAConfig] = field( - default=None, - metadata={"help": "LoRA config. None means full fine-tuning."}, - ) + Args: + model_fn (Callable[[], nn.Module]): Model factory for training. + strategy (str): Training strategy (seq, sft, dpo, grpo, online_*). + dataset (Dataset): Dataset for training. + optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training. + scheduler_fn (Callable[[Optimizer], LRScheduler]): Scheduler factory for training. + n_epoch (int): Number of epochs for training. Defaults to 1. + batch_per_device (int): Batch size per device. Defaults to 4. + grad_accum_steps (int): Number of iterations between optimizer steps. Defaults to 1. + max_grad_norm (Optional[float]): Maximum gradient norm. None disables clipping. Defaults to 1.0. + gradient_checkpointing_modules (List[type]): Module types to enable activation checkpointing for. Defaults to []. + compile_mode (Optional[str]): torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None. Defaults to None. + start_epoch (int): Start epoch for training. Defaults to 0. + start_samples (int): Start samples count (per rank). Superseded by checkpoint consumed_samples. Defaults to 0. + ckpt_dir (str): Checkpoint directory. Defaults to "./checkpoint". + ckpt_interval (int): Number of optimizer steps between checkpoints. Defaults to 5000. + lora (Optional[LoRAConfig]): LoRA config. None means full fine-tuning. Defaults to None. + metrics (List[str]): Metrics to record during training. Defaults to ["loss", "lr", "grad_norm"]. + random_seed (int): Random seed. Defaults to 3407. + num_workers (int): Number of workers for dataloader. Defaults to 0. + prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None. + pin_memory (bool): Pin memory for dataloader. Defaults to False. + collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None. + nprocs (int): Number of processes for distributed training. Defaults to 1. + backend (str): Distributed training backend. Defaults to "nccl". + master_addr (str): Master address for distributed training. Defaults to "localhost". + master_port (str): Master port for distributed training. Defaults to "29500". + parallel_mode (str): Parallel strategy: none, ddp, fsdp. Defaults to "none". + start_method (str): Multiprocessing start method: spawn/fork/forkserver. Defaults to "spawn". + device_type (str): Device type for distributed training. Defaults to "cuda". + val_dataset (Optional[Dataset]): Dataset for validation. Defaults to None. + val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None. + val_step (int): Number of optimizer steps between validation runs. Defaults to 1000. + neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0. + rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512. + rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7. + rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0. + rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9. + rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024. + reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None. + executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}. + extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}. + """ - # metric setting - metrics: List[str] = field( - default_factory=lambda: ["loss", "lr", "grad_norm"], - metadata={"help": "Metrics to record during training."}, - ) + model_fn: Callable[[], nn.Module] + strategy: str + dataset: Dataset + optimizer_fn: Callable[[nn.Module], Optimizer] + scheduler_fn: Callable[[Optimizer], LRScheduler] + n_epoch: int = 1 + batch_per_device: int = 4 + grad_accum_steps: int = 1 + max_grad_norm: Optional[float] = 1.0 + gradient_checkpointing_modules: List[type] = field(default_factory=list) + compile_mode: Optional[str] = None - # dataloader setting - random_seed: int = field(default=3407, metadata={"help": "Random seed."}) - num_workers: int = field( - default=0, metadata={"help": "Number of workers for dataloader."} - ) - prefetch_factor: Optional[int] = field( - default=None, metadata={"help": "Prefetch factor for dataloader."} - ) - pin_memory: bool = field( - default=False, metadata={"help": "Pin memory for dataloader."} - ) - collate_fn: Optional[Callable[[List[Any]], Any]] = field( - default=None, - metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."}, - ) + start_epoch: int = 0 + start_samples: int = 0 + ckpt_dir: str = "./checkpoint" + ckpt_interval: int = 5000 - # distributed training - nprocs: int = field( - default=1, metadata={"help": "Number of processes for distributed training."} - ) - backend: str = field( - default="nccl", metadata={"help": "Distributed training backend."} - ) - master_addr: str = field( - default="localhost", - metadata={"help": "Master address for distributed training."}, - ) - master_port: str = field( - default="29500", metadata={"help": "Master port for distributed training."} - ) - parallel_mode: str = field( - default="none", - metadata={"help": "Parallel strategy: none, ddp, fsdp."}, - ) - start_method: str = field( - default="spawn", - metadata={"help": "Multiprocessing start method (spawn/fork/forkserver)."}, - ) + lora: Optional[LoRAConfig] = None - # others - device_type: str = field( - default="cuda", metadata={"help": "Device type for distributed training."} - ) - val_dataset: Optional[Dataset] = field( - default=None, metadata={"help": "Dataset for validation."} - ) - val_split: Optional[float] = field( - default=None, - metadata={ - "help": "Ratio to split from training dataset for validation (e.g. 0.05). Ignored if val_dataset is set." - }, - ) - val_step: int = field( - default=1000, - metadata={"help": "Number of optimizer steps between validation runs."}, - ) - neftune_alpha: float = field( - default=0.0, - metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."}, - ) + metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"]) - # online rollout - rollout_interval: int = field( - default=512, - metadata={"help": "Number of optimizer steps between online rollouts."}, - ) - rollout_temperature: float = field( - default=0.7, metadata={"help": "Sampling temperature for online rollout."} - ) - rollout_top_k: int = field( - default=0, metadata={"help": "Top-k filtering for online rollout (0=disable)."} - ) - rollout_top_p: float = field( - default=0.9, - metadata={"help": "Top-p (nucleus) filtering for online rollout."}, - ) - rollout_max_tokens: int = field( - default=1024, - metadata={"help": "Maximum generated tokens per response in rollout."}, - ) - reward_model_fn: Optional[Callable] = field( - default=None, - metadata={ - "help": "Factory for reward model (required for online RL strategies)." - }, - ) + random_seed: int = 3407 + num_workers: int = 0 + prefetch_factor: Optional[int] = None + pin_memory: bool = False + collate_fn: Optional[Callable[[List[Any]], Any]] = None - executor_kwargs: Dict[str, Any] = field( - default_factory=dict, - metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."}, - ) - extra_kwargs: Dict[str, Any] = field( - default_factory=dict, metadata={"help": "Other arguments."} - ) + nprocs: int = 1 + backend: str = "nccl" + master_addr: str = "localhost" + master_port: str = "29500" + parallel_mode: str = "none" + start_method: str = "spawn" - def __post_init__(self): - self.validate() + device_type: str = "cuda" + val_dataset: Optional[Dataset] = None + val_split: Optional[float] = None + val_step: int = 1000 + neftune_alpha: float = 0.0 - def validate(self): - for fld in fields(self): - if fld.metadata.get("required") and getattr(self, fld.name) is None: - raise ValueError(f"TrainConfig.{fld.name} is required but got None.") + rollout_interval: int = 512 + rollout_temperature: float = 0.7 + rollout_top_k: int = 0 + rollout_top_p: float = 0.9 + rollout_max_tokens: int = 1024 + reward_model_fn: Optional[Callable] = None + + executor_kwargs: Dict[str, Any] = field(default_factory=dict) + extra_kwargs: Dict[str, Any] = field(default_factory=dict) diff --git a/astrai/inference/api/tool_parser.py b/astrai/inference/api/tool_parser.py index 3ed72b8..2e871a4 100644 --- a/astrai/inference/api/tool_parser.py +++ b/astrai/inference/api/tool_parser.py @@ -22,13 +22,10 @@ class BaseToolParser(ABC): Maintains streaming state internally so that each call to :meth:`feed` can diff against previously emitted content. - Parameters - ---------- - tools : list of dict, optional - Tool definitions from the request. - tool_choice : str - ``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice - dict. + Args: + tools (list of dict, optional): Tool definitions from the request. + tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named + tool choice dict. """ def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"): @@ -51,14 +48,12 @@ class BaseToolParser(ABC): Returns an empty list when nothing new should be emitted. - Parameters - ---------- - body : str - The complete accumulated generated text so far. - current_token_ids : list of int, optional - All token IDs decoded into *body* (cumulative). - delta_token_ids : list of int, optional - Only the token IDs for this chunk. + Args: + body (str): The complete accumulated generated text so far. + current_token_ids (list of int, optional): All token IDs decoded + into *body* (cumulative). + delta_token_ids (list of int, optional): Only the token IDs for + this chunk. """ @abstractmethod diff --git a/astrai/model/components/lora.py b/astrai/model/components/lora.py index 9c470a5..3b02cec 100644 --- a/astrai/model/components/lora.py +++ b/astrai/model/components/lora.py @@ -1,11 +1,12 @@ import logging -from dataclasses import asdict, dataclass +from dataclasses import asdict from pathlib import Path from typing import Optional, Set import torch import torch.nn as nn import torch.nn.functional as F +from pydantic.dataclasses import dataclass from astrai.model.components.linear import Linear from astrai.serialization import ( diff --git a/pyproject.toml b/pyproject.toml index 2da7f86..dfacc57 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,6 +15,7 @@ dependencies = [ "safetensors==0.5.3", "huggingface-hub==0.34.3", "jinja2>=3.0.0", + "pydantic>=2.0", "fastapi", "uvicorn[standard]", "click>=8.0",