canopy-258m-r3 / modeling_canopy.py
psikosen's picture
Deploy Canopy-258M-R3 Stage 13 Unified Flagship Model
cb87b3d verified
Raw History Blame Contribute Delete
23.3 kB
"""
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