Download code/train.py from ukung/semantic-lite-2-decoder-smoke-test: direct link, hf CLI and curl.
- Browser
- Download file 4.65 kB
-
https://huggingface.co/ukung/semantic-lite-2-decoder-smoke-test/resolve/main/code/train.py
- Command line
-
hf download hf://ukung/semantic-lite-2-decoder-smoke-test/code/train.py
-
curl -L -o train.py https://huggingface.co/ukung/semantic-lite-2-decoder-smoke-test/resolve/main/code/train.py
4.65 kB
| """ | |
| 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() | |