coder-fake / tinker /tinker_data.py
halle01's picture
Add files using upload-large-folder tool
78cb49e verified
Raw History Blame Contribute Delete
5.31 kB
"""
Shared dataloader for Tinker training of Qwen/Qwen3.6-27B on our threejs data.
Builds MULTIMODAL Datums (image + text) — verified working on Tinker:
a 1024x1024 image = 1024 image tokens (patch_size 16 x merge_size 2 -> factor
32; tokens = (W/32)*(H/32)). Our reference PNGs are 1024x1024, so 1024 each.
Sequence layout per SFT example:
[system + "\n<image>\n"] [ImageChunk(1024)] ["\n"+user_text+"\n"] [assistant code + EOS]
SFT loss (cross_entropy) is put ONLY on the completion tokens (weight 1),
everything else weight 0. Uses the cookbook shift: input = full[:-1],
target = full[1:] (see tinker_cookbook.supervised.common.datum_from_tokens_weights).
"""
from __future__ import annotations
import json, os, signal
from PIL import Image
import tinker
from tinker import types as tt
PATCH, MERGE = 16, 2
FACTOR = PATCH * MERGE # 32
TRAIN_RATE_PER_1M = 4.103 # $/1M tokens, Qwen3.6-27B train
def image_expected_tokens(path: str) -> int:
"""Qwen3-VL image-token count. All our renders are 1024x1024 -> 1024.
General: round each side to a multiple of FACTOR, then (W/FACTOR)*(H/FACTOR).
(Assumes within [min,max]_pixels — true for our 1024x1024 set.)"""
with Image.open(path) as im:
w, h = im.size
gw = max(1, round(w / FACTOR))
gh = max(1, round(h / FACTOR))
return gw * gh
def _texts(row):
sys_txt = row["prompt"][0]["content"]
user_txt = ""
for c in row["prompt"][1]["content"]:
if isinstance(c, dict) and c.get("type") == "text":
user_txt = c["text"]
return sys_txt, user_txt
def _build_multimodal(tok, root, img_rel, pre_txt, post_txt, comp_txt, max_length):
"""Return (model_input, target_tokens, weights, n_tokens) for one example.
Weight=1 only on the completion (comp_txt) tokens."""
img_path = os.path.join(root, img_rel)
itoks = image_expected_tokens(img_path)
img_bytes = open(img_path, "rb").read()
pre = tok.encode(pre_txt, add_special_tokens=False)
post = tok.encode(post_txt, add_special_tokens=False)
eos = getattr(tok, "eos_token_id", None)
comp = tok.encode(comp_txt, add_special_tokens=False) + ([eos] if eos is not None else [])
n_prefix = len(pre) + itoks + len(post)
total = n_prefix + len(comp)
# budget: if over max_length, truncate the COMPLETION from the end (keep image+prompt)
if total > max_length:
keep = max(1, max_length - n_prefix)
comp = comp[:keep]
total = n_prefix + len(comp)
flat_targets = [0] * n_prefix + comp
flat_weights = [0.0] * n_prefix + [1.0] * len(comp)
tgt = flat_targets[1:]
w = flat_weights[1:]
mi = tinker.ModelInput.empty()
for t in pre:
mi = mi.append_int(t)
mi = mi.append(tt.ImageChunk(data=img_bytes, format="png", expected_tokens=itoks, type="image"))
for t in (post + comp)[:-1]: # full[:-1]
mi = mi.append_int(t)
return mi, tgt, w, total
def build_sft_datum(row, tok, root, max_length=16384):
sys_txt, user_txt = _texts(row)
mi, tgt, w, ntok = _build_multimodal(
tok, root, row["images"][0],
pre_txt=sys_txt + "\n<image>\n", post_txt="\n" + user_txt + "\n",
comp_txt=row["completion"][0]["content"], max_length=max_length,
)
datum = tinker.Datum(
model_input=mi,
loss_fn_inputs={
"weights": tt.TensorData(data=[float(x) for x in w], dtype="float32", shape=[len(w)]),
"target_tokens": tt.TensorData(data=[int(x) for x in tgt], dtype="int64", shape=[len(tgt)]),
},
)
return datum, ntok
def build_dpo_sequences(row, tok, root, max_length=16384):
"""For DPO: same prompt+image, two completions (chosen, rejected).
Returns two (model_input, comp_token_ids, comp_start_index, n_tokens) tuples —
the caller computes per-token logprobs of the completion under policy+ref."""
sys_txt, user_txt = _texts(row)
out = {}
for key, msgs in (("chosen", row["chosen"]), ("rejected", row["rejected"])):
mi, tgt, w, ntok = _build_multimodal(
tok, root, row["images"][0],
pre_txt=sys_txt + "\n<image>\n", post_txt="\n" + user_txt + "\n",
comp_txt=msgs[0]["content"], max_length=max_length,
)
out[key] = dict(model_input=mi, target_tokens=tgt, weights=w, n_tokens=ntok)
return out
def load_rows(path, limit=None):
rows = []
for line in open(path):
line = line.strip()
if line:
rows.append(json.loads(line))
if limit and len(rows) >= limit:
break
return rows
def install_stop_handler():
"""Graceful kill: first SIGINT/SIGTERM sets a flag so the loop saves a
resume checkpoint after the current step, then exits (lossless). A second
Ctrl-C force-quits. Lets you watch cost on the dashboard and kill safely."""
state = {"stop": False}
def _h(signum, frame):
if state["stop"]:
print("\n[signal] second interrupt — force quit", flush=True); os._exit(130)
state["stop"] = True
print("\n[signal] caught — saving a resume checkpoint after the current step "
"(Ctrl-C again to force-quit)", flush=True)
signal.signal(signal.SIGINT, _h)
try: signal.signal(signal.SIGTERM, _h)
except Exception: pass
return state