""" Ablation: does the semantic conditioning actually contribute? Arm A (model.py) : memory = proj_b(Data B) + Data A prefix Arm B (model.py) : memory = learned constant + no prefix Trainable parameter counts are identical. The only variable that changes is whether the cross-attention memory carries information about the input. Run: python ablation.py """ import gc import json 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 DecoderNoConditioning, SemanticConditionedDecoder from train import (ACCUM_STEPS, BATCH_SIZE, DIM_FEEDFORWARD, DROPOUT, EPOCHS, GRAD_CLIP, LR, NHEAD, NUM_DECODER_LAYERS, WARMUP_STEPS, D_MODEL, evaluate) def train_arm(model, train_loader, eval_loader, criterion, device, bos_id, label): trainable = [p for p in model.parameters() if p.requires_grad] optimizer = torch.optim.AdamW(trainable, lr=LR, weight_decay=0.01) total_steps = (len(train_loader) // ACCUM_STEPS) * EPOCHS scheduler = get_linear_schedule_with_warmup(optimizer, WARMUP_STEPS, total_steps) curve = [] 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() avg = epoch_loss / len(train_loader) eval_loss = evaluate(model, eval_loader, criterion, device, bos_id) curve.append((epoch + 1, avg, eval_loss)) print(f"[{label}] Epoch {epoch+1}/{EPOCHS} | Train {avg:.4f} | Eval {eval_loss:.4f} " f"| {time.time()-start:.1f}s") set_train_mode(model) return curve def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") encoder, tokenizer = load_encoder(device=str(device)) _, _, 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) criterion = nn.CrossEntropyLoss(ignore_index=tokenizer.pad_token_id) bos_id = tokenizer.bos_token_id # --- Arm A: conditioned --- arm_a = 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) curve_a = train_arm(arm_a, train_loader, eval_loader, criterion, device, bos_id, "ARM A") eval_a = curve_a[-1][2] # --- Arm B: unconditioned baseline --- gc.collect() torch.cuda.empty_cache() arm_b = DecoderNoConditioning( encoder=encoder, tokenizer=tokenizer, d_model=D_MODEL, nhead=NHEAD, num_decoder_layers=NUM_DECODER_LAYERS, dim_feedforward=DIM_FEEDFORWARD, dropout=DROPOUT, ).to(device) curve_b = train_arm(arm_b, train_loader, eval_loader, criterion, device, bos_id, "ARM B") eval_b = curve_b[-1][2] results = {"eval_A_final": eval_a, "curve_A": curve_a, "eval_B_final": eval_b, "curve_B": curve_b} with open("ablation_results.json", "w") as f: json.dump(results, f, indent=2) diff = eval_b - eval_a print("\n" + "=" * 70) print(f"Arm A (with semantics) : Eval Loss = {eval_a:.4f}") print(f"Arm B (without semantics) : Eval Loss = {eval_b:.4f}") print(f"Difference (B - A) : {diff:+.4f}") if diff > 0.1: print("=> Conditioning CONTRIBUTES (A clearly better)") elif abs(diff) <= 0.1: print("=> NO evidence conditioning helps (difference within noise)") else: print("=> Baseline is BETTER (needs investigation)") print("=" * 70) if __name__ == "__main__": main()