Text Generation
Transformers
Safetensors
English
canopy
browser-use
web-agent
recurrent-moe
edge-llm
lightpanda
obscura
multi-agent
robotics-web
conversational
custom_code
Instructions to use psikosen/canopy-258m-r3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use psikosen/canopy-258m-r3 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="psikosen/canopy-258m-r3", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("psikosen/canopy-258m-r3", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use psikosen/canopy-258m-r3 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "psikosen/canopy-258m-r3" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "psikosen/canopy-258m-r3", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/psikosen/canopy-258m-r3
- SGLang
How to use psikosen/canopy-258m-r3 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "psikosen/canopy-258m-r3" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "psikosen/canopy-258m-r3", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "psikosen/canopy-258m-r3" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "psikosen/canopy-258m-r3", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use psikosen/canopy-258m-r3 with Docker Model Runner:
docker model run hf.co/psikosen/canopy-258m-r3
Download modeling_canopy.py from psikosen/canopy-258m-r3: direct link, hf CLI and curl.
- Browser
- Download file 23.3 kB
-
https://huggingface.co/psikosen/canopy-258m-r3/resolve/main/modeling_canopy.py
- Command line
-
hf download hf://psikosen/canopy-258m-r3/modeling_canopy.py
-
curl -L -o modeling_canopy.py https://huggingface.co/psikosen/canopy-258m-r3/resolve/main/modeling_canopy.py
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, | |
| } | |
| 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 | |