File size: 4,647 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
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
"""
Training loop for the semantic-conditioned decoder.

Run:  python train.py

MEMORY NOTE (Tesla T4, 16 GB)
-----------------------------
The output head produces logits of shape (B, L, 131072). At L=768 that is
~100M floats per sample, so batch_size=1 is not a stylistic choice — batch>1
OOMs. Gradient accumulation recovers an effective batch of 4.
"""

import gc
import time

import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from transformers import get_linear_schedule_with_warmup

from data import SPLIT, build_dataset, make_collate_fn, split_data
from encoder_loader import load_encoder, set_train_mode
from model import SemanticConditionedDecoder

# --- Hyperparameters ---------------------------------------------------------
BATCH_SIZE = 1
ACCUM_STEPS = 4
EPOCHS = 10
LR = 1e-4
WEIGHT_DECAY = 0.01
WARMUP_STEPS = 100
GRAD_CLIP = 1.0

D_MODEL = 1024
NHEAD = 16
NUM_DECODER_LAYERS = 1
DIM_FEEDFORWARD = 1024
DROPOUT = 0.1


def evaluate(model, loader, criterion, device, bos_id):
    model.eval()
    total = 0.0
    with torch.no_grad():
        for src, mask, tgt in loader:
            src, mask, tgt = src.to(device), mask.to(device), tgt.to(device)
            bos = torch.full((tgt.shape[0], 1), bos_id, dtype=torch.long, device=device)
            dec_input = torch.cat([bos, tgt[:, :-1]], dim=1)
            logits = model(input_ids=src, attention_mask=mask, decoder_input_ids=dec_input)
            total += criterion(logits.reshape(-1, logits.shape[-1]), tgt.reshape(-1)).item()
    return total / len(loader)


def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Device: {device}")

    gc.collect()
    torch.cuda.empty_cache()

    encoder, tokenizer = load_encoder(device=str(device))
    problems, targets, src_ids, src_masks, tgt_ids = build_dataset(tokenizer)
    train_data, eval_data = split_data(src_ids, src_masks, tgt_ids, SPLIT)

    collate_fn = make_collate_fn(tokenizer.pad_token_id)
    train_loader = DataLoader(train_data, batch_size=BATCH_SIZE, shuffle=True, collate_fn=collate_fn)
    eval_loader = DataLoader(eval_data, batch_size=BATCH_SIZE, shuffle=False, collate_fn=collate_fn)

    model = SemanticConditionedDecoder(
        encoder=encoder, tokenizer=tokenizer, d_model=D_MODEL, nhead=NHEAD,
        num_decoder_layers=NUM_DECODER_LAYERS, dim_feedforward=DIM_FEEDFORWARD,
        dropout=DROPOUT,
    ).to(device)

    trainable = [p for p in model.parameters() if p.requires_grad]
    print(f"Trainable params: {sum(p.numel() for p in trainable):,}")

    optimizer = torch.optim.AdamW(trainable, lr=LR, weight_decay=WEIGHT_DECAY)
    total_steps = (len(train_loader) // ACCUM_STEPS) * EPOCHS
    scheduler = get_linear_schedule_with_warmup(optimizer, WARMUP_STEPS, total_steps)
    criterion = nn.CrossEntropyLoss(ignore_index=tokenizer.pad_token_id)

    bos_id = tokenizer.bos_token_id
    print(f"BOS={bos_id} EOS={tokenizer.eos_token_id} PAD={tokenizer.pad_token_id}")

    set_train_mode(model)
    for epoch in range(EPOCHS):
        epoch_loss = 0.0
        start = time.time()
        optimizer.zero_grad()

        for step, (src, mask, tgt) in enumerate(train_loader):
            src, mask, tgt = src.to(device), mask.to(device), tgt.to(device)

            bos = torch.full((tgt.shape[0], 1), bos_id, dtype=torch.long, device=device)
            dec_input = torch.cat([bos, tgt[:, :-1]], dim=1)

            logits = model(input_ids=src, attention_mask=mask, decoder_input_ids=dec_input)
            loss = criterion(logits.reshape(-1, logits.shape[-1]), tgt.reshape(-1))
            (loss / ACCUM_STEPS).backward()

            if (step + 1) % ACCUM_STEPS == 0:
                torch.nn.utils.clip_grad_norm_(trainable, GRAD_CLIP)
                optimizer.step()
                scheduler.step()
                optimizer.zero_grad()

            epoch_loss += loss.item()
            if (step + 1) % 50 == 0:
                print(f"  Epoch {epoch+1}/{EPOCHS} | Step {step+1}/{len(train_loader)} "
                      f"| Loss {loss.item():.4f}")

        avg = epoch_loss / len(train_loader)
        eval_loss = evaluate(model, eval_loader, criterion, device, bos_id)
        print(f"Epoch {epoch+1}/{EPOCHS} | Train {avg:.4f} | Eval {eval_loss:.4f} "
              f"| {time.time()-start:.1f}s")

        set_train_mode(model)

    torch.save(
        {k: v for k, v in model.state_dict().items() if not k.startswith("encoder.")},
        "decoder_weights.pt",
    )
    print("Saved decoder_weights.pt (encoder excluded — load it from ukung/semantic-lite-2).")


if __name__ == "__main__":
    main()