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:]