Text Classification
Transformers
Safetensors
edlm
feature-extraction
decision-model
system-one
custom_code
Instructions to use nace-ai/drex-dlm with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use nace-ai/drex-dlm with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="nace-ai/drex-dlm", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("nace-ai/drex-dlm", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download code/kev/model.py from nace-ai/drex-dlm: direct link, hf CLI and curl.
- Browser
- Download file 28.8 kB
-
https://huggingface.co/nace-ai/drex-dlm/resolve/main/code/kev/model.py
- Command line
-
hf download hf://nace-ai/drex-dlm/code/kev/model.py
-
curl -L -o model.py https://huggingface.co/nace-ai/drex-dlm/resolve/main/code/kev/model.py
28.8 kB
| """Decision model: causal LM backbone + block-causal branch mask + pointer readout.""" | |
| import copy, math, os, re | |
| from pathlib import Path | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer, DynamicCache | |
| # Reuse existing rarely-used Qwen special tokens as delimiters (state, q, opt, /opt, decide) so no | |
| # embedding rows need to be added/trained; LoRA adapts their meaning. | |
| SPECIAL = ["<|fim_prefix|>", "<|fim_middle|>", "<|box_start|>", "<|box_end|>", "<|fim_suffix|>"] | |
| # training context: state tokens, tokens per question branch, and the whole packed record. Frozen suites are admitted with | |
| # this rule (kev.suite) and training applies it to records built on the fly, so train and eval see the same population. | |
| MAX_STATE, MAX_BRANCH, MAX_PACKED = 384, 1024, 2048 | |
| # Position window of the weights. Serving recommends 16,384 and can be raised to this. | |
| CONTEXT_LENGTH = 32768 | |
| RECOMMENDED_CONTEXT = 16384 | |
| def _serve_limit(name, default): | |
| """A serving cap from the environment. Missing or invalid keeps `default`. Never above the position window.""" | |
| raw = os.environ.get(name) | |
| if raw is None or raw == "": | |
| return default | |
| try: | |
| value = int(raw) | |
| except ValueError: | |
| return default | |
| if value < 1 or value > CONTEXT_LENGTH: | |
| return default | |
| return value | |
| # KEV_CONTEXT selects the window. KEV_SERVE_MAX_STATE, KEV_SERVE_MAX_BRANCH, and KEV_SERVE_MAX_PACKED override one cap. | |
| _context = _serve_limit("KEV_CONTEXT", RECOMMENDED_CONTEXT) | |
| SERVE_MAX_STATE = _serve_limit("KEV_SERVE_MAX_STATE", _context) | |
| SERVE_MAX_BRANCH = _serve_limit("KEV_SERVE_MAX_BRANCH", _context) | |
| SERVE_MAX_PACKED = _serve_limit("KEV_SERVE_MAX_PACKED", _context) | |
| # Training admission stays on the historical 8,192-token row, so a wider serving window does not widen training. | |
| MAX_TRAIN_STATE = 8192 - (MAX_BRANCH - MAX_STATE) | |
| def training_context(max_state=MAX_STATE): | |
| """The encoder limits for training with the state limit lifted to `max_state`: the row (state + one branch) and | |
| packed limits grow by the same amount, so every question keeps its token budget. training_context() is the default | |
| training context (kev.suite.CONTEXT without `truncate`).""" | |
| if not MAX_STATE <= max_state <= MAX_TRAIN_STATE: | |
| raise ValueError(f"max_state must be in [{MAX_STATE}, {MAX_TRAIN_STATE}]") | |
| extra = max_state - MAX_STATE | |
| return {"max_state": max_state, "max_branch": MAX_BRANCH + extra, "max_packed": MAX_PACKED + extra} | |
| def rows_per_pass(rows, prefix_len=0, budget=SERVE_MAX_PACKED): | |
| """How many causal rows one inference forward pass takes: as many as fit `budget` tokens counting the cached state | |
| each row carries (prefix_len) plus its own tokens, at least one. Memory per pass is bounded by one maximal row however | |
| many questions a request has, and the rows are independent, so the answers do not depend on the split. Kev-4B on MLX, | |
| a 4.8k-token state with 64 questions: 24.7 GB peak in one pass, 8.9 GB one row at a time (and faster: the cost was | |
| 64 copies of the state cache).""" | |
| return max(1, budget // (prefix_len + max(len(r) for r in rows))) | |
| class ContextOverflow(ValueError): | |
| """A record does not encode within its context (state, branch or packed limit). Serving turns it into a 422; the | |
| benchmark counts it as a rejected record for suites scored as published (skip_overlong).""" | |
| def load_tokenizer(name, revision=None): | |
| return AutoTokenizer.from_pretrained(name, revision=revision, trust_remote_code=False) | |
| def load_backbone(name, revision, dtype, attn): | |
| """Load the encoder. Efficient-DLM is bidirectional unless the caller passes a mask. | |
| The shards use Qwen3 block names. The config does not. Loading them as | |
| Qwen3ForCausalLM would build a causal mask whenever the caller omits one. | |
| """ | |
| root = Path(name) | |
| if (root / "modeling_edlm.py").is_file(): | |
| import importlib.util | |
| import sys | |
| root_s = str(root.resolve()) | |
| if root_s not in sys.path: | |
| sys.path.insert(0, root_s) | |
| def load_mod(mod_name, filename): | |
| spec = importlib.util.spec_from_file_location(mod_name, root / filename) | |
| mod = importlib.util.module_from_spec(spec) | |
| sys.modules[mod_name] = mod | |
| spec.loader.exec_module(mod) | |
| return mod | |
| configuration = load_mod("configuration_edlm", "configuration_edlm.py") | |
| modeling = load_mod("modeling_edlm", "modeling_edlm.py") | |
| cfg = configuration.EfficientDLMConfig.from_pretrained(str(root), revision=revision) | |
| return modeling.EfficientDLM.from_pretrained( | |
| str(root), config=cfg, revision=revision, dtype=dtype, attn_implementation=attn | |
| ).model | |
| cfg = AutoConfig.from_pretrained(name, revision=revision, trust_remote_code=True) | |
| if getattr(cfg, "model_type", None) == "edlm" or (getattr(cfg, "architectures", [""]) or [""])[0] == "EfficientDLM": | |
| return AutoModelForCausalLM.from_pretrained( | |
| name, revision=revision, dtype=dtype, attn_implementation=attn, trust_remote_code=True | |
| ).model | |
| return AutoModelForCausalLM.from_pretrained(name, revision=revision, dtype=dtype, attn_implementation=attn).model | |
| def pad_id(tok): | |
| """The id used to right-pad token rows (never attended to); Qwen tokenizers define one, others fall back to 0.""" | |
| return tok.pad_token_id if tok.pad_token_id is not None else 0 | |
| def is_hybrid(config): | |
| """Whether a (text) config has Gated DeltaNet layers (Qwen3.5). Such backbones cannot honour the block-causal mask and | |
| run the row form; on Apple Silicon they are what the MLX backend is for.""" | |
| return "linear_attention" in set(getattr(config, "layer_types", None) or []) | |
| _SPECIAL_RE = re.compile(r"<\|([A-Za-z0-9_]+)\|>") | |
| def user_tokens(tok, text): | |
| """Tokenize caller-supplied text so it can never produce delimiter/control tokens (option boundaries are unforgeable). | |
| The fast tokenizer ignores split_special_tokens, so `<|name|>` is rewritten to `<¦name¦>` before tokenizing.""" | |
| return tok(_SPECIAL_RE.sub(r"<¦\1¦>", text), add_special_tokens=False).input_ids | |
| OPT_NONE, OPT_DECIDE = -1, -2 # values of enc["opt"]: instruction/state tokens, and the <decide> token | |
| def encode(tok, rec, max_state=MAX_STATE, max_branch=MAX_BRANCH, strict=False, option_isolation=False): | |
| """Pack one record: [<state> ...] then per-question [<q> instr <opt> o </opt>... <decide>]. | |
| Returns ids, seg (0 = state, k = question k), pos (branch positions restart after state), | |
| decide_idx [Q], opt_idx [Q][K] (index of </opt> token for each option), opt (per-token option index within its | |
| question: OPT_NONE for state/instruction, 0..K-1 for option spans, OPT_DECIDE for <decide>). | |
| option_isolation=True: every option span is its own sub-branch (it sees state + instruction + itself only), all | |
| option spans share the same position ids, and <decide> sits at one fixed position after the longest span. Then the | |
| per-option representations and <decide>'s attention over them are permutation-invariant by construction. | |
| """ | |
| state_tokens = user_tokens(tok, rec["state"]) | |
| if strict and len(state_tokens) + 1 > max_state: | |
| raise ContextOverflow(f"state exceeds {max_state} tokens: {len(state_tokens) + 1}") | |
| S = [tok.convert_tokens_to_ids(SPECIAL[0])] + state_tokens[: max_state - 1] | |
| ids, seg, pos, opt = list(S), [0] * len(S), list(range(len(S))), [OPT_NONE] * len(S) | |
| q_id, o_id, c_id, d_id = (tok.convert_tokens_to_ids(t) for t in SPECIAL[1:]) | |
| decide_idx, opt_idx = [], [] | |
| for k, q in enumerate(rec["questions"], start=1): | |
| instr = [q_id] + user_tokens(tok, q["instr"]) | |
| spans = [[o_id] + user_tokens(tok, o) + [c_id] for o in q["options"]] | |
| br = instr + [t for sp in spans for t in sp] + [d_id] | |
| if len(br) > max_branch - len(S): | |
| raise ContextOverflow(f"branch too long: {len(br)} tokens with a {len(S)}-token state (row limit {max_branch})") | |
| base = len(ids); p0 = len(S) | |
| br_opt = [OPT_NONE] * len(instr) + [j for j, sp in enumerate(spans) for _ in sp] + [OPT_DECIDE] | |
| if option_isolation: | |
| longest = max(len(sp) for sp in spans) | |
| br_pos = list(range(p0, p0 + len(instr))) + [p0 + len(instr) + i for sp in spans for i in range(len(sp))] + [p0 + len(instr) + longest] | |
| else: | |
| br_pos = list(range(p0, p0 + len(br))) | |
| ends, cursor = [], len(instr) | |
| for sp in spans: | |
| cursor += len(sp); ends.append(cursor - 1) | |
| ids += br; seg += [k] * len(br); pos += br_pos; opt += br_opt | |
| decide_idx.append(base + len(br) - 1); opt_idx.append([base + e for e in ends]) | |
| return {"ids": ids, "seg": seg, "pos": pos, "opt": opt, "option_isolation": option_isolation, "decide_idx": decide_idx, "opt_idx": opt_idx, | |
| "labels": [q["label"] for q in rec["questions"]], "state_truncated": len(state_tokens) + 1 > max_state} | |
| def fits(rec, *tokenizers, max_state=MAX_STATE, max_branch=MAX_BRANCH, max_packed=MAX_PACKED): | |
| """True when the internal record encodes strictly (no truncation) within the training context under every tokenizer | |
| given (frozen suites are admitted against the tokenizers of all their pinned bases).""" | |
| try: | |
| return all(len(encode(tok, rec, max_state=max_state, max_branch=max_branch, strict=True)["ids"]) <= max_packed for tok in tokenizers) | |
| except ValueError: | |
| return False | |
| def branch_mask(seg, device, dtype=torch.float32): | |
| """attend(i,j) iff j<=i and (seg[j]==0 or seg[j]==seg[i]). Returns additive [1,1,L,L].""" | |
| return branch_mask_batch([seg], device, dtype) | |
| def branch_mask_batch(segs, device, dtype=torch.float32, opts=None, length=None, state_bidir=False): | |
| """Batched block-causal mask, additive [B,1,L,L], right-padded to the longest sequence. | |
| Padded key positions are masked for every query; padded query rows keep the diagonal so no row is fully | |
| masked (finfo.min, not -inf, so softmax stays finite either way). Real tokens never see pads because pads sit | |
| after them (causal) and belong to no segment (-1). | |
| opts (option isolation): within a question, an option-span token may attend to state, the instruction, and its own | |
| span only; <decide> attends to everything in its question. Instruction tokens never see option spans (causal). | |
| state_bidir: state tokens attend to every state token (both directions), for backbones trained bidirectionally | |
| (diffusion LMs). The state still never sees the branches, so the state prefix cache stays exact.""" | |
| L = max(max(len(s) for s in segs), length or 0) | |
| s = torch.full((len(segs), L), -1, device=device) | |
| for b, seg in enumerate(segs): | |
| s[b, : len(seg)] = torch.tensor(seg, device=device) | |
| causal = torch.tril(torch.ones(L, L, dtype=torch.bool, device=device)) | |
| same = (s[:, None, :] == s[:, :, None]) | (s[:, None, :] == 0) | |
| valid_key = (s != -1)[:, None, :] | |
| allow = causal[None] & same & valid_key | |
| if state_bidir: | |
| allow = allow | ((s[:, :, None] == 0) & (s[:, None, :] == 0)) | |
| if opts is not None: | |
| o = torch.full((len(segs), L), OPT_NONE, device=device) | |
| for b, op in enumerate(opts): | |
| o[b, : len(op)] = torch.tensor(op, device=device) | |
| key_is_option = (o[:, None, :] >= 0) | |
| query_is_decide = (o[:, :, None] == OPT_DECIDE) | |
| same_option = o[:, None, :] == o[:, :, None] | |
| allow = allow & (~key_is_option | query_is_decide | same_option) | |
| allow = allow | torch.eye(L, dtype=torch.bool, device=device)[None] | |
| return torch.zeros(len(segs), L, L, dtype=dtype, device=device).masked_fill(~allow, torch.finfo(dtype).min)[:, None] | |
| def rows_of(enc): | |
| """Split a packed encoding into its state and per-question branch rows. | |
| Returns (state_ids, state_pos, rows) with rows[k] = {"ids", "pos", "decide", "opts"}: the branch tokens of question | |
| k with their (already state-continuing) positions, and the readout offsets *within the branch*. Feeding | |
| state + rows[k] as one causal row is equivalent to the packed block-causal form for that question, on any | |
| architecture: the row contains exactly the tokens question k may attend to, in the same positions.""" | |
| seg = enc["seg"]; Ls = seg.count(0) | |
| rows, start = [], Ls | |
| for k, (d, oi) in enumerate(zip(enc["decide_idx"], enc["opt_idx"]), start=1): | |
| end = d + 1 # <decide> is the last token of its branch | |
| if seg[start] != k or seg[end - 1] != k: raise ValueError("branch layout mismatch") | |
| rows.append({"ids": enc["ids"][start:end], "pos": enc["pos"][start:end], "decide": d - start, "opts": [o - start for o in oi]}) | |
| start = end | |
| return enc["ids"][:Ls], enc["pos"][:Ls], rows | |
| class PointerHead(nn.Module): | |
| def __init__(self, d, dp=256): | |
| """dp = pointer dimension (head capacity knob).""" | |
| super().__init__() | |
| self.q, self.k = nn.Linear(d, dp), nn.Linear(d, dp) | |
| self.scale = 1 / math.sqrt(dp) | |
| # calibration: logits are divided by this at inference (eval mode) only. 1.0 = raw. A checkpoint carries the value fitted on | |
| # its in-distribution development rows (scripts/calibrate_checkpoint.py -> head.pt["temperature"]); training always sees T=1 so | |
| # a fitted value stays meaningful, and the argmax is unchanged by construction. | |
| self.temperature = 1.0 | |
| def forward(self, h_decide, h_opts): # [d], [K,d] -> logits [K] | |
| z = (self.k(h_opts) @ self.q(h_decide)) * self.scale | |
| return z if self.training or self.temperature == 1.0 else z / self.temperature | |
| # What a loaded model exposes to kev.serve, kev.predictors and the Space: the scoring interface both DecisionModel (torch) | |
| # and kev.mlx_model.MLXDecisionModel implement. tests/test_mlx.py checks the MLX class against this list. | |
| SCORING_INTERFACE = ("encode", "forward", "probs", "probs_and_prefix", "probs_with_prefix", "eval", | |
| "head", "backend", "dtype", "device", "hybrid", "option_isolation", "prefix_min_tokens") | |
| class DecisionModel(nn.Module): | |
| def __init__(self, name, tok, device, lora=None, revision=None, attn=None, head_dim=256, option_isolation=False, special_embeddings=False, lora_targets="all", dtype=torch.float32, state_bidir=False): | |
| super().__init__() | |
| # backbone only (no vocab head): we never generate text. | |
| # eager on MPS/CPU (known-good with our float 4D mask); SDPA on CUDA (accepts arbitrary additive masks). | |
| attn = attn or ("sdpa" if str(device).startswith("cuda") else "eager") | |
| # dtype: fp32 for training and exact evaluation; bf16 is a serving option for large backbones (8B on a 32 GB Mac) | |
| self.lm = load_backbone(name, revision, dtype, attn) | |
| self.pad_id = pad_id(tok) | |
| # hybrid backbones (Qwen3.5: Gated DeltaNet layers, recurrent) cannot honour the block-causal mask, so every | |
| # question runs as its own causal row continuing from the state (rows_of). Attention-only backbones keep the | |
| # packed form; the two agree to fp32 noise (tests/test_model.py::test_rows_match_packed). | |
| self.hybrid = is_hybrid(self.lm.config) | |
| if self.hybrid and option_isolation: raise ValueError("option_isolation needs the packed mask; not available on hybrid backbones") | |
| if self.hybrid and state_bidir: raise ValueError("state_bidir needs attention masks; hybrid backbones ignore them") | |
| self.option_isolation = option_isolation | |
| self.state_bidir = state_bidir | |
| if lora: | |
| from peft import LoraConfig, get_peft_model | |
| extra = {"trainable_token_indices": {"embed_tokens": [tok.convert_tokens_to_ids(t) for t in SPECIAL]}} if special_embeddings else {} | |
| targets = {"all": ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"], | |
| "dense": ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"], # "all" minus the DeltaNet projections on hybrids (retention ablation) | |
| "attn": ["q_proj", "k_proj", "v_proj", "o_proj"], "qv": ["q_proj", "v_proj"]}[lora_targets] | |
| if self.hybrid and lora_targets in ("all", "attn"): | |
| # Gated DeltaNet projections (transformers 5 names, verified on Qwen3_5TextModel); the mixer's out_proj too | |
| targets = targets + ["in_proj_qkv", "in_proj_z", "in_proj_a", "in_proj_b", "out_proj"] | |
| cfg = LoraConfig(task_type="FEATURE_EXTRACTION", r=lora, lora_alpha=2 * lora, lora_dropout=0.05, target_modules=targets, **extra) | |
| self.lm = get_peft_model(self.lm, cfg) | |
| self.head = PointerHead(self.lm.config.hidden_size, dp=head_dim) | |
| self.device = device | |
| self.to(device) | |
| backend = "torch" # kev.mlx_model.MLXDecisionModel is the other implementation of this scoring interface | |
| def prefix_min_tokens(self): | |
| """kev.serve caches the state prefix from this many state tokens. Attention-only backbones: 384, below which the | |
| branch-only pass is not faster than one packed pass on MPS (per-op overhead). Hybrid backbones: always, because | |
| their miss path otherwise recomputes the state once per question (Kev-0.8B bf16 on MPS, 5 questions: 1011 -> 413 ms).""" | |
| return 0 if self.hybrid else 384 | |
| def dtype(self): | |
| return str(next(self.lm.parameters()).dtype).removeprefix("torch.") | |
| def encode(self, tok, rec, **kw): | |
| """encode() with this model's option-isolation setting; use this from serving/eval code.""" | |
| return encode(tok, rec, option_isolation=self.option_isolation, **kw) | |
| def hidden(self, enc): | |
| return self.hidden_batch([enc])[0, : len(enc["ids"])] | |
| SHAPE_BUCKET = int(os.environ.get("KEV_SHAPE_BUCKET", "64")) # MPS: pad the sequence to a multiple of this (per-shape kernel warm-up); 1 disables | |
| def _pad_rows(self, rows): | |
| """Right-pad (ids, pos) token rows into [N, L] id / position tensors and a [N, L] attention mask (1 = real token). | |
| Pads sit after every real token and are masked keys, so they never change a real token's hidden state (parity | |
| measured exact). On MPS in eval mode L is rounded up to a SHAPE_BUCKET multiple so kernels are warmed per bucket.""" | |
| L = max(len(ids) for ids, _ in rows) | |
| if str(self.device) == "mps" and not self.training: L = -(-L // self.SHAPE_BUCKET) * self.SHAPE_BUCKET | |
| ids = torch.full((len(rows), L), self.pad_id, device=self.device) | |
| pos = torch.zeros((len(rows), L), dtype=torch.long, device=self.device) | |
| att = torch.zeros((len(rows), L), dtype=torch.long, device=self.device) | |
| for i, (rid, rpos) in enumerate(rows): | |
| ids[i, : len(rid)] = torch.tensor(rid, device=self.device); pos[i, : len(rpos)] = torch.tensor(rpos, device=self.device); att[i, : len(rid)] = 1 | |
| return ids, pos, att | |
| def hidden_batch(self, encs): | |
| """[B, L_max, d] hidden states for a right-padded batch of encoded records under the packed block-causal mask.""" | |
| ids, pos, _ = self._pad_rows([(e["ids"], e["pos"]) for e in encs]) | |
| isolate = any(e.get("option_isolation") for e in encs) | |
| if isolate and not all(e.get("option_isolation") for e in encs): | |
| raise ValueError("cannot mix option-isolated and plain encodings in one batch") | |
| lm_dtype = next(self.lm.parameters()).dtype | |
| mask = branch_mask_batch([e["seg"] for e in encs], self.device, dtype=lm_dtype, opts=[e["opt"] for e in encs] if isolate else None, | |
| length=ids.shape[1], state_bidir=self.state_bidir) | |
| return self.lm(input_ids=ids, position_ids=pos, attention_mask=mask).last_hidden_state.float() # head stays fp32 | |
| def _readout(self, h, enc): | |
| return [self.head(h[d], h[torch.tensor(oi, device=self.device)]) for d, oi in zip(enc["decide_idx"], enc["opt_idx"])] | |
| def rows_form(self, encs): | |
| """Whether these records run as causal rows: hybrid backbones always (the recurrent layers cannot honour the packed | |
| mask), attention-only ones when the packed sequence would exceed one serving row (its L x L mask grows without | |
| bound with the number of questions). The two forms agree (tests/test_model.py::test_rows_match_packed).""" | |
| return self.hybrid or any(len(e["ids"]) > SERVE_MAX_PACKED for e in encs) | |
| def _row_mask(self, att, state_lens): | |
| """Additive [N,1,L,L] mask for self-contained rows (state + one branch) when the state reads bidirectionally.""" | |
| N, L = att.shape | |
| allow = torch.tril(torch.ones(L, L, dtype=torch.bool, device=self.device))[None].repeat(N, 1, 1) | |
| for i, sl in enumerate(state_lens): | |
| allow[i, :sl, :sl] = True | |
| allow = allow & att.bool()[:, None, :] | |
| allow = allow | torch.eye(L, dtype=torch.bool, device=self.device)[None] | |
| dt = next(self.lm.parameters()).dtype | |
| return torch.zeros(N, L, L, dtype=dt, device=self.device).masked_fill(~allow, torch.finfo(dt).min)[:, None] | |
| def _rows_hidden(self, rows, cache=None, prefix_len=0, state_lens=None): | |
| """Hidden states of causal token rows, one [L_i, d] tensor per row. In eval mode the rows go through the backbone | |
| rows_per_pass at a time; training keeps one batch (its batches are small and autograd needs the whole graph anyway). | |
| With `cache`, the rows are branches continuing the cached state: the cache is replicated once per chunk (a copy, | |
| so the caller's prefix stays pristine) and the cached tokens are marked real in the attention mask.""" | |
| chunk = len(rows) if self.training else rows_per_pass([ids for ids, _ in rows], prefix_len) | |
| out = [] | |
| for start in range(0, len(rows), chunk): | |
| part = rows[start:start + chunk] | |
| ids, pos, att = self._pad_rows(part) | |
| past = {} | |
| if cache is not None: | |
| replica = copy.deepcopy(cache); replica.reorder_cache(torch.zeros(len(part), dtype=torch.long, device=self.device)) | |
| att = torch.cat([torch.ones((len(part), prefix_len), dtype=torch.long, device=self.device), att], 1) | |
| past = {"past_key_values": replica, "use_cache": True} | |
| elif self.state_bidir and state_lens is not None: | |
| att = self._row_mask(att, state_lens[start:start + chunk]) | |
| h = self.lm(input_ids=ids, position_ids=pos, attention_mask=att, **past).last_hidden_state.float() | |
| out += [h[i, : len(row_ids)] for i, (row_ids, _) in enumerate(part)] | |
| return out | |
| def forward_rows_batch(self, encs): | |
| """Row form: every question of every record is one causal row = state tokens + its branch tokens. Returns the same | |
| nested logits as forward_batch. Exact isolation by construction (rows are independent); the state is recomputed | |
| per row (Q x state tokens), which training accepts; serving uses the prefix cache instead.""" | |
| rows, readouts, state_lens = [], [], [] # one causal row per question; readouts[i] = (record, <decide> offset, option offsets) | |
| for b, e in enumerate(encs): | |
| S, Sp, brs = rows_of(e) | |
| for r in brs: | |
| rows.append((S + r["ids"], Sp + r["pos"])); readouts.append((b, len(S) + r["decide"], [len(S) + o for o in r["opts"]])) | |
| state_lens.append(len(S)) | |
| out = [[] for _ in encs] | |
| for h, (b, d, oi) in zip(self._rows_hidden(rows, state_lens=state_lens), readouts): | |
| out[b].append(self.head(h[d], h[torch.tensor(oi, device=self.device)])) | |
| return out | |
| def forward(self, enc): | |
| """Returns list of logits tensors, one per question.""" | |
| return self.forward_batch([enc])[0] | |
| def forward_batch(self, encs): | |
| """List (per record) of lists (per question) of logits. Row form or packed block-causal mask, see rows_form.""" | |
| if self.rows_form(encs): return self.forward_rows_batch(encs) | |
| hs = self.hidden_batch(encs) | |
| return [self._readout(hs[b], e) for b, e in enumerate(encs)] | |
| def probs(self, enc): | |
| return [F.softmax(z, -1).cpu() for z in self.forward(enc)] | |
| # --- state-prefix reuse (serving): the state is encoded once, question branches attend to its cached keys/values. | |
| # Exact by construction: branch tokens never attend to each other across questions (block-causal mask) and the state | |
| # never sees the branches (causal), so the state's hidden states and KV are identical with or without the branches. | |
| def _branch_rows_from_prefix(self, enc, cache): | |
| """Row-form serving: the branches run as causal rows continuing the cached state (exactly the forward_rows_batch | |
| layout, minus the recomputed state).""" | |
| S, _, rows = rows_of(enc) | |
| hs = self._rows_hidden([(r["ids"], r["pos"]) for r in rows], cache=cache, prefix_len=len(S)) | |
| return [F.softmax(self.head(h[r["decide"]], h[torch.tensor(r["opts"], device=self.device)]), -1).cpu() for h, r in zip(hs, rows)] | |
| def prefix(self, enc): | |
| """Run the state tokens only. Returns (n_state_tokens, kv cache, state hidden states [Ls, d]).""" | |
| Ls = enc["seg"].count(0) | |
| ids = torch.tensor([enc["ids"][:Ls]], device=self.device); pos = torch.tensor([enc["pos"][:Ls]], device=self.device) | |
| # the cache must know the layer types (hybrid backbones keep recurrent + conv states per DeltaNet layer) | |
| extra = {} | |
| if self.state_bidir: | |
| extra["attention_mask"] = torch.zeros(1, 1, Ls, Ls, dtype=next(self.lm.parameters()).dtype, device=self.device) | |
| out = self.lm(input_ids=ids, position_ids=pos, past_key_values=DynamicCache(config=self.lm.config), use_cache=True, **extra) | |
| return Ls, out.past_key_values, out.last_hidden_state[0].float() | |
| def probs_and_prefix(self, enc): | |
| """One full pass that also returns the state prefix (KV cropped to the state, state hidden states): a cache miss | |
| costs a single forward pass, not two.""" | |
| Ls = enc["seg"].count(0) | |
| if self.rows_form([enc]): | |
| # recurrent layers cannot be cropped back to the state (and an over-long packed pass is what the row form avoids), | |
| # so a miss here is a state pass (kept as the prefix) plus the branch rows | |
| Ls, cache, h_state = self.prefix(enc) | |
| return self._branch_rows_from_prefix(enc, cache), (Ls, cache, h_state) | |
| ids = torch.tensor([enc["ids"]], device=self.device); pos = torch.tensor([enc["pos"]], device=self.device) | |
| dt = next(self.lm.parameters()).dtype | |
| mask = branch_mask_batch([enc["seg"]], self.device, dtype=dt, opts=[enc["opt"]] if enc.get("option_isolation") else None, state_bidir=self.state_bidir) | |
| out = self.lm(input_ids=ids, position_ids=pos, attention_mask=mask, past_key_values=DynamicCache(config=self.lm.config), use_cache=True) | |
| h = out.last_hidden_state[0].float() | |
| out.past_key_values.crop(-(len(enc["ids"]) - Ls)) # keep the state only (negative = drop that many trailing tokens; positive form deprecated in transformers 5) | |
| return [F.softmax(z, -1).cpu() for z in self._readout(h, enc)], (Ls, out.past_key_values, h[:Ls].clone()) | |
| def probs_with_prefix(self, enc, prefix): | |
| """probs() for a record whose state tokens equal the cached prefix's; only the branches run. The cache is cropped | |
| back to the state afterwards so it can be reused.""" | |
| Ls, cache, h_state = prefix | |
| if enc["seg"].count(0) != Ls: raise ValueError("prefix does not match this record's state") | |
| if self.rows_form([enc]): | |
| return self._branch_rows_from_prefix(enc, cache) | |
| ids = torch.tensor([enc["ids"][Ls:]], device=self.device); pos = torch.tensor([enc["pos"][Ls:]], device=self.device) | |
| dt = next(self.lm.parameters()).dtype | |
| mask = branch_mask_batch([enc["seg"]], self.device, dtype=dt, opts=[enc["opt"]] if enc.get("option_isolation") else None, | |
| state_bidir=self.state_bidir)[:, :, Ls:, :] | |
| try: | |
| out = self.lm(input_ids=ids, position_ids=pos, past_key_values=cache, attention_mask=mask, use_cache=True) | |
| h = torch.cat([h_state, out.last_hidden_state[0].float()], 0) | |
| finally: | |
| cache.crop(-(len(enc["ids"]) - Ls)) | |
| return [F.softmax(z, -1).cpu() for z in self._readout(h, enc)] | |
| def trainable_parameters(self): | |
| return [p for p in self.parameters() if p.requires_grad] | |