Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Binary file added docs/figures/neuron/phase15/loss_curves.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file added docs/figures/neuron/phase15/sparsity_tradeoff.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
391 changes: 391 additions & 0 deletions notebooks/02-function-level/14-phase15-sparsity-prune.ipynb
Original file line number Diff line number Diff line change
@@ -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
}
Loading
Loading