ukung's picture
Add source code (encoder_loader, model, data, train, ablation, generate) + NOTES
fd3090c verified
Raw History Blame Contribute Delete
4.83 kB
"""
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()