#!/usr/bin/env python3 """Phase 7.4-EWC — Three-way forgetting benchmark. Compares three continual learning strategies on identical 5-domain × 10-task stream: KAIZEN : per-task episodic adapters — structural isolation, BWT=0 by construction SHARED : single shared adapter, naive fine-tuning — catastrophic forgetting baseline EWC : single shared adapter + EWC diagonal Laplace (Kirkpatrick 2017, λ=1000) Expected: BWT(KAIZEN) ≈ 0 > BWT(EWC) > BWT(SHARED) If EWC barely improves over SHARED, λ is too small — see LAMBDA_EWC diagnostic output. """ import time import torch from tokenizers import Tokenizer from huggingface_hub import hf_hub_download from lora import KaizenWithLoRA, LoRAAdapter from task_memory import TaskMemory from forgetting_benchmark import ( eval_retention, kaizen_train_domain, shared_train_domain, cl_metrics, build_domains, DEFAULT_CKPT, ) from online_learner import ( HF_TOKEN, HF_TOK_REPO, TOP_K, LORA_RANK, LORA_ALPHA, ATTEMPT_THRESHOLD, LOSS_EARLY_STOP, ONLINE_STEPS, ONLINE_LR, build_update_seq, ) from eval_benchmark import build_prompt_ids, clean_ids, generate, token_f1, BLOCK_SIZE, MAX_GEN from ewc import EWC KAIZEN_STORE_EWC = os.path.join(os.path.expanduser('~'), '.kaizen', 'forgetting_ewc') LAMBDA_EWC = 1000.0 def ewc_online_update(model: KaizenWithLoRA, adapter: LoRAAdapter, x: torch.Tensor, y: torch.Tensor, ewc: EWC) -> float: """Adam update with EWC penalty. Returns final task loss (not total loss).""" for p in model.parameters(): p.requires_grad_(False) for p in adapter.parameters(): p.requires_grad_(True) optimizer = torch.optim.Adam(adapter.parameters(), lr=ONLINE_LR) last_task_loss = float('inf') for _ in range(ONLINE_STEPS): optimizer.zero_grad() _, task_loss = model(x, targets=y, adapter=adapter) total_loss = task_loss + ewc.penalty(adapter) total_loss.backward() optimizer.step() last_task_loss = task_loss.item() if last_task_loss < LOSS_EARLY_STOP: break return last_task_loss def ewc_train_domain(model: KaizenWithLoRA, tokenizer, shared_adapter: LoRAAdapter, tasks: list, ewc: EWC) -> None: """Train shared adapter on domain tasks with EWC penalty against prior domains.""" for question, answer, _ in tasks: prompt_ids = build_prompt_ids(tokenizer, question)[:BLOCK_SIZE - MAX_GEN] ref_ids = clean_ids(tokenizer, answer) with torch.no_grad(): gen_ids = generate(model, tokenizer, prompt_ids, adapter=shared_adapter) if token_f1(gen_ids, ref_ids) >= ATTEMPT_THRESHOLD: continue x_upd, y_upd = build_update_seq(tokenizer, question, answer) ewc_online_update(model, shared_adapter, x_upd, y_upd, ewc) def main(): t0 = time.time() print('Phase 7.4-EWC — Three-way forgetting benchmark') print('Conditions: KAIZEN / SHARED / EWC-SHARED (Kirkpatrick 2017)') print('=' * 70) tok_file = hf_hub_download(HF_TOK_REPO, 'tokenizer.json', token=HF_TOKEN, cache_dir=None) tokenizer = Tokenizer.from_file(tok_file) model = KaizenWithLoRA() model.load_base(DEFAULT_CKPT) model.eval() print(f'Model loaded ({time.time()-t0:.0f}s)') names, domains = build_domains() N = len(names) print(f'{N} domains × 10 tasks: {names}') # ── KAIZEN ──────────────────────────────────────────────────────────────── print('\n── KAIZEN (per-task episodic adapters, structural isolation) ──') memory = TaskMemory(KAIZEN_STORE_EWC, top_k=TOP_K) R_k = [[0.0]*N for _ in range(N)] for j, (dname, tasks) in enumerate(zip(names, domains)): n_stored = kaizen_train_domain(model, tokenizer, memory, tasks) for i in range(j+1): R_k[j][i] = eval_retention(model, tokenizer, domains[i], memory=memory) row = ' '.join(f'{names[i][:6]}={R_k[j][i]:.3f}' for i in range(j+1)) print(f' d{j+1} {dname}: +{n_stored} stored | {row}') kaizen_aa, kaizen_bwt, kaizen_fgt = cl_metrics(R_k, N) print(f' → AA={kaizen_aa:.4f} BWT={kaizen_bwt:+.4f} Fgt={kaizen_fgt:.4f} ' f'mem={len(memory)}') # ── SHARED ──────────────────────────────────────────────────────────────── print('\n── SHARED (single adapter, no regularization) ──') shared = LoRAAdapter(model.N_LAYERS, model.D_MODEL, LORA_RANK, LORA_ALPHA) R_s = [[0.0]*N for _ in range(N)] for j, (dname, tasks) in enumerate(zip(names, domains)): shared_train_domain(model, tokenizer, shared, tasks) for i in range(j+1): R_s[j][i] = eval_retention(model, tokenizer, domains[i], shared_adapter=shared) row = ' '.join(f'{names[i][:6]}={R_s[j][i]:.3f}' for i in range(j+1)) print(f' d{j+1} {dname}: | {row}') shared_aa, shared_bwt, shared_fgt = cl_metrics(R_s, N) print(f' → AA={shared_aa:.4f} BWT={shared_bwt:+.4f} Fgt={shared_fgt:.4f}') # ── EWC-SHARED ──────────────────────────────────────────────────────────── print(f'\n── EWC-SHARED (single adapter + EWC λ={LAMBDA_EWC:.0f}, Kirkpatrick 2017) ──') ewc = EWC(lambda_ewc=LAMBDA_EWC) shared_ewc = LoRAAdapter(model.N_LAYERS, model.D_MODEL, LORA_RANK, LORA_ALPHA) R_e = [[0.0]*N for _ in range(N)] for j, (dname, tasks) in enumerate(zip(names, domains)): ewc_train_domain(model, tokenizer, shared_ewc, tasks, ewc) for i in range(j+1): R_e[j][i] = eval_retention(model, tokenizer, domains[i], shared_adapter=shared_ewc) row = ' '.join(f'{names[i][:6]}={R_e[j][i]:.3f}' for i in range(j+1)) # Compute Fisher AFTER domain j training, anchor = post-training params fisher_norm = ewc.update(model, shared_ewc, tasks, tokenizer) print(f' d{j+1} {dname}: Fisher_norm={fisher_norm:.4f} | {row}') ewc_aa, ewc_bwt, ewc_fgt = cl_metrics(R_e, N) print(f' → AA={ewc_aa:.4f} BWT={ewc_bwt:+.4f} Fgt={ewc_fgt:.4f}') # ── Summary table ───────────────────────────────────────────────────────── print() print('=' * 70) print(f'{"Condition":24s} {"AA":>8s} {"BWT":>8s} {"Fgt":>8s}') print('-' * 54) print(f'{"KAIZEN episodic":24s} {kaizen_aa:8.4f} {kaizen_bwt:+8.4f} {kaizen_fgt:8.4f}') print(f'{"EWC-SHARED":24s} {ewc_aa:8.4f} {ewc_bwt:+8.4f} {ewc_fgt:8.4f}') print(f'{"SHARED adapter":24s} {shared_aa:8.4f} {shared_bwt:+8.4f} {shared_fgt:8.4f}') print() ewc_reduces = ewc_bwt > shared_bwt structural_wins = kaizen_bwt > ewc_bwt print(f'EWC reduces forgetting vs SHARED : {"YES" if ewc_reduces else "NO "} ' f'(EWC BWT {ewc_bwt:+.4f} vs SHARED BWT {shared_bwt:+.4f})') print(f'Structural isolation beats EWC : {"YES" if structural_wins else "NO "} ' f'(KAIZEN BWT {kaizen_bwt:+.4f} vs EWC BWT {ewc_bwt:+.4f})') if not ewc_reduces: print(f' NOTE: EWC did not reduce forgetting. ' f'Fisher norms suggest λ={LAMBDA_EWC:.0f} may be too small.') print(f' Retry with λ=10000 or λ=100000 to see EWC effect.') print(f'\nRuntime: {time.time()-t0:.0f}s') if __name__ == '__main__': main()