| |
| """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}') |
|
|
| |
| 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)}') |
|
|
| |
| 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}') |
|
|
| |
| 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)) |
| |
| 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}') |
|
|
| |
| 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() |
|
|