| """glint-router-1m: standalone inference. torch + tokenizers + safetensors, nothing else. |
| |
| one forward pass over the prompt gives you every field at once: domain, |
| complexity, code, math, reasoning, long_output, route, and a 64-d projection |
| used for user-defined categories. no decoding loop, no output parsing. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import json |
| import math |
| from dataclasses import dataclass |
| from pathlib import Path |
|
|
| import torch |
| from torch import Tensor, nn |
| from torch.nn import functional |
|
|
| HERE = Path(__file__).parent |
| PAD_ID = 0 |
| BOS_ID = 1 |
| BINARY_FIELDS = ("code", "math", "reasoning", "long_output") |
| COMPLEXITY_LEVELS = 5 |
|
|
| DOMAINS = ( |
| "programming", "web_dev", "databases", "devops_sysadmin", "security", |
| "data_science_ml", "math", "science", "engineering", "reasoning_puzzle", |
| "creative_writing", "professional_writing", "editing_grammar", "translation", |
| "factual_qa", "education", "business", "finance", "legal", "health", |
| "travel", "food", "entertainment", "lifestyle", "other", |
| ) |
|
|
|
|
| @dataclass(frozen=True) |
| class RouterConfig: |
| vocab_size: int = 4096 |
| dim: int = 128 |
| n_heads: int = 8 |
| layers: int = 3 |
| ffn_hidden: int = 208 |
| max_len: int = 256 |
| rope_base: float = 10_000.0 |
| proj_dim: int = 64 |
| n_domains: int = len(DOMAINS) |
|
|
|
|
| def build_rope_cache(config: RouterConfig) -> tuple[Tensor, Tensor]: |
| head_dim = config.dim // config.n_heads |
| positions = torch.arange(config.max_len, dtype=torch.float32) |
| inv_freq = 1.0 / (config.rope_base ** (torch.arange(0, head_dim, 2).float() / head_dim)) |
| angles = torch.outer(positions, inv_freq) |
| return torch.cos(angles), torch.sin(angles) |
|
|
|
|
| def apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor: |
| x_even, x_odd = x[..., 0::2], x[..., 1::2] |
| rotated_even = x_even * cos - x_odd * sin |
| rotated_odd = x_even * sin + x_odd * cos |
| return torch.stack((rotated_even, rotated_odd), dim=-1).flatten(-2) |
|
|
|
|
| class SwiGlu(nn.Module): |
| def __init__(self, dim: int, hidden: int) -> None: |
| super().__init__() |
| self.gate_up = nn.Linear(dim, 2 * hidden, bias=False) |
| self.down = nn.Linear(hidden, dim, bias=False) |
|
|
| def forward(self, x: Tensor) -> Tensor: |
| gate, up = self.gate_up(x).chunk(2, dim=-1) |
| return self.down(functional.silu(gate) * up) |
|
|
|
|
| class RouterAttention(nn.Module): |
| def __init__(self, config: RouterConfig) -> None: |
| super().__init__() |
| self.n_heads = config.n_heads |
| self.head_dim = config.dim // config.n_heads |
| self.qkv = nn.Linear(config.dim, 3 * config.dim, bias=False) |
| self.out = nn.Linear(config.dim, config.dim, bias=False) |
|
|
| def forward(self, x: Tensor, cos: Tensor, sin: Tensor, attn_mask: Tensor) -> Tensor: |
| batch, seq_len, dim = x.shape |
| q, k, v = self.qkv(x).split(dim, dim=-1) |
| shape = (batch, seq_len, self.n_heads, self.head_dim) |
| q = apply_rope(q.view(shape).transpose(1, 2), cos, sin) |
| k = apply_rope(k.view(shape).transpose(1, 2), cos, sin) |
| v = v.view(shape).transpose(1, 2) |
| attended = functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask) |
| return self.out(attended.transpose(1, 2).reshape(batch, seq_len, dim)) |
|
|
|
|
| class RouterBlock(nn.Module): |
| def __init__(self, config: RouterConfig) -> None: |
| super().__init__() |
| self.attn_norm = nn.RMSNorm(config.dim) |
| self.attn = RouterAttention(config) |
| self.ffn_norm = nn.RMSNorm(config.dim) |
| self.ffn = SwiGlu(config.dim, config.ffn_hidden) |
|
|
| def forward(self, x: Tensor, cos: Tensor, sin: Tensor, attn_mask: Tensor) -> Tensor: |
| x = x + self.attn(self.attn_norm(x), cos, sin, attn_mask) |
| return x + self.ffn(self.ffn_norm(x)) |
|
|
|
|
| class GlintRouter(nn.Module): |
| def __init__(self, config: RouterConfig | None = None) -> None: |
| super().__init__() |
| config = config or RouterConfig() |
| self.config = config |
| self.embed = nn.Embedding(config.vocab_size, config.dim, padding_idx=PAD_ID) |
| self.blocks = nn.ModuleList(RouterBlock(config) for _ in range(config.layers)) |
| self.final_norm = nn.RMSNorm(config.dim) |
| cos, sin = build_rope_cache(config) |
| self.register_buffer("rope_cos", cos, persistent=False) |
| self.register_buffer("rope_sin", sin, persistent=False) |
|
|
| pooled = 2 * config.dim |
| self.domain_head = nn.Linear(pooled, config.n_domains) |
| self.complexity_dir = nn.Linear(pooled, 1, bias=False) |
| self.complexity_bias = nn.Parameter(torch.zeros(COMPLEXITY_LEVELS - 1)) |
| self.binary_head = nn.Linear(pooled, len(BINARY_FIELDS)) |
| self.route_head = nn.Linear(pooled, 1) |
| self.proj_head = nn.Linear(pooled, config.proj_dim) |
| self.register_buffer("temperature", torch.ones(3), persistent=True) |
|
|
| def encode(self, tokens: Tensor) -> Tensor: |
| valid = tokens != PAD_ID |
| seq_len = tokens.shape[1] |
| cos = self.rope_cos[:seq_len].to(self.embed.weight.dtype) |
| sin = self.rope_sin[:seq_len].to(self.embed.weight.dtype) |
| attn_mask = torch.zeros(tokens.shape, dtype=self.embed.weight.dtype, |
| device=tokens.device) |
| attn_mask = attn_mask.masked_fill(~valid, float("-inf"))[:, None, None, :] |
| x = self.embed(tokens) |
| for block in self.blocks: |
| x = block(x, cos, sin, attn_mask) |
| x = self.final_norm(x) |
| mask = valid.unsqueeze(-1) |
| mean = (x * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1) |
| maximum = torch.nan_to_num(x.masked_fill(~mask, float("-inf")).max(dim=1).values, |
| neginf=0.0) |
| return torch.cat((mean, maximum), dim=-1) |
|
|
| def forward(self, tokens: Tensor) -> dict[str, Tensor]: |
| pooled = self.encode(tokens) |
| return { |
| "domain": self.domain_head(pooled), |
| "complexity": self.complexity_dir(pooled) + self.complexity_bias, |
| "binary": self.binary_head(pooled), |
| "route": self.route_head(pooled).squeeze(-1), |
| "proj": functional.normalize(self.proj_head(pooled), dim=-1), |
| } |
|
|
| def calibrated(self, tokens: Tensor) -> dict[str, Tensor]: |
| """probabilities with the fitted temperatures applied. the routing |
| arithmetic in policy.py runs on these numbers, so they have to mean |
| something. temperatures were fitted on held-out data after training.""" |
| out = self.forward(tokens) |
| route_t, complexity_t, binary_t = self.temperature.unbind() |
| return { |
| "domain": out["domain"].softmax(dim=-1), |
| "complexity": (out["complexity"] / complexity_t).sigmoid(), |
| "binary": (out["binary"] / binary_t).sigmoid(), |
| "route": (out["route"] / route_t).sigmoid(), |
| "proj": out["proj"], |
| } |
|
|
|
|
| def complexity_from_cumulative(probabilities: Tensor) -> Tensor: |
| """coral decode. level = 1 + how many thresholds the prompt clears.""" |
| return 1 + (probabilities > 0.5).sum(dim=-1) |
|
|
|
|
| def encode_batch(tokenizer, texts: list[str], max_len: int) -> Tensor: |
| """bos-prefixed, right-padded. long prompts lose their tail, because the |
| instruction verb lives at the front and the pasted context lives at the back.""" |
| out = torch.full((len(texts), max_len), PAD_ID, dtype=torch.long) |
| for row, text in enumerate(texts): |
| ids = [BOS_ID] + tokenizer.encode(text).ids[: max_len - 1] |
| out[row, : len(ids)] = torch.tensor(ids, dtype=torch.long) |
| return out |
|
|
|
|
| def load_router(directory: Path = HERE, device: str = "cpu"): |
| """returns (model, tokenizer). reads config.json + model.safetensors + tokenizer.json.""" |
| from safetensors.torch import load_file |
| from tokenizers import Tokenizer |
|
|
| directory = Path(directory) |
| config = RouterConfig(**json.loads((directory / "config.json").read_text())["model"]) |
| model = GlintRouter(config).to(device) |
| model.load_state_dict(load_file(directory / "model.safetensors", device=device)) |
| model.eval() |
| return model, Tokenizer.from_file(str(directory / "tokenizer.json")) |
|
|
|
|
| def count_parameters(model: nn.Module) -> int: |
| return sum(p.numel() for p in model.parameters()) |
|
|