""" Dataset preparation. Source : ukung/Opus-4.6-Reasoning Columns: id, problem, thinking, solution, difficulty, category, timestamp, hash Target : thinking + "\\n\\n" + solution (chain-of-thought, then the answer) TOKEN ALIGNMENT (easy to get wrong — read this before changing anything) ------------------------------------------------------------------------ tgt_ids = [t0, t1, ..., tN, EOS] <- no BOS dec_input = [BOS, t0, t1, ..., tN] <- shifted right, BOS prepended labels = [t0, t1, ..., tN, EOS] Inference starts from [BOS] and generates autoregressively. Note on EOS: for samples shorter than max_len, EOS lands in dec_input at a position whose label is PAD (ignored by the loss). That is mathematically correct — the token preceding the first PAD really is EOS — and no gradient flows through those positions. """ import torch from datasets import load_dataset DATASET_ID = "ukung/Opus-4.6-Reasoning" N_SAMPLES = 1000 SHUFFLE_SEED = 42 MAX_SRC_LEN = 512 MAX_TGT_LEN = 768 SEP = "\n\n" SPLIT = 900 # 900 train / 100 eval def load_samples(n=N_SAMPLES, seed=SHUFFLE_SEED): """Deterministic subset of the dataset.""" raw = load_dataset(DATASET_ID, split="train") return raw.shuffle(seed=seed).select(range(n)) def format_target(example): return example["thinking"].strip() + SEP + example["solution"].strip() def tokenize_src(tokenizer, text): enc = tokenizer(text, padding="max_length", truncation=True, max_length=MAX_SRC_LEN, return_tensors="pt") return enc["input_ids"][0], enc["attention_mask"][0] def tokenize_tgt(tokenizer, text): """[tokens..., EOS] — no BOS; BOS is added by the training loop.""" enc = tokenizer(text, truncation=True, max_length=MAX_TGT_LEN - 1, return_tensors="pt", add_special_tokens=False) ids = enc["input_ids"][0] return torch.cat([ids, torch.tensor([tokenizer.eos_token_id])]) def build_dataset(tokenizer, samples=None): """Return (problems, targets, src_ids, src_masks, tgt_ids).""" if samples is None: samples = load_samples() problems = samples["problem"] targets = [format_target(samples[i]) for i in range(len(samples))] src_ids, src_masks, tgt_ids = [], [], [] for i in range(len(problems)): si, sm = tokenize_src(tokenizer, problems[i]) src_ids.append(si) src_masks.append(sm) tgt_ids.append(tokenize_tgt(tokenizer, targets[i])) return problems, targets, src_ids, src_masks, tgt_ids def make_collate_fn(pad_id): """Pad targets to a uniform length within each batch.""" def collate_fn(batch): src = torch.stack([b[0] for b in batch]) mask = torch.stack([b[1] for b in batch]) tgts = [b[2] for b in batch] max_len = max(t.shape[0] for t in tgts) padded = torch.full((len(tgts), max_len), pad_id, dtype=torch.long) for i, t in enumerate(tgts): padded[i, :t.shape[0]] = t return src, mask, padded return collate_fn def split_data(src_ids, src_masks, tgt_ids, split=SPLIT): all_data = list(zip(src_ids, src_masks, tgt_ids)) return all_data[:split], all_data[split:]