""" 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\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\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\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