Instructions to use Duke-CEI-SVD/traj-mc with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Duke-CEI-SVD/traj-mc with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("Duke-CEI-SVD/traj-mc", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download code/eval/eval.py from Duke-CEI-SVD/traj-mc: direct link, hf CLI and curl.
- Browser
- Download file 21.5 kB
-
https://huggingface.co/Duke-CEI-SVD/traj-mc/resolve/main/code/eval/eval.py
- Command line
-
hf download hf://Duke-CEI-SVD/traj-mc/code/eval/eval.py
-
curl -L -o eval.py https://huggingface.co/Duke-CEI-SVD/traj-mc/resolve/main/code/eval/eval.py
21.5 kB
| """ | |
| [SUPERSEDED — diagnostic only] The authoritative evaluation engine is now the | |
| official lm-eval harness: eval/llada_harness.py driven by eval/run_lmeval.py | |
| (README section 0). This self-contained file is kept only as a fast smoke tool; | |
| it does NOT implement the official per-task fewshot/cfg/mc_num, so its numbers | |
| are NOT reported. Use run_lmeval.py + analysis/lmeval_to_items.py for all | |
| REF/BASE/OURS results. | |
| eval.py -- ONE shared evaluation harness for REF / BASE / OURS. | |
| Same code, same protocol, same seed for every arm. Arm identity = which weights | |
| dir is loaded (REF = none / dense). Per-item results are dumped as JSONL so the | |
| paired McNemar test (analysis/mcnemar.py) can read them directly. | |
| PROTOCOL (aligned with the two reference papers): | |
| - SVD-LLM (ICLR 2025) evaluates MCQ with LM-Evaluation-Harness defaults, i.e. | |
| CONDITIONAL LOG-LIKELIHOOD over the choice -- not single-letter generation. | |
| - Sink-Aware Pruning (arXiv 2602.17664) LLaDA command: | |
| eval_llada.py --num_fewshot 0 --model llada_dist \ | |
| --model_args cfg=0.5,is_check_greedy=False,mc_num=128 | |
| -> we adopt num_fewshot=0, cfg=0.5, mc_num=128, is_check_greedy=False. | |
| So EVERY multiple-choice benchmark (mmlu, arc_c, arc_e, piqa, winogrande, | |
| hellaswag) goes through the SAME conditional-likelihood estimator with | |
| mc_num=128 and cfg=0.5. There is no single-token/mc_num=1 special case and no | |
| single-token assert: that exactness argument only held under the letter- | |
| GENERATION framing, which we no longer use. | |
| gsm8k : generation, gen_length=256 steps=256, 0-shot (strided sharding) | |
| Both acc (total log-likelihood) and acc_norm (length-normalized) are recorded | |
| per item, since the harness reports both and different papers cite different ones. | |
| Per-item record: {item_id, prompt_hash, correct, correct_norm, pred, pred_norm, gold}. | |
| is_check_greedy=False: we never run the greedy-decoding consistency check, so | |
| this is structurally satisfied (documented here for the protocol record). | |
| """ | |
| import os | |
| import re | |
| import sys | |
| import json | |
| import argparse | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) | |
| import common as C | |
| LETTERS = ["A", "B", "C", "D", "E", "F", "G", "H"] | |
| MAX_CTX = 1900 | |
| # ══ faithful LLaDA generation (from generate.py) ══════════════════════════════ | |
| def add_gumbel_noise(logits, temperature): | |
| if temperature == 0: | |
| return logits | |
| logits = logits.to(torch.float64) | |
| noise = torch.rand_like(logits, dtype=torch.float64) | |
| return logits.exp() / ((-torch.log(noise)) ** temperature) | |
| def get_num_transfer_tokens(mask_index, steps): | |
| mask_num = mask_index.sum(dim=1, keepdim=True) | |
| base = mask_num // steps | |
| remainder = mask_num % steps | |
| out = torch.zeros(mask_num.size(0), steps, device=mask_index.device, | |
| dtype=torch.int64) + base | |
| for i in range(mask_num.size(0)): | |
| out[i, : remainder[i]] += 1 | |
| return out | |
| def llada_generate(model, prompt, attention_mask=None, steps=256, gen_length=256, | |
| block_length=8, temperature=0.0, cfg_scale=0.0, mask_id=C.MASK_ID): | |
| x = torch.full((prompt.shape[0], prompt.shape[1] + gen_length), mask_id, | |
| dtype=torch.long, device=model.device) | |
| x[:, : prompt.shape[1]] = prompt.clone() | |
| prompt_index = x != mask_id | |
| if attention_mask is not None: | |
| attention_mask = torch.cat( | |
| [attention_mask, torch.ones((prompt.shape[0], gen_length), | |
| dtype=attention_mask.dtype, device=model.device)], | |
| dim=-1) | |
| assert gen_length % block_length == 0 | |
| num_blocks = gen_length // block_length | |
| assert steps % num_blocks == 0 | |
| steps_pb = steps // num_blocks | |
| for nb in range(num_blocks): | |
| blk = (x[:, prompt.shape[1] + nb * block_length: | |
| prompt.shape[1] + (nb + 1) * block_length] == mask_id) | |
| ntt = get_num_transfer_tokens(blk, steps_pb) | |
| for i in range(steps_pb): | |
| mask_index = x == mask_id | |
| if cfg_scale > 0.0: | |
| un_x = x.clone() | |
| un_x[prompt_index] = mask_id | |
| x_ = torch.cat([x, un_x], dim=0) | |
| am_ = torch.cat([attention_mask, attention_mask], dim=0) \ | |
| if attention_mask is not None else None | |
| logits = model(x_, attention_mask=am_).logits | |
| logits, un_logits = torch.chunk(logits, 2, dim=0) | |
| logits = un_logits + (cfg_scale + 1) * (logits - un_logits) | |
| else: | |
| logits = model(x, attention_mask=attention_mask).logits | |
| x0 = torch.argmax(add_gumbel_noise(logits, temperature), dim=-1) | |
| p = F.softmax(logits.to(torch.float64), dim=-1) | |
| x0_p = torch.gather(p, -1, x0.unsqueeze(-1)).squeeze(-1) | |
| x0_p[:, prompt.shape[1] + (nb + 1) * block_length:] = -np.inf | |
| x0 = torch.where(mask_index, x0, x) | |
| conf = torch.where(mask_index, x0_p, | |
| torch.tensor(-np.inf, device=x.device, dtype=x0_p.dtype)) | |
| transfer = torch.zeros_like(x0, dtype=torch.bool) | |
| for j in range(conf.shape[0]): | |
| _, sel = torch.topk(conf[j], k=int(ntt[j, i])) | |
| transfer[j, sel] = True | |
| x[transfer] = x0[transfer] | |
| return x | |
| # ══ conditional likelihood with CFG (from get_log_likelihood.py) ══════════════ | |
| def _forward_process(batch, prompt_index, mask_id): | |
| b, l = batch.shape | |
| target_len = (l - prompt_index.sum()).item() | |
| k = torch.randint(1, target_len + 1, (), device=batch.device) | |
| x = torch.round(torch.linspace(float(k), k + (b - 1) * (target_len / b), | |
| steps=b, device=batch.device)).long() | |
| x = ((x - 1) % target_len) + 1 | |
| indices = torch.arange(target_len, device=batch.device).repeat(b, 1) | |
| is_mask = indices < x.unsqueeze(1) | |
| for i in range(b): | |
| is_mask[i] = is_mask[i][torch.randperm(target_len, device=batch.device)] | |
| is_mask = torch.cat( | |
| (torch.zeros(b, prompt_index.sum(), dtype=torch.bool, device=batch.device), | |
| is_mask), dim=1) | |
| noisy = torch.where(is_mask, mask_id, batch) | |
| return noisy, (x / target_len).unsqueeze(1).repeat(1, l) | |
| def _get_logits_cfg(model, batch, prompt_index, cfg_scale, mask_id): | |
| """Unsupervised classifier-free guidance (cfg=0.5 per Sink-Aware protocol).""" | |
| if cfg_scale > 0.0: | |
| pi = prompt_index.unsqueeze(0).repeat(batch.shape[0], 1) | |
| un_batch = batch.clone() | |
| un_batch[pi] = mask_id | |
| batch = torch.cat([batch, un_batch], dim=0) | |
| logits = model(batch).logits | |
| if cfg_scale > 0.0: | |
| logits, un_logits = torch.chunk(logits, 2, dim=0) | |
| logits = un_logits + (cfg_scale + 1) * (logits - un_logits) | |
| return logits | |
| def conditional_loglik(model, prompt, answer, mc_num=128, batch_size=16, | |
| cfg_scale=0.5, mask_id=C.MASK_ID): | |
| """Returns (total_loglik, per_token_loglik). Monte-Carlo over mask patterns.""" | |
| seq = torch.cat([prompt, answer])[None, :].repeat(batch_size, 1).to(model.device) | |
| prompt_index = torch.arange(seq.shape[1], device=model.device) < len(prompt) | |
| losses = [] | |
| for _ in range(max(1, mc_num // batch_size)): | |
| perturbed, p_mask = _forward_process(seq, prompt_index, mask_id) | |
| mask_index = perturbed == mask_id | |
| logits = _get_logits_cfg(model, perturbed, prompt_index, cfg_scale, mask_id) | |
| loss = F.cross_entropy(logits[mask_index].to(torch.float32), seq[mask_index], | |
| reduction="none") / p_mask[mask_index] | |
| losses.append((loss.sum() / batch_size).item()) | |
| total = -(sum(losses) / len(losses)) | |
| return total, total / max(1, len(answer)) | |
| def score_choices(model, tok, context, continuations, mc_num, batch_size, cfg): | |
| """Score each continuation; return (pred_acc, pred_accnorm, totals, norms).""" | |
| pids = tok(context, add_special_tokens=False, | |
| return_tensors="pt")["input_ids"][0][-MAX_CTX:].to(model.device) | |
| totals, norms = [], [] | |
| for cont in continuations: | |
| a = tok(cont, add_special_tokens=False, | |
| return_tensors="pt")["input_ids"][0].to(model.device) | |
| if a.numel() == 0: # degenerate empty continuation | |
| totals.append(-1e30) | |
| norms.append(-1e30) | |
| continue | |
| t, n = conditional_loglik(model, pids, a, mc_num, batch_size, cfg) | |
| totals.append(t) | |
| norms.append(n) | |
| return int(np.argmax(totals)), int(np.argmax(norms)), totals, norms | |
| # ══ arm weight loading (head guard) ═══════════════════════════════════════════ | |
| def replace_with_lowrank(model, weights_path): | |
| n = 0 | |
| for name, mod in list(C.iter_target_linears(model, "all")): | |
| prefix = name.replace(".", "_") | |
| pa = os.path.join(weights_path, f"{prefix}_A.pt") | |
| pb = os.path.join(weights_path, f"{prefix}_B.pt") | |
| if os.path.exists(pa) and os.path.exists(pb): | |
| A = torch.load(pa, map_location="cpu") | |
| B = torch.load(pb, map_location="cpu") | |
| parent, attr = C.get_parent_attr(model, name) | |
| setattr(parent, attr, C.LowRankLinear(A, B).to(model.device)) | |
| n += 1 | |
| return n | |
| def load_arm_model(model_path, weights_path): | |
| model, tok = C.load_model(model_path) | |
| if weights_path: | |
| n = replace_with_lowrank(model, weights_path) | |
| print(f"[eval] replaced {n} linears from {weights_path}") | |
| C.assert_head_dense(model) | |
| return model, tok | |
| # ══ task builders (LM-Eval-Harness style, 0-shot) ═════════════════════════════ | |
| def _mmlu_task(row): | |
| subj = row["subject"].replace("_", " ") | |
| ctx = (f"The following are multiple choice questions (with answers) about " | |
| f"{subj}.\n\n{row['question'].strip()}\n") | |
| for i, c in enumerate(row["choices"]): | |
| ctx += f"{LETTERS[i]}. {c}\n" | |
| ctx += "Answer:" | |
| conts = [f" {LETTERS[i]}" for i in range(len(row["choices"]))] | |
| return ctx, conts, int(row["answer"]) | |
| def _arc_task(row): | |
| texts = row["choices"]["text"] | |
| labels = row["choices"]["label"] | |
| key = row["answerKey"] | |
| if key not in labels: | |
| return None | |
| ctx = f"Question: {row['question'].strip()}\nAnswer:" | |
| conts = [f" {t}" for t in texts] | |
| return ctx, conts, labels.index(key) | |
| def _piqa_task(row): | |
| ctx = f"Question: {row['goal'].strip()}\nAnswer:" | |
| conts = [f" {row['sol1'].strip()}", f" {row['sol2'].strip()}"] | |
| return ctx, conts, int(row["label"]) | |
| def _wino_task(row): | |
| """Harness style: context varies with the option, continuation is shared.""" | |
| sent = row["sentence"] | |
| idx = sent.index("_") | |
| opts = [row["option1"], row["option2"]] | |
| cont = sent[idx + 1:] | |
| ctxs = [sent[:idx] + o for o in opts] | |
| return ctxs, cont, int(row["answer"]) - 1 | |
| def _hellaswag_preprocess(text): | |
| text = text.strip().replace(" [title]", ". ") | |
| text = re.sub(r"\[.*?\]", "", text) | |
| return text.replace(" ", " ") | |
| def _hellaswag_task(row): | |
| ctx_a = row["ctx_a"] | |
| ctx_b = row["ctx_b"].capitalize() | |
| ctx = _hellaswag_preprocess(row["activity_label"] + ": " + ctx_a + " " + ctx_b) | |
| conts = [" " + _hellaswag_preprocess(e) for e in row["endings"]] | |
| return ctx, conts, int(row["label"]) | |
| # ══ evaluators ════════════════════════════════════════════════════════════════ | |
| def _dump(f, rec): | |
| f.write(json.dumps(rec) + "\n") | |
| f.flush() | |
| def _run_mcq(model, tok, out_f, name, rows, task_fn, mc_num, batch_size, cfg, | |
| shard, num_shards, limit): | |
| idxs = list(range(len(rows)))[shard::num_shards] | |
| if limit: | |
| idxs = idxs[:limit] | |
| corr = corr_n = n = 0 | |
| for idx in idxs: | |
| t = task_fn(rows[idx]) | |
| if t is None: | |
| continue | |
| if name == "winogrande": | |
| ctxs, cont, gold = t | |
| totals, norms = [], [] | |
| for cx in ctxs: | |
| p, q, tt, nn_ = score_choices(model, tok, cx, [cont], mc_num, | |
| batch_size, cfg) | |
| totals.append(tt[0]) | |
| norms.append(nn_[0]) | |
| pred, pred_n = int(np.argmax(totals)), int(np.argmax(norms)) | |
| ctx_hash = C.sha256_text(ctxs[0]) | |
| else: | |
| ctx, conts, gold = t | |
| pred, pred_n, totals, norms = score_choices(model, tok, ctx, conts, | |
| mc_num, batch_size, cfg) | |
| ctx_hash = C.sha256_text(ctx) | |
| ok, ok_n = (pred == gold), (pred_n == gold) | |
| corr += int(ok) | |
| corr_n += int(ok_n) | |
| n += 1 | |
| _dump(out_f, {"item_id": f"{name}-{idx}", "prompt_hash": ctx_hash, | |
| "correct": bool(ok), "correct_norm": bool(ok_n), | |
| "pred": pred, "pred_norm": pred_n, "gold": gold}) | |
| return corr, corr_n, n | |
| def eval_gsm8k(model, tok, out_f, shard, num_shards, limit, cfg, use_chat, | |
| subset_n=None, seed=42): | |
| from datasets import load_dataset | |
| ds = load_dataset("gsm8k", "main", split="test") | |
| all_idx = list(range(len(ds))) | |
| if subset_n and subset_n < len(all_idx): | |
| # seeded RANDOM subsample, applied BEFORE sharding so shards stay disjoint | |
| rng = np.random.default_rng(seed) | |
| all_idx = sorted(int(i) for i in rng.permutation(len(all_idx))[:subset_n]) | |
| idxs = all_idx[shard::num_shards] | |
| if limit: | |
| idxs = idxs[:limit] | |
| corr = 0 | |
| for idx in idxs: | |
| row = ds[idx] | |
| prompt = f"Question: {row['question'].strip()}\nAnswer:" | |
| text_in = tok.apply_chat_template([{"role": "user", "content": prompt}], | |
| add_generation_prompt=True, tokenize=False) \ | |
| if use_chat else prompt | |
| enc = tok(text_in, return_tensors="pt", add_special_tokens=False) | |
| ids = enc["input_ids"][:, -MAX_CTX:].to(model.device) | |
| am = enc["attention_mask"][:, -MAX_CTX:].to(model.device) | |
| out = llada_generate(model, ids, am, steps=256, gen_length=256, | |
| block_length=8, cfg_scale=cfg) | |
| text = tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True) | |
| pred = extract_number(text) | |
| gold = row["answer"].split("####")[-1].strip().replace(",", "") | |
| ok = numbers_equal(pred, gold) if pred is not None else False | |
| corr += int(ok) | |
| _dump(out_f, {"item_id": f"gsm8k-{idx}", "prompt_hash": C.sha256_text(prompt), | |
| "correct": bool(ok), "correct_norm": bool(ok), | |
| "pred": pred, "pred_norm": pred, "gold": gold}) | |
| return corr, corr, len(idxs) | |
| def extract_number(text): | |
| m = re.search(r"####\s*(-?[\d,]+)", text) | |
| if m: | |
| return m.group(1).replace(",", "") | |
| m = re.search(r"[Tt]he answer is[^\d-]*(-?[\d,]+)", text) | |
| if m: | |
| return m.group(1).replace(",", "") | |
| nums = re.findall(r"-?[\d,]+\.?\d*", text) | |
| return nums[-1].replace(",", "") if nums else None | |
| def numbers_equal(a, b): | |
| try: | |
| return abs(float(a) - float(b)) < 1e-4 | |
| except Exception: | |
| return str(a) == str(b) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--benchmark", required=True, | |
| choices=["gsm8k", "mmlu", "arc_c", "arc_e", "piqa", | |
| "winogrande", "hellaswag"]) | |
| ap.add_argument("--arm", required=True, choices=["ref", "base", "ours"]) | |
| ap.add_argument("--weights", type=str, default=None) | |
| ap.add_argument("--model_path", type=str, default=C.DEFAULT_MODEL_PATH) | |
| ap.add_argument("--shard", type=int, default=0) | |
| ap.add_argument("--num_shards", type=int, default=1) | |
| ap.add_argument("--limit", type=int, default=None, | |
| help="take the first N (schema checks only; NOT unbiased)") | |
| ap.add_argument("--subset_n", type=int, default=None, | |
| help="seeded RANDOM subsample of N items (use for sanity runs)") | |
| ap.add_argument("--mmlu_n", type=int, default=2000) | |
| ap.add_argument("--num_fewshot", type=int, default=0) # Sink-Aware protocol | |
| ap.add_argument("--cfg", type=float, default=0.5) # Sink-Aware protocol | |
| ap.add_argument("--mc_num", type=int, default=128) # Sink-Aware protocol | |
| ap.add_argument("--batch_size", type=int, default=16) | |
| ap.add_argument("--seed", type=int, default=42) | |
| ap.add_argument("--chat_template", choices=["auto", "yes", "no"], default="auto") | |
| ap.add_argument("--tag", type=str, default="") | |
| ap.add_argument("--out_dir", type=str, default=None) | |
| args = ap.parse_args() | |
| if args.num_fewshot != 0: | |
| print(f"[eval] WARNING: num_fewshot={args.num_fewshot}; Sink-Aware protocol is 0.") | |
| torch.manual_seed(args.seed) | |
| np.random.seed(args.seed) | |
| if args.arm == "ref" and args.weights: | |
| sys.exit("ref arm must NOT load weights (it is dense).") | |
| if args.arm in ("base", "ours") and not args.weights: | |
| sys.exit(f"{args.arm} arm requires --weights.") | |
| use_chat = (args.chat_template == "yes") or ( | |
| args.chat_template == "auto" and "Instruct" in args.model_path) | |
| ghash = C.git_hash() | |
| root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| out_dir = args.out_dir or os.path.join(root, "results", "eval", args.arm) | |
| os.makedirs(out_dir, exist_ok=True) | |
| shard_tag = f"_shard{args.shard}of{args.num_shards}" if args.num_shards > 1 else "" | |
| tag = f"_{args.tag}" if args.tag else "" | |
| out_path = os.path.join( | |
| out_dir, f"{ghash}_{args.arm}_{args.benchmark}{tag}{shard_tag}_items.jsonl") | |
| model, tok = load_arm_model(args.model_path, args.weights) | |
| from datasets import load_dataset | |
| def subsample(rows): | |
| """Seeded RANDOM subsample -- unbiased, unlike taking a prefix.""" | |
| if not args.subset_n or args.subset_n >= len(rows): | |
| return rows | |
| rng = np.random.default_rng(args.seed) | |
| sel = rng.permutation(len(rows))[: args.subset_n] | |
| return [rows[int(i)] for i in sorted(sel)] | |
| with open(out_path, "w") as f: | |
| b = args.benchmark | |
| if b == "gsm8k": | |
| c, cn, tot = eval_gsm8k(model, tok, f, args.shard, args.num_shards, | |
| args.limit, args.cfg, use_chat, | |
| args.subset_n, args.seed) | |
| elif b == "mmlu": | |
| rows = load_dataset("cais/mmlu", "all", split="test") | |
| rng = np.random.default_rng(args.seed) | |
| sel = rng.permutation(len(rows))[: args.mmlu_n] | |
| rows = subsample([rows[int(i)] for i in sel]) | |
| c, cn, tot = _run_mcq(model, tok, f, "mmlu", rows, _mmlu_task, | |
| args.mc_num, args.batch_size, args.cfg, | |
| args.shard, args.num_shards, args.limit) | |
| elif b in ("arc_c", "arc_e"): | |
| cfgname = "ARC-Challenge" if b == "arc_c" else "ARC-Easy" | |
| rows = subsample(load_dataset("allenai/ai2_arc", cfgname, split="test")) | |
| c, cn, tot = _run_mcq(model, tok, f, b, rows, _arc_task, | |
| args.mc_num, args.batch_size, args.cfg, | |
| args.shard, args.num_shards, args.limit) | |
| elif b == "piqa": | |
| rows = subsample(load_dataset("lighteval/piqa", "plain_text", split="validation")) | |
| c, cn, tot = _run_mcq(model, tok, f, b, rows, _piqa_task, | |
| args.mc_num, args.batch_size, args.cfg, | |
| args.shard, args.num_shards, args.limit) | |
| elif b == "winogrande": | |
| rows = subsample(load_dataset("winogrande", "winogrande_xl", split="validation")) | |
| c, cn, tot = _run_mcq(model, tok, f, b, rows, _wino_task, | |
| args.mc_num, args.batch_size, args.cfg, | |
| args.shard, args.num_shards, args.limit) | |
| elif b == "hellaswag": | |
| rows = subsample(load_dataset("Rowan/hellaswag", split="validation")) | |
| c, cn, tot = _run_mcq(model, tok, f, b, rows, _hellaswag_task, | |
| args.mc_num, args.batch_size, args.cfg, | |
| args.shard, args.num_shards, args.limit) | |
| summ = {"arm": args.arm, "benchmark": args.benchmark, "weights": args.weights, | |
| "model_path": args.model_path, "shard": args.shard, | |
| "num_shards": args.num_shards, "correct": c, "correct_norm": cn, | |
| "total": tot, "acc": (c / tot) if tot else None, | |
| "acc_norm": (cn / tot) if tot else None, | |
| "num_fewshot": args.num_fewshot, "cfg": args.cfg, "mc_num": args.mc_num, | |
| "is_check_greedy": False, "chat_template": use_chat, | |
| "seed": args.seed, "git_hash": ghash, "items_file": out_path} | |
| C.dump_json(summ, out_path.replace("_items.jsonl", "_summary.json")) | |
| print(f"[eval] {args.arm}/{args.benchmark}{shard_tag}: " | |
| f"acc={summ['acc']} acc_norm={summ['acc_norm']} ({c}/{tot}) -> {out_path}") | |
| if __name__ == "__main__": | |
| main() | |