# -*- coding: utf-8 -*- """FalconDec — Falcon Decision model: single-pass, typed, calibrated closed-set decisions. Layout : [CLS] question [SEP] [MASK] opt_1 ... [MASK] opt_k [SEP] state [SEP] Head : marker vectors + CLS context + question-type embedding -> permutation-equivariant set transformer (options attend to each other; no positions) -> MLP -> one logit per option -> softmax over this question's options Calib. : temperature per (question type, option-count bucket), stored in the checkpoint """ from __future__ import annotations import contextlib import json import shutil from pathlib import Path import numpy as np import torch import torch.nn as nn QTYPES = {"choice": 0, "noul": 1, "score": 2} N_BUCKETS = 4 CONFIG_FILE = "falcondec_config.json" FP16_FILE = "model.safetensors" INT8_FILE = "model_int8.safetensors" MODEL_KEYS = ("input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype") def n_bucket(n: int) -> int: return 0 if n <= 2 else 1 if n <= 5 else 2 if n <= 12 else 3 def special_ids(tok) -> dict: sp = {"cls": tok.cls_token_id, "sep": tok.sep_token_id, "mask": tok.mask_token_id, "pad": tok.pad_token_id} if sp["cls"] is None: sp["cls"] = tok.bos_token_id if sp["sep"] is None: sp["sep"] = tok.eos_token_id if sp["pad"] is None: sp["pad"] = sp["sep"] if sp["mask"] is None: raise ValueError("FalconDec needs a tokenizer with a mask token (used as the option marker).") return {k: int(v) for k, v in sp.items()} def _as_list(x): return x.tolist() if hasattr(x, "tolist") else list(x) def assemble(q_ids, opt_ids, s_ids, sp, max_len=512, head_max_len=192, max_tok_per_opt=24, long_max_len=2048, long_opts_threshold=24): """Build one input sequence. Returns (ids, marker_positions).""" n = len(opt_ids) eff = max_len if n <= long_opts_threshold else max(max_len, long_max_len) q = _as_list(q_ids)[:96] need = len(q) + 2 + n * (max_tok_per_opt + 1) head = min(eff - 64, max(head_max_len, need)) per = max(2, min(max_tok_per_opt, (head - len(q) - 2) // max(n, 1) - 1)) ids = [sp["cls"]] + q + [sp["sep"]] markers = [] for o in opt_ids: markers.append(len(ids)) ids.append(sp["mask"]) ids.extend(_as_list(o[:per])) ids.append(sp["sep"]) room = eff - len(ids) - 1 if room > 0 and s_ids is not None and len(s_ids) > 0: ids.extend(_as_list(s_ids[:room])) ids.append(sp["sep"]) if len(ids) > eff: ids = ids[:eff] if markers and markers[-1] >= len(ids): raise ValueError("Too many options for one pass; reduce options or raise long_max_len.") return ids, markers def collate_features(feats, pad_id, device=None): """feats: list of (ids, markers, qtype_index) -> dict of padded tensors matching FalconDec.forward.""" B = len(feats) T = max(len(f[0]) for f in feats) K = max(len(f[1]) for f in feats) ids = torch.full((B, T), pad_id, dtype=torch.long) att = torch.zeros((B, T), dtype=torch.long) mpos = torch.zeros((B, K), dtype=torch.long) mmask = torch.zeros((B, K), dtype=torch.bool) qt = torch.zeros(B, dtype=torch.long) for i, (x, mk, q) in enumerate(feats): ids[i, : len(x)] = torch.as_tensor(x, dtype=torch.long) att[i, : len(x)] = 1 mpos[i, : len(mk)] = torch.as_tensor(mk, dtype=torch.long) mmask[i, : len(mk)] = True qt[i] = int(q) out = dict(input_ids=ids, attention_mask=att, marker_pos=mpos, marker_mask=mmask, qtype=qt) if device is not None: out = {k: v.to(device, non_blocking=True) for k, v in out.items()} return out class OptionInteraction(nn.Module): """Set transformer over the options of one question (no positional encoding => order-equivariant).""" def __init__(self, d, n_layers=2, n_heads=8, dropout=0.1): super().__init__() layer = nn.TransformerEncoderLayer(d, n_heads, dim_feedforward=2 * d, dropout=dropout, activation="gelu", batch_first=True, norm_first=True) self.enc = nn.TransformerEncoder(layer, n_layers, enable_nested_tensor=False) def forward(self, x, mask): return self.enc(x, src_key_padding_mask=~mask) class FalconDec(nn.Module): def __init__(self, encoder, fcfg: dict): super().__init__() self.encoder = encoder self.fcfg = dict(fcfg) d = encoder.config.hidden_size self.qtype_emb = nn.Embedding(len(QTYPES), d) self.ctx_proj = nn.Linear(d, d) self.opt_norm = nn.LayerNorm(d) self.interact = OptionInteraction(d, int(self.fcfg.get("interact_layers", 2)), int(self.fcfg.get("interact_heads", 8))) self.scorer = nn.Sequential(nn.Linear(d, d), nn.GELU(), nn.Dropout(0.1), nn.Linear(d, 1)) self.register_buffer("temperature", torch.ones(len(QTYPES), N_BUCKETS)) @property def device(self): return next(self.parameters()).device def num_parameters(self): return sum(p.numel() for p in self.parameters()) def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype): h = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state B, K = marker_pos.shape idx = marker_pos.unsqueeze(-1).expand(B, K, h.size(-1)) opt = h.gather(1, idx) ctx = self.ctx_proj(h[:, 0]).unsqueeze(1) x = self.opt_norm(opt + ctx + self.qtype_emb(qtype).unsqueeze(1)) x = self.interact(x, marker_mask) logits = self.scorer(x).squeeze(-1).float() return logits.masked_fill(~marker_mask, -1e4) # ------------------------------------------------------------------ storage def clean_state_dict(sd): return {k.replace("_orig_mod.", ""): v.detach().cpu().contiguous() for k, v in sd.items()} def quantize_int8(sd, min_numel=4096): """Per-output-channel symmetric int8 for every matrix; fp16 for everything else.""" out = {} for k, v in sd.items(): if v.is_floating_point() and v.ndim == 2 and v.numel() >= min_numel: w = v.float() s = (w.abs().amax(dim=1, keepdim=True) / 127.0).clamp_min(1e-12) out[k] = torch.round(w / s).clamp_(-127, 127).to(torch.int8).contiguous() out[k + "::scale"] = s.contiguous() elif v.is_floating_point(): out[k] = v.half().contiguous() else: out[k] = v.contiguous() return out def dequantize_int8(sd): out = {} for k, v in sd.items(): if k.endswith("::scale"): continue s = sd.get(k + "::scale") out[k] = (v.float() * s).half() if s is not None else v return out def save_falcondec(model, tok, out_dir, int8=False, extra_files=None): from safetensors.torch import save_file out = Path(out_dir) out.mkdir(parents=True, exist_ok=True) enc_cfg = model.encoder.config if hasattr(enc_cfg, "reference_compile"): enc_cfg.reference_compile = False enc_cfg.save_pretrained(str(out / "encoder")) tok.save_pretrained(str(out / "tokenizer")) sd = clean_state_dict(model.state_dict()) fc = dict(model.fcfg) fc["temperature"] = model.temperature.detach().float().cpu().tolist() fc["weights"] = INT8_FILE if int8 else FP16_FILE if int8: save_file(quantize_int8(sd), str(out / INT8_FILE), metadata={"format": "falcondec-int8"}) else: save_file({k: (v.half() if v.is_floating_point() else v) for k, v in sd.items()}, str(out / FP16_FILE), metadata={"format": "falcondec-fp16"}) (out / CONFIG_FILE).write_text(json.dumps(fc, indent=2, default=str), encoding="utf-8") try: shutil.copy(__file__, out / "falcondec_modeling.py") except Exception: pass for name, content in (extra_files or {}).items(): (out / name).write_text(content, encoding="utf-8") return out def load_falcondec(path, device=None, dtype=None, attn_implementation="sdpa"): """Load a FalconDec directory (fp16 or int8) or Hub repo. Returns (model, tokenizer).""" from safetensors.torch import load_file from transformers import AutoConfig, AutoModel, AutoTokenizer p = Path(path) if not p.exists(): from huggingface_hub import snapshot_download p = Path(snapshot_download(str(path))) fc = json.loads((p / CONFIG_FILE).read_text(encoding="utf-8")) ecfg = AutoConfig.from_pretrained(str(p / "encoder")) if hasattr(ecfg, "reference_compile"): ecfg.reference_compile = False try: enc = AutoModel.from_config(ecfg, attn_implementation=attn_implementation) except Exception: enc = AutoModel.from_config(ecfg) model = FalconDec(enc, fc) wf = p / fc.get("weights", FP16_FILE) if not wf.exists(): wf = p / (INT8_FILE if (p / INT8_FILE).exists() else FP16_FILE) sd = load_file(str(wf)) if wf.name == INT8_FILE: sd = dequantize_int8(sd) missing, unexpected = model.load_state_dict(sd, strict=False) if missing or unexpected: print(f"[FalconDec] load warning: missing={list(missing)[:5]} unexpected={list(unexpected)[:5]}") tok = AutoTokenizer.from_pretrained(str(p / "tokenizer")) dev = torch.device(device) if device is not None else torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(dev) if dtype is not None: model.to(dtype) model.eval() return model, tok # ------------------------------------------------------------------ inference def _amp(model): if model.device.type == "cuda" and next(model.parameters()).dtype == torch.float32: dt = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 return torch.autocast("cuda", dtype=dt) return contextlib.nullcontext() def _state_text(state): if state is None: return "" return state if isinstance(state, str) else json.dumps(state, ensure_ascii=False) @torch.no_grad() def score_items(model, tok, items, batch_size=32): """items: [{"state", "question", "options", "type"?, "option_tokens"?, "seq_len"?}] -> list of prob arrays.""" fc = model.fcfg M = int(fc.get("max_opts_single_pass", 96)) out = [None] * len(items) small = [i for i, it in enumerate(items) if len(it["options"]) <= M] for i, it in enumerate(items): if len(it["options"]) > M: out[i] = _score_large(model, tok, it, batch_size) if not small: return out sp = fc["special"] prepared = {} for i in small: it = items[i] if len(it["options"]) < 2: raise ValueError("Every question needs at least two options.") budget = int(it.get("option_tokens", fc["max_tok_per_opt"])) q_ids = tok(it.get("question", ""), add_special_tokens=False, truncation=True, max_length=96)["input_ids"] o_ids = tok([str(o) for o in it["options"]], add_special_tokens=False, truncation=True, max_length=budget)["input_ids"] s = _state_text(it.get("state")) s_ids = tok(s, add_special_tokens=False, truncation=True, max_length=int(fc["long_max_len"]))["input_ids"] if s else [] ids, mk = assemble(q_ids, o_ids, s_ids, sp, int(it.get("seq_len", fc["max_len"])), fc["head_max_len"], budget, fc["long_max_len"], fc["long_opts_threshold"]) prepared[i] = (ids, mk, QTYPES[it.get("type", "choice")]) order = sorted(small, key=lambda i: len(prepared[i][0])) model.eval() for b0 in range(0, len(order), batch_size): idxs = order[b0: b0 + batch_size] batch = collate_features([prepared[i] for i in idxs], sp["pad"], model.device) with _amp(model): logits = model(**batch) for j, i in enumerate(idxs): n = len(prepared[i][1]) T = model.temperature[prepared[i][2], n_bucket(n)].float().clamp_min(1e-3) out[i] = torch.softmax(logits[j, :n].float() / T, -1).cpu().numpy() return out def _score_large(model, tok, it, batch_size): """Tournament for > max_opts_single_pass options: keep the best of each chunk, then one final pass.""" M = int(model.fcfg.get("max_opts_single_pass", 96)) opts = list(it["options"]) cand = list(range(len(opts))) while len(cand) > M: groups = [cand[i: i + M] for i in range(0, len(cand), M)] keep = max(1, M // len(groups)) ps = score_items(model, tok, [dict(it, options=[opts[c] for c in g]) for g in groups], batch_size) cand = [g[t] for g, p in zip(groups, ps) for t in np.argsort(-p)[:keep]] p = score_items(model, tok, [dict(it, options=[opts[c] for c in cand])], batch_size)[0] full = np.zeros(len(opts), dtype=np.float32) full[cand] = p return full def _normalize_question(q): qtype = q.get("type", "choice") text = q.get("question") or q.get("instructions") or "" if qtype == "noul": lab = q.get("labels") or {} return qtype, text, [True, False], [str(lab.get("true", "Yes")), str(lab.get("false", "No"))] crit = q.get("criteria", q.get("options")) if qtype == "score": opts = [str(c) for c in crit] return qtype, text, list(range(len(opts))), opts if isinstance(crit, dict): return qtype, text, list(crit), [f"{k}: {v}" if v else str(k) for k, v in crit.items()] return qtype, text, list(crit), [str(c) for c in crit] @torch.no_grad() def decide(model, tok, state, questions, defer_threshold=None, batch_size=32): """Answer typed questions about one state. questions: list of dicts, or a Jev/Laya-style dict {key: question}. Each question: {"type": "choice"|"noul"|"score", "question"/"instructions": str, "options": [...] or "criteria": {key: description} / [levels], "option_tokens"?: int} """ if isinstance(questions, dict): questions = [dict(q, key=k) for k, q in questions.items()] items, metas = [], [] for q in questions: qtype, text, keys, opts = _normalize_question(q) item = {"state": state, "question": text, "options": opts, "type": qtype} for extra in ("option_tokens", "seq_len"): if extra in q: item[extra] = q[extra] items.append(item) metas.append((q, qtype, keys, opts)) probs = score_items(model, tok, items, batch_size) thr = model.fcfg.get("defer_threshold", 0.0) if defer_threshold is None else defer_threshold results = [] for (q, qtype, keys, opts), p in zip(metas, probs): i = int(np.argmax(p)) r = {"key": q.get("key"), "type": qtype, "question": items[len(results)]["question"], "choice": keys[i], "choice_text": opts[i], "confidence": float(p[i]), "probs": {str(k): float(v) for k, v in zip(keys, p)}, "defer": float(p[i]) < thr} if qtype == "noul": r["p_true"] = float(p[0]) if qtype == "score": r["expected_level"] = float(np.dot(p, np.arange(len(p)))) results.append(r) return {"results": results, "answers": {r["key"]: r for r in results if r["key"] is not None}}