kaizen-42m / forgetting_benchmark_ewc.py
qoa's picture
Add KAIZEN inference code, benchmarks, semantic head, example memory, README, requirements
4700286 verified
Raw
History Blame Contribute Delete
7.91 kB
#!/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()