""" Canopy-R3 Model Architecture: Causal GQA Transformer, Looped MoE, Tokenwise Thought Bus, and Visit-Conditioned Low-Rank Adapters. """ import math from typing import Optional, Tuple, List, Dict, Any import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint try: from .configuration_canopy import CanopyConfig except ImportError: from configuration_canopy import CanopyConfig from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutputWithPast class CanopyPreTrainedModel(PreTrainedModel): config_class = CanopyConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["CanopyBlock"] def _init_weights(self, module): std = getattr(self.config, "initializer_range", 0.02) if isinstance(module, (nn.Linear, nn.Embedding)): module.weight.data.normal_(mean=0.0, std=std) if hasattr(module, "bias") and module.bias is not None: module.bias.data.zero_() # Optional quantization try: from .quant import convert_to_bitlinear, BitLinear except Exception: pass class RMSNorm(nn.Module): """Root Mean Square Layer Normalization.""" def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: variance = x.pow(2).mean(dim=-1, keepdim=True) return x * torch.rsqrt(variance + self.eps) * self.weight class RotaryEmbedding(nn.Module): """Rotary Position Embedding (RoPE) cache.""" def __init__(self, dim: int, max_seq_len: int = 2048, theta: float = 10000.0): super().__init__() self.dim = dim self.max_seq_len = max_seq_len self.theta = theta inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) self._build_cache(max_seq_len) def _build_cache(self, seq_len: int): t = torch.arange(seq_len, dtype=torch.float32, device=self.inv_freq.device) freqs = torch.outer(t, self.inv_freq) emb = torch.cat((freqs, freqs), dim=-1) self.register_buffer("cos_cached", emb.cos(), persistent=False) self.register_buffer("sin_cached", emb.sin(), persistent=False) def forward(self, x: torch.Tensor, seq_len: int, offset: int = 0) -> Tuple[torch.Tensor, torch.Tensor]: total_len = offset + seq_len if total_len > self.cos_cached.shape[0]: self._build_cache(total_len) return self.cos_cached[offset:total_len], self.sin_cached[offset:total_len] def rotate_half(x: torch.Tensor) -> torch.Tensor: x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) def apply_rotary_pos_emb(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: # q: (batch, heads, seq, dim) # cos, sin: (seq, dim) -> (1, 1, seq, dim) cos = cos.unsqueeze(0).unsqueeze(0).to(dtype=q.dtype) sin = sin.unsqueeze(0).unsqueeze(0).to(dtype=q.dtype) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed.to(dtype=q.dtype), k_embed.to(dtype=k.dtype) class CausalSelfAttention(nn.Module): """Grouped Query Attention (GQA) with SDPA and KV Caching.""" def __init__(self, config: CanopyConfig): super().__init__() self.d_model = config.d_model self.num_heads = config.num_heads self.num_kv_heads = config.num_kv_heads self.head_dim = config.head_dim self.num_kv_groups = self.num_heads // self.num_kv_heads self.q_proj = nn.Linear(self.d_model, self.num_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(self.d_model, self.num_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(self.d_model, self.num_kv_heads * self.head_dim, bias=False) self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.d_model, bias=False) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, kv_cache: Optional[Dict[str, torch.Tensor]] = None, layer_idx: int = 0, ) -> torch.Tensor: batch_size, seq_len, _ = x.shape q = self.q_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(x).view(batch_size, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(x).view(batch_size, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) q, k = apply_rotary_pos_emb(q, k, cos, sin) if kv_cache is not None: cache_k_key = f"k_{layer_idx}" cache_v_key = f"v_{layer_idx}" if cache_k_key in kv_cache: k = torch.cat([kv_cache[cache_k_key], k], dim=2) v = torch.cat([kv_cache[cache_v_key], v], dim=2) kv_cache[cache_k_key] = k kv_cache[cache_v_key] = v # Repeat KV heads for GQA if self.num_kv_groups > 1: k_expanded = k.repeat_interleave(self.num_kv_groups, dim=1) v_expanded = v.repeat_interleave(self.num_kv_groups, dim=1) else: k_expanded = k v_expanded = v # Scaled dot-product attention is_causal = (kv_cache is None or kv_cache.get("is_prefill", False)) and seq_len > 1 attn_output = F.scaled_dot_product_attention( q, k_expanded.to(dtype=q.dtype), v_expanded.to(dtype=q.dtype), is_causal=is_causal ) attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) return self.o_proj(attn_output) class SwiGLU(nn.Module): """Dense SwiGLU Feed-Forward Network.""" def __init__(self, d_model: int, intermediate_size: int): super().__init__() self.gate_proj = nn.Linear(d_model, intermediate_size, bias=False) self.up_proj = nn.Linear(d_model, intermediate_size, bias=False) self.down_proj = nn.Linear(intermediate_size, d_model, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class ThoughtBus(nn.Module): """ Tokenwise Thought Bus: fuses selected expert outputs at token t via cross-attention. Query: token hidden state. Keys/Values: top-K selected expert outputs at token t. """ def __init__(self, d_model: int, bus_width: int): super().__init__() self.d_model = d_model self.bus_width = bus_width self.q_proj = nn.Linear(d_model, bus_width, bias=False) self.k_proj = nn.Linear(d_model, bus_width, bias=False) self.v_proj = nn.Linear(d_model, bus_width, bias=False) self.out_proj = nn.Linear(bus_width, d_model, bias=False) self.bus_norm = RMSNorm(d_model) def forward(self, hidden_state: torch.Tensor, expert_outputs: torch.Tensor) -> torch.Tensor: # hidden_state: (B, S, D) # expert_outputs: (B, S, K, D) b, s, k, d = expert_outputs.shape q = self.q_proj(self.bus_norm(hidden_state)).unsqueeze(2) # (B, S, 1, bus_width) k_vec = self.k_proj(expert_outputs) # (B, S, K, bus_width) v_vec = self.v_proj(expert_outputs) # (B, S, K, bus_width) scale = 1.0 / math.sqrt(self.bus_width) attn_scores = torch.matmul(q, k_vec.transpose(-1, -2)) * scale # (B, S, 1, K) attn_weights = F.softmax(attn_scores, dim=-1) context = torch.matmul(attn_weights, v_vec).squeeze(2) # (B, S, bus_width) return self.out_proj(context) # (B, S, D) class VisitAdapter(nn.Module): """Visit-Conditioned Low-Rank Adapter for Recurrent Transformer Blocks.""" def __init__(self, d_model: int, rank: int = 8, num_visits: int = 2): super().__init__() self.d_model = d_model self.rank = rank self.num_visits = num_visits # (num_visits, d_model, rank) and (num_visits, rank, d_model) self.down_proj = nn.Parameter(torch.empty(num_visits, d_model, rank)) self.up_proj = nn.Parameter(torch.empty(num_visits, rank, d_model)) self.visit_norm = RMSNorm(d_model) self.adapter_scale = nn.Parameter(torch.zeros(1)) self.reset_parameters() def reset_parameters(self): # Kaiming uniform for down, zero for up to warm-start at identity nn.init.kaiming_uniform_(self.down_proj, a=math.sqrt(5)) nn.init.zeros_(self.up_proj) def forward(self, x: torch.Tensor, visit_idx: int) -> torch.Tensor: v_idx = min(visit_idx, self.num_visits - 1) x_norm = self.visit_norm(x) down = self.down_proj[v_idx] # (D, R) up = self.up_proj[v_idx] # (R, D) delta = torch.matmul(torch.matmul(x_norm, down), up) * self.adapter_scale return delta class MoESwiGLU(nn.Module): """Sparse Mixture of Experts with Top-2 Softmax Routing and Thought Bus.""" def __init__(self, config: CanopyConfig): super().__init__() self.d_model = config.d_model self.num_experts = config.moe_num_experts self.top_k = config.moe_top_k self.router_aux_loss_coef = config.router_aux_loss_coef self.router_z_loss_coef = config.router_z_loss_coef self.use_thought_bus = config.use_thought_bus # Top-2 router self.router = nn.Linear(self.d_model, self.num_experts, bias=False) # Experts self.experts = nn.ModuleList([ SwiGLU(config.d_model, config.moe_intermediate_size) for _ in range(self.num_experts) ]) # Tokenwise Thought Bus if self.use_thought_bus: self.thought_bus = ThoughtBus(config.d_model, config.thought_bus_width) else: self.thought_bus = None def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: batch_size, seq_len, d_model = x.shape flat_x = x.view(-1, d_model) # (N, D) # Router logits router_logits = self.router(flat_x) # (N, E) router_probs = F.softmax(router_logits, dim=-1) # Top-K selection topk_weights, topk_indices = torch.topk(router_probs, self.top_k, dim=-1) topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) # normalize # Compute auxiliary load-balancing loss and router z-loss # Switch aux loss: E * sum_i (f_i * P_i) if self.training: tokens_per_expert = torch.zeros(self.num_experts, device=x.device, dtype=torch.float32) for k_idx in range(self.top_k): tokens_per_expert.scatter_add_( 0, topk_indices[:, k_idx], torch.ones_like(topk_indices[:, k_idx], dtype=torch.float32) ) fraction_tokens = tokens_per_expert / (flat_x.shape[0] * self.top_k) fraction_prob = router_probs.mean(dim=0) aux_loss = self.num_experts * torch.sum(fraction_tokens * fraction_prob) * self.router_aux_loss_coef # Router z-loss: mean(log(sum(exp(logits)))^2) log_sum_exp = torch.logsumexp(router_logits, dim=-1) z_loss = torch.mean(log_sum_exp ** 2) * self.router_z_loss_coef total_router_loss = aux_loss + z_loss else: total_router_loss = torch.tensor(0.0, device=x.device) # Execute experts out = torch.zeros_like(flat_x) expert_outputs = torch.zeros(flat_x.shape[0], self.top_k, d_model, device=x.device, dtype=x.dtype) for expert_id, expert in enumerate(self.experts): for k_idx in range(self.top_k): mask = (topk_indices[:, k_idx] == expert_id) if mask.any(): expert_in = flat_x[mask] e_out = expert(expert_in) expert_outputs[mask, k_idx] = e_out weight = topk_weights[mask, k_idx].unsqueeze(-1) out[mask] = out[mask] + weight * e_out out = out.view(batch_size, seq_len, d_model) # Thought Bus latent collaboration if self.thought_bus is not None: expert_outputs_3d = expert_outputs.view(batch_size, seq_len, self.top_k, d_model) bus_residual = self.thought_bus(x, expert_outputs_3d) out = out + bus_residual return out, total_router_loss class CanopyBlock(nn.Module): """Transformer Block supporting Dense Prelude/Coda and Recurrent MoE execution.""" def __init__(self, config: CanopyConfig, is_moe: bool = False): super().__init__() self.is_moe = is_moe self.loop_residual_scale = config.loop_residual_scale self.use_visit_adapter = config.use_visit_adapter and is_moe self.attn_norm = RMSNorm(config.d_model, eps=config.norm_eps) self.attn = CausalSelfAttention(config) self.ffn_norm = RMSNorm(config.d_model, eps=config.norm_eps) if is_moe: self.ffn = MoESwiGLU(config) if self.use_visit_adapter: self.visit_adapter = VisitAdapter(config.d_model, rank=config.visit_adapter_rank, num_visits=config.recurrent_visits) else: self.visit_adapter = None else: self.ffn = SwiGLU(config.d_model, config.dense_intermediate_size) self.visit_adapter = None def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, visit_idx: int = 0, kv_cache: Optional[Dict[str, torch.Tensor]] = None, layer_idx: int = 0, ) -> Tuple[torch.Tensor, torch.Tensor]: # Visit adapter delta if recurrent if self.visit_adapter is not None: x = x + self.visit_adapter(x, visit_idx) # Attention sublayer attn_out = self.attn(self.attn_norm(x), cos, sin, kv_cache=kv_cache, layer_idx=layer_idx) # Residual connection (scaled by 1/r if in second visit of looped recurrence) res_scale = self.loop_residual_scale if (self.is_moe and visit_idx > 0) else 1.0 x = x + (attn_out * res_scale) # FFN sublayer if self.is_moe: ffn_out, router_loss = self.ffn(self.ffn_norm(x)) else: ffn_out = self.ffn(self.ffn_norm(x)) router_loss = torch.tensor(0.0, device=x.device) x = x + (ffn_out * res_scale) return x, router_loss class CanopyForCausalLM(CanopyPreTrainedModel): """ Canopy-R3 Causal Language Model. Exact parameter count: 258,555,654 parameters. """ def __init__(self, config: CanopyConfig): super().__init__() self.config = config # Token embeddings self.embed_tokens = nn.Embedding(config.vocab_size, config.d_model) # Rotary embeddings self.rotary_emb = RotaryEmbedding(config.head_dim, max_seq_len=config.max_seq_len, theta=config.rope_theta) # Transformer blocks: 3 prelude dense + 6 recurrent MoE + 3 coda dense = 12 physical blocks self.layers = nn.ModuleList() for idx in range(config.num_layers): is_recurrent_moe = (config.prelude_layers <= idx < config.prelude_layers + config.recurrent_layers) self.layers.append(CanopyBlock(config, is_moe=is_recurrent_moe)) # Final RMSNorm self.final_norm = RMSNorm(config.d_model, eps=config.norm_eps) # Tied LM Head self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) if config.tie_word_embeddings: self.lm_head.weight = self.embed_tokens.weight # Apply quantization mode if requested if config.quant_mode in ("ternary", "binary"): bits = 1.58 if config.quant_mode == "ternary" else 1.0 convert_to_bitlinear(self, bits=bits) self.apply(self._init_weights) def _init_weights(self, module: nn.Module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) def num_parameters(self, trainable_only: bool = True) -> int: if trainable_only: return sum(p.numel() for p in self.parameters() if p.requires_grad) return sum(p.numel() for p in self.parameters()) def forward( self, input_ids: torch.Tensor, labels: Optional[torch.Tensor] = None, kv_cache: Optional[Dict[str, torch.Tensor]] = None, use_checkpointing: bool = False, ptrm_stochastic_scale: float = 0.0, ) -> Dict[str, torch.Tensor]: batch_size, seq_len = input_ids.shape x = self.embed_tokens(input_ids) offset = 0 if kv_cache is not None and not kv_cache.get("is_prefill", False): if "k_0" in kv_cache: offset = kv_cache["k_0"].shape[2] cos, sin = self.rotary_emb(x, seq_len, offset=offset) total_router_loss = torch.tensor(0.0, device=x.device) logical_layer_idx = 0 # 1. Prelude blocks (dense, unlooped) for idx in range(self.config.prelude_layers): block = self.layers[idx] if use_checkpointing and self.training: x, r_loss = checkpoint(block, x, cos, sin, 0, None, logical_layer_idx, use_reentrant=False) else: x, r_loss = block(x, cos, sin, visit_idx=0, kv_cache=kv_cache, layer_idx=logical_layer_idx) total_router_loss = total_router_loss + r_loss logical_layer_idx += 1 # 2. Recurrent MoE blocks (looped recurrent_visits times) recurrent_start = self.config.prelude_layers recurrent_end = recurrent_start + self.config.recurrent_layers for visit in range(self.config.recurrent_visits): # PTRM: Stochastic perturbation on recurrent transitions to escape deterministic loop attractors if ptrm_stochastic_scale > 0.0 and not self.training and visit > 0: x = x + torch.randn_like(x) * ptrm_stochastic_scale for idx in range(recurrent_start, recurrent_end): block = self.layers[idx] if use_checkpointing and self.training: x, r_loss = checkpoint(block, x, cos, sin, visit, None, logical_layer_idx, use_reentrant=False) else: x, r_loss = block(x, cos, sin, visit_idx=visit, kv_cache=kv_cache, layer_idx=logical_layer_idx) total_router_loss = total_router_loss + r_loss logical_layer_idx += 1 # 3. Coda blocks (dense, unlooped) for idx in range(recurrent_end, self.config.num_layers): block = self.layers[idx] if use_checkpointing and self.training: x, r_loss = checkpoint(block, x, cos, sin, 0, None, logical_layer_idx, use_reentrant=False) else: x, r_loss = block(x, cos, sin, visit_idx=0, kv_cache=kv_cache, layer_idx=logical_layer_idx) total_router_loss = total_router_loss + r_loss logical_layer_idx += 1 x = self.final_norm(x) logits = self.lm_head(x) loss = None if labels is not None: # Shift tokens for next-token prediction shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = F.cross_entropy(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1)) total_loss = loss + total_router_loss else: total_loss = None return { "logits": logits, "loss": total_loss, "lm_loss": loss, "router_loss": total_router_loss, } @torch.no_grad() def generate( self, input_ids: torch.Tensor, max_new_tokens: int = 64, temperature: float = 0.8, top_k: int = 50, top_p: float = 0.95, repetition_penalty: float = 1.15, ptrm_stochastic_scale: float = 0.0, eos_token_id: Optional[int] = None, ) -> torch.Tensor: """Autoregressive generation with KV caching, repetition penalty, and PTRM stochastic exploration.""" self.eval() batch_size, prompt_len = input_ids.shape kv_cache = {"is_prefill": True} # Prefill prompt outputs = self.forward(input_ids, kv_cache=kv_cache, ptrm_stochastic_scale=ptrm_stochastic_scale) next_token_logits = outputs["logits"][:, -1, :].clone() generated = input_ids.clone() kv_cache["is_prefill"] = False for _ in range(max_new_tokens): logits = next_token_logits.clone() # Apply repetition penalty to already generated token IDs if repetition_penalty > 1.0: for b_idx in range(batch_size): unique_tokens = torch.unique(generated[b_idx]) token_logits = logits[b_idx, unique_tokens] logits[b_idx, unique_tokens] = torch.where( token_logits > 0, token_logits / repetition_penalty, token_logits * repetition_penalty, ) if temperature > 0: logits = logits / temperature if top_k > 0: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[:, [-1]]] = -float("Inf") if top_p < 1.0: sorted_logits, sorted_indices = torch.sort(logits, descending=True) cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) sorted_indices_to_remove = cumulative_probs > top_p sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() sorted_indices_to_remove[..., 0] = 0 indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove) logits[indices_to_remove] = -float("Inf") probs = F.softmax(logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) else: next_token = torch.argmax(logits, dim=-1, keepdim=True) generated = torch.cat([generated, next_token], dim=-1) if eos_token_id is not None and (next_token == eos_token_id).all(): break # Forward single token with cached KV and PTRM stochastic exploration step_outputs = self.forward(next_token, kv_cache=kv_cache, ptrm_stochastic_scale=ptrm_stochastic_scale) next_token_logits = step_outputs["logits"][:, -1, :] return generated