Download tinker/tinker_data.py from halle01/coder-fake: direct link, hf CLI and curl.
- Browser
- Download file 5.31 kB
-
https://huggingface.co/halle01/coder-fake/resolve/main/tinker/tinker_data.py
- Command line
-
hf download hf://halle01/coder-fake/tinker/tinker_data.py
-
curl -L -o tinker_data.py https://huggingface.co/halle01/coder-fake/resolve/main/tinker/tinker_data.py
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 | |