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}")