""" Canopy-R3 Architecture and Training Configuration. """ from dataclasses import dataclass, field, asdict from typing import Optional, Dict, Any import yaml import json from pathlib import Path @dataclass class CanopyConfig: """Configuration for Canopy-R3 and matched baseline models.""" # Model dimensions vocab_size: int = 49152 max_seq_len: int = 2048 d_model: int = 768 num_heads: int = 12 num_kv_heads: int = 4 head_dim: int = 64 # Layer layout (12 physical blocks: 3 prelude + 6 recurrent MoE + 3 coda) num_layers: int = 12 prelude_layers: int = 3 recurrent_layers: int = 6 coda_layers: int = 3 # Recurrence / looping recurrent_visits: int = 2 loop_residual_scale: float = 0.5 # FFN parameters dense_intermediate_size: int = 2048 moe_num_experts: int = 8 moe_top_k: int = 2 moe_intermediate_size: int = 1536 # Thought Bus use_thought_bus: bool = True thought_bus_width: int = 192 # Visit Adapter use_visit_adapter: bool = True visit_adapter_rank: int = 8 # Routing & Stability router_aux_loss_coef: float = 0.01 router_z_loss_coef: float = 0.001 # Precision & Normalization norm_eps: float = 1e-6 rope_theta: float = 10000.0 tie_word_embeddings: bool = True dropout: float = 0.0 initializer_range: float = 0.02 # Quantization mode: "none", "ternary" (W1.58A8), "binary" (W1A8) quant_mode: str = "none" # Adaptive Recurrent Inference (Dynamic Early-Exit) adaptive_early_exit: bool = False early_exit_entropy_threshold: float = 2.5 # v4 Innovations: Prefix Sliding & SMELT enable_prefix_sliding: bool = False prefix_tokens_len: int = 128 sliding_window_len: int = 1024 use_smelt_residual: bool = True smelt_residual_scale: float = 0.7071 # 1 / sqrt(2) per SMELT (ByteDance / Tsinghua) suppress_attention_sink: bool = True @property def effective_layers(self) -> int: """Executed depth in blocks.""" return self.prelude_layers + (self.recurrent_layers * self.recurrent_visits) + self.coda_layers @property def kv_cache_elements_per_token(self) -> int: """KV cache elements stored per token across all executed layers.""" return 2 * self.num_kv_heads * self.head_dim * self.effective_layers def to_dict(self) -> Dict[str, Any]: return asdict(self) @classmethod def from_dict(cls, data: Dict[str, Any]) -> "CanopyConfig": valid_keys = {f.name for f in cls.__dataclass_fields__.values()} filtered = {k: v for k, v in data.items() if k in valid_keys} return cls(**filtered) def to_yaml(self, path: str | Path) -> None: with open(path, "w", encoding="utf-8") as f: yaml.dump(self.to_dict(), f, sort_keys=False) @classmethod def from_yaml(cls, path: str | Path) -> "CanopyConfig": with open(path, "r", encoding="utf-8") as f: data = yaml.safe_load(f) return cls.from_dict(data) def to_json(self, path: str | Path) -> None: with open(path, "w", encoding="utf-8") as f: json.dump(self.to_dict(), f, indent=2) @classmethod def from_json(cls, path: str | Path) -> "CanopyConfig": with open(path, "r", encoding="utf-8") as f: data = json.load(f) return cls.from_dict(data) @dataclass class TrainingConfig: """Training hyperparameters and run settings.""" # Batch sizing global_batch_tokens: int = 65536 micro_batch_size: int = 2 max_seq_len: int = 2048 gradient_accumulation_steps: int = 16 # auto-computed from global_batch_tokens # Optimization learning_rate: float = 3e-3 min_learning_rate: float = 3e-4 warmup_steps: int = 500 total_steps: int = 4578 # 300M tokens @ 65536 tokens/step weight_decay: float = 0.1 beta1: float = 0.9 beta2: float = 0.95 max_grad_norm: float = 1.0 adam_eps: float = 1e-8 # Features activation_checkpointing: bool = True mixed_precision: str = "bf16" torch_compile: bool = False # Checkpointing & logging log_interval: int = 10 eval_interval: int = 250 save_interval: int = 1000 keep_last_k_checkpoints: int = 3 output_dir: str = "checkpoints/canopy_258m" seed: int = 42 def compute_accumulation_steps(self) -> None: tokens_per_micro = self.micro_batch_size * self.max_seq_len self.gradient_accumulation_steps = max(1, self.global_batch_tokens // tokens_per_micro) def to_dict(self) -> Dict[str, Any]: return asdict(self) @classmethod def from_dict(cls, data: Dict[str, Any]) -> "TrainingConfig": valid_keys = {f.name for f in cls.__dataclass_fields__.values()} filtered = {k: v for k, v in data.items() if k in valid_keys} cfg = cls(**filtered) cfg.compute_accumulation_steps() return cfg @classmethod def from_yaml(cls, path: str | Path) -> "TrainingConfig": with open(path, "r", encoding="utf-8") as f: data = yaml.safe_load(f) return cls.from_dict(data)