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()
|