Team Ai
Modelpublic

IFM/K2-Type-0.9B

sourceHugging Faceapache-2.0updated 8d agoView on Hugging Face
33likes1.8kdownloads
encode.py148 linesDownload Raw Back to jev
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