| """Central config. Scales from K=10k (pilot) to K=1,000,000 clusters by changing CLUSTERS.""" |
| from dataclasses import dataclass, field |
|
|
| MODEL = "Qwen/Qwen3-8B" |
| D_MODEL = 4096 |
| READ_LAYER = 27 |
| INJECT_LAYER = 1 |
| STEER_COEFF = 1.0 |
| EMBED_MODEL = "BAAI/bge-large-en-v1.5" |
| CORPUS = "openbmb/Ultra-FineWeb" |
|
|
|
|
| @dataclass |
| class ClusterConfig: |
| corpus: str = CORPUS |
| embed_model: str = EMBED_MODEL |
| n_docs: int = 4_000_000 |
| clusters: int = 10_000 |
| min_chars: int = 200 |
| max_chars: int = 2000 |
| embed_batch: int = 1024 |
| shard_size: int = 500_000 |
| out_dir: str = "data/clusters" |
| seed: int = 0 |
|
|
|
|
| @dataclass |
| class ProbeCacheConfig: |
| members_per_cluster: int = 64 |
| pool: str = "mean" |
| batch: int = 128 |
| out_dir: str = "data/resid_cache" |
|
|
|
|
| @dataclass |
| class BuildDataConfig: |
| n_examples: int = 8_000_000 |
| targets_per_example: int = 5 |
| probe_c: float = 1.0 |
| negatives_per_probe: int = 512 |
| pair_mode: str = "AvsB" |
| shard_examples: int = 250_000 |
| out_dir: str = "data/pretrain" |
| seed: int = 0 |
|
|
|
|
| @dataclass |
| class TrainConfig: |
| init_adapter: str | None = None |
| lora_r: int = 64 |
| lora_alpha: int = 16 |
| lora_dropout: float = 0.0 |
| lr: float = 3e-5 |
| batch_size: int = 64 |
| epochs: int = 1 |
| max_seq: int = 192 |
| warmup_frac: float = 0.02 |
| save_dir: str = "checkpoints/pretrain" |
| run_name: str = "mxf-pretrain" |
|
|
|
|
| @dataclass |
| class RLConfig: |
| """Dr. GRPO — no /std, no KL, global-token normalizer.""" |
| init_adapter: str = "checkpoints/pretrain/final" |
| groups_per_step: int = 256 |
| group_size: int = 8 |
| lr: float = 1e-6 |
| clip_eps: float = 0.2 |
| tis_cap: float = 2.0 |
| entropy_coef: float = 0.0 |
| max_new_tokens: int = 96 |
| min_new_tokens: int = 16 |
| temperature: float = 1.0 |
| total_steps: int = 30_000 |
| sync_every: int = 10 |
| fluency_floor: float | None = -4.5 |
| distinct_floor: float | None = 0.5 |
| gate_penalty: float = 25.0 |
| len_penalty_start: int | None = 64 |
| len_penalty_per_tok: float = 0.5 |
| direction_source: str = "cluster" |
| save_dir: str = "checkpoints/rl" |
| run_name: str = "mxf-rl-drgrpo" |
|
|
|
|
| @dataclass |
| class Config: |
| cluster: ClusterConfig = field(default_factory=ClusterConfig) |
| cache: ProbeCacheConfig = field(default_factory=ProbeCacheConfig) |
| build: BuildDataConfig = field(default_factory=BuildDataConfig) |
| train: TrainConfig = field(default_factory=TrainConfig) |
| rl: RLConfig = field(default_factory=RLConfig) |
|
|