"""Pydantic configuration schemas for TaoTrain.""" from enum import Enum from typing import Optional, Literal from pathlib import Path import json from pydantic import BaseModel as PydanticBaseModel, Field, validator import yaml # ============================================================================ # Enums # ============================================================================ class DataTypeEnum(str, Enum): """Data types for training.""" FLOAT32 = "float32" FLOAT16 = "float16" BFLOAT16 = "bfloat16" class OptimizerEnum(str, Enum): """Supported optimizers.""" ADAM = "adam" ADAMW = "adamw" SGD = "sgd" HYBRID_MUON_ADAMW = "hybrid_muon_adamw" class ModelArchitectureEnum(str, Enum): """Built-in model architectures.""" TRANSFORMER = "transformer" TAONET = "taonet" GAMMA_NET = "gamma_net" MULTIMODAL_WRAPPER = "multimodal_wrapper" class SchedulerEnum(str, Enum): """Supported learning rate schedulers.""" LINEAR_WARMUP = "linearWarmup" COSINE_WARMUP = "cosineWarmup" CONSTANT = "constant" class RLMethodEnum(str, Enum): """Supported RL training methods.""" PPO = "ppo" DPO = "dpo" class TrainingModeEnum(str, Enum): """Training stages.""" PRETRAIN = "pretrain" SFT = "sft" RL = "rl" VLM = "vlm" VLM_SFT = "vlm_sft" # ============================================================================ # Base Configs # ============================================================================ class BaseConfig(PydanticBaseModel): """Base Pydantic model with utility methods.""" class Config: """Pydantic config.""" arbitrary_types_allowed = True def to_dict(self) -> dict: """Convert to dictionary.""" data = self.model_dump(mode='json') # Enums -> strings return data def to_json_str(self) -> str: """Convert to JSON string.""" return json.dumps(self.to_dict(), indent=2) def save_yaml(self, path: str | Path) -> None: """Save config to YAML file.""" path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) with open(path, 'w') as f: yaml.dump(self.to_dict(), f, default_flow_style=False, sort_keys=False) def save_json(self, path: str | Path) -> None: """Save config to JSON file.""" path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) with open(path, 'w') as f: f.write(self.to_json_str()) @classmethod def load_yaml(cls, path: str | Path) -> "BaseConfig": """Load config from YAML file.""" with open(path) as f: data = yaml.safe_load(f) return cls(**data) @classmethod def load_json(cls, path: str | Path) -> "BaseConfig": """Load config from JSON file.""" with open(path) as f: data = json.load(f) return cls(**data) # ============================================================================ # Model Config # ============================================================================ class ModelConfig(BaseConfig): """Configuration for model architecture.""" architecture_type: ModelArchitectureEnum = Field( default=ModelArchitectureEnum.TRANSFORMER, description="Type of model architecture" ) llm_architecture_type: Optional[ModelArchitectureEnum] = Field( default=None, description="Base language-model architecture when using a multimodal wrapper." ) # Transformer-specific vocab_size: int = Field(default=50257, description="Vocabulary size") hidden_dim: int = Field(default=768, description="Hidden dimension") num_layers: int = Field(default=12, description="Number of transformer blocks") num_heads: int = Field(default=12, description="Number of attention heads") head_dim: Optional[int] = Field( default=None, description="Head dimension (defaults to hidden_dim // num_heads)" ) intermediate_dim: Optional[int] = Field( default=None, description="FFN intermediate dimension (defaults to 4 * hidden_dim)" ) dropout: float = Field(default=0.1, description="Dropout rate") max_seq_length: int = Field(default=2048, description="Maximum sequence length") # TaoNet (DeepSeek MLA) specific d_latent_kv: Optional[int] = Field( default=None, description="KV compression dimension for MLA (defaults to 3/4 * hidden_dim). Only used for taonet architecture." ) d_rope: Optional[int] = Field( default=None, description="RoPE dimension per head (defaults to hidden_dim // num_heads). Only used for taonet architecture." ) gqa_groups: int = Field( default=1, description="Grouped Query Attention groups (1 = standard MLA, >1 = GQA). Only used for taonet architecture." ) hidden_dim_ff: Optional[int] = Field( default=None, description="Feed-forward intermediate dimension (defaults to 4 * hidden_dim)." ) use_factorized_embedding: bool = Field( default=False, description="Use low-rank factorized embedding instead of standard embedding (reduces params). Only for taonet." ) d_embed_rank: int = Field( default=96, description="Rank dimension for factorized embedding. Only used if use_factorized_embedding=True." ) vision_encoder_type: str = Field( default="cnn", description="Vision encoder type for multimodal models." ) image_size: int = Field( default=224, description="Square image size for multimodal image preprocessing." ) vision_output_dim: int = Field( default=256, description="Output feature dimension from the vision encoder before projection." ) vision_prefix_tokens: int = Field( default=10, description="Number of visual prefix tokens projected into the LLM sequence." ) image_token: str = Field( default="", description="Tokenizer special token reserved for image placeholders." ) cnn_channels: list[int] = Field( default_factory=lambda: [32, 64, 128], description="Per-stage channel sizes for the in-repo CNN encoder." ) cnn_kernel_size: int = Field( default=3, description="Kernel size for CNN encoder convolutions." ) # GammaSpaceModel-specific gamma_hidden_dim: int = Field( default=1536, description="GammaSpaceBlock hidden_dim (state-space hidden size). Only used for gamma_net." ) gamma_dt_min: float = Field( default=1e-3, description="GammaSpaceBlock dt_min. Only used for gamma_net." ) gamma_dt_max: float = Field( default=1e-1, description="GammaSpaceBlock dt_max. Only used for gamma_net." ) gamma_dt_init: float = Field( default=1e-2, description="GammaSpaceBlock dt_init. Only used for gamma_net." ) gamma_discretization: str = Field( default="bilinear", description="GammaSpaceBlock discretization mode: bilinear, zoh, or euler. Only used for gamma_net." ) gamma_prenorm: bool = Field( default=True, description="GammaSpaceBlock prenorm flag. Only used for gamma_net." ) gamma_residual_scale: float = Field( default=1.0, description="GammaSpaceBlock residual_scale. Only used for gamma_net." ) gamma_activation: str = Field( default="gelu", description="GammaSpaceBlock activation. Only used for gamma_net." ) gamma_gate: bool = Field( default=True, description="GammaSpaceBlock gate flag. Only used for gamma_net." ) gamma_use_D: bool = Field( default=True, description="GammaSpaceBlock use_D flag. Only used for gamma_net." ) gamma_kernel_mode: str = Field( default="auto", description="GammaSpaceBlock kernel_mode: auto, recurrent, or conv. Only used for gamma_net." ) gamma_kernel_threshold: int = Field( default=64, description="GammaSpaceBlock kernel_threshold. Only used for gamma_net." ) gamma_use_output_linear: bool = Field( default=True, description="GammaSpaceBlock use_output_linear flag. Only used for gamma_net." ) gamma_gate_bias: float = Field( default=2.0, description="GammaSpaceBlock gate_bias. Only used for gamma_net." ) gamma_input_gate: bool = Field( default=True, description="GammaSpaceBlock input_gate flag. Only used for gamma_net." ) gamma_input_gate_bias: float = Field( default=2.0, description="GammaSpaceBlock input_gate_bias. Only used for gamma_net." ) gamma_layer_scale_init: float = Field( default=0.1, description="GammaSpaceBlock layer_scale_init. Only used for gamma_net." ) # YaRN (Yet another RoPE eXtension) for context length extension rope_scale: float = Field( default=40.0, description="Base RoPE scale factor (default: 40.0). Controls position frequency base." ) yarn_enabled: bool = Field( default=False, description="Enable YaRN (Yet another RoPE eXtension) for context length interpolation." ) yarn_original_max_seq_length: Optional[int] = Field( default=None, description="Original trained context length that YaRN extends from. Defaults to max_seq_length when unset." ) yarn_alpha: float = Field( default=1.0, description="YaRN interpolation smoothness (1.0=smooth, <1.0=aggressive, >1.0=conservative). Only used if yarn_enabled=True." ) # Initializations init_std: float = Field(default=0.02, description="Weight initialization standard deviation") @validator("head_dim", always=True) def validate_head_dim(cls, v, values): """Validate head dimension.""" if v is None and 'hidden_dim' in values: return values['hidden_dim'] // values.get('num_heads', 12) return v @validator("intermediate_dim", always=True) def validate_intermediate_dim(cls, v, values): """Validate intermediate dimension.""" if v is None and 'hidden_dim' in values: return 4 * values['hidden_dim'] return v # ============================================================================ # Dataset Config # ============================================================================ class DatasetConfig(BaseConfig): """Configuration for dataset loading.""" # Local vs HuggingFace dataset selection local: bool = Field(default=False, description="Use local JSONL dataset instead of HuggingFace") # HuggingFace dataset fields dataset_name: Optional[str] = Field(default=None, description="HuggingFace dataset name (e.g., 'wikitext', 'openwebtext')") split: str = Field(default="train", description="Dataset split to use") config: Optional[str] = Field(default=None, description="Dataset config if multi-config (e.g., 'wikitext-103')") # Local JSONL dataset fields jsonl_path: Optional[str] = Field(default=None, description="Path to local JSONL dataset file") text_field: str = Field(default="text", description="Name of text field in JSONL") image_path_column: str = Field( default="image_path", description="Column containing a local image path for multimodal JSONL datasets." ) image_path_aliases: list[str] = Field( default_factory=lambda: ["image", "image_path", "image_file", "file_name"], description="Fallback column names to try when a multimodal JSONL record does not contain image_path_column." ) caption_prompt: str = Field( default="Describe the image.", description="Prompt paired with caption-style multimodal records that only provide image + text." ) # Text column name varies by dataset text_column: str = Field(default="text", description="Name of text column in dataset") # Preprocessing max_samples: Optional[int] = Field( default=None, description="Limit dataset to N samples (useful for debugging)" ) cache_dir: str = Field(default=".cache/datasets", description="HuggingFace cache directory") # For SFT/RL datasets with instruction-response format instruction_column: Optional[str] = Field(default=None, description="Instruction column for SFT") response_column: Optional[str] = Field(default=None, description="Response column for SFT") prompt_column: Optional[str] = Field(default=None, description="Prompt column for RL") # Instruction template instruction_template: Optional[str] = Field( default=None, description="Template for combining instruction and response. E.g., '{instruction}\\n{response}'" ) # Tokenizer configuration tokenizer_type: Optional[str] = Field( default=None, description="Tokenizer type: 'huggingface' or 'sentencepiece'. If None, defaults based on tokenizer_path." ) tokenizer_path: Optional[str] = Field( default=None, description="Path to saved tokenizer (for SentencePiece: .model file, for HuggingFace: model name or local path)" ) # Chunked loading for large JSONL files enable_streaming: bool = Field( default=True, description="Enable streaming/chunked loading for large JSONL files to reduce memory usage" ) chunk_size_gb: float = Field( default=5.0, description="Approximate chunk size in GB (ignored if samples_per_chunk is set)" ) samples_per_chunk: Optional[int] = Field( default=1000, description="Number of samples per chunk (takes precedence over chunk_size_gb). Default: 1000 samples" ) # Chunk caching enable_chunk_metadata_cache: bool = Field( default=True, description="Enable caching of chunk metadata (file scan results) to avoid re-scanning large JSONL files" ) enable_chunk_data_cache: bool = Field( default=False, description="Enable caching of actual chunk data as separate files for faster loading (uses more disk space)" ) chunk_cache_dir: str = Field( default=".cache/chunks", description="Directory to store chunk metadata and data cache files" ) # Tokenization parallelization tokenizer_threads: int = Field( default=1, description="Number of background threads for tokenization (1-32 recommended). Higher values speed up tokenization but increase memory usage." ) @validator('jsonl_path', always=True) def validate_dataset_source(cls, v, values): """Validate that either local JSONL or HuggingFace dataset is specified.""" local = values.get('local', False) dataset_name = values.get('dataset_name') if local and not v: raise ValueError("jsonl_path must be provided when local=True") if not local and not dataset_name: raise ValueError("dataset_name must be provided when local=False (HuggingFace dataset)") return v @validator('tokenizer_threads') def validate_tokenizer_threads(cls, v): """Validate tokenizer_threads is a positive integer.""" if v < 1: raise ValueError("tokenizer_threads must be at least 1") if v > 128: raise ValueError("tokenizer_threads should not exceed 128 (recommended: 1-32)") return v # ============================================================================ # Tokenizer Config # ============================================================================ class TokenizerConfig(BaseConfig): """Configuration for tokenizer training.""" # Dataset source jsonl_path: str = Field(description="Path to JSONL file containing training data") text_field: str = Field(default="text", description="Field name in JSONL for text data") # Training configuration vocab_size: int = Field(default=50000, description="Vocabulary size") model_type: str = Field(default="unigram", description="SentencePiece model type (unigram, bpe, char, word)") character_coverage: float = Field( default=0.9995, description="Character coverage for SentencePiece training" ) output_dir: str = Field(default="tokenizers", description="Directory to save trained tokenizer") tokenizer_prefix: Optional[str] = Field( default=None, description="Prefix for tokenizer output files (default: model_type)" ) # Custom special tokens registered as SentencePiece user-defined symbols. special_tokens: Optional[list[str]] = Field( default=None, description=( "Custom special tokens such as , , , or . " "Built-in SentencePiece tokens are managed by the tokenizer itself and should not be listed here." ) ) # Data sampling max_samples: Optional[int] = Field( default=None, description="Limit training to first N samples from JSONL (useful for quick testing)" ) # Tokenizer metadata tokenizer_name: Optional[str] = Field( default=None, description="Optional name for the tokenizer" ) @validator("special_tokens", pre=True) def normalize_special_tokens(cls, value): """Normalize special token config to a list of custom token strings.""" if value is None: return None if isinstance(value, dict): return [str(token) for token in value.keys() if str(token) not in {"", "", "", ""}] if isinstance(value, (list, tuple)): return [str(token) for token in value] raise ValueError("special_tokens must be a list of token strings") @validator("special_tokens") def validate_special_tokens(cls, value): """Validate custom special token registry.""" if value is None: return value builtin_tokens = {"", "", "", ""} seen: set[str] = set() normalized: list[str] = [] for token in value: if token in builtin_tokens: continue if token in seen: continue seen.add(token) normalized.append(token) return normalized # ============================================================================ # Training Config # ============================================================================ class OptimizerConfig(BaseConfig): """Optimizer configuration.""" optimizer_type: OptimizerEnum = Field(default=OptimizerEnum.ADAMW, description="Optimizer type") learning_rate: float = Field(default=1e-4, description="Peak learning rate (for Muon 2D weights)") adamw_lr: Optional[float] = Field( default=None, description="Learning rate for AdamW (1D parameters). If None, defaults to learning_rate / 10. Used in hybrid_muon_adamw optimizer." ) weight_decay: float = Field(default=1e-2, description="Weight decay (L2 regularization)") betas: tuple[float, float] = Field(default=(0.9, 0.999), description="Adam betas") eps: float = Field(default=1e-8, description="Optimizer epsilon") @validator('adamw_lr', always=True) def set_default_adamw_lr(cls, v, values): """Set default adamw_lr as 1/10 of learning_rate if not specified.""" if v is None and 'learning_rate' in values: return values['learning_rate'] / 10 return v class SchedulerConfig(BaseConfig): """Learning rate scheduler configuration.""" scheduler_type: SchedulerEnum = Field(default=SchedulerEnum.LINEAR_WARMUP, description="Scheduler type") warmup_steps: int = Field(default=0, description="Number of warmup steps (takes precedence over warmup_ratio)") warmup_ratio: float = Field(default=0.1, description="Warmup as fraction of total steps (used if warmup_steps=0)") # Cosine scheduler specific num_cycles: float = Field(default=0.5, description="Number of cycles for cosine schedule") last_epoch: int = Field(default=-1, description="Last epoch for scheduler") # TaoNet 3-phase scheduler (warmup -> steady -> cosine decay) steady_ratio: float = Field( default=0.0, description="Fraction of training steps at peak LR before cosine decay (0.0 = no steady phase). Only for cosineWarmup." ) min_lr_ratio: float = Field( default=0.0, description="Minimum LR as fraction of peak LR at end of training (0.0 = decay to 0). Only for cosineWarmup." ) @validator('warmup_ratio') def validate_warmup_ratio(cls, v): """Validate warmup ratio is between 0 and 1.""" if not 0 <= v <= 1: raise ValueError("warmup_ratio must be between 0 and 1") return v @validator('steady_ratio') def validate_steady_ratio(cls, v): """Validate steady ratio is between 0 and 1.""" if not 0 <= v <= 1: raise ValueError("steady_ratio must be between 0 and 1") return v @validator('min_lr_ratio') def validate_min_lr_ratio(cls, v): """Validate min_lr_ratio is between 0 and 1.""" if not 0 <= v <= 1: raise ValueError("min_lr_ratio must be between 0 and 1") return v @validator('warmup_steps') def validate_warmup_steps(cls, v): """Validate warmup steps is non-negative.""" if v < 0: raise ValueError("warmup_steps must be non-negative") return v class TrainingConfig(BaseConfig): """Base training configuration shared across all modes.""" # Data and model model: ModelConfig = Field(default_factory=ModelConfig, description="Model configuration") dataset: DatasetConfig = Field(description="Dataset configuration") # Training hyperparameters batch_size: int = Field(default=32, description="Batch size per device") num_epochs: int = Field(default=3, description="Number of training epochs") max_steps: Optional[int] = Field( default=None, description="Maximum steps (overrides num_epochs if set)" ) gradient_accumulation_steps: int = Field( default=1, description="Gradient accumulation steps" ) max_grad_norm: float = Field(default=1.0, description="Gradient clipping max norm") # Optimizer optimizer: OptimizerConfig = Field( default_factory=OptimizerConfig, description="Optimizer configuration" ) # Scheduler scheduler: SchedulerConfig = Field( default_factory=SchedulerConfig, description="Learning rate scheduler configuration" ) # Data type and device dtype: DataTypeEnum = Field( default=DataTypeEnum.BFLOAT16, description="Training data type" ) device: str = Field(default="cuda", description="Device to train on (cuda, cpu)") seed: int = Field(default=42, description="Random seed") # Checkpointing checkpoint_dir: str = Field(default="checkpoints", description="Directory to save checkpoints") checkpoint_path: Optional[str] = Field( default=None, description="Path to load pretrained checkpoint (for SFT/RL). If provided, loads weights before training starts." ) save_every_steps: int = Field(default=500, description="Save checkpoint every N steps") keep_last_n_checkpoints: int = Field(default=3, description="Keep only last N checkpoints") save_best_model: bool = Field(default=True, description="Save best model based on validation loss") # Validation eval_every_steps: int = Field(default=500, description="Evaluate every N steps") eval_samples: int = Field(default=1000, description="Number of validation samples") # Logging log_every_steps: int = Field(default=10, description="Log metrics every N steps") aim_repo: str = Field(default=".aim", description="AimStack repository path") # Misc num_workers: int = Field(default=0, description="Number of DataLoader workers") pin_memory: bool = Field(default=True, description="Pin memory for DataLoader") use_compile: bool = Field(default=False, description="Use torch.compile (experimental)") # Mode mode: TrainingModeEnum = Field(default=TrainingModeEnum.PRETRAIN, description="Training mode") # ============================================================================ # Stage-Specific Configs # ============================================================================ class PretrainConfig(TrainingConfig): """Configuration for pretraining.""" mode: Literal[TrainingModeEnum.PRETRAIN] = TrainingModeEnum.PRETRAIN # Pretraining-specific sequence_length: int = Field(default=1024, description="Sequence length for pretraining") class SFTConfig(TrainingConfig): """Configuration for supervised fine-tuning.""" mode: Literal[TrainingModeEnum.SFT] = TrainingModeEnum.SFT # SFT-specific response_loss_only: bool = Field( default=True, description="Only compute loss on response/assistant tokens (not instruction/user tokens). Uses -100 label masking." ) # Multi-turn conversation role tokens user_token: str = Field( default="", description="Special token representing user/instruction role in conversations" ) assistant_token: str = Field( default="", description="Special token representing assistant/response role in conversations" ) class RLConfig(TrainingConfig): """Configuration for reinforcement learning training.""" mode: Literal[TrainingModeEnum.RL] = TrainingModeEnum.RL # RL-specific rl_method: RLMethodEnum = Field( default=RLMethodEnum.PPO, description="RL training method (PPO or DPO)" ) # Reward model reward_model_path: str = Field(description="Path to trained reward model checkpoint") # PPO-specific ppo_epochs: int = Field(default=4, description="PPO inner epochs") ppo_clip_ratio: float = Field(default=0.2, description="PPO clipping ratio") entropy_coeff: float = Field(default=0.01, description="Entropy bonus coefficient") value_loss_coeff: float = Field(default=1.0, description="Value function loss coefficient") # DPO-specific (Direct Preference Optimization) dpo_beta: float = Field(default=0.1, description="DPO inverse temperature (beta)") # Prompt distribution prompt_dataset: Optional[DatasetConfig] = Field( default=None, description="Separate dataset for prompts (if different from main dataset)" ) generation_max_length: int = Field( default=256, description="Maximum length for generated responses during RL" ) class VLMConfig(TrainingConfig): """Configuration for multimodal vision-language connector training.""" mode: Literal[TrainingModeEnum.VLM] = TrainingModeEnum.VLM response_loss_only: bool = Field( default=True, description="Only compute loss on assistant/response tokens for multimodal samples." ) user_token: str = Field( default="", description="Special token representing user/instruction role in multimodal conversations." ) assistant_token: str = Field( default="", description="Special token representing assistant/response role in multimodal conversations." ) freeze_llm: bool = Field( default=True, description="Freeze the LLM except for any explicitly unfrozen trailing layers." ) unfreeze_last_n_layers: int = Field( default=0, description="Number of final LLM blocks to unfreeze during VLM connector training." ) vision_learning_rate: float = Field( default=1e-4, description="Learning rate for the vision encoder and multimodal projector." ) llm_learning_rate: float = Field( default=5e-5, description="Learning rate for the unfrozen LLM parameter subset." ) vision_prefix_tokens: int = Field( default=10, description="Number of visual prefix tokens to inject in place of a single image placeholder token." ) image_token: str = Field( default="", description="Tokenizer token used as the multimodal image placeholder." ) image_size: int = Field( default=224, description="Square image size used for multimodal preprocessing." ) @validator("image_token") def validate_image_token_matches_model_config(cls, v, values): """Validate duplicated image token settings stay aligned.""" model = values.get("model") if model is not None and getattr(model, "image_token", None) != v: raise ValueError( f"image_token must match model.image_token for VLM configs. " f"Got top-level `{v}` and model `{model.image_token}`." ) return v @validator("vision_prefix_tokens") def validate_vision_prefix_tokens_match_model_config(cls, v, values): """Validate duplicated visual prefix settings stay aligned.""" model = values.get("model") if model is not None and getattr(model, "vision_prefix_tokens", None) != v: raise ValueError( "vision_prefix_tokens must match model.vision_prefix_tokens for VLM configs. " f"Got top-level `{v}` and model `{model.vision_prefix_tokens}`." ) return v @validator("image_size") def validate_image_size_matches_model_config(cls, v, values): """Validate duplicated image size settings stay aligned.""" model = values.get("model") if model is not None and getattr(model, "image_size", None) != v: raise ValueError( f"image_size must match model.image_size for VLM configs. " f"Got top-level `{v}` and model `{model.image_size}`." ) return v @validator("unfreeze_last_n_layers") def validate_unfreeze_last_n_layers(cls, v): """Validate unfrozen layer count.""" if v < 0: raise ValueError("unfreeze_last_n_layers must be non-negative") return v @validator("vision_learning_rate", "llm_learning_rate") def validate_multimodal_learning_rates(cls, v): """Validate multimodal learning rates.""" if v <= 0: raise ValueError("multimodal learning rates must be greater than 0") return v @validator("vision_prefix_tokens") def validate_vision_prefix_tokens(cls, v): """Validate visual prefix length.""" if v < 1: raise ValueError("vision_prefix_tokens must be at least 1") return v @validator("image_size") def validate_image_size(cls, v): """Validate image size.""" if v < 8: raise ValueError("image_size must be at least 8") return v class VLMSFTConfig(VLMConfig): """Configuration for end-to-end multimodal supervised fine-tuning.""" mode: Literal[TrainingModeEnum.VLM_SFT] = TrainingModeEnum.VLM_SFT freeze_llm: bool = Field( default=False, description="Whether to freeze the LLM during end-to-end multimodal SFT." ) # ============================================================================ # Factory function # ============================================================================ def load_config(path: str | Path, mode: TrainingModeEnum | str) -> TrainingConfig: """Load config file and return appropriate config class.""" if isinstance(mode, str): mode = TrainingModeEnum(mode) config_map = { TrainingModeEnum.PRETRAIN: PretrainConfig, TrainingModeEnum.SFT: SFTConfig, TrainingModeEnum.RL: RLConfig, TrainingModeEnum.VLM: VLMConfig, TrainingModeEnum.VLM_SFT: VLMSFTConfig, } config_class = config_map[mode] path = Path(path) if path.suffix == '.yaml' or path.suffix == '.yml': return config_class.load_yaml(path) elif path.suffix == '.json': return config_class.load_json(path) else: raise ValueError(f"Unsupported config file format: {path.suffix}") def load_tokenizer_config(path: str | Path) -> TokenizerConfig: """Load tokenizer config from YAML or JSON file.""" path = Path(path) if path.suffix == '.yaml' or path.suffix == '.yml': return TokenizerConfig.load_yaml(path) elif path.suffix == '.json': return TokenizerConfig.load_json(path) else: raise ValueError(f"Unsupported config file format: {path.suffix}")