IFM/K2-Type-0.9B
331.8k
1"""Turn a jev record into one packed sequence: the state once, then every question after it.2 3 <state> state tokens | <q> instr <opt> option </opt> ... <decide> | <q> ... <decide> | ...4 5Isolation: a question token may attend to the state and to earlier tokens of its own question, never to another6question. Position ids restart at the end of the state for every question, so each question sees exactly the7positions it would see if it were asked alone.8 9Marker tokens are unused reserved ids of the K2 tokenizer, inserted by id (never through text tokenization).10"""11import json12 13import torch14 15MARKERS = {"state": "reserved_special_token_100", "q": "reserved_special_token_101", "opt": "reserved_special_token_102",16 "end_opt": "reserved_special_token_103", "decide": "reserved_special_token_104"}17 18 19def render(v, indent=0):20 pad = " " * indent21 if isinstance(v, dict):22 return "\n".join(f"{pad}{k}:\n{render(x, indent + 1)}" if isinstance(x, (dict, list)) else f"{pad}{k}: {x}"23 for k, x in v.items())24 if isinstance(v, list):25 return "\n".join(f"{pad}- {render(x, indent + 1).strip() if isinstance(x, (dict, list)) else x}" for x in v)26 return f"{pad}{v}"27 28 29def options_and_target(q):30 """Option texts and the target distribution for one question."""31 t = q["type"]32 if t == "choice":33 keys = list(q["criteria"])34 opts = [f"{k}: {d}" if d else str(k) for k, d in q["criteria"].items()]35 if q.get("soft") is not None:36 target = [float(q["soft"][k]) for k in keys]37 else:38 target = [1.0 if k == q["label"] else 0.0 for k in keys]39 elif t == "score":40 opts = [str(x) for x in q["criteria"]]41 target = [float(x) for x in q["soft"]] if q.get("soft") is not None else \42 [1.0 if i == q["label"] else 0.0 for i in range(len(opts))]43 else:44 crit = q.get("criteria") or {} # optional definitions of the two outcomes: {"false": ..., "true": ...}45 opts = [f"{k}: {render(crit[k]).strip()}" if crit.get(k) else k for k in ("false", "true")]46 p = float(q["soft"]) if q.get("soft") is not None else float(bool(q["label"]))47 target = [1.0 - p, p]48 s = sum(target)49 return opts, [x / s for x in target]50 51 52def teacher_target(q):53 """The question's `teacher` field as a distribution in canonical option order (None if absent)."""54 t = q.get("teacher")55 if t is None:56 return None57 if q["type"] == "choice":58 return [float(t[k]) for k in q["criteria"]]59 if q["type"] == "score":60 return [float(x) for x in t]61 return [1.0 - float(t), float(t)]62 63 64class Encoder:65 def __init__(self, tok, max_len=4096, max_state=3072, teacher_mix=None):66 """teacher_mix = alpha: target = alpha * gold + (1 - alpha) * teacher when a question has `teacher`."""67 self.tok, self.max_len, self.max_state, self.teacher_mix = tok, max_len, max_state, teacher_mix68 vocab = tok.get_vocab()69 self.ids = {k: vocab[v] for k, v in MARKERS.items()}70 self.bos = tok.bos_token_id71 72 def text(self, s):73 return self.tok(s, add_special_tokens=False).input_ids74 75 def encode(self, rec, rng=None):76 """Returns a dict of python lists, or None when not even one question fits.77 With `rng`, choice and noul options are shown in a random order (targets permuted to match)."""78 state = self.text(render(rec["state"]))[: self.max_state]79 ids = [self.bos, self.ids["state"]] + state80 seg = [0] * len(ids)81 pos = list(range(len(ids)))82 S = len(ids)83 decide, opt_ends, targets, qtypes, qkeys, perms = [], [], [], [], [], []84 for j, (k, q) in enumerate(rec["questions"].items()):85 opts, target = options_and_target(q)86 tt = teacher_target(q) if self.teacher_mix is not None else None87 if tt is not None:88 a = self.teacher_mix89 target = [a * g + (1 - a) * t for g, t in zip(target, tt)]90 order = list(range(len(opts)))91 if rng is not None and q["type"] != "score":92 rng.shuffle(order)93 opts, target = [opts[i] for i in order], [target[i] for i in order]94 instr = q["instructions"] if isinstance(q["instructions"], str) else render(q["instructions"]) # Kev uses dicts too95 qi = [self.ids["q"]] + self.text(instr)96 ends = []97 for o in opts:98 qi += [self.ids["opt"]] + self.text(o)[:128] + [self.ids["end_opt"]]99 ends.append(len(qi) - 1)100 qi.append(self.ids["decide"])101 if len(ids) + len(qi) > self.max_len:102 continue # skip questions that do not fit; the others still train103 base = len(ids)104 ids += qi105 seg += [j + 1] * len(qi)106 pos += list(range(S, S + len(qi)))107 decide.append(base + len(qi) - 1)108 opt_ends.append([base + e for e in ends])109 targets.append(target)110 qtypes.append(q["type"])111 qkeys.append(k)112 perms.append(order) # shown position -> canonical option index113 if not decide:114 return None115 return {"ids": ids, "seg": seg, "pos": pos, "decide": decide, "opt_ends": opt_ends,116 "targets": targets, "qtypes": qtypes, "qkeys": qkeys, "perms": perms}117 118 119def collate(items, pad_id):120 """Right-pad a batch and build the [B, 1, L, L] boolean isolation mask (True = may attend)."""121 B, L = len(items), max(len(x["ids"]) for x in items)122 ids = torch.full((B, L), pad_id, dtype=torch.long)123 pos = torch.zeros((B, L), dtype=torch.long)124 seg = torch.full((B, L), -1, dtype=torch.long)125 for b, x in enumerate(items):126 n = len(x["ids"])127 ids[b, :n] = torch.tensor(x["ids"])128 pos[b, :n] = torch.tensor(x["pos"])129 seg[b, :n] = torch.tensor(x["seg"])130 causal = torch.ones(L, L, dtype=torch.bool).tril()131 sq, sk = seg[:, :, None], seg[:, None, :]132 mask = causal[None] & (sk >= 0) & (sq >= 0) & ((sk == 0) | (sk == sq))133 mask |= torch.eye(L, dtype=torch.bool)[None] # padding rows attend to themselves only (avoids NaN softmax)134 # flat question index: (batch row, decide position, option end positions, target)135 qs = [(b, d, e, t, ty) for b, x in enumerate(items)136 for d, e, t, ty in zip(x["decide"], x["opt_ends"], x["targets"], x["qtypes"])]137 return {"input_ids": ids, "position_ids": pos, "mask": mask[:, None], "questions": qs}138 139 140def load_jsonl(path, limit=0):141 out = []142 with open(path) as fh:143 for i, line in enumerate(fh):144 if limit and i >= limit:145 break146 out.append(json.loads(line))147 return out148 