diff --git a/docs/figures/neuron/phase16b/loss_curves.png b/docs/figures/neuron/phase16b/loss_curves.png new file mode 100644 index 0000000..21d9a99 Binary files /dev/null and b/docs/figures/neuron/phase16b/loss_curves.png differ diff --git a/docs/figures/neuron/phase16b/params_vs_loss.png b/docs/figures/neuron/phase16b/params_vs_loss.png new file mode 100644 index 0000000..4f80dcd Binary files /dev/null and b/docs/figures/neuron/phase16b/params_vs_loss.png differ diff --git a/notebooks/02-function-level/16-phase16b-net2net-grow.ipynb b/notebooks/02-function-level/16-phase16b-net2net-grow.ipynb new file mode 100644 index 0000000..4d8eb57 --- /dev/null +++ b/notebooks/02-function-level/16-phase16b-net2net-grow.ipynb @@ -0,0 +1,399 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "# 16-phase16b-net2net-grow\n", + "\n", + "**neuron Phase 16b** — paradigm 의 dynamic phase 두 번째 갈래. Phase 16a 가 *within-shape DST* (sparsity 재할당, parameter 수 동일) 였다면, Phase 16b 는 **cross-shape expansion** — 학습 중 `ffn_dim` 을 확장하여 *parameter 수 자체 증가*. **function preservation** 보장 (Net2Net-style: 새 input weight = 0).\n", + "\n", + "핵심 가설:\n", + "1. **function preservation** — grow 직후 forward output 이 grow 전과 정확히 동일?\n", + "2. **grown ≤ large-baseline + threshold** — 작게 시작 + 중간에 grow 한 모델이 처음부터 큰 모델 근방까지 학습?\n", + "3. **grown < small-baseline** — grow 가 capacity 부족 모델보다 명확히 좋음?\n", + "4. **all-finite** — grow 후 학습 안정성?\n", + "\n", + "설계: 3 mode × 2 seed = 6 run.\n", + "- `small_baseline`: ffn=128 끝까지 (small capacity)\n", + "- `grown`: ffn=128 시작 → step 750 에서 256 으로 grow\n", + "- `large_baseline`: ffn=256 끝까지 (large capacity)\n", + "\n", + "arch 고정: `hybrid_around_one_around_one` + `use_full_graph=True` (Phase 14 최저 loss 구조)\n", + "데이터: TinyShakespeare (char-LM, block_size=64)\n", + "시드: [42, 123]\n", + "작성일: 2026-05-27\n", + "연관: Issue [#77](https://github.com/EinSofINTEREST/GraphLM/issues/77) (Phase 16 main [#75](https://github.com/EinSofINTEREST/GraphLM/issues/75)) / Phase 16a PR [#78](https://github.com/EinSofINTEREST/GraphLM/pull/78)" + ] + }, + { + "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", + " HybridTransformerTrainConfig,\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", + "HIDDEN_DIM = 128\n", + "N_HEADS = 4\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", + "GROW_AT_STEP = MAX_STEPS // 2 # 750\n", + "SEEDS = [42, 123]\n", + "ARCH = \"hybrid_around_one_around_one\"\n", + "\n", + "FFN_SMALL = 128\n", + "FFN_LARGE = 256\n", + "\n", + "# 3 mode 정의\n", + "MODES = [\n", + " {\n", + " \"name\": \"small_baseline\",\n", + " \"ffn_dim\": FFN_SMALL,\n", + " \"grow_at_step\": None,\n", + " \"grow_ffn_target\": None,\n", + " },\n", + " {\n", + " \"name\": \"grown\",\n", + " \"ffn_dim\": FFN_SMALL,\n", + " \"grow_at_step\": GROW_AT_STEP,\n", + " \"grow_ffn_target\": FFN_LARGE,\n", + " },\n", + " {\n", + " \"name\": \"large_baseline\",\n", + " \"ffn_dim\": FFN_LARGE,\n", + " \"grow_at_step\": None,\n", + " \"grow_ffn_target\": None,\n", + " },\n", + "]\n", + "print(f\"\\nGrow 시점: step {GROW_AT_STEP} / {MAX_STEPS}, ffn_dim {FFN_SMALL} → {FFN_LARGE}\")" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. Sweep 실행 (3 mode × 2 seed = 6 run)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "results = {}\n", + "for mode in MODES:\n", + " for seed in SEEDS:\n", + " key = (mode[\"name\"], seed)\n", + " print(f\"\\n== mode={mode['name']} 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=mode[\"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", + " grow_at_step=mode[\"grow_at_step\"],\n", + " grow_ffn_target=mode[\"grow_ffn_target\"],\n", + " seed=seed,\n", + " device=device,\n", + " )\n", + " out = train_hybrid_transformer_lm(cfg)\n", + " results[key] = out\n", + " ev = out[\"grow_event\"]\n", + " print(\n", + " f\" final_loss = {out['final_loss']:.4f} (ppl = {safe_perplexity(out['final_loss']):.2f})\"\n", + " f\" params = {out['final_param_count']:,}\"\n", + " f\" grow_event = {ev}\"\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 3. 결과 표 + 자동 verdict" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "print(f\"{'mode':>16s} {'seed':>6s} {'final_loss':>12s} {'perplexity':>12s} {'params':>10s}\")\n", + "print(\"-\" * 65)\n", + "for (name, seed), out in results.items():\n", + " fl = out[\"final_loss\"]\n", + " print(\n", + " f\"{name:>16s} {seed:>6d} {fl:>12.4f} {safe_perplexity(fl):>12.2f} \"\n", + " f\"{out['final_param_count']:>10,}\"\n", + " )\n", + "\n", + "# mode 별 평균\n", + "print(\"\\n== Mode summary (mean ± σ across seeds) ==\")\n", + "summary = {}\n", + "for mode in MODES:\n", + " name = mode[\"name\"]\n", + " vals = [results[(name, s)][\"final_loss\"] for s in SEEDS]\n", + " params = [results[(name, s)][\"final_param_count\"] for s in SEEDS]\n", + " m = statistics.mean(vals)\n", + " sd = statistics.stdev(vals) if len(vals) > 1 else 0.0\n", + " summary[name] = (m, sd, statistics.mean(params))\n", + " print(\n", + " f\" {name:>16s} {m:.4f} ± {sd:.4f} (ppl ≈ {safe_perplexity(m):.2f}) \"\n", + " f\"params={statistics.mean(params):,.0f}\"\n", + " )\n", + "\n", + "# 자동 verdict\n", + "print(\"\\n== Verdict ==\")\n", + "small_loss = summary[\"small_baseline\"][0]\n", + "grown_loss = summary[\"grown\"][0]\n", + "large_loss = summary[\"large_baseline\"][0]\n", + "\n", + "# 1. all-finite\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 (grow stability): {all_finite} [{verdict_1}]\")\n", + "\n", + "# 2. grown < small_baseline — grow 가 capacity 부족 모델보다 명확히 좋음\n", + "diff_grown_small = grown_loss - small_loss\n", + "verdict_2 = \"PASS\" if diff_grown_small < 0 else \"FAIL\"\n", + "print(f\"2. grown < small_baseline: diff = {diff_grown_small:+.4f} [{verdict_2}]\")\n", + "\n", + "# 3. grown ≤ large_baseline + 0.05 — grown 이 처음부터 큰 모델 근방까지 학습\n", + "diff_grown_large = grown_loss - large_loss\n", + "verdict_3 = \"PASS\" if diff_grown_large <= 0.05 else \"FAIL\"\n", + "print(f\"3. grown ≤ large_baseline + 0.05: diff = {diff_grown_large:+.4f} [{verdict_3}]\")\n", + "\n", + "# 4. grow 직후 function preservation 검증 — 학습 curve 의 step 750 근방에 spike 없음\n", + "# (정량 검증은 별도 unit test 가 atol=1e-5 로 보장; 여기서는 학습 curve smoothness 만 추정)\n", + "for seed in SEEDS:\n", + " losses = results[(\"grown\", seed)][\"losses\"]\n", + " # step 750 직전 100step 평균 vs step 750 직후 1 step 비교 (drift 측정)\n", + " before = sum(losses[max(0, GROW_AT_STEP - 100) : GROW_AT_STEP]) / 100\n", + " after = losses[GROW_AT_STEP] if len(losses) > GROW_AT_STEP else losses[-1]\n", + " spike = after - before\n", + " print(\n", + " f\" seed={seed}: grow 직전 100 step 평균 = {before:.4f}, grow 직후 1 step = {after:.4f}, spike = {spike:+.4f}\"\n", + " )\n", + "\n", + "# grow 이벤트 정보\n", + "print(\"\\n== Grow events ==\")\n", + "for (name, seed), out in results.items():\n", + " if out[\"grow_event\"] is not None:\n", + " ev = out[\"grow_event\"]\n", + " print(\n", + " f\" {name} seed={seed}: step={ev['step']}, layers_grown={ev['n_layers_grown']}, \"\n", + " f\"ffn {ev['old_ffn_dim']} → {ev['new_ffn_dim']}, n_new_groups={ev['n_new_groups']}\"\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## 4. Loss curve 시각화 — grow step 표시" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "10", + "metadata": {}, + "outputs": [], + "source": [ + "fig, ax = plt.subplots(1, 1, figsize=(12, 6))\n", + "colors = {\n", + " \"small_baseline\": \"tab:red\",\n", + " \"grown\": \"tab:blue\",\n", + " \"large_baseline\": \"tab:green\",\n", + "}\n", + "window = 50\n", + "\n", + "for mode in MODES:\n", + " name = mode[\"name\"]\n", + " losses_per_seed = [results[(name, s)][\"losses\"] for s in SEEDS]\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", + " ax.plot(steps, mean, label=name, color=colors[name], linewidth=1.5)\n", + " ax.fill_between(steps, mean - std, mean + std, color=colors[name], alpha=0.15)\n", + "\n", + "# grow step 수직선\n", + "ax.axvline(GROW_AT_STEP, color=\"black\", linestyle=\":\", alpha=0.5, label=f\"grow @ {GROW_AT_STEP}\")\n", + "\n", + "ax.set_xlabel(\"step\")\n", + "ax.set_ylabel(f\"loss (rolling mean w={window})\")\n", + "ax.set_title(\n", + " f\"Phase 16b — Net2Net grow (ffn {FFN_SMALL} → {FFN_LARGE}): small / grown / large (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-phase16b\")\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. Parameter count 비교" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "12", + "metadata": {}, + "outputs": [], + "source": [ + "fig, ax = plt.subplots(1, 1, figsize=(8, 5))\n", + "names = [m[\"name\"] for m in MODES]\n", + "params = [summary[n][2] for n in names]\n", + "losses = [summary[n][0] for n in names]\n", + "stds = [summary[n][1] for n in names]\n", + "\n", + "color_list = [colors[n] for n in names]\n", + "ax.errorbar(params, losses, yerr=stds, marker=\"o\", capsize=4, linewidth=1.5, color=\"tab:gray\")\n", + "for p, lo, n in zip(params, losses, names, strict=True):\n", + " ax.annotate(n, (p, lo), xytext=(8, -8), textcoords=\"offset points\", fontsize=10)\n", + "\n", + "ax.set_xlabel(\"final parameter count\")\n", + "ax.set_ylabel(\"final loss (last 100 mean ± σ)\")\n", + "ax.set_title(\"Phase 16b — params vs loss: grown vs static baselines\")\n", + "ax.grid(alpha=0.3)\n", + "plt.tight_layout()\n", + "fig.savefig(out_dir / \"params_vs_loss.png\", dpi=150, bbox_inches=\"tight\")\n", + "plt.show()\n", + "print(f\"saved: {out_dir / 'params_vs_loss.png'}\")" + ] + }, + { + "cell_type": "markdown", + "id": "13", + "metadata": {}, + "source": [ + "## 6. 결론 / 다음 단계\n", + "\n", + "(셀 출력 보고 사용자가 채울 영역)\n", + "\n", + "- grow 직후 function preservation 확인 (spike 없음)?\n", + "- grown 의 final loss 가 large_baseline 근방까지 회복?\n", + "- 16a (DST) 와 16b (grow) 의 의미 비교 — 16b 의 capacity 증가가 더 효과적?\n", + "\n", + "**Phase 17 후보**:\n", + "- 16a + 16b 결합 — grow + shrink 동시 dynamic\n", + "- layer-wise 차등 (attention vs FFN 별 다른 grow / prune 정책)\n", + "- hidden_dim 까지 grow (downstream layer 모두 영향 — 더 invasive)" + ] + } + ], + "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 89b8499..1519779 100644 --- a/src/graphlm/neuron/graph_hybrid.py +++ b/src/graphlm/neuron/graph_hybrid.py @@ -373,6 +373,127 @@ def _activate_edges(self, indices_4d: Tensor, *, reset_weight: bool) -> None: if reset_weight: self.weight[i0, i1, i2, i3] = 0.0 + # ── Phase 16b: shape expansion (Net2Net-style) ────────────── + + @staticmethod + def _grow_param(old: nn.Parameter, new_data: Tensor, dim: int) -> nn.Parameter: + """Concat old + new_data along ``dim``, preserving ``requires_grad``. + + Copilot #3308323611 — ``.data`` 비권장 + freeze (requires_grad=False) 한 layer 가 + Parameter replace 후 unfreeze 되는 buf 회피. ``detach().clone()`` 사용 + 기존 flag 복원. + """ + combined = torch.cat([old.detach().clone(), new_data], dim=dim) + new_param = nn.Parameter(combined) + new_param.requires_grad_(old.requires_grad) + return new_param + + def grow_out(self, n_new_groups: int) -> None: + """G_out 차원을 ``n_new_groups`` 만큼 확장 (out_features 증가). + + function preservation 의 책임은 **downstream layer 의 ``grow_in``** 에 있음: + - 본 layer 의 새 G_out row 가 어떤 weight 든, downstream 이 새 input column 의 weight 를 + 0 으로 받으면 전체 forward 는 변경 없음. + - 본 layer 자체의 기존 G_out 출력은 영향 없음 (새 row 가 다른 row 와 독립). + + 새 weight: 작은 random init (기존 와 같은 fan_in 기반). adj_outer/inner = 1, bias = 0, + edge_mask = 1. + + Args: + n_new_groups: 추가할 G_out group 수 (positive int). + + Side effects (Parameter / Buffer **replace**): + optimizer state 가 기존 weight/adj_* 의 id 에 묶여 있어 grow 후 새 Parameter 에 + 대한 state 가 없음. **caller (train loop) 가 optimizer 를 재생성** 해야 함. + """ + if not isinstance(n_new_groups, int) or n_new_groups < 1: + raise ValueError(f"n_new_groups must be a positive int, got {n_new_groups!r}") + with torch.no_grad(): + k = self.group_size + G_in = self.n_groups_in + # 새 weight rows — 같은 fan_in (in_features) 기반 uniform + bound = 1.0 / math.sqrt(self.in_features) + new_w = torch.empty( + n_new_groups, + G_in, + k, + k, + device=self.weight.device, + dtype=self.weight.dtype, + ).uniform_(-bound, bound) + self.weight = self._grow_param(self.weight, new_w, dim=0) + # adj_outer / adj_inner / edge_mask 새 rows 는 1.0 + new_outer = torch.ones( + n_new_groups, G_in, device=self.adj_outer.device, dtype=self.adj_outer.dtype + ) + self.adj_outer = self._grow_param(self.adj_outer, new_outer, dim=0) + new_inner = torch.ones( + n_new_groups, + G_in, + k, + k, + device=self.adj_inner.device, + dtype=self.adj_inner.dtype, + ) + self.adj_inner = self._grow_param(self.adj_inner, new_inner, dim=0) + new_mask = torch.ones( + n_new_groups, G_in, k, k, device=self.edge_mask.device, dtype=self.edge_mask.dtype + ) + self.edge_mask = torch.cat([self.edge_mask, new_mask], dim=0) + # bias 추가 (0 init) + if self.bias is not None: + new_bias = torch.zeros( + n_new_groups * k, device=self.bias.device, dtype=self.bias.dtype + ) + self.bias = self._grow_param(self.bias, new_bias, dim=0) + # 메타데이터 갱신 + self.n_groups_out += n_new_groups + self.out_features += n_new_groups * k + + def grow_in(self, n_new_groups: int) -> None: + """G_in 차원을 ``n_new_groups`` 만큼 확장 (in_features 증가). + + **function preservation 의 핵심 layer** — 새 input column 의 weight 를 0 으로 초기화하여 + upstream layer 가 어떤 새 output 을 보내든 본 layer 의 기존 output 은 변경 없음. + + 새 weight: **0** (function preservation). adj_outer/inner = 1, edge_mask = 1, bias 영향 없음. + + Args: + n_new_groups: 추가할 G_in group 수 (positive int). + + Side effects: ``grow_out`` 과 동일 (Parameter replace, optimizer 재생성 필요). + """ + if not isinstance(n_new_groups, int) or n_new_groups < 1: + raise ValueError(f"n_new_groups must be a positive int, got {n_new_groups!r}") + with torch.no_grad(): + k = self.group_size + G_out = self.n_groups_out + # 새 weight columns = 0 (function preservation 의 핵심) + new_w = torch.zeros( + G_out, n_new_groups, k, k, device=self.weight.device, dtype=self.weight.dtype + ) + self.weight = self._grow_param(self.weight, new_w, dim=1) + # adj_outer / adj_inner / edge_mask 새 columns 는 1.0 + new_outer = torch.ones( + G_out, n_new_groups, device=self.adj_outer.device, dtype=self.adj_outer.dtype + ) + self.adj_outer = self._grow_param(self.adj_outer, new_outer, dim=1) + new_inner = torch.ones( + G_out, + n_new_groups, + k, + k, + device=self.adj_inner.device, + dtype=self.adj_inner.dtype, + ) + self.adj_inner = self._grow_param(self.adj_inner, new_inner, dim=1) + new_mask = torch.ones( + G_out, n_new_groups, k, k, device=self.edge_mask.device, dtype=self.edge_mask.dtype + ) + self.edge_mask = torch.cat([self.edge_mask, new_mask], dim=1) + # 메타데이터 갱신 + self.n_groups_in += n_new_groups + self.in_features += n_new_groups * k + 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 54fc80e..4508cd4 100644 --- a/src/graphlm/neuron/hybrid_transformer_demo.py +++ b/src/graphlm/neuron/hybrid_transformer_demo.py @@ -26,6 +26,7 @@ from graphlm.neuron.hybrid_transformer import ( Arch, FullGraphTransformerBlock, + HybridGraphFFN, HybridGraphTransformerBlock, make_block, make_full_block, @@ -82,6 +83,13 @@ class HybridTransformerTrainConfig: # DST cycle 종료 step (default None = max_steps 까지). 마지막 200 step 정도는 stabilize. dst_end_step: int | None = None + # Phase 16b: Net2Net-style FFN expansion (capacity 증가) + # 동작: grow_at_step + grow_ffn_target 둘 다 설정 시 해당 step 에서 FFN 의 ffn_dim 을 + # target 값으로 확장 (function preservation 보장). plain arch / non-hybrid FFN 무효. + # 확장 직후 optimizer 재생성 (parameter 객체 replace 됐으므로). + grow_at_step: int | None = None + grow_ffn_target: int | None = None + # runtime seed: int = 0 device: str = "cpu" @@ -108,6 +116,25 @@ def __post_init__(self) -> None: raise ValueError( f"dst_end_step must be in [1, max_steps={self.max_steps}], got {self.dst_end_step}" ) + # Phase 16b grow 검증 + if self.grow_at_step is not None and not 1 <= self.grow_at_step <= self.max_steps: + raise ValueError( + f"grow_at_step must be in [1, max_steps={self.max_steps}], got {self.grow_at_step}" + ) + if self.grow_ffn_target is not None: + if self.grow_ffn_target <= self.ffn_dim: + raise ValueError( + f"grow_ffn_target ({self.grow_ffn_target}) must be > current ffn_dim " + f"({self.ffn_dim})" + ) + if self.grow_ffn_target % self.group_size != 0: + raise ValueError( + f"grow_ffn_target ({self.grow_ffn_target}) must be divisible by " + f"group_size ({self.group_size})" + ) + # grow_at_step / grow_ffn_target 의 일관성 (둘 다 None 이거나 둘 다 set) + if (self.grow_at_step is None) != (self.grow_ffn_target is None): + raise ValueError("grow_at_step 과 grow_ffn_target 은 함께 설정 (또는 함께 None)") class HybridGraphTransformerLM(nn.Module): @@ -314,6 +341,47 @@ def _dst_swap_step( return summary +# ── Phase 16b: FFN shape expansion (Net2Net) ──────────────── + + +def _grow_ffn_in_model(model: nn.Module, target_ffn_dim: int, group_size: int) -> dict: + """모델의 모든 HybridGraphFFN 의 ffn_dim 을 target 으로 확장 (function-preserving). + + fc1.grow_out(n) + fc2.grow_in(n) 동시 적용. n = (target - current) / group_size. + + Returns: + {"n_layers_grown": int, "old_ffn_dim": int, "new_ffn_dim": int, "n_new_groups": int} + plain arch 등 HybridGraphFFN 이 없는 경우 n_layers_grown=0. + """ + n_layers = 0 + old_ffn_dim = None + n_new_groups = 0 + for ffn in model.modules(): + if not isinstance(ffn, HybridGraphFFN): + continue + current = ffn.fc1.out_features + if old_ffn_dim is None: + old_ffn_dim = current + if target_ffn_dim <= current: + continue # 이미 충분 + delta = target_ffn_dim - current + if delta % group_size != 0: + raise RuntimeError( + f"target_ffn_dim - current ({delta}) not divisible by group_size ({group_size})" + ) + n_new = delta // group_size + ffn.fc1.grow_out(n_new) # ffn_dim 차원 증가 + ffn.fc2.grow_in(n_new) # downstream — function preservation 보장 + n_layers += 1 + n_new_groups = n_new + return { + "n_layers_grown": n_layers, + "old_ffn_dim": old_ffn_dim if old_ffn_dim is not None else 0, + "new_ffn_dim": target_ffn_dim if n_layers > 0 else (old_ffn_dim or 0), + "n_new_groups": n_new_groups, + } + + def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: """1 run 학습 — Phase 13/14/15 sweep unit. @@ -343,6 +411,7 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: losses: list[float] = [] prune_event: dict | None = None dst_cycles: list[dict] = [] + grow_event: dict | None = None model.train() for step in range(1, config.max_steps + 1): x, y = next(data_iter) @@ -394,6 +463,20 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: } ) + # Phase 16b: FFN expansion (function-preserving grow) + if ( + config.grow_at_step is not None + and step == config.grow_at_step + and config.grow_ffn_target is not None + ): + grow_info = _grow_ffn_in_model(model, config.grow_ffn_target, config.group_size) + # plain arch / HybridGraphFFN 없는 모델은 실제로 확장된 layer 0개 — optimizer 재생성 + # + grow_event 기록 모두 건너뜀 (Copilot #3308323642, 주석과 동작 일치). + if grow_info["n_layers_grown"] > 0: + # Parameter 객체 replace 됐으므로 optimizer 재생성 (옛 state 손실) + optimizer = torch.optim.AdamW(model.parameters(), lr=config.lr) + grow_event = {"step": step, **grow_info} + n_last = min(100, len(losses)) final_loss = sum(losses[-n_last:]) / n_last if n_last > 0 else 0.0 @@ -404,6 +487,8 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: "final_sparsity": _model_sparsity(model), "prune_event": prune_event, "dst_cycles": dst_cycles, + "grow_event": grow_event, + "final_param_count": count_parameters(model), } diff --git a/tests/neuron/test_graph_hybrid.py b/tests/neuron/test_graph_hybrid.py index ec9ca40..22f4be7 100644 --- a/tests/neuron/test_graph_hybrid.py +++ b/tests/neuron/test_graph_hybrid.py @@ -421,6 +421,129 @@ def test_constant_sparsity_prune_then_regrow(): assert lin.n_alive_edges() == alive_before_cycle +# ── Phase 16b: shape expansion (Net2Net) ───────────────────── + + +def test_grow_out_shape_increased(): + """grow_out 후 out_features / n_groups_out 가 정확히 증가.""" + lin = HybridGraphLinear(16, 24, group_size=4) # G_out=6, G_in=4 + lin.grow_out(2) # G_out: 6 → 8 + assert lin.n_groups_out == 8 + assert lin.out_features == 32 + assert lin.weight.shape == (8, 4, 4, 4) + assert lin.adj_outer.shape == (8, 4) + assert lin.adj_inner.shape == (8, 4, 4, 4) + assert lin.edge_mask.shape == (8, 4, 4, 4) + assert lin.bias.shape == (32,) + + +def test_grow_in_shape_increased(): + """grow_in 후 in_features / n_groups_in 가 정확히 증가.""" + lin = HybridGraphLinear(16, 24, group_size=4) # G_in=4 + lin.grow_in(3) # G_in: 4 → 7 + assert lin.n_groups_in == 7 + assert lin.in_features == 28 + assert lin.weight.shape == (6, 7, 4, 4) + assert lin.adj_outer.shape == (6, 7) + assert lin.adj_inner.shape == (6, 7, 4, 4) + assert lin.edge_mask.shape == (6, 7, 4, 4) + + +def test_grow_in_function_preservation(): + """grow_in 후 기존 in_features 만큼의 input 에 대한 forward 동일 (새 input=anything).""" + torch.manual_seed(0) + lin = HybridGraphLinear(16, 24, group_size=4) + x_old = torch.randn(2, 8, 16) + lin.eval() + with torch.no_grad(): + y_before = lin(x_old) + + lin.grow_in(2) # in_features: 16 → 24 + # 새 input position 에 임의 값을 채워서 확장된 input 만들기 + x_extra = torch.randn(2, 8, 8) # 새 8 차원 (G_in 2 × group_size 4) + x_new = torch.cat([x_old, x_extra], dim=-1) + with torch.no_grad(): + y_after = lin(x_new) + # 새 input 의 forward 기여 = 0 (weight=0) → 기존 input 부분의 forward 결과만 남음 → y_after == y_before + assert torch.allclose(y_after, y_before, atol=1e-5), ( + f"grow_in function preservation 깨짐: max |diff| = {(y_after - y_before).abs().max().item()}" + ) + + +def test_grow_out_then_grow_in_downstream_preserves_forward(): + """FFN-style: fc1.grow_out(n) + fc2.grow_in(n) 동시 적용 → 2-layer chain forward 동일.""" + torch.manual_seed(0) + fc1 = HybridGraphLinear(16, 32, group_size=4) # 16 → 32 + fc2 = HybridGraphLinear(32, 16, group_size=4) # 32 → 16 + x = torch.randn(2, 8, 16) + fc1.eval() + fc2.eval() + with torch.no_grad(): + # GELU 없이 단순 chain — function preservation 의 핵심은 fc2 의 새 input weight=0 + y_before = fc2(fc1(x)) + + # 확장: fc1 의 G_out 늘리고 fc2 의 G_in 동일 증가 + fc1.grow_out(2) # 32 → 40 + fc2.grow_in(2) # 32 → 40 + assert fc1.out_features == 40 + assert fc2.in_features == 40 + + with torch.no_grad(): + y_after = fc2(fc1(x)) + # function preservation — 새 ffn unit 의 weight 가 0 이라 forward 결과 정확히 동일 + assert torch.allclose(y_after, y_before, atol=1e-5), ( + f"FFN grow function preservation 깨짐: max |diff| = {(y_after - y_before).abs().max().item()}" + ) + + +def test_grow_out_then_train_step_runs(): + """grow_out 후 학습 가능 — gradient 흐름 정상.""" + torch.manual_seed(0) + lin = HybridGraphLinear(16, 24, group_size=4) + lin.grow_out(2) + # parameter 가 replace 되었으므로 새 parameter 가 list 에 있는지 + params = list(lin.parameters()) + assert any(p.shape == (8, 4, 4, 4) for p in params if p.dim() == 4) + # forward + backward + x = torch.randn(2, 16) + lin(x).sum().backward() + # 새 weight 의 grad 도 흐름 + assert lin.weight.grad is not None + assert lin.weight.grad.shape == lin.weight.shape + + +def test_grow_in_then_train_step_runs(): + torch.manual_seed(0) + lin = HybridGraphLinear(16, 24, group_size=4) + lin.grow_in(2) + x = torch.randn(2, 24) # 확장된 in_features + lin(x).sum().backward() + assert lin.weight.grad is not None + + +def test_grow_negative_rejected(): + lin = HybridGraphLinear(16, 16, group_size=4) + with pytest.raises(ValueError, match="positive int"): + lin.grow_out(0) + with pytest.raises(ValueError, match="positive int"): + lin.grow_in(-1) + + +def test_grow_preserves_existing_prune_mask(): + """grow 가 기존 edge_mask 의 pruned 위치를 보존 (확장된 새 위치만 mask=1 추가).""" + torch.manual_seed(0) + lin = HybridGraphLinear(16, 16, group_size=4) + lin.prune_bottom_fraction(0.5) + pruned_count_before = lin.n_pruned_edges() + + lin.grow_out(2) + # 기존 pruned 위치는 그대로, 새 row 는 모두 alive (mask=1) + pruned_count_after = lin.n_pruned_edges() + assert pruned_count_after == pruned_count_before, ( + f"기존 pruned 손실: before={pruned_count_before}, after={pruned_count_after}" + ) + + def test_edge_mask_in_state_dict(): """edge_mask 가 state_dict 에 포함되어 save/load 보존.""" lin1 = HybridGraphLinear(16, 16, group_size=4)