psikosen's picture
Update to v5: Prefix Sliding KV-cache, SMELT scaling, CMA checkpoint, sPTC tool caller
2976bd9 verified
Raw History Blame Contribute Delete
5.2 kB
"""
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)