diff --git a/docs/figures/neuron/phase15/loss_curves.png b/docs/figures/neuron/phase15/loss_curves.png new file mode 100644 index 0000000..818bc2a Binary files /dev/null and b/docs/figures/neuron/phase15/loss_curves.png differ diff --git a/docs/figures/neuron/phase15/sparsity_tradeoff.png b/docs/figures/neuron/phase15/sparsity_tradeoff.png new file mode 100644 index 0000000..a0ad243 Binary files /dev/null and b/docs/figures/neuron/phase15/sparsity_tradeoff.png differ diff --git a/notebooks/02-function-level/14-phase15-sparsity-prune.ipynb b/notebooks/02-function-level/14-phase15-sparsity-prune.ipynb new file mode 100644 index 0000000..cbfee85 --- /dev/null +++ b/notebooks/02-function-level/14-phase15-sparsity-prune.ipynb @@ -0,0 +1,391 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "# 14-phase15-sparsity-prune\n", + "\n", + "**neuron Phase 15** — paradigm 의 **정적 → 동적 위상 진입** 의 첫 단계. Phase 14 까지는 dense topology + learned magnitude 였고, Phase 15 는 학습 중 **edge 자체를 영구 제거** (sparsity-driven prune).\n", + "\n", + "핵심 가설:\n", + "1. **prune 후에도 학습 지속** — pruned edge 의 gradient 가 차단됨 (resurrection 방지) 으로 학습 자체가 깨지지 않음?\n", + "2. **sparsity vs final_loss trade-off** — 30% / 50% / 70% sparsity 에서 loss 가 얼마나 악화?\n", + "3. **soft degradation** — sparsity 가 증가해도 loss 가 catastrophic 하게 폭발하지 않고 gradual?\n", + "4. **post-prune recovery** — prune 직후의 loss spike 가 이후 학습으로 어느 정도 복구?\n", + "\n", + "설계: full graph block (`hybrid_around_one_around_one`, Phase 14 최저 loss 구조) 위에서 **prune fraction 4 단계** × 2 seed = 8 run.\n", + "- mode: `dense` (no prune) / `prune_0.3` / `prune_0.5` / `prune_0.7`\n", + "- prune 시점: max_steps / 2 = 750 (학습 중간)\n", + "- 그 외 hyperparameter 는 Phase 14 와 동일 (공정 비교)\n", + "\n", + "데이터: TinyShakespeare (char-LM, block_size=64)\n", + "시드: [42, 123]\n", + "작성일: 2026-05-27\n", + "연관: Issue [#71](https://github.com/EinSofINTEREST/GraphLM/issues/71) / Phase 14 baseline PR [#70](https://github.com/EinSofINTEREST/GraphLM/pull/70)" + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "## 0. 환경 / 의존성" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "from __future__ import annotations\n", + "\n", + "import math\n", + "import statistics\n", + "from pathlib import Path\n", + "\n", + "import matplotlib.pyplot as plt\n", + "import torch\n", + "\n", + "from graphlm.data.tinyshakespeare import (\n", + " CharTokenizer,\n", + " TinyShakespeareDataset,\n", + " load_tinyshakespeare_text,\n", + ")\n", + "from graphlm.neuron.hybrid_transformer_demo import (\n", + " HybridGraphTransformerLM,\n", + " HybridTransformerTrainConfig,\n", + " count_parameters,\n", + " train_hybrid_transformer_lm,\n", + ")\n", + "from graphlm.utils import safe_perplexity\n", + "\n", + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "print(f\"device: {device}\")\n", + "print(f\"torch: {torch.__version__}\")" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. Config + 데이터" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "text = load_tinyshakespeare_text()\n", + "tokenizer = CharTokenizer(text)\n", + "dataset = TinyShakespeareDataset(text, tokenizer)\n", + "vocab_size = tokenizer.vocab_size\n", + "print(f\"vocab_size = {vocab_size}, dataset size = {len(dataset)}\")\n", + "\n", + "# Phase 14 와 동일 hyperparameter — full graph 위에서 prune 검증\n", + "HIDDEN_DIM = 128\n", + "N_HEADS = 4\n", + "FFN_DIM = 256\n", + "N_LAYERS = 4\n", + "GROUP_SIZE = 16\n", + "BLOCK_SIZE = 64\n", + "BATCH_SIZE = 32\n", + "LR = 3e-4\n", + "MAX_STEPS = 1500\n", + "PRUNE_AT_STEP = MAX_STEPS // 2 # 750 = 학습 중간\n", + "SEEDS = [42, 123]\n", + "ARCH = \"hybrid_around_one_around_one\" # Phase 14 최저 loss 구조 고정\n", + "PRUNE_FRACTIONS = [0.0, 0.3, 0.5, 0.7] # dense, 30%, 50%, 70%\n", + "\n", + "# 모델 파라미터 수 (prune 전)\n", + "m = HybridGraphTransformerLM(\n", + " vocab_size=vocab_size,\n", + " hidden_dim=HIDDEN_DIM,\n", + " n_heads=N_HEADS,\n", + " ffn_dim=FFN_DIM,\n", + " n_layers=N_LAYERS,\n", + " max_seq_len=BLOCK_SIZE,\n", + " arch=ARCH,\n", + " group_size=GROUP_SIZE,\n", + " use_full_graph=True,\n", + ")\n", + "print(f\"\\nfull graph 모델 (prune 전): params = {count_parameters(m):,}\")\n", + "print(f\"prune 시점: step {PRUNE_AT_STEP} / {MAX_STEPS} (학습 중간)\")" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. Sweep 실행 (4 prune fraction × 2 seed = 8 run)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "results = {}\n", + "for frac in PRUNE_FRACTIONS:\n", + " for seed in SEEDS:\n", + " key = (frac, seed)\n", + " mode = \"dense\" if frac == 0.0 else f\"prune_{frac}\"\n", + " print(f\"\\n== mode={mode} seed={seed} ==\")\n", + " cfg = HybridTransformerTrainConfig(\n", + " dataset=dataset,\n", + " vocab_size=vocab_size,\n", + " hidden_dim=HIDDEN_DIM,\n", + " n_heads=N_HEADS,\n", + " ffn_dim=FFN_DIM,\n", + " n_layers=N_LAYERS,\n", + " group_size=GROUP_SIZE,\n", + " arch=ARCH,\n", + " use_full_graph=True,\n", + " block_size=BLOCK_SIZE,\n", + " batch_size=BATCH_SIZE,\n", + " lr=LR,\n", + " max_steps=MAX_STEPS,\n", + " prune_at_step=PRUNE_AT_STEP if frac > 0 else None,\n", + " prune_fraction=frac,\n", + " seed=seed,\n", + " device=device,\n", + " )\n", + " out = train_hybrid_transformer_lm(cfg)\n", + " results[key] = out\n", + " sparsity = out[\"final_sparsity\"]\n", + " print(\n", + " f\" final_loss = {out['final_loss']:.4f} (perplexity = {safe_perplexity(out['final_loss']):.2f})\"\n", + " f\" sparsity = {sparsity:.3f}\"\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 3. 결과 표 + 자동 verdict" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "print(f\"{'mode':>10s} {'seed':>6s} {'final_loss':>12s} {'perplexity':>12s} {'sparsity':>10s}\")\n", + "print(\"-\" * 60)\n", + "for (frac, seed), out in results.items():\n", + " fl = out[\"final_loss\"]\n", + " mode = \"dense\" if frac == 0.0 else f\"prune_{frac}\"\n", + " print(\n", + " f\"{mode:>10s} {seed:>6d} {fl:>12.4f} {safe_perplexity(fl):>12.2f} {out['final_sparsity']:>10.3f}\"\n", + " )\n", + "\n", + "# fraction 별 평균\n", + "print(\"\\n== Fraction summary (mean ± σ across seeds) ==\")\n", + "summary = {}\n", + "for frac in PRUNE_FRACTIONS:\n", + " vals = [results[(frac, s)][\"final_loss\"] for s in SEEDS]\n", + " sparsities = [results[(frac, s)][\"final_sparsity\"] for s in SEEDS]\n", + " mean = statistics.mean(vals)\n", + " std = statistics.stdev(vals) if len(vals) > 1 else 0.0\n", + " sp_mean = statistics.mean(sparsities)\n", + " summary[frac] = (mean, std, sp_mean)\n", + " mode = \"dense\" if frac == 0.0 else f\"prune_{frac}\"\n", + " print(\n", + " f\" {mode:>10s} {mean:.4f} ± {std:.4f} (perplexity ≈ {safe_perplexity(mean):.2f}) sparsity={sp_mean:.3f}\"\n", + " )\n", + "\n", + "# 자동 verdict\n", + "print(\"\\n== Verdict ==\")\n", + "dense_loss = summary[0.0][0]\n", + "\n", + "# 1. prune 후에도 학습 지속 — 모든 prune mode 가 finite loss 로 종료\n", + "all_finite = all(math.isfinite(out[\"final_loss\"]) for out in results.values())\n", + "verdict_1 = \"PASS\" if all_finite else \"FAIL\"\n", + "print(f\"1. all-finite after prune: {all_finite} [{verdict_1}]\")\n", + "\n", + "# 2. soft degradation — 70% prune 도 dense + 1.0 이내\n", + "p70_loss = summary[0.7][0]\n", + "diff_70 = p70_loss - dense_loss\n", + "verdict_2 = \"PASS\" if diff_70 < 1.0 else \"FAIL\"\n", + "print(f\"2. soft degradation (70% prune ≤ dense + 1.0): diff = {diff_70:+.4f} [{verdict_2}]\")\n", + "\n", + "# 3. monotonic degradation — sparsity 증가에 따라 loss 비감소\n", + "losses_by_frac = [summary[f][0] for f in PRUNE_FRACTIONS]\n", + "monotone = all(\n", + " losses_by_frac[i] <= losses_by_frac[i + 1] + 0.05 for i in range(len(losses_by_frac) - 1)\n", + ")\n", + "verdict_3 = \"PASS\" if monotone else \"FAIL\"\n", + "print(f\"3. monotonic loss vs sparsity (±0.05 tolerance): {monotone} [{verdict_3}]\")\n", + "\n", + "# 4. moderate prune (30%) 은 거의 무손실 — dense + 0.1 이내\n", + "p30_loss = summary[0.3][0]\n", + "diff_30 = p30_loss - dense_loss\n", + "verdict_4 = \"PASS\" if diff_30 < 0.1 else \"FAIL\"\n", + "print(f\"4. moderate prune (30%) ≤ dense + 0.1: diff = {diff_30:+.4f} [{verdict_4}]\")\n", + "\n", + "# prune 발생 정보\n", + "print(\"\\n== Prune events ==\")\n", + "for (frac, seed), out in results.items():\n", + " if out[\"prune_event\"] is not None:\n", + " ev = out[\"prune_event\"]\n", + " print(\n", + " f\" frac={frac} seed={seed}: step={ev['step']}, total_pruned={ev['total_pruned']}, sparsity_after={ev['sparsity_after']:.3f}\"\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## 4. Loss curve 시각화 — prune step 표시" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "10", + "metadata": {}, + "outputs": [], + "source": [ + "fig, ax = plt.subplots(1, 1, figsize=(12, 6))\n", + "colors = {\n", + " 0.0: \"tab:gray\",\n", + " 0.3: \"tab:blue\",\n", + " 0.5: \"tab:orange\",\n", + " 0.7: \"tab:red\",\n", + "}\n", + "window = 50\n", + "\n", + "for frac in PRUNE_FRACTIONS:\n", + " losses_per_seed = [results[(frac, s)][\"losses\"] for s in SEEDS]\n", + " # rolling mean — slice 시작점 +1 시프트로 window 와 divisor 일치\n", + " smoothed = []\n", + " for losses in losses_per_seed:\n", + " smoothed.append(\n", + " [\n", + " sum(losses[max(0, i - window + 1) : i + 1]) / min(i + 1, window)\n", + " for i in range(len(losses))\n", + " ]\n", + " )\n", + " arr = torch.tensor(smoothed)\n", + " mean = arr.mean(dim=0)\n", + " std = arr.std(dim=0)\n", + " steps = list(range(len(mean)))\n", + " mode = \"dense\" if frac == 0.0 else f\"prune {int(frac * 100)}%\"\n", + " ax.plot(steps, mean, label=mode, color=colors[frac], linewidth=1.5)\n", + " ax.fill_between(steps, mean - std, mean + std, color=colors[frac], alpha=0.15)\n", + "\n", + "# prune step 수직선\n", + "ax.axvline(\n", + " PRUNE_AT_STEP, color=\"black\", linestyle=\":\", alpha=0.5, label=f\"prune @ step {PRUNE_AT_STEP}\"\n", + ")\n", + "\n", + "ax.set_xlabel(\"step\")\n", + "ax.set_ylabel(f\"loss (rolling mean w={window})\")\n", + "ax.set_title(\n", + " f\"Phase 15 — sparsity-driven prune ({ARCH}, full graph): dense vs 30/50/70% prune (mean ± σ over 2 seeds)\"\n", + ")\n", + "ax.legend(loc=\"upper right\")\n", + "ax.grid(alpha=0.3)\n", + "plt.tight_layout()\n", + "\n", + "out_dir = Path(\"../../runs/notebook-neuron-phase15\")\n", + "out_dir.mkdir(parents=True, exist_ok=True)\n", + "fig.savefig(out_dir / \"loss_curves.png\", dpi=150, bbox_inches=\"tight\")\n", + "plt.show()\n", + "print(f\"saved: {out_dir / 'loss_curves.png'}\")" + ] + }, + { + "cell_type": "markdown", + "id": "11", + "metadata": {}, + "source": [ + "## 5. Sparsity vs final_loss trade-off" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "12", + "metadata": {}, + "outputs": [], + "source": [ + "fig, ax = plt.subplots(1, 1, figsize=(8, 5))\n", + "fracs = [f for f in PRUNE_FRACTIONS]\n", + "means = [summary[f][0] for f in fracs]\n", + "stds = [summary[f][1] for f in fracs]\n", + "sparsities = [summary[f][2] for f in fracs]\n", + "\n", + "ax.errorbar(sparsities, means, yerr=stds, marker=\"o\", capsize=4, linewidth=1.5, color=\"tab:blue\")\n", + "for sp, m, frac in zip(sparsities, means, fracs, strict=True):\n", + " label = \"dense\" if frac == 0.0 else f\"{int(frac * 100)}%\"\n", + " ax.annotate(label, (sp, m), xytext=(8, -8), textcoords=\"offset points\", fontsize=10)\n", + "\n", + "ax.set_xlabel(\"effective sparsity (mask=0 비율)\")\n", + "ax.set_ylabel(\"final loss (last 100 mean ± σ)\")\n", + "ax.set_title(\"Phase 15 — sparsity vs final_loss trade-off\")\n", + "ax.grid(alpha=0.3)\n", + "plt.tight_layout()\n", + "fig.savefig(out_dir / \"sparsity_tradeoff.png\", dpi=150, bbox_inches=\"tight\")\n", + "plt.show()\n", + "print(f\"saved: {out_dir / 'sparsity_tradeoff.png'}\")" + ] + }, + { + "cell_type": "markdown", + "id": "13", + "metadata": {}, + "source": [ + "## 6. 결론 / 다음 단계\n", + "\n", + "(셀 출력 보고 사용자가 채울 영역)\n", + "\n", + "- soft degradation 곡선의 형태 (linear / convex / cliff)?\n", + "- 30% prune 의 무손실 / 70% 의 cliff 여부?\n", + "- prune 직후 loss spike 와 회복 양상?\n", + "\n", + "**Phase 16 후보**:\n", + "- **Net2Net / LiGO 식 grow** — pruned slot 에 신규 edge 추가 (function preservation 유지), dynamic grow + shrink 동시\n", + "- **RigL / SET 식 dynamic sparse training** — 매 prune step 후 같은 수의 edge 를 다른 위치에서 재할당 (constant sparsity 유지하며 topology 만 변화)\n", + "- **layer-wise 차등 prune** — attention vs FFN 별로 다른 sparsity target" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "GraphLM (uv .venv)", + "language": "python", + "name": "graphlm-uv-venv" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/src/graphlm/neuron/graph_hybrid.py b/src/graphlm/neuron/graph_hybrid.py index cafcacc..16f56f0 100644 --- a/src/graphlm/neuron/graph_hybrid.py +++ b/src/graphlm/neuron/graph_hybrid.py @@ -160,6 +160,14 @@ def __init__( else: self.register_parameter("bias", None) + # Phase 15: edge_mask — buffer (학습 X, state_dict 에 포함). 초기값 모두 1 (no prune). + # prune_by_magnitude() 호출 시 일부 위치가 영구 0 으로 전환 → 해당 edge 의 forward + # 기여도 0 + gradient chain 도 mask=0 위치에서 끊김 → resurrection 방지. + self.register_buffer( + "edge_mask", + torch.ones(self.n_groups_out, self.n_groups_in, group_size, group_size), + ) + def forward(self, x: Tensor) -> Tensor: # x: (..., in_features) → (..., n_groups_in, group_size) *batch, in_f = x.shape @@ -169,10 +177,13 @@ def forward(self, x: Tensor) -> Tensor: # 메모리 효율 최적화 (gemini #3302293739): adj_outer 와 adj_inner 를 weight 수준에서 # 미리 결합 → (*batch, G_out, G_in, k) 의 큰 intermediate tensor 회피. - # 수학적 등치: contrib[..., go, gi, ko] = Σ_ki adj_inner[go,gi,ki,ko] · W[go,gi,ki,ko] · x_g[..., gi, ki] - # y[..., go, ko] = Σ_gi adj_outer[go, gi] · contrib[..., go, gi, ko] - # = Σ_gi Σ_ki (adj_outer · adj_inner · W)[...] · x_g[..., gi, ki] - eff_w = self.adj_outer.unsqueeze(-1).unsqueeze(-1) * self.adj_inner * self.weight + # Phase 15: edge_mask 도 같은 위치에 곱함 → pruned edge 의 forward 0 + gradient 차단. + eff_w = ( + self.adj_outer.unsqueeze(-1).unsqueeze(-1) + * self.adj_inner + * self.weight + * self.edge_mask + ) # single einsum — output (*batch, G_out, k), 중간 (*batch, G_out, G_in, k) 텐서 없음 y_g = torch.einsum("...gi,Ggik->...Gk", x_g, eff_w) # flatten back to (..., out_features) @@ -201,6 +212,80 @@ def freeze_adj_outer(self) -> None: def freeze_adj_inner(self) -> None: self.adj_inner.requires_grad_(False) + # ── Phase 15: edge prune ──────────────────────────────────── + + def effective_edge_magnitude(self) -> Tensor: + """현재 forward 에 적용되는 edge magnitude (shape: (G_out, G_in, k, k)). + + = ``|adj_outer · adj_inner · W| · edge_mask`` — 이미 pruned 된 edge 는 0. + prune 결정 / sparsity 측정에 사용. + """ + with torch.no_grad(): + return ( + self.adj_outer.unsqueeze(-1).unsqueeze(-1) * self.adj_inner * self.weight + ).abs() * self.edge_mask + + def prune_by_magnitude(self, threshold: float) -> int: + """edge magnitude < threshold 인 위치를 영구 prune (mask=0). + + gradient resurrection 방지: forward 의 ``eff_w *= edge_mask`` 곱셈으로 + 해당 위치의 gradient 가 ``adj_outer``, ``adj_inner``, ``weight`` 모두에서 0 으로 차단됨. + + Args: + threshold: |adj_outer · adj_inner · W| 의 절대값 임계값. + + Returns: + 이번 호출에서 신규로 prune 된 edge 수. + """ + if threshold < 0: + raise ValueError(f"threshold must be >= 0, got {threshold}") + with torch.no_grad(): + mag = self.effective_edge_magnitude() + new_dead = (mag < threshold) & (self.edge_mask > 0) + n_pruned = int(new_dead.sum().item()) + self.edge_mask[new_dead] = 0.0 + return n_pruned + + def prune_bottom_fraction(self, fraction: float) -> int: + """살아있는 edge 중 magnitude 하위 ``fraction`` 비율을 prune. + + threshold 기반보다 target sparsity 제어가 용이. fraction=0.3 → 살아있는 edge 의 30% 추가 prune. + + Args: + fraction: 0.0 ~ 1.0, 살아있는 edge 중 prune 할 비율. + + Returns: + 이번 호출에서 신규로 prune 된 edge 수. + """ + if not 0.0 <= fraction <= 1.0: + raise ValueError(f"fraction must be in [0, 1], got {fraction}") + with torch.no_grad(): + mag = self.effective_edge_magnitude() + alive_mask = self.edge_mask > 0 + alive_count = int(alive_mask.sum().item()) + n_to_prune = int(alive_count * fraction) + if n_to_prune == 0: + return 0 + # 살아있는 edge 의 4D 인덱스 + magnitude + alive_indices = torch.nonzero(alive_mask, as_tuple=False) # (alive_count, 4) + alive_mag = mag[alive_mask] # (alive_count,) + # torch.topk(largest=False) 로 하위 n 개의 정확한 인덱스 추출 — tie-breaking 결정론적 + # (gemini #3307531740): kthvalue + mag<=kth 방식은 동률 시 의도보다 많이 prune 위험 + _, topk_idx = torch.topk(alive_mag, n_to_prune, largest=False) + prune_4d = alive_indices[topk_idx] # (n_to_prune, 4) + self.edge_mask[prune_4d[:, 0], prune_4d[:, 1], prune_4d[:, 2], prune_4d[:, 3]] = 0.0 + return n_to_prune + + def effective_sparsity(self) -> float: + """전체 edge 중 영구 prune (mask=0) 비율.""" + with torch.no_grad(): + return float((self.edge_mask == 0).float().mean().item()) + + def n_alive_edges(self) -> int: + """살아있는 edge 수 (mask=1).""" + with torch.no_grad(): + return int((self.edge_mask > 0).sum().item()) + def extra_repr(self) -> str: return ( f"in_features={self.in_features}, out_features={self.out_features}, " diff --git a/src/graphlm/neuron/hybrid_transformer_demo.py b/src/graphlm/neuron/hybrid_transformer_demo.py index 8e95892..1fd6415 100644 --- a/src/graphlm/neuron/hybrid_transformer_demo.py +++ b/src/graphlm/neuron/hybrid_transformer_demo.py @@ -21,6 +21,7 @@ from torch import Tensor, nn from graphlm.data.tinyshakespeare import TinyShakespeareDataset, iter_random_batches +from graphlm.neuron.graph_hybrid import HybridGraphLinear from graphlm.neuron.hybrid_transformer import ( Arch, FullGraphTransformerBlock, @@ -61,10 +62,26 @@ class HybridTransformerTrainConfig: lr: float = 3e-4 max_steps: int = 1500 + # Phase 15: edge prune (one-shot at midpoint by default) + # prune_at_step=None → no prune. 0 < step ≤ max_steps 면 해당 step 끝에서 prune 실행. + prune_at_step: int | None = None + # 살아있는 edge 중 하위 magnitude 비율 (prune_bottom_fraction 사용). 0.0 → no-op. + prune_fraction: float = 0.0 + # runtime seed: int = 0 device: str = "cpu" + def __post_init__(self) -> None: + # Phase 15 prune 인자 입력 검증 (Copilot #3307536553) — 잘못된 값의 silent no-op 회피. + if not 0.0 <= self.prune_fraction <= 1.0: + raise ValueError(f"prune_fraction must be in [0, 1], got {self.prune_fraction}") + if self.prune_at_step is not None and not 1 <= self.prune_at_step <= self.max_steps: + raise ValueError( + f"prune_at_step must be in [1, max_steps={self.max_steps}], " + f"got {self.prune_at_step}" + ) + class HybridGraphTransformerLM(nn.Module): """Small char-LM Transformer with arch-dispatched FFN. @@ -160,11 +177,40 @@ def _snapshot_adj(model: HybridGraphTransformerLM) -> list[dict[str, dict[str, T return snapshots +def _prune_model(model: nn.Module, fraction: float) -> dict[str, int]: + """모델의 모든 HybridGraphLinear 에 prune_bottom_fraction 적용. + + Returns: + per-layer pruned count (디버깅용 — layer 이름 → 신규 prune edge 수). + """ + counts: dict[str, int] = {} + for name, mod in model.named_modules(): + if isinstance(mod, HybridGraphLinear): + counts[name] = mod.prune_bottom_fraction(fraction) + return counts + + +def _model_sparsity(model: nn.Module) -> float: + """모든 HybridGraphLinear edge 전체에 대한 평균 sparsity.""" + total = 0 + dead = 0 + for mod in model.modules(): + if isinstance(mod, HybridGraphLinear): + total += mod.edge_mask.numel() + dead += int((mod.edge_mask == 0).sum().item()) + if total == 0: + return 0.0 + return dead / total + + def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: - """1 run 학습 — Phase 13 sweep unit. + """1 run 학습 — Phase 13/14/15 sweep unit. - Returns: ``losses``, ``final_loss`` (last 100 mean), ``final_adj`` - (hybrid arch 인 경우 block 별 fc1/fc2 outer/inner snapshot list). + Phase 15: ``config.prune_at_step`` 에 도달 시 ``config.prune_fraction`` 만큼 prune. + plain arch 는 HybridGraphLinear 가 없어 prune 무효 (HybridGraphTransformerLM 의 plain 도 동일). + + Returns: ``losses``, ``final_loss`` (last 100 mean), ``final_adj``, ``final_sparsity``, + ``prune_event`` (prune 실행 시점의 step + 신규 prune edge 수). """ set_seed(config.seed) model = HybridGraphTransformerLM( @@ -184,8 +230,9 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: ) optimizer = torch.optim.AdamW(model.parameters(), lr=config.lr) losses: list[float] = [] + prune_event: dict | None = None model.train() - for _step in range(1, config.max_steps + 1): + for step in range(1, config.max_steps + 1): x, y = next(data_iter) x, y = x.to(config.device), y.to(config.device) optimizer.zero_grad() @@ -195,6 +242,19 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: optimizer.step() losses.append(loss.item()) + # Phase 15: prune at midpoint (or configured step) + if ( + config.prune_at_step is not None + and step == config.prune_at_step + and config.prune_fraction > 0 + ): + per_layer = _prune_model(model, config.prune_fraction) + prune_event = { + "step": step, + "total_pruned": sum(per_layer.values()), + "sparsity_after": _model_sparsity(model), + } + n_last = min(100, len(losses)) final_loss = sum(losses[-n_last:]) / n_last if n_last > 0 else 0.0 @@ -202,6 +262,8 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: "losses": losses, "final_loss": final_loss, "final_adj": _snapshot_adj(model), + "final_sparsity": _model_sparsity(model), + "prune_event": prune_event, } diff --git a/tests/neuron/test_graph_hybrid.py b/tests/neuron/test_graph_hybrid.py index b0c5907..f1d343a 100644 --- a/tests/neuron/test_graph_hybrid.py +++ b/tests/neuron/test_graph_hybrid.py @@ -140,3 +140,151 @@ def test_freeze_helpers(): assert lin.weight.requires_grad lin.freeze_adj_inner() assert not lin.adj_inner.requires_grad + + +# ── Phase 15: edge prune ──────────────────────────────────────── + + +def test_edge_mask_initial_all_ones(): + """초기 edge_mask 는 전부 1 (no prune).""" + lin = HybridGraphLinear(16, 16, group_size=4) + assert lin.edge_mask.shape == (4, 4, 4, 4) + assert torch.all(lin.edge_mask == 1.0) + assert lin.effective_sparsity() == 0.0 + assert lin.n_alive_edges() == 16 * 16 # 4·4·4·4 + + +def test_forward_with_initial_mask_unchanged(): + """edge_mask=1 (초기) 일 때 forward 가 mask 없이 계산한 값과 정확히 동일 (function preservation). + + Copilot #3307536521 — 이름과 검증 일치: mask=1 의 forward 가 mask 적용 안 한 직접 계산과 같은지 직접 비교. + """ + torch.manual_seed(0) + lin = HybridGraphLinear(16, 24, group_size=4) + x = torch.randn(2, 8, 16) + y = lin(x) + + # mask 없는 forward 직접 계산 (mask=1 이므로 결과 동일해야 함) + with torch.no_grad(): + x_g = x.reshape(2, 8, lin.n_groups_in, lin.group_size) + eff_w_no_mask = lin.adj_outer.unsqueeze(-1).unsqueeze(-1) * lin.adj_inner * lin.weight + y_g = torch.einsum("...gi,Ggik->...Gk", x_g, eff_w_no_mask) + expected = y_g.reshape(2, 8, lin.out_features) + lin.bias + assert torch.allclose(y, expected, atol=1e-6), ( + f"mask=1 forward 가 mask 없는 계산과 달라짐, max |diff| = {(y - expected).abs().max().item()}" + ) + + # 추가 sanity: mask 를 모두 0 으로 만들면 출력은 bias 만 남음 + with torch.no_grad(): + lin.edge_mask.zero_() + y_zero = lin(x) + bias_expected = lin.bias.expand(2, 8, 24) + assert torch.allclose(y_zero, bias_expected, atol=1e-6) + + +def test_prune_by_magnitude_basic(): + """threshold 이하 edge 가 mask=0 으로 prune.""" + torch.manual_seed(0) + lin = HybridGraphLinear(16, 16, group_size=4) + initial_mag = lin.effective_edge_magnitude() + # threshold = 50 percentile 로 잡으면 절반 정도 prune + threshold = float(initial_mag.median().item()) + n_pruned = lin.prune_by_magnitude(threshold) + assert n_pruned > 0 + sparsity = lin.effective_sparsity() + assert 0.4 < sparsity < 0.6, f"median threshold 면 ~50% sparsity, got {sparsity}" + + +def test_prune_idempotent_below_threshold(): + """같은 threshold 로 두 번 호출 시 두 번째는 0 신규 prune.""" + torch.manual_seed(0) + lin = HybridGraphLinear(16, 16, group_size=4) + threshold = float(lin.effective_edge_magnitude().median().item()) + n1 = lin.prune_by_magnitude(threshold) + n2 = lin.prune_by_magnitude(threshold) + assert n1 > 0 + assert n2 == 0, f"두 번째 prune 은 신규 0 이어야 함, got {n2}" + + +def test_pruned_edges_do_not_resurrect_via_gradient(): + """pruned edge 의 weight/adj gradient 가 정확히 0 — optimizer 가 살릴 수 없음.""" + torch.manual_seed(0) + lin = HybridGraphLinear( + 16, + 16, + group_size=4, + adj_outer_init="uniform_around_one", + adj_inner_init="uniform_around_one", + ) + threshold = float(lin.effective_edge_magnitude().median().item()) + lin.prune_by_magnitude(threshold) + dead_positions = lin.edge_mask == 0 + + x = torch.randn(2, 16) + lin(x).sum().backward() + # weight gradient: pruned 위치는 0 + assert torch.all(lin.weight.grad[dead_positions] == 0), "weight grad at pruned == 0" + # adj_inner gradient: pruned 위치는 0 + assert torch.all(lin.adj_inner.grad[dead_positions] == 0), "adj_inner grad at pruned == 0" + + +def test_prune_bottom_fraction(): + """fraction=0.3 → 살아있는 edge 의 정확히 30% prune (topk 기반 deterministic). + + gemini #3307531745 — topk 로 정확 n개 prune 보장되므로 tolerance 불필요. + """ + torch.manual_seed(0) + lin = HybridGraphLinear(16, 16, group_size=4) + alive_before = lin.n_alive_edges() + n_pruned = lin.prune_bottom_fraction(0.3) + expected = int(alive_before * 0.3) + assert n_pruned == expected, f"target {expected}, got {n_pruned}" + + +def test_prune_bottom_fraction_tie_breaking_deterministic(): + """모든 magnitude 가 동률일 때도 정확히 n_to_prune 만큼만 prune (no over-prune). + + gemini #3307531740 의 시나리오 — uniform 같은 magnitude → kthvalue 방식은 100% prune 위험. + """ + lin = HybridGraphLinear(16, 16, group_size=4) + # 모든 magnitude 가 정확히 같도록 설정 + with torch.no_grad(): + lin.weight.fill_(1.0) + lin.adj_outer.fill_(1.0) + lin.adj_inner.fill_(1.0) + alive_before = lin.n_alive_edges() + n_pruned = lin.prune_bottom_fraction(0.3) + expected = int(alive_before * 0.3) + assert n_pruned == expected + # over-prune 안 됨 — 살아있는 edge 수가 expected 만큼 줄음 + assert lin.n_alive_edges() == alive_before - expected + + +def test_prune_bottom_fraction_zero_fraction_noop(): + lin = HybridGraphLinear(16, 16, group_size=4) + assert lin.prune_bottom_fraction(0.0) == 0 + + +def test_prune_negative_threshold_rejected(): + lin = HybridGraphLinear(16, 16, group_size=4) + with pytest.raises(ValueError, match="must be >= 0"): + lin.prune_by_magnitude(-0.1) + + +def test_prune_invalid_fraction_rejected(): + lin = HybridGraphLinear(16, 16, group_size=4) + with pytest.raises(ValueError, match=r"must be in \[0, 1\]"): + lin.prune_bottom_fraction(1.5) + + +def test_edge_mask_in_state_dict(): + """edge_mask 가 state_dict 에 포함되어 save/load 보존.""" + lin1 = HybridGraphLinear(16, 16, group_size=4) + lin1.prune_by_magnitude(float(lin1.effective_edge_magnitude().median().item())) + sparsity_before = lin1.effective_sparsity() + assert sparsity_before > 0 + + lin2 = HybridGraphLinear(16, 16, group_size=4) + lin2.load_state_dict(lin1.state_dict()) + assert lin2.effective_sparsity() == sparsity_before + assert torch.all(lin2.edge_mask == lin1.edge_mask) diff --git a/tests/neuron/test_hybrid_transformer_demo.py b/tests/neuron/test_hybrid_transformer_demo.py new file mode 100644 index 0000000..e265b47 --- /dev/null +++ b/tests/neuron/test_hybrid_transformer_demo.py @@ -0,0 +1,59 @@ +"""Tests for HybridTransformerTrainConfig validation (Phase 15).""" + +from __future__ import annotations + +import pytest + +from graphlm.data.tinyshakespeare import CharTokenizer, TinyShakespeareDataset +from graphlm.neuron.hybrid_transformer_demo import HybridTransformerTrainConfig + + +def _dummy_dataset() -> TinyShakespeareDataset: + text = "abcdefghij" * 100 + tok = CharTokenizer(text) + return TinyShakespeareDataset(text, tok) + + +def _base_kwargs() -> dict: + ds = _dummy_dataset() + return dict( + dataset=ds, + vocab_size=10, + hidden_dim=16, + n_heads=4, + ffn_dim=32, + n_layers=2, + group_size=4, + arch="hybrid_full_full", + block_size=8, + batch_size=4, + lr=1e-3, + max_steps=100, + ) + + +def test_default_config_valid(): + """Phase 15 인자 default (prune_at_step=None, prune_fraction=0.0) 가 유효.""" + cfg = HybridTransformerTrainConfig(**_base_kwargs()) + assert cfg.prune_at_step is None + assert cfg.prune_fraction == 0.0 + + +@pytest.mark.parametrize("frac", [-0.1, 1.1, 2.0]) +def test_invalid_prune_fraction_rejected(frac): + """prune_fraction ∉ [0, 1] 는 __post_init__ 에서 거부 (Copilot #3307536553).""" + with pytest.raises(ValueError, match=r"prune_fraction must be in \[0, 1\]"): + HybridTransformerTrainConfig(**_base_kwargs(), prune_fraction=frac) + + +@pytest.mark.parametrize("step", [0, -1, 101, 9999]) +def test_invalid_prune_at_step_rejected(step): + """prune_at_step ∉ [1, max_steps] 는 거부.""" + with pytest.raises(ValueError, match="prune_at_step must be in"): + HybridTransformerTrainConfig(**_base_kwargs(), prune_at_step=step) + + +def test_valid_prune_config_accepted(): + cfg = HybridTransformerTrainConfig(**_base_kwargs(), prune_at_step=50, prune_fraction=0.3) + assert cfg.prune_at_step == 50 + assert cfg.prune_fraction == 0.3