LightDec / falcondec_modeling.py
RealFalconsAI's picture
Upload 14 files
3e2212d verified
Raw History Blame Contribute Delete
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))
@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}}