File size: 3,213 Bytes
fd3090c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 | """
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:]
|