ukung's picture
Add source code (encoder_loader, model, data, train, ablation, generate) + NOTES
fd3090c verified
Raw History Blame Contribute Delete
3.21 kB
"""
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:]