"""Byte-level (ByT5) input for the same Qwen-token dataset. Qwen3.8 uses byte-level BPE, so every Qwen token id maps to an exact byte string; the answer's bytes are the concatenation of its tokens' bytes. ByT5 ids are byte + 3 (0 pad, 1 eos, 2 unk). Labels move from tokens to bytes: a BAD token gives B on its first byte and I on the rest, TAIL bytes are I, -100 stays -100 (grouped scheme: +2 for grammar classes, as for token models). For evaluation byte probabilities are averaged back to Qwen tokens, so every metric is computed on exactly the same Qwen tokens as for the ModernBERT / XLM-R taggers. """ from __future__ import annotations import numpy as np import torch from transformers import AutoTokenizer from modeling import _byte_decoder PROMPT_BYTES = 200 SEP = "\n### ответ:\n".encode("utf-8") class QwenBytes: def __init__(self, qwen_tokenizer_dir: str): qtok = AutoTokenizer.from_pretrained(qwen_tokenizer_dir, local_files_only=True) dec = _byte_decoder() # one get_vocab() call: per-id convert_ids_to_tokens is very slow in some transformers versions inv = {i: s for s, i in qtok.get_vocab().items()} self.tb = [] for i in range(248320): s = inv.get(i) if s is None or s.startswith("<|") or any(c not in dec for c in s): self.tb.append(b"") else: self.tb.append(bytes(dec[c] for c in s)) def prompt(self, query: str) -> list[int]: q = query.encode("utf-8")[:PROMPT_BYTES] return [b + 3 for b in q + SEP] def answer(self, token_ids, token_labels=None): """Returns byte ids, byte labels, and for each token its [start, end) in the byte sequence.""" ids, labs, spans = [], [], [] for k, t in enumerate(token_ids): bs = self.tb[t] s = len(ids) ids.extend(b + 3 for b in bs) if token_labels is not None: l = token_labels[k] if l < 0: labs.extend([-100] * len(bs)) elif l == 0: labs.extend([0] * len(bs)) elif l % 2 == 1: # BAD-like: B then I labs.extend([l] + [l + 1] * (len(bs) - 1)) else: labs.extend([l] * len(bs)) spans.append((s, len(ids))) return ids, labs, spans def byte_windows(n: int, a0: int, max_len: int, stride: int): span = max_len - a0 - 1 # room for eos starts, s = [], 0 while True: starts.append(s) if s + span >= n: break s += span - stride return [(s, min(n, s + span)) for s in starts] def expand_windows_bytes(qb: QwenBytes, rows, max_len, stride, synth_weight, labels_fn, include_synth=True, include_corrected=True): out = {"input_ids": [], "labels": [], "weight": []} for r in rows: v = r["variant"] if (v.startswith("synthetic") and not include_synth) or (v == "corrected" and not include_corrected): continue a0 = r["answer_start"] lab = labels_fn(r) ids, blabs, _ = qb.answer(r["input_ids"][a0:], lab[a0:]) if not ids: continue p = qb.prompt(r["query"]) w = synth_weight if v.startswith("synthetic") else 1.0 for s, e in byte_windows(len(ids), len(p), max_len, stride): out["input_ids"].append(p + ids[s:e] + [1]) out["labels"].append([-100] * len(p) + blabs[s:e] + [-100]) out["weight"].append(w) return out @torch.no_grad() def predict_row_bytes(qb: QwenBytes, model, row, max_len, stride, device, batch=8): a0 = row["answer_start"] ids, _, spans = qb.answer(row["input_ids"][a0:]) C = model.config.num_labels n_tok = len(spans) if not ids: return np.zeros((n_tok, C), dtype=np.float32) p0 = qb.prompt(row["query"]) ws = byte_windows(len(ids), len(p0), max_len, stride) probs = np.zeros((len(ids), C), dtype=np.float32) best = np.full(len(ids), -1.0) for b in range(0, len(ws), batch): chunk = ws[b:b + batch] seqs = [p0 + ids[s:e] + [1] for s, e in chunk] L = max(len(x) for x in seqs) x = torch.zeros((len(seqs), L), dtype=torch.long) m = torch.zeros((len(seqs), L), dtype=torch.long) for i, sq in enumerate(seqs): x[i, :len(sq)] = torch.as_tensor(sq) m[i, :len(sq)] = 1 logits = model(input_ids=x.to(device), attention_mask=m.to(device)).logits.float() p = torch.softmax(logits, -1).cpu().numpy() for i, (s, e) in enumerate(chunk): pos = np.arange(s, e) centr = np.minimum(pos - s, e - 1 - pos).astype(float) upd = centr > best[pos] probs[pos[upd]] = p[i, len(p0) + (pos[upd] - s)] best[pos[upd]] = centr[upd] out = np.zeros((n_tok, C), dtype=np.float32) for k, (s, e) in enumerate(spans): if e > s: out[k] = probs[s:e].mean(0) else: out[k, 0] = 1.0 # special tokens have no bytes return out