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