Text Generation
Transformers
Safetensors
taonet
trust-remote-code
sentencepiece
custom-architecture
custom_code
Instructions to use TaoTern/TaoNet-mini-A2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use TaoTern/TaoNet-mini-A2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="TaoTern/TaoNet-mini-A2", trust_remote_code=True)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("TaoTern/TaoNet-mini-A2", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use TaoTern/TaoNet-mini-A2 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "TaoTern/TaoNet-mini-A2" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "TaoTern/TaoNet-mini-A2", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/TaoTern/TaoNet-mini-A2
- SGLang
How to use TaoTern/TaoNet-mini-A2 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "TaoTern/TaoNet-mini-A2" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "TaoTern/TaoNet-mini-A2", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "TaoTern/TaoNet-mini-A2" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "TaoTern/TaoNet-mini-A2", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use TaoTern/TaoNet-mini-A2 with Docker Model Runner:
docker model run hf.co/TaoTern/TaoNet-mini-A2
File size: 33,786 Bytes
fd448dd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 | """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}")
|