StandardOne-3B-SH / code /render.py
MyeongHoJeong's picture
Standard One 3B SH v1
43e5946 verified
Raw History Blame Contribute Delete
21.9 kB
"""Template `schema-v1`: render a schema request into one user turn and locate every question/option span.
Design: request rendering for the joint schema head.
A *schema request* (the internal, normalised form used everywhere in this package):
{"state": str | dict | list,
"images": [data URL, ...], # optional
"questions": {key: {"type": "choice" | "noul" | "score",
"instructions": str | dict | list,
"options": [{"name": str, "description": str | dict | list | None}, ...]}}}
Options are always stored in CANONICAL order: choice = caller order, noul = [true, false], score = levels 0..K-1.
Rendering may show them in another order (augmentation); spans are always returned per canonical index.
Encoding (``Encoder.encode``) returns token ids plus token spans. Spans are half-open [start, end) token indices.
"""
import hashlib
import os
import re
import json
import math
import random
SYSTEM_PROMPT = os.environ.get("SH_SYSTEM_PROMPT", "none") # this model was trained without the template's default system prompt; "default" keeps it
TEMPLATE_ID = "schema-v1" if SYSTEM_PROMPT == "default" else "schema-v1-nosys"
MAX_LENGTH = 32768
MAX_OPTIONS = 255
MAX_QUESTIONS = 256
MAX_SCORE_LEVELS = 10
QTYPES = ("choice", "noul", "score")
# Token roles used by the head's role embedding.
ROLE_OTHER, ROLE_STATE, ROLE_PREVIEW, ROLE_QUESTION, ROLE_OPTION = 0, 1, 2, 3, 4
N_ROLES = 5
# Wording variants. W0 is canonical (eval always uses W0); W1..W3 are augmentation only.
WORDINGS = {
"W0": {"preamble": "Read the state and answer every question. Each question lists its possible answers.",
"preview": "Questions to answer:", "state": "State:", "questions": "Questions:",
"choice": "(choose one)", "noul": "(true or false)", "score": "(score 0 to {top})", "bullet": "- "},
"W1": {"preamble": "Answer each question below using only the information provided. Every question lists "
"the answers it allows.",
"preview": "You will be asked:", "state": "Context:", "questions": "Answer these:",
"choice": "[pick one]", "noul": "[true/false]", "score": "[rate 0-{top}]", "bullet": "* "},
"W2": {"preamble": "Use the input to decide every question. Pick exactly one of the listed answers for each.",
"preview": "Questions:", "state": "Input:", "questions": "Decide:",
"choice": "(one of)", "noul": "(true/false)", "score": "(scale 0..{top})", "bullet": "- "},
"W3": {"preamble": "Below is some material followed by questions. Choose one listed answer per question.",
"preview": "Asked below:", "state": "Material:", "questions": "Questions and answers:",
"choice": "[choose one]", "noul": "[true or false]", "score": "[score from 0 to {top}]", "bullet": "* "},
}
class EncodeError(ValueError):
"""A request that cannot be encoded; `reason` is a short machine-readable code (drop statistics)."""
def __init__(self, reason, message=""):
super().__init__(f"{reason}: {message}" if message else reason)
self.reason = reason
def require(condition, reason, message=""):
if not condition:
raise EncodeError(reason, message)
def render_value(value):
"""Plain strings verbatim; structured values as sorted, indented JSON (= jev-adapter render_native_state_text)."""
if value is None:
return ""
if isinstance(value, str):
return value
return json.dumps(value, sort_keys=True, ensure_ascii=False, indent=2, allow_nan=False)
# ----------------------------------------------------------------------------------------------- validation
def validate_request(req):
"""Structural checks of a schema request (raises EncodeError). Returns the list of question keys."""
require(isinstance(req, dict) and isinstance(req.get("questions"), dict), "bad_request", "questions missing")
keys = list(req["questions"])
require(1 <= len(keys) <= MAX_QUESTIONS, "bad_question_count", str(len(keys)))
for key in keys:
require(isinstance(key, str) and key.strip() and "\n" not in key and "[" not in key and "]" not in key,
"bad_question_key", repr(key))
q = req["questions"][key]
t = q.get("type")
require(t in QTYPES, "bad_type", repr(t))
opts = q.get("options")
require(isinstance(opts, list), "bad_options", key)
names = [o.get("name") for o in opts]
require(all(isinstance(n, str) and n.strip() and "\n" not in n for n in names), "bad_option_name", key)
require(len(set(names)) == len(names), "duplicate_option_name", key)
if t == "choice":
require(2 <= len(opts) <= MAX_OPTIONS, "bad_option_count", f"{key}: {len(opts)}")
elif t == "noul":
require(names == ["true", "false"], "bad_noul_options", f"{key}: {names}")
else:
require(2 <= len(opts) <= MAX_SCORE_LEVELS, "bad_score_levels", f"{key}: {len(opts)}")
require(names == [str(i) for i in range(len(opts))], "bad_score_names", key)
return keys
# ----------------------------------------------------------------------------------------------- rendering
def render(req, wording="W0", question_order=None, option_orders=None, preview=True):
"""Render the user text. Returns (content, segments).
question_order: list of keys in display order (default: request order).
option_orders: {key: [canonical index shown at display position 0, 1, ...]} (default identity; score must be
identity because levels are ordinal).
segments: {"state": (c0, c1), "preview": {key: (c0, c1)}, "question": {key: (c0, c1)},
"option": {key: [(c0, c1) per CANONICAL option index]}} as character offsets into `content`.
The state segment includes its header line so it is never empty.
"""
keys = validate_request(req)
order = list(question_order) if question_order is not None else keys
require(sorted(order) == sorted(keys), "bad_question_order")
w = WORDINGS[wording]
parts, pos = [], 0
segs = {"state": None, "preview": {}, "question": {}, "option": {}}
def emit(text):
nonlocal pos
start = pos
parts.append(text)
pos += len(text)
return start, pos
emit(w["preamble"] + "\n\n")
if preview:
emit(w["preview"] + "\n")
for key in order:
text = f"[{key}] {render_value(req['questions'][key]['instructions'])}"
segs["preview"][key] = emit(text)
emit("\n")
emit("\n")
s0, _ = emit(w["state"] + "\n")
_, s1 = emit(render_value(req.get("state", "")))
segs["state"] = (s0, s1)
emit("\n\n" + w["questions"])
for key in order:
q = req["questions"][key]
n = len(q["options"])
tag = w[q["type"]].format(top=n - 1)
emit("\n")
segs["question"][key] = emit(f"[{key}] {tag} {render_value(q['instructions'])}")
perm = list(range(n)) if not option_orders or key not in option_orders else list(option_orders[key])
require(sorted(perm) == list(range(n)), "bad_option_order", key)
require(q["type"] != "score" or perm == list(range(n)), "score_order_must_be_identity", key)
spans = [None] * n
for canonical in perm:
o = q["options"][canonical]
desc = render_value(o.get("description"))
text = o["name"] if not desc else f"{o['name']}: {desc}"
emit("\n" + w["bullet"])
spans[canonical] = emit(text)
segs["option"][key] = spans
return "".join(parts), segs
# ----------------------------------------------------------------------------------------------- augmentation
def _draw(row_id, epoch, what):
return int(hashlib.sha256(f"{row_id}|{epoch}|{what}|schema-aug-v1".encode()).hexdigest()[:16], 16)
def augmentation(req, row_id, epoch, p_choice=1.0, p_noul=0.5, p_field=1.0, p_wording=0.5):
"""Deterministic per (row id, epoch) augmentation plan: wording, question order, option orders."""
keys = list(req["questions"])
frac = lambda what: _draw(row_id, epoch, what) / float(1 << 64) # noqa: E731
wording = "W0"
if frac("wording") < p_wording:
wording = ("W1", "W2", "W3")[_draw(row_id, epoch, "wording-pick") % 3]
order = list(keys)
if len(keys) > 1 and frac("fields") < p_field:
random.Random(_draw(row_id, epoch, "field-perm")).shuffle(order)
option_orders = {}
for key in keys:
q = req["questions"][key]
n = len(q["options"])
p = p_choice if q["type"] == "choice" else p_noul if q["type"] == "noul" else 0.0
if p > 0 and frac("opt-gate|" + key) < p:
perm = list(range(n))
random.Random(_draw(row_id, epoch, "opt-perm|" + key)).shuffle(perm)
option_orders[key] = perm
return {"wording": wording, "question_order": order, "option_orders": option_orders}
# ----------------------------------------------------------------------------------------------- encoding
class Encoder:
"""Request -> token ids + token spans, with the trainer's chat-template call shape.
tokenizer: the canonical tokenizer (AutoTokenizer as the trainer loads it); used for apply_chat_template and
the round-trip check.
fast: a `tokenizers.Tokenizer` loaded from the same tokenizer.json (character offsets).
processor: AutoProcessor, only needed for requests with images.
"""
def __init__(self, tokenizer, fast, processor=None, max_length=MAX_LENGTH, image_token_id=None,
image_end_id=None, check_roundtrip=True):
self.tokenizer, self.fast, self.processor = tokenizer, fast, processor
self.max_length = max_length
self.check_roundtrip = check_roundtrip
self.image_token_id = image_token_id
self.image_end_id = image_end_id
self.special_ids = set(getattr(tokenizer, "all_special_ids", []) or [])
self.chat_template = None
if SYSTEM_PROMPT == "none":
tpl = tokenizer.chat_template
new, n = re.subn(r"set default_system_message = '(?:[^'\\]|\\.)*'", "set default_system_message = ''", tpl)
require(n == 1, "no_default_system_message_in_template")
self.chat_template = new
@classmethod
def from_snapshot(cls, snapshot, with_processor=True, max_length=MAX_LENGTH, mistral_regex_fix=True):
import transformers
from tokenizers import Tokenizer
kwargs = {"fix_mistral_regex": True} if mistral_regex_fix else {}
tok = transformers.AutoTokenizer.from_pretrained(str(snapshot), local_files_only=True,
trust_remote_code=False, token=False, **kwargs)
fast = Tokenizer.from_file(str(snapshot) + "/tokenizer.json")
proc = img = img_end = None
if with_processor:
proc = transformers.AutoProcessor.from_pretrained(str(snapshot), local_files_only=True,
trust_remote_code=False, token=False, **kwargs)
img = tok.convert_tokens_to_ids(proc.image_token)
img_end = tok.convert_tokens_to_ids(proc.image_end_token)
enc = cls(tok, fast, proc, max_length, img, img_end)
enc.snapshot = str(snapshot)
return enc
# -- helpers
def _chat_text(self, content, n_images):
if n_images:
messages = [{"role": "user", "content": [{"type": "image"} for _ in range(n_images)]
+ [{"type": "text", "text": content}]}]
# no enable_thinking here: the Ministral template has no such variable, and the processor's
# apply_chat_template treats unknown kwargs as processor kwargs (warning on every image encode). The text
# is identical either way (tests/test_image_batch.py::test_processor_kwargs_accepted).
kw = {"chat_template": self.chat_template} if self.chat_template else {}
return self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, **kw), messages
messages = [{"role": "user", "content": content}]
kw = {"chat_template": self.chat_template} if self.chat_template else {}
return self.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True,
enable_thinking=False, **kw), messages
def encode(self, req, wording="W0", question_order=None, option_orders=None, preview=True, pil_images=None):
"""Returns a dict (all plain Python ints/lists):
input_ids, n_tokens, keys (rendered question order), qtype, n_opt,
q_span[i] = (s, e) question line, p_span[i] = (s, e) preview line or None,
o_span[i] = [(s, e) per CANONICAL option], o_display[i] = canonical index per display position,
state_spans = [(s, e), ...] (state text incl. header, plus the image block), roles = run-length
[(role, start, end)], image info."""
content, segs = render(req, wording, question_order, option_orders, preview)
images = req.get("images") or []
chat, _messages = self._chat_text(content, len(images))
require(chat.count(content) == 1, "content_not_unique_in_chat")
c0 = chat.index(content)
enc = self.fast.encode(chat, add_special_tokens=False)
ids, offsets = list(enc.ids), list(enc.offsets)
if self.check_roundtrip and not images:
ref = self.tokenizer.apply_chat_template([{"role": "user", "content": content}], tokenize=True,
add_generation_prompt=True, enable_thinking=False,
return_dict=False,
**({"chat_template": self.chat_template} if self.chat_template else {}))
require(list(ref) == ids, "roundtrip_mismatch")
image_span = None
pixel = None
if images:
require(self.processor is not None, "no_processor")
img, img_end = self.image_token_id, self.image_end_id
require(ids.count(img) == len(images), "image_placeholder_count")
if pil_images is None:
pil_images = decode_images(images)
# flat add_special_tokens: _merge_kwargs routes it to the tokenizer (no extra BOS; checked below and in
# tests/test_image_batch.py)
out = self.processor(text=chat, images=pil_images, return_tensors="pt", add_special_tokens=False)
expanded = out["input_ids"][0].tolist()
plain_ids, plain_off = ids, offsets
first_plain = plain_ids.index(img)
last_plain = len(plain_ids) - 1 - plain_ids[::-1].index(img)
require(plain_off[last_plain][1] <= c0, "images_not_before_text")
require(img in expanded and img_end in expanded, "no_image_tokens_after_processor")
first_exp = expanded.index(img)
last_exp = len(expanded) - 1 - expanded[::-1].index(img_end)
require(expanded[:first_exp] == plain_ids[:first_plain], "image_prefix_mismatch")
require(expanded[last_exp + 1:] == plain_ids[last_plain + 1:], "image_suffix_mismatch")
shift = last_exp - last_plain
image_span = (first_exp, last_exp + 1)
ids = expanded
offsets = [None] * len(expanded)
for i in range(first_plain):
offsets[i] = plain_off[i]
for i in range(last_plain + 1, len(plain_ids)):
offsets[i + shift] = plain_off[i]
pixel = {"pixel_values": out["pixel_values"], "image_sizes": out.get("image_sizes")}
n = len(ids)
require(1 <= n <= self.max_length, "too_long", str(n))
# char -> segment id map over the content region
seg_names = []
seg_ranges = []
def add(name, rng):
seg_names.append(name)
seg_ranges.append((rng[0] + c0, rng[1] + c0))
return len(seg_names) - 1
keys = list(question_order) if question_order is not None else list(req["questions"])
sid_state = add(("state",), segs["state"])
sid_prev, sid_q, sid_o = {}, {}, {}
for key in keys:
if key in segs["preview"]:
sid_prev[key] = add(("preview", key), segs["preview"][key])
sid_q[key] = add(("question", key), segs["question"][key])
sid_o[key] = [add(("option", key, j), r) for j, r in enumerate(segs["option"][key])]
# sort segment ranges for binary search
order = sorted(range(len(seg_ranges)), key=lambda i: seg_ranges[i][0])
starts = [seg_ranges[i][0] for i in order]
import bisect
def seg_of(ch):
k = bisect.bisect_right(starts, ch) - 1
if k < 0:
return -1
sid = order[k]
a, b = seg_ranges[sid]
return sid if a <= ch < b else -1
token_seg = [-1] * n
content_end = c0 + len(content)
for t, off in enumerate(offsets):
if off is None:
continue
a, b = off
if b <= c0 or a >= content_end:
continue
ch = a
while ch < b and chat[ch].isspace():
ch += 1
if ch >= b:
ch = a
token_seg[t] = seg_of(ch)
# special tokens must not appear inside the caller's content (e.g. a key spelled like a control token)
require(ids[t] not in self.special_ids, "special_token_in_content", str(ids[t]))
# spans per segment + coverage check (every non-space char of a segment lies in one of its tokens)
spans = {}
for t, sid in enumerate(token_seg):
if sid < 0:
continue
s = spans.get(sid)
spans[sid] = (t, t + 1) if s is None else (s[0], t + 1)
for sid, (a, b) in spans.items():
require(all(token_seg[t] == sid for t in range(a, b)), "span_not_contiguous", str(seg_names[sid]))
for sid, (ca, cb) in enumerate(seg_ranges):
require(sid in spans, "empty_span", str(seg_names[sid]))
a, b = spans[sid]
covered_lo = offsets[a][0]
covered_hi = offsets[b - 1][1]
first = ca
while first < cb and chat[first].isspace():
first += 1
last = cb
while last > first and chat[last - 1].isspace():
last -= 1
require(covered_lo <= first and covered_hi >= last, "span_straddle", str(seg_names[sid]))
state_spans = [spans[sid_state]]
if image_span is not None:
state_spans.insert(0, image_span)
out = {"template": TEMPLATE_ID, "wording": wording, "input_ids": ids, "n_tokens": n, "keys": keys,
"qtype": [req["questions"][k]["type"] for k in keys],
"n_opt": [len(req["questions"][k]["options"]) for k in keys],
"q_span": [spans[sid_q[k]] for k in keys],
"p_span": [spans[sid_prev[k]] if k in sid_prev else None for k in keys],
"o_span": [[spans[s] for s in sid_o[k]] for k in keys],
"o_display": [list(option_orders[k]) if option_orders and k in option_orders
else list(range(len(req["questions"][k]["options"]))) for k in keys],
"state_spans": state_spans, "n_images": len(images), "image_tokens": (image_span[1] - image_span[0]
if image_span else 0)}
if pixel is not None:
out["_pixel"] = pixel
return out
def decode_images(data_urls):
import base64
import io
from PIL import Image
out = []
for url in data_urls:
require(isinstance(url, str) and url.startswith("data:image/") and "," in url, "bad_image")
out.append(Image.open(io.BytesIO(base64.b64decode(url.split(",", 1)[1]))).convert("RGB"))
return out
def token_roles(encoded):
"""Per-token role ids (list of length n_tokens)."""
roles = [ROLE_OTHER] * encoded["n_tokens"]
for a, b in encoded["state_spans"]:
for t in range(a, b):
roles[t] = ROLE_STATE
for i in range(len(encoded["keys"])):
if encoded["p_span"][i]:
a, b = encoded["p_span"][i]
for t in range(a, b):
roles[t] = ROLE_PREVIEW
a, b = encoded["q_span"][i]
for t in range(a, b):
roles[t] = ROLE_QUESTION
for a, b in encoded["o_span"][i]:
for t in range(a, b):
roles[t] = ROLE_OPTION
return roles
def request_from_systemone(body):
"""/v1/systemone request JSON (jev-adapter protocol) -> schema request (canonical option order)."""
questions = {}
for key, q in body["questions"].items():
t = q["type"]
if t == "choice":
opts = [{"name": n, "description": d} for n, d in q["criteria"].items()]
elif t == "noul":
crit = q.get("criteria") or {}
opts = [{"name": "true", "description": crit.get("true")}, {"name": "false", "description": crit.get("false")}]
elif t == "score":
opts = [{"name": str(i), "description": d} for i, d in enumerate(q["criteria"])]
else:
raise EncodeError("bad_type", repr(t))
questions[key] = {"type": t, "instructions": q["instructions"], "options": opts}
return {"state": body.get("state", ""), "images": list(body.get("images") or []), "questions": questions}
def entropy_confidence(p):
h = -sum(x * math.log(x) for x in p if x > 0)
return min(1.0, max(0.0, 1.0 - h / math.log(len(p))))