TaoNet-mini-A2 / src /taoTrain /config.py
Lobakkang's picture
Upload folder using huggingface_hub
fd448dd verified
Raw
History Blame
33.8 kB
"""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="<image>",
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 <think>, <user>, <assistant>, or <image>. "
"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 {"<UNK>", "<BOS>", "<EOS>", "<PAD>"}]
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 = {"<UNK>", "<BOS>", "<EOS>", "<PAD>"}
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="<user>",
description="Special token representing user/instruction role in conversations"
)
assistant_token: str = Field(
default="<assistant>",
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="<user>",
description="Special token representing user/instruction role in multimodal conversations."
)
assistant_token: str = Field(
default="<assistant>",
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="<image>",
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}")