hfvladkon's picture
Upload main defect tagger (weights, tokenizer, inference code, card)
22ca93a
Raw History Blame Contribute Delete
5.13 kB
"""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