Zero-Shot Classification
Transformers
Safetensors
GGUF
English
decision-model
system-one
falcondec
lightdec
calibrated-decisions
multiple-choice
intent-classification
customer-support
natural-language-inference
code
guardrails
agents
selective-prediction
falconsai
model-surgeon
attested-lineage
Instructions to use Falconsai/LightDec with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Falconsai/LightDec with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("zero-shot-classification", model="Falconsai/LightDec")# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("Falconsai/LightDec", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download falcondec_modeling.py from Falconsai/LightDec: direct link, hf CLI and curl.
- Browser
- Download file 15.6 kB
-
https://huggingface.co/Falconsai/LightDec/resolve/main/falcondec_modeling.py
- Command line
-
hf download hf://Falconsai/LightDec/falcondec_modeling.py
-
curl -L -o falcondec_modeling.py https://huggingface.co/Falconsai/LightDec/resolve/main/falcondec_modeling.py
15.6 kB
| # -*- 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)) | |
| 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) | |
| 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] | |
| 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}} | |