File size: 4,833 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
"""
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()