Decision-1.0-Route-0.6B / decision1_vela.py
Xunzhuo's picture
Load with 🤗 Transformers (trust_remote_code) (#3)
a5b21df
Raw History Blame Contribute Delete
22.5 kB
# Copyright 2026 The vLLM Semantic Router Authors.
# Copyright 2024 Answer.AI, LightOn, and contributors, and the HuggingFace Inc. team.
# SPDX-License-Identifier: Apache-2.0
#
# The ModernBERT encoder below is adapted from Hugging Face Transformers 4.57.6
# (models/modernbert/modeling_modernbert.py and modeling_rope_utils.py, Apache-2.0):
# the SDPA path with YaRN rotary embeddings, reduced to inference.
"""Kai / Lex / Route: one ModernBERT encoder with Choice, Noul and Score paths.
Inference follows the native Decision 1.0 runtime: FP32 weights and math (no TF32,
no fused attention fast path), one marker per candidate, complete inputs only (no
truncation), and rows sorted by question type into physical batches of eight.
"""
from __future__ import annotations
import copy
import json
import math
from contextlib import contextmanager
from pathlib import Path
from typing import Any
import torch
from torch import nn
from torch.nn import functional
from .decision1_system_one import DecisionInputTooLongError, Row, content_text
KINDS = ("choice", "noul", "score")
PHYSICAL_BATCH = 8
NOUL_DEFAULT_FALSE = "No. The statement or question is not satisfied."
NOUL_DEFAULT_TRUE = "Yes. The statement or question is satisfied."
ENCODER_GEOMETRY = {
"model_type": "modernbert",
"hidden_size": 768,
"num_hidden_layers": 22,
"num_attention_heads": 12,
"max_position_embeddings": 32768,
}
def _yarn_inverse_frequencies(
config: dict[str, Any], base: float
) -> tuple[torch.Tensor, float]:
scaling = config["rope_scaling"]
dim = config["hidden_size"] // config["num_attention_heads"]
factor = scaling["factor"]
original = (
scaling.get("original_max_position_embeddings")
or config["max_position_embeddings"]
)
def get_mscale(scale, mscale=1):
if scale <= 1:
return 1.0
return 0.1 * mscale * math.log(scale) + 1.0
attention_factor = scaling.get("attention_factor")
mscale, mscale_all_dim = scaling.get("mscale"), scaling.get("mscale_all_dim")
if attention_factor is None:
if mscale and mscale_all_dim:
attention_factor = float(
get_mscale(factor, mscale) / get_mscale(factor, mscale_all_dim)
)
else:
attention_factor = get_mscale(factor)
beta_fast = scaling.get("beta_fast") or 32
beta_slow = scaling.get("beta_slow") or 1
def correction_dim(rotations):
return (dim * math.log(original / (rotations * 2 * math.pi))) / (
2 * math.log(base)
)
low, high = correction_dim(beta_fast), correction_dim(beta_slow)
if scaling.get("truncate", True):
low, high = math.floor(low), math.ceil(high)
low, high = max(low, 0), min(high, dim - 1)
if low == high:
high += 0.001
ramp = torch.clamp(
(torch.arange(dim // 2, dtype=torch.float32) - low) / (high - low), 0, 1
)
pos_freqs = base ** (torch.arange(0, dim, 2).to(dtype=torch.float) / dim)
extrapolation = 1.0 / pos_freqs
interpolation = 1.0 / (factor * pos_freqs)
extrapolation_factor = 1 - ramp.to(dtype=torch.float)
inverse = (
interpolation * (1 - extrapolation_factor)
+ extrapolation * extrapolation_factor
)
return inverse, attention_factor
class RotaryEmbedding(nn.Module):
def __init__(self, config: dict[str, Any], base: float):
super().__init__()
scaling = config.get("rope_scaling")
if (
not isinstance(scaling, dict)
or scaling.get("rope_type", scaling.get("type")) != "yarn"
):
raise ValueError("Decision 1.0 encoders use YaRN rotary embeddings")
inverse, self.attention_scaling = _yarn_inverse_frequencies(config, base)
self.register_buffer("inv_freq", inverse, persistent=False)
@torch.no_grad()
def forward(self, x: torch.Tensor, position_ids: torch.Tensor):
inverse = (
self.inv_freq[None, :, None]
.float()
.expand(position_ids.shape[0], -1, 1)
.to(x.device)
)
positions = position_ids[:, None, :].float()
device_type = x.device.type if x.device.type != "mps" else "cpu"
with torch.autocast(device_type=device_type, enabled=False):
freqs = (inverse.float() @ positions.float()).transpose(1, 2)
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos() * self.attention_scaling
sin = emb.sin() * self.attention_scaling
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
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(q, k, cos, sin):
cos, sin = cos.unsqueeze(1), sin.unsqueeze(1)
return (q * cos) + (_rotate_half(q) * sin), (k * cos) + (_rotate_half(k) * sin)
class Attention(nn.Module):
def __init__(self, config: dict[str, Any], layer_id: int):
super().__init__()
hidden = config["hidden_size"]
self.num_heads = config["num_attention_heads"]
self.head_dim = hidden // self.num_heads
self.all_head_size = self.head_dim * self.num_heads
self.Wqkv = nn.Linear(
hidden, 3 * self.all_head_size, bias=config["attention_bias"]
)
if layer_id % config["global_attn_every_n_layers"] != 0:
self.local = True
base = config["local_rope_theta"]
if base is None:
base = config["global_rope_theta"]
else:
self.local = False
base = config["global_rope_theta"]
self.rotary_emb = RotaryEmbedding(config, base)
self.Wo = nn.Linear(hidden, hidden, bias=config["attention_bias"])
def forward(self, hidden_states, attention_mask, sliding_window_mask, position_ids):
batch = hidden_states.shape[0]
qkv = self.Wqkv(hidden_states).view(batch, -1, 3, self.num_heads, self.head_dim)
cos, sin = self.rotary_emb(qkv, position_ids=position_ids)
query, key, value = qkv.transpose(3, 1).unbind(dim=2)
query, key = _apply_rotary(query, key, cos, sin)
mask = sliding_window_mask if self.local else attention_mask
if (
torch.version.hip is not None
and hidden_states.device.type == "cuda"
and torch.backends.cuda.mem_efficient_sdp_enabled()
):
# ROCm's efficient SDPA kernel needs contiguous post-RoPE inputs.
query, key, value = query.contiguous(), key.contiguous(), value.contiguous()
output = functional.scaled_dot_product_attention(
query, key, value, attn_mask=mask
)
output = output.transpose(1, 2).contiguous().view(batch, -1, self.all_head_size)
return self.Wo(output)
class MLP(nn.Module):
def __init__(self, config: dict[str, Any]):
super().__init__()
hidden, intermediate = config["hidden_size"], int(config["intermediate_size"])
if config["hidden_activation"] != "gelu":
raise ValueError("Decision 1.0 encoders use exact GELU")
self.Wi = nn.Linear(hidden, intermediate * 2, bias=config["mlp_bias"])
self.Wo = nn.Linear(intermediate, hidden, bias=config["mlp_bias"])
def forward(self, hidden_states):
values, gate = self.Wi(hidden_states).chunk(2, dim=-1)
return self.Wo(functional.gelu(values) * gate)
class EncoderLayer(nn.Module):
def __init__(self, config: dict[str, Any], layer_id: int):
super().__init__()
hidden, eps, bias = (
config["hidden_size"],
config["norm_eps"],
config["norm_bias"],
)
self.attn_norm = (
nn.Identity() if layer_id == 0 else nn.LayerNorm(hidden, eps=eps, bias=bias)
)
self.attn = Attention(config, layer_id)
self.mlp_norm = nn.LayerNorm(hidden, eps=eps, bias=bias)
self.mlp = MLP(config)
def forward(self, hidden_states, attention_mask, sliding_window_mask, position_ids):
hidden_states = hidden_states + self.attn(
self.attn_norm(hidden_states),
attention_mask,
sliding_window_mask,
position_ids,
)
return hidden_states + self.mlp(self.mlp_norm(hidden_states))
class Embeddings(nn.Module):
def __init__(self, config: dict[str, Any]):
super().__init__()
hidden = config["hidden_size"]
self.tok_embeddings = nn.Embedding(
config["vocab_size"], hidden, padding_idx=config["pad_token_id"]
)
self.norm = nn.LayerNorm(
hidden, eps=config["norm_eps"], bias=config["norm_bias"]
)
def forward(self, input_ids):
return self.norm(self.tok_embeddings(input_ids))
class Encoder(nn.Module):
"""Parameter names match transformers' ModernBertModel."""
def __init__(self, config: dict[str, Any]):
super().__init__()
self.config = config
self.embeddings = Embeddings(config)
self.layers = nn.ModuleList(
[
EncoderLayer(config, index)
for index in range(config["num_hidden_layers"])
]
)
self.final_norm = nn.LayerNorm(
config["hidden_size"], eps=config["norm_eps"], bias=config["norm_bias"]
)
def masks(self, attention_mask: torch.Tensor, dtype: torch.dtype):
batch, length = attention_mask.shape
expanded = (
attention_mask[:, None, None, :].expand(batch, 1, length, length).to(dtype)
)
inverted = torch.tensor(1.0, dtype=dtype) - expanded
global_mask = inverted.masked_fill(
inverted.to(torch.bool), torch.finfo(dtype).min
)
rows = torch.arange(length).unsqueeze(0)
window = (
(torch.abs(rows - rows.T) <= self.config["local_attention"] // 2)
.unsqueeze(0)
.unsqueeze(0)
.to(attention_mask.device)
)
sliding = global_mask.masked_fill(window.logical_not(), torch.finfo(dtype).min)
return global_mask, sliding
class VelaDecision(nn.Module):
"""Shared encoder (Noul path), private Choice and Score encoder copies, typed heads."""
def __init__(self, encoder_config: dict[str, Any], head: dict[str, Any]):
super().__init__()
for key, expected in ENCODER_GEOMETRY.items():
if encoder_config.get(key) != expected:
raise ValueError(
f"Unsupported Decision 1.0 encoder configuration: {key}"
)
if head.get("head_layers") != 2 or head.get("head_heads") != 12:
raise ValueError("Unsupported Decision 1.0 head geometry")
hidden = encoder_config["hidden_size"]
self.encoder = Encoder(encoder_config)
self.type_embedding = nn.Embedding(3, hidden)
self.heads = nn.ModuleDict(
{
kind: nn.ModuleList(
[
nn.TransformerEncoderLayer(
hidden,
12,
4 * hidden,
0.1,
activation="relu",
batch_first=True,
norm_first=True,
)
for _ in range(2)
]
)
for kind in KINDS
}
)
self.scorers = nn.ModuleDict(
{
kind: nn.Sequential(
nn.LayerNorm(hidden),
nn.Linear(hidden, hidden),
nn.GELU(),
nn.Linear(hidden, 1),
)
for kind in KINDS
}
)
self.choice_blocks = nn.ModuleList(
copy.deepcopy(layer) for layer in self.encoder.layers
)
self.choice_final_norm = copy.deepcopy(self.encoder.final_norm)
self.score_blocks = nn.ModuleList(
copy.deepcopy(layer) for layer in self.encoder.layers
)
self.score_final_norm = copy.deepcopy(self.encoder.final_norm)
@staticmethod
def _path(hidden, layers, final_norm, masks, positions):
for layer in layers:
hidden = layer(hidden, masks[0], masks[1], positions)
return final_norm(hidden)
def forward(
self, input_ids, attention_mask, kind_ids, marker_positions, valid_candidates
):
positions = torch.arange(input_ids.shape[1], device=input_ids.device).unsqueeze(
0
)
masks = self.encoder.masks(attention_mask, torch.float32)
embedded = self.encoder.embeddings(input_ids)
present = set(kind_ids.tolist())
type_offset = self.type_embedding(kind_ids)[:, None, :]
hidden_by_kind = {}
if 0 in present:
hidden = self._path(
embedded, self.choice_blocks, self.choice_final_norm, masks, positions
)
hidden_by_kind["choice"] = hidden + type_offset.to(hidden.dtype)
if 1 in present:
hidden = self._path(
embedded, self.encoder.layers, self.encoder.final_norm, masks, positions
)
hidden_by_kind["noul"] = hidden + type_offset.to(hidden.dtype)
if 2 in present:
hidden = self._path(
embedded, self.score_blocks, self.score_final_norm, masks, positions
)
hidden_by_kind["score"] = hidden + type_offset.to(hidden.dtype)
pad = ~attention_mask.bool()
output = torch.empty(
marker_positions.shape, device=input_ids.device, dtype=torch.float32
)
for index, kind in enumerate(KINDS):
rows = torch.nonzero(kind_ids == index, as_tuple=False).flatten()
if rows.numel() == 0:
continue
hidden = hidden_by_kind[kind].index_select(0, rows)
branch_pad = pad.index_select(0, rows)
for layer in self.heads[kind]:
hidden = layer(hidden, src_key_padding_mask=branch_pad)
where = marker_positions.index_select(0, rows)
markers = torch.gather(
hidden, 1, where[:, :, None].expand(-1, -1, hidden.shape[-1])
)
output = output.index_copy(
0, rows, self.scorers[kind](markers).squeeze(-1).float()
)
return output.masked_fill(~valid_candidates, torch.finfo(torch.float32).min)
@contextmanager
def native_flags():
"""The published runtime's settings: no fused attention fast path and no TF32, restored afterwards."""
fastpath = torch.backends.mha.get_fastpath_enabled()
matmul, cudnn = (
torch.backends.cuda.matmul.allow_tf32,
torch.backends.cudnn.allow_tf32,
)
torch.backends.mha.set_fastpath_enabled(False)
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
try:
yield
finally:
torch.backends.mha.set_fastpath_enabled(fastpath)
torch.backends.cuda.matmul.allow_tf32 = matmul
torch.backends.cudnn.allow_tf32 = cudnn
class VelaRuntime:
"""Loaded weights, tokenizer and the complete-input limit of one Decision 1.0 encoder."""
noul_default_false = NOUL_DEFAULT_FALSE
noul_default_true = NOUL_DEFAULT_TRUE
noul_explicit_null = "use_default"
def __init__(self, model: VelaDecision, tokenizer: Any, max_input_tokens: int):
self.model = model
self.tokenizer = tokenizer
self.max_input_tokens = max_input_tokens
special = (
(
tokenizer.cls_token_id
if tokenizer.cls_token_id is not None
else tokenizer.bos_token_id
),
(
tokenizer.sep_token_id
if tokenizer.sep_token_id is not None
else tokenizer.eos_token_id
),
tokenizer.pad_token_id,
tokenizer.mask_token_id,
)
if any(value is None for value in special):
raise ValueError(
"The tokenizer must define CLS/BOS, SEP/EOS, PAD and MASK tokens"
)
self.bos, self.sep, self.pad, self.marker = special
@classmethod
def load(
cls, root: Path, descriptor: dict[str, Any], *, max_input_tokens: int, device
):
from safetensors.torch import load_file
from transformers import AutoTokenizer
model_config = json.loads(
(root / descriptor["model_config"]).read_text(encoding="utf-8")
)
if (
model_config.get("arm") != "all22"
or model_config.get("training_arm") != "S22"
or model_config.get("type_order") != list(KINDS)
or model_config.get("packing", {}).get("state_truncation") != "error"
):
raise ValueError(
"Only the three-path all22 / S22 Decision 1.0 encoder is supported"
)
encoder_config = json.loads(
(root / descriptor["backbone"]["config"]).read_text(encoding="utf-8")
)
model = VelaDecision(encoder_config, model_config["head"])
state = {
f"encoder.{name}": tensor
for name, tensor in load_file(
str(root / descriptor["backbone"]["weights"][0])
).items()
}
for role in ("decision_heads", "choice_encoder", "score_encoder"):
part = load_file(str(root / descriptor["decision_weights"][role]))
if set(part) & set(state):
raise ValueError("Decision 1.0 weight files overlap")
state.update(part)
if any(tensor.dtype != torch.float32 for tensor in state.values()):
raise ValueError("Decision 1.0 encoder weights must be FP32")
model.load_state_dict(state, strict=True)
expected = model_config.get("parameters")
loaded = sum(parameter.numel() for parameter in model.parameters())
if expected is not None and loaded != expected:
raise ValueError(
f"Loaded {loaded:,} parameters; the model declares {expected:,}"
)
model.to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(
str((root / descriptor["tokenizer"]["json"]).parent),
trust_remote_code=False,
)
return cls(model, tokenizer, max_input_tokens)
def _tokens(self, text: str) -> list[int]:
ids = self.tokenizer(text, add_special_tokens=False, truncation=False)[
"input_ids"
]
if not ids:
raise ValueError("A candidate, question or state renders to no tokens")
return list(ids)
def encode(self, row: Row, cache: dict[str, list[int]]) -> dict[str, Any]:
def tokens(text):
if text not in cache:
cache[text] = self._tokens(text)
return cache[text]
ids = [
self.bos,
*tokens(f"{row.type} question: {content_text(row.instructions)}"),
self.sep,
]
positions = []
for index, candidate in enumerate(row.candidates):
if candidate.description is None:
text = candidate.key
else:
text = content_text(candidate.description)
if row.type == "choice":
text = f"{candidate.key}: {text}"
if row.type == "score":
text = f"level {index}: {text}"
positions.append(len(ids))
ids.extend((self.marker, *tokens(text), self.sep))
state = tokens(row.state)
room = self.max_input_tokens - len(ids) - 1
if room < 1 or len(state) > room:
raise DecisionInputTooLongError(
f"{row.question_id}: the complete input exceeds {self.max_input_tokens} "
"tokens; no truncation allowed"
)
ids.extend((*state, self.sep))
return {"ids": ids, "positions": positions, "kind": KINDS.index(row.type)}
def predict(self, rows: list[Row]) -> tuple[list[list[float]], list[int]]:
"""Probabilities per row in request order, and input tokens per row."""
cache: dict[str, list[int]] = {}
encoded = [self.encode(row, cache) for row in rows]
order = sorted(
range(len(rows)), key=lambda index: rows[index].type.capitalize()
)
device = next(self.model.parameters()).device
results: list[list[float] | None] = [None] * len(rows)
with torch.inference_mode(), native_flags():
for start in range(0, len(order), PHYSICAL_BATCH):
chunk = order[start : start + PHYSICAL_BATCH]
items = [encoded[index] for index in chunk]
length = max(len(item["ids"]) for item in items)
width = max(len(item["positions"]) for item in items)
input_ids = torch.full((len(items), length), self.pad, dtype=torch.long)
mask = torch.zeros((len(items), length), dtype=torch.bool)
markers = torch.zeros((len(items), width), dtype=torch.long)
valid = torch.zeros((len(items), width), dtype=torch.bool)
for slot, item in enumerate(items):
input_ids[slot, : len(item["ids"])] = torch.tensor(item["ids"])
mask[slot, : len(item["ids"])] = True
markers[slot, : len(item["positions"])] = torch.tensor(
item["positions"]
)
valid[slot, : len(item["positions"])] = True
kinds = torch.tensor([item["kind"] for item in items])
logits = self.model(
input_ids.to(device),
mask.to(device),
kinds.to(device),
markers.to(device),
valid.to(device),
)
if not torch.isfinite(logits).all():
raise FloatingPointError("Non-finite Decision logits")
for slot, index in enumerate(chunk):
count = len(encoded[index]["positions"])
results[index] = logits[slot, :count].softmax(-1).cpu().tolist()
return results, [len(item["ids"]) for item in encoded] # type: ignore[return-value]