/
Mihaham
/
CNN-NEAT
Обзор
Документация
Войти
/
Mihaham
/
CNN-NEAT
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
scripts/bench_wave_cache_ab.py
532 строки
17 KB
MihahamYT
feat(scripts): OVA holdout, hash studies, and cache benches
04 авг 2026, 20:36
04 авг 2026, 20:36
e256d0f
Код
Авторство
О чём код?
#!/usr/bin/env python3 """A/B: naive per-genome forward vs planned wave-cache (same init). Builds one GELU (or chosen) CIFAR population, then measures: A) baseline — each genome ``forward`` separately (no shared cache) B) hash-cache per-genome ``execute`` after elite commit (sequential hits) C) planned unique-miss wave (``plan`` + ``execute_planned_population``) Checks output correctness (max abs err / fitness agreement) and speedup. Usage ----- :: python scripts/bench_wave_cache_ab.py python scripts/bench_wave_cache_ab.py --pop 48 --images 512 --activation gelu python scripts/bench_wave_cache_ab.py --pop 64 --images 1024 --warmup 1 """ from __future__ import annotations import argparse import json import random import sys import time from dataclasses import asdict, dataclass, field from pathlib import Path from typing import Any, Dict, List, Optional, Sequence, Tuple import torch ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from cnn_neat.binary_head import output_to_score # noqa: E402 from cnn_neat.forward_cache import ( # noqa: E402 ForwardActivationCache, data_id_from_tensor, ) from cnn_neat.mutations import ConvolutionalNetworkMutator, MutationType # noqa: E402 # Reuse real-bench population / CIFAR helpers. import importlib.util _real_path = ROOT / "scripts" / "bench_wave_cache_real.py" _spec = importlib.util.spec_from_file_location("bench_wave_cache_real", _real_path) assert _spec is not None and _spec.loader is not None _real = importlib.util.module_from_spec(_spec) sys.modules["bench_wave_cache_real"] = _real _spec.loader.exec_module(_real) _seed_genome = _real._seed_genome _breed = _real._breed _count_nodes = _real._count_nodes _count_edges = _real._count_edges load_cifar_neat_batch = _real.load_cifar_neat_batch GenomeId = Tuple[int, int] PopRec = Tuple[Any, GenomeId] @dataclass class GenomeRow: genome_id: List[int] nodes: int edges: int wall_s: float fitness: float max_abs_err_vs_baseline: float = 0.0 @dataclass class ModeResult: name: str wall_s: float peak_vram_mb: float genomes: List[GenomeRow] = field(default_factory=list) plan: Dict[str, int] = field(default_factory=dict) plan_s: float = 0.0 gpu_s: float = 0.0 max_abs_err: float = 0.0 mean_fitness_abs_delta: float = 0.0 speedup_vs_baseline: Optional[float] = None def _sync(device: torch.device) -> None: if device.type == "cuda": torch.cuda.synchronize(device) def _peak_vram_mb(device: torch.device) -> float: if device.type != "cuda": return 0.0 return float(torch.cuda.max_memory_allocated(device)) / (1024.0 * 1024.0) def _reset_peak(device: torch.device) -> None: if device.type == "cuda": torch.cuda.reset_peak_memory_stats(device) torch.cuda.empty_cache() def _fitness(out: torch.Tensor, labels: torch.Tensor, activation: str) -> float: scores = output_to_score(out.float(), output_readout="mean", activation=activation) preds = (scores > 0).long() tp = ((labels == 1) & (preds == 1)).sum().item() tn = ((labels == 0) & (preds == 0)).sum().item() fp = ((labels == 0) & (preds == 1)).sum().item() fn = ((labels == 1) & (preds == 0)).sum().item() tpr = tp / max(1, tp + fn) tnr = tn / max(1, tn + fp) return 0.5 * (tpr + tnr) def _max_err(a: torch.Tensor, b: torch.Tensor) -> float: return float((a.float() - b.float()).abs().max().item()) def build_population( *, pop: int, n_elite: int, activation: str, seed: int, device: torch.device, ) -> List[PopRec]: rng = random.Random(seed) torch.manual_seed(seed) if device.type == "cuda": torch.cuda.manual_seed_all(seed) stem = _seed_genome(activation, rng) elites: List[Any] = [] for i in range(n_elite): e = ConvolutionalNetworkMutator(stem).apply( MutationType.MUTATE_WEIGHTS, rate=1.0, scale=0.02 + 0.01 * i ) e.to(device) e.eval() elites.append(e) pop_recs = _breed(elites, pop_size=pop, generation=1, rng=rng) for g, _ in pop_recs: g.to(device) g.eval() return pop_recs @torch.inference_mode() def run_baseline_naive( items: Sequence[PopRec], images: torch.Tensor, labels: torch.Tensor, *, activation: str, ) -> Tuple[ModeResult, Dict[GenomeId, torch.Tensor], Dict[GenomeId, float]]: device = images.device _reset_peak(device) outs: Dict[GenomeId, torch.Tensor] = {} fits: Dict[GenomeId, float] = {} rows: List[GenomeRow] = [] _sync(device) t0 = time.perf_counter() for genome, gid in items: _sync(device) tg0 = time.perf_counter() out = genome(images) _sync(device) dt = time.perf_counter() - tg0 fit = _fitness(out, labels, activation) outs[gid] = out fits[gid] = fit rows.append( GenomeRow( genome_id=list(gid), nodes=_count_nodes(genome), edges=_count_edges(genome), wall_s=dt, fitness=fit, ) ) wall = time.perf_counter() - t0 return ( ModeResult( name="A_naive_per_genome_forward", wall_s=wall, peak_vram_mb=_peak_vram_mb(device), genomes=rows, ), outs, fits, ) @torch.inference_mode() def run_hash_per_genome( items: Sequence[PopRec], elite_items: Sequence[PopRec], images: torch.Tensor, labels: torch.Tensor, *, activation: str, max_bytes: int, ) -> Tuple[ModeResult, Dict[GenomeId, torch.Tensor], Dict[GenomeId, float]]: device = images.device data_id = data_id_from_tensor(images) _reset_peak(device) cache = ForwardActivationCache(max_bytes=max_bytes, data_id=data_id) cache.precompute_population(list(elite_items), register_overlay=False) for genome, gid in elite_items: cache.execute(genome, images, gid, pin=True, store_activations=True) cache.precompute_population(list(items), register_overlay=False) outs: Dict[GenomeId, torch.Tensor] = {} fits: Dict[GenomeId, float] = {} rows: List[GenomeRow] = [] _sync(device) t0 = time.perf_counter() for genome, gid in items: _sync(device) tg0 = time.perf_counter() out = cache.execute(genome, images, gid, store_activations=False) _sync(device) dt = time.perf_counter() - tg0 fit = _fitness(out, labels, activation) outs[gid] = out fits[gid] = fit rows.append( GenomeRow( genome_id=list(gid), nodes=_count_nodes(genome), edges=_count_edges(genome), wall_s=dt, fitness=fit, ) ) wall = time.perf_counter() - t0 st = cache.stats return ( ModeResult( name="B_hash_per_genome_execute", wall_s=wall, peak_vram_mb=_peak_vram_mb(device), genomes=rows, plan={ "hits": int(st.hits), "misses": int(st.misses), "entries": int(st.entries), }, ), outs, fits, ) @torch.inference_mode() def run_planned_wave( items: Sequence[PopRec], elite_items: Sequence[PopRec], images: torch.Tensor, labels: torch.Tensor, *, activation: str, max_bytes: int, baseline_outs: Dict[GenomeId, torch.Tensor], baseline_fits: Dict[GenomeId, float], ) -> Tuple[ModeResult, Dict[GenomeId, torch.Tensor], Dict[GenomeId, float]]: device = images.device data_id = data_id_from_tensor(images) _reset_peak(device) cache = ForwardActivationCache(max_bytes=max_bytes, data_id=data_id) cache.precompute_population(list(elite_items), register_overlay=False) for genome, gid in elite_items: cache.execute(genome, images, gid, pin=True, store_activations=True) # Precompute outside timed section (same Merkle setup cost for A/B fairness # as hash path; naive path has no Merkle — noted in report). cache.precompute_population(list(items), register_overlay=True) _sync(device) t_plan0 = time.perf_counter() work = cache.plan_population_misses(list(items)) plan_s = time.perf_counter() - t_plan0 t_gpu0 = time.perf_counter() outs = cache.execute_planned_population( images, work, store_activations=False, pin=False ) _sync(device) gpu_s = time.perf_counter() - t_gpu0 wall = plan_s + gpu_s fits: Dict[GenomeId, float] = {} rows: List[GenomeRow] = [] max_err = 0.0 fit_deltas: List[float] = [] for genome, gid in items: out = outs[gid] fit = _fitness(out, labels, activation) fits[gid] = fit err = _max_err(baseline_outs[gid], out) if gid in baseline_outs else float("inf") max_err = max(max_err, err) fit_deltas.append(abs(fit - baseline_fits.get(gid, fit))) rows.append( GenomeRow( genome_id=list(gid), nodes=_count_nodes(genome), edges=_count_edges(genome), wall_s=0.0, # wave is collective; per-genome wall not separable fitness=fit, max_abs_err_vs_baseline=err, ) ) summary = work.summary() return ( ModeResult( name="C_planned_unique_miss_wave", wall_s=wall, peak_vram_mb=_peak_vram_mb(device), genomes=rows, plan=dict(summary), plan_s=plan_s, gpu_s=gpu_s, max_abs_err=max_err, mean_fitness_abs_delta=(sum(fit_deltas) / max(1, len(fit_deltas))), ), outs, fits, ) def _attach_baseline_err( mode: ModeResult, baseline_outs: Dict[GenomeId, torch.Tensor], outs: Dict[GenomeId, torch.Tensor], baseline_fits: Dict[GenomeId, float], fits: Dict[GenomeId, float], ) -> None: max_err = 0.0 deltas: List[float] = [] for row in mode.genomes: gid = (int(row.genome_id[0]), int(row.genome_id[1])) err = _max_err(baseline_outs[gid], outs[gid]) row.max_abs_err_vs_baseline = err max_err = max(max_err, err) deltas.append(abs(fits[gid] - baseline_fits[gid])) mode.max_abs_err = max_err mode.mean_fitness_abs_delta = sum(deltas) / max(1, len(deltas)) def build_parser() -> argparse.ArgumentParser: p = argparse.ArgumentParser(description=__doc__) p.add_argument("--pop", type=int, default=48) p.add_argument("--images", type=int, default=512) p.add_argument("--elite-percent", type=float, default=0.2) p.add_argument("--activation", default="gelu") p.add_argument("--seed", type=int, default=42) p.add_argument("--warmup", type=int, default=1) p.add_argument("--cache-gb", type=float, default=8.0) p.add_argument( "--cifar-root", type=Path, default=ROOT / "data" / "cifar10", ) p.add_argument( "--out", type=Path, default=ROOT / "runs" / "cache_wave_sim" / "ab_correctness", ) return p def main(argv: Optional[Sequence[str]] = None) -> int: args = build_parser().parse_args(argv) if not torch.cuda.is_available(): print("CUDA required", file=sys.stderr) return 2 device = torch.device("cuda") args.out.mkdir(parents=True, exist_ok=True) max_bytes = int(args.cache_gb * 1024**3) n_elite = max(1, int(round(args.pop * float(args.elite_percent)))) print( f"Building pop={args.pop} elites={n_elite} images={args.images} " f"activation={args.activation} seed={args.seed}", flush=True, ) images, labels = load_cifar_neat_batch( n_images=args.images, device=device, cifar_root=args.cifar_root, seed=args.seed, ) items = build_population( pop=args.pop, n_elite=n_elite, activation=args.activation, seed=args.seed, device=device, ) elite_items = items[:n_elite] # Warmup (not timed): touch CUDA paths. for _ in range(max(0, args.warmup)): with torch.inference_mode(): _ = elite_items[0][0](images[: min(16, images.shape[0])]) _sync(device) print("\n=== A) Naive per-genome forward ===", flush=True) mode_a, outs_a, fits_a = run_baseline_naive( items, images, labels, activation=args.activation ) print( f" wall={mode_a.wall_s:.3f}s peak_vram={mode_a.peak_vram_mb:.0f}MB " f"mean_genome={mode_a.wall_s / max(1, len(items)):.4f}s", flush=True, ) print("\n=== B) Hash-cache per-genome execute (elites committed) ===", flush=True) mode_b, outs_b, fits_b = run_hash_per_genome( items, elite_items, images, labels, activation=args.activation, max_bytes=max_bytes, ) _attach_baseline_err(mode_b, outs_a, outs_b, fits_a, fits_b) mode_b.speedup_vs_baseline = ( mode_a.wall_s / mode_b.wall_s if mode_b.wall_s > 1e-9 else None ) print( f" wall={mode_b.wall_s:.3f}s peak_vram={mode_b.peak_vram_mb:.0f}MB " f"speedup={mode_b.speedup_vs_baseline:.3f}x " f"max_err={mode_b.max_abs_err:.3e} " f"stats={mode_b.plan}", flush=True, ) print("\n=== C) Planned unique-miss wave ===", flush=True) mode_c, _outs_c, _fits_c = run_planned_wave( items, elite_items, images, labels, activation=args.activation, max_bytes=max_bytes, baseline_outs=outs_a, baseline_fits=fits_a, ) mode_c.speedup_vs_baseline = ( mode_a.wall_s / mode_c.wall_s if mode_c.wall_s > 1e-9 else None ) print( f" wall={mode_c.wall_s:.3f}s (plan={mode_c.plan_s:.3f}s gpu={mode_c.gpu_s:.3f}s) " f"peak_vram={mode_c.peak_vram_mb:.0f}MB " f"speedup={mode_c.speedup_vs_baseline:.3f}x " f"max_err={mode_c.max_abs_err:.3e} " f"plan={mode_c.plan}", flush=True, ) # Per-genome table (baseline times + planned err) print("\n=== Per-genome (baseline wall | planned abs-err) ===", flush=True) print( f"{'gid':>10} {'nodes':>5} {'edges':>5} {'A_s':>8} {'B_s':>8} " f"{'A_fit':>7} {'err_B':>10} {'err_C':>10}", flush=True, ) b_by = {tuple(r.genome_id): r for r in mode_b.genomes} c_by = {tuple(r.genome_id): r for r in mode_c.genomes} for row in mode_a.genomes: key = tuple(row.genome_id) br = b_by[key] cr = c_by[key] print( f"{str(key):>10} {row.nodes:5d} {row.edges:5d} " f"{row.wall_s:8.4f} {br.wall_s:8.4f} " f"{row.fitness:7.3f} {br.max_abs_err_vs_baseline:10.3e} " f"{cr.max_abs_err_vs_baseline:10.3e}", flush=True, ) report = { "device": torch.cuda.get_device_name(0), "pop": args.pop, "images": args.images, "n_elite": n_elite, "activation": args.activation, "seed": args.seed, "modes": [asdict(mode_a), asdict(mode_b), asdict(mode_c)], "summary": { "baseline_wall_s": mode_a.wall_s, "hash_per_genome_wall_s": mode_b.wall_s, "planned_wall_s": mode_c.wall_s, "speedup_hash_vs_naive": mode_b.speedup_vs_baseline, "speedup_planned_vs_naive": mode_c.speedup_vs_baseline, "speedup_planned_vs_hash": ( mode_b.wall_s / mode_c.wall_s if mode_c.wall_s > 1e-9 else None ), "max_abs_err_hash": mode_b.max_abs_err, "max_abs_err_planned": mode_c.max_abs_err, "mean_fitness_abs_delta_hash": mode_b.mean_fitness_abs_delta, "mean_fitness_abs_delta_planned": mode_c.mean_fitness_abs_delta, "mean_nodes": sum(r.nodes for r in mode_a.genomes) / max(1, len(mode_a.genomes)), "mean_edges": sum(r.edges for r in mode_a.genomes) / max(1, len(mode_a.genomes)), "correct": bool( mode_b.max_abs_err < 1e-4 and mode_c.max_abs_err < 1e-4 ), }, } out_path = args.out / "ab_report.json" out_path.write_text(json.dumps(report, indent=2), encoding="utf-8") s = report["summary"] print("\n=== SUMMARY ===", flush=True) print( f"Naive A: {s['baseline_wall_s']:.3f}s\n" f"Hash B: {s['hash_per_genome_wall_s']:.3f}s " f"({s['speedup_hash_vs_naive']:.2f}x vs A) max_err={s['max_abs_err_hash']:.3e}\n" f"Planned C: {s['planned_wall_s']:.3f}s " f"({s['speedup_planned_vs_naive']:.2f}x vs A, " f"{s['speedup_planned_vs_hash']:.2f}x vs B) " f"max_err={s['max_abs_err_planned']:.3e}\n" f"Correct (err<1e-4): {s['correct']}", flush=True, ) print(f"Wrote {out_path}", flush=True) return 0 if s["correct"] else 1 if __name__ == "__main__": raise SystemExit(main())