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/phase13/hybrid_adj.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/phase13/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.
282 changes: 282 additions & 0 deletions notebooks/02-function-level/12-phase13-hybrid-transformer.ipynb
Original file line number Diff line number Diff line change
@@ -0,0 +1,282 @@
{
"cells": [
{
"cell_type": "markdown",
"id": "0",
"metadata": {},
"source": [
"# 12-phase13-hybrid-transformer\n",
"\n",
"**neuron Phase 13** — HybridGraphLinear 를 **실제 Transformer 의 FFN 위치** 에 통합 + **RMSNorm** 도입.\n",
"Phase 8~12 는 모두 MLP-LM baseline 이었고, Phase 13 부터 standard pre-norm Transformer block 위에서 graph hidden layer 의 paradigm 이 동작하는지 검증.\n",
"\n",
"핵심 가설:\n",
"1. **function preservation (Transformer 위)** — `hybrid_full_full` ≈ `plain` (RMSNorm + 표준 FFN)?\n",
"2. **dual routing 우위 재현** — Phase 12 에서 hybrid_full_around_one 이 최저 final_loss 였음. Transformer 에서도 유지?\n",
"3. **full scale-corrected (around_one × around_one)** — outer/inner 둘 다 학습 활성화 init 의 효과?\n",
"4. **RMSNorm 안정성** — LayerNorm 대비 학습 동등 또는 우위 (모든 arch 에서 NaN 없음)?\n",
"\n",
"설계: 4 arch × 2 seed = 8 run, max_steps=1500.\n",
"\n",
"데이터: TinyShakespeare (char-LM, block_size=64)\n",
"시드: [42, 123]\n",
"작성일: 2026-05-26\n",
"연관: Issue [#67](https://github.com/EinSofINTEREST/GraphLM/issues/67) / Phase 12 baseline PR [#66](https://github.com/EinSofINTEREST/GraphLM/pull/66)\n",
"\n",
"Phase 13 의 ``identity`` outer 미지원 → 4 arch 는 plain / hybrid_full_full / hybrid_full_around_one / hybrid_around_one_around_one."
]
},
{
"cell_type": "markdown",
"id": "1",
"metadata": {},
"source": [
"## 0. 환경 / 의존성"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "2",
"metadata": {},
"outputs": [],
"source": "from __future__ import annotations\n\nimport math\nimport statistics\nfrom pathlib import Path\n\nimport matplotlib.pyplot as plt\nimport torch\n\nfrom graphlm.data.tinyshakespeare import (\n CharTokenizer,\n TinyShakespeareDataset,\n load_tinyshakespeare_text,\n)\nfrom graphlm.neuron.hybrid_transformer_demo import (\n HybridGraphTransformerLM,\n HybridTransformerTrainConfig,\n count_parameters,\n train_hybrid_transformer_lm,\n)\nfrom graphlm.utils import safe_perplexity\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"device: {device}\")\nprint(f\"torch: {torch.__version__}\")"
},
{
"cell_type": "markdown",
"id": "3",
"metadata": {},
"source": [
"## 1. Config + 데이터"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "4",
"metadata": {},
"outputs": [],
"source": [
"# data\n",
"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",
"# model + train hyperparameters\n",
"HIDDEN_DIM = 128\n",
"N_HEADS = 4\n",
"FFN_DIM = 256 # 2x hidden (작게 — sweep 속도 위해)\n",
"N_LAYERS = 4\n",
"GROUP_SIZE = 16 # hidden_dim/group_size = 8, ffn_dim/group_size = 16\n",
"BLOCK_SIZE = 64\n",
"BATCH_SIZE = 32\n",
"LR = 3e-4\n",
"MAX_STEPS = 1500\n",
"SEEDS = [42, 123]\n",
"ARCHS = [\n",
" \"plain\",\n",
" \"hybrid_full_full\",\n",
" \"hybrid_full_around_one\",\n",
" \"hybrid_around_one_around_one\",\n",
"]\n",
"\n",
"# 모델 파라미터 수 비교 (arch 별)\n",
"print(\"\\n== Parameter count by arch ==\")\n",
"for arch in ARCHS:\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",
" )\n",
" print(f\" {arch:32s} params = {count_parameters(m):,}\")"
]
},
{
"cell_type": "markdown",
"id": "5",
"metadata": {},
"source": [
"## 2. Sweep 실행 (4 arch × 2 seed = 8 run)"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "6",
"metadata": {},
"outputs": [],
"source": "results = {}\nfor arch in ARCHS:\n for seed in SEEDS:\n key = (arch, seed)\n print(f\"\\n== arch={arch} 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 block_size=BLOCK_SIZE,\n batch_size=BATCH_SIZE,\n lr=LR,\n max_steps=MAX_STEPS,\n seed=seed,\n device=device,\n )\n out = train_hybrid_transformer_lm(cfg)\n results[key] = out\n print(\n f\" final_loss = {out['final_loss']:.4f} (perplexity = {safe_perplexity(out['final_loss']):.2f})\"\n )"
},
{
"cell_type": "markdown",
"id": "7",
"metadata": {},
"source": [
"## 3. 결과 표 + 자동 verdict"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "8",
"metadata": {},
"outputs": [],
"source": "print(f\"{'arch':32s} {'seed':>6s} {'final_loss':>12s} {'perplexity':>12s}\")\nprint(\"-\" * 70)\nfor (arch, seed), out in results.items():\n fl = out[\"final_loss\"]\n print(f\"{arch:32s} {seed:>6d} {fl:>12.4f} {safe_perplexity(fl):>12.2f}\")\n\n# arch 별 평균 / 표준편차\nprint(\"\\n== Arch-level summary (mean ± σ across seeds) ==\")\nsummary = {}\nfor arch in ARCHS:\n vals = [results[(arch, s)][\"final_loss\"] for s in SEEDS]\n summary[arch] = (statistics.mean(vals), statistics.stdev(vals) if len(vals) > 1 else 0.0)\n m, s = summary[arch]\n print(f\" {arch:32s} {m:.4f} ± {s:.4f} (perplexity ≈ {safe_perplexity(m):.2f})\")\n\n# 자동 verdict\nprint(\"\\n== Verdict ==\")\nplain_loss = summary[\"plain\"][0]\nff_loss = summary[\"hybrid_full_full\"][0]\nfa_loss = summary[\"hybrid_full_around_one\"][0]\naa_loss = summary[\"hybrid_around_one_around_one\"][0]\n\n# 1. function preservation (학습 후 hybrid_full_full ≈ plain — 학습 dynamics 가 동일 sweet spot 유지)\ndiff_ff = abs(ff_loss - plain_loss)\nverdict_1 = \"PASS\" if diff_ff < 0.15 else \"FAIL\"\nprint(f\"1. function preservation: |hybrid_full_full - plain| = {diff_ff:.4f} [{verdict_1}]\")\n\n# 2. scale-corrected 우위 (hybrid_full_around_one 이 plain 보다 우수 또는 동등)\nverdict_2 = \"PASS\" if fa_loss <= plain_loss + 0.05 else \"FAIL\"\ndiff_fa = fa_loss - plain_loss\nprint(f\"2. inner around_one ≤ plain + 0.05: diff = {diff_fa:+.4f} [{verdict_2}]\")\n\n# 3. full scale-corrected (around_one × around_one) 안정성 (유한 + plain 근방)\n# CodeRabbit #3304306127 — isfinite 가 inf/-inf 도 거부 (isnan 만 쓰면 inf 가 PASS 됨)\nverdict_3 = \"PASS\" if (math.isfinite(aa_loss) and aa_loss < plain_loss + 0.5) else \"FAIL\"\ndiff_aa = aa_loss - plain_loss\nprint(f\"3. full around_one stable: diff = {diff_aa:+.4f} [{verdict_3}]\")\n\n# 4. 모두 finite 로 학습 종료 (RMSNorm 안정성, inf/nan 모두 거부)\nall_finite = all(math.isfinite(out[\"final_loss\"]) for out in results.values())\nverdict_4 = \"PASS\" if all_finite else \"FAIL\"\nprint(f\"4. RMSNorm stability (all finite): {all_finite} [{verdict_4}]\")"
},
{
"cell_type": "markdown",
"id": "9",
"metadata": {},
"source": [
"## 4. Loss curve 시각화"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "10",
"metadata": {},
"outputs": [],
"source": [
"fig, ax = plt.subplots(1, 1, figsize=(10, 5))\n",
"colors = {\n",
" \"plain\": \"tab:gray\",\n",
" \"hybrid_full_full\": \"tab:blue\",\n",
" \"hybrid_full_around_one\": \"tab:orange\",\n",
" \"hybrid_around_one_around_one\": \"tab:green\",\n",
"}\n",
"window = 50 # rolling mean for smoothing\n",
"\n",
"for arch in ARCHS:\n",
" losses_per_seed = [results[(arch, s)][\"losses\"] for s in SEEDS]\n",
" # rolling mean\n",
" smoothed = []\n",
" for losses in losses_per_seed:\n",
" smoothed.append(\n",
" [\n",
" sum(losses[max(0, i - window) : i + 1]) / min(i + 1, window)\n",
" for i in range(len(losses))\n",
" ]\n",
" )\n",
" # mean ± σ across seeds\n",
" arr = torch.tensor(smoothed)\n",
" mean = arr.mean(dim=0)\n",
" std = arr.std(dim=0)\n",
" steps = list(range(len(mean)))\n",
" color = colors[arch]\n",
" ax.plot(steps, mean, label=arch, color=color, linewidth=1.5)\n",
" ax.fill_between(steps, mean - std, mean + std, color=color, alpha=0.15)\n",
"\n",
"ax.set_xlabel(\"step\")\n",
"ax.set_ylabel(f\"loss (rolling mean w={window})\")\n",
"ax.set_title(\"Phase 13 — Transformer FFN: 4 arch loss curves (mean ± σ over 2 seeds)\")\n",
"ax.legend(loc=\"upper right\")\n",
"ax.grid(alpha=0.3)\n",
"plt.tight_layout()\n",
"\n",
"out_dir = Path(\"../../runs/notebook-neuron-phase13\")\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. adj_outer / adj_inner heatmap (hybrid arch 만)\n",
"\n",
"각 hybrid arch 의 첫 block 의 fc1 adj_outer 와 adj_inner (block-aggregated) 학습 후 모습."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "12",
"metadata": {},
"outputs": [],
"source": [
"hybrid_archs = [a for a in ARCHS if a != \"plain\"]\n",
"n_arch = len(hybrid_archs)\n",
"fig, axes = plt.subplots(n_arch, 2, figsize=(10, 3 * n_arch))\n",
"if n_arch == 1:\n",
" axes = axes.reshape(1, -1)\n",
"\n",
"for row, arch in enumerate(hybrid_archs):\n",
" # seed 42 의 첫 block 의 fc1\n",
" snap = results[(arch, 42)][\"final_adj\"][0][\"fc1\"]\n",
" outer = snap[\"outer\"].numpy() # (G_out, G_in)\n",
" inner = (\n",
" snap[\"inner\"].abs().mean(dim=(-1, -2)).numpy()\n",
" ) # (G_out, G_in) — block-aggregated magnitude\n",
"\n",
" ax_o = axes[row, 0]\n",
" im_o = ax_o.imshow(outer, cmap=\"RdBu_r\", vmin=-2, vmax=2)\n",
" ax_o.set_title(f\"{arch} — adj_outer (block 0, fc1)\")\n",
" ax_o.set_xlabel(\"G_in\")\n",
" ax_o.set_ylabel(\"G_out\")\n",
" plt.colorbar(im_o, ax=ax_o, fraction=0.046)\n",
"\n",
" ax_i = axes[row, 1]\n",
" im_i = ax_i.imshow(inner, cmap=\"viridis\")\n",
" ax_i.set_title(f\"{arch} — |adj_inner| block-mean (block 0, fc1)\")\n",
" ax_i.set_xlabel(\"G_in\")\n",
" ax_i.set_ylabel(\"G_out\")\n",
" plt.colorbar(im_i, ax=ax_i, fraction=0.046)\n",
"\n",
"plt.tight_layout()\n",
"fig.savefig(out_dir / \"hybrid_adj.png\", dpi=150, bbox_inches=\"tight\")\n",
"plt.show()\n",
"print(f\"saved: {out_dir / 'hybrid_adj.png'}\")"
]
},
{
"cell_type": "markdown",
"id": "13",
"metadata": {},
"source": [
"## 6. 결론 / 다음 단계\n",
"\n",
"(셀 출력 보고 사용자가 채울 영역)\n",
"\n",
"- function preservation: hybrid_full_full vs plain — Transformer 위에서도 동등?\n",
"- dual routing 우위: Phase 12 패턴 재현?\n",
"- adj 학습 패턴: outer / inner 가 다른 영역을 cover?\n",
"\n",
"**Phase 14 후보**:\n",
"- attention 의 qkv / out 을 HybridGraphLinear 로 교체 — function-level graph 가 attention 까지 확장\n",
"- Net2Net / LiGO 식 growable transformer block — 학습 중 hidden_dim 또는 n_layers 증가"
]
}
],
"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
}
12 changes: 12 additions & 0 deletions src/graphlm/neuron/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,18 +16,30 @@
from graphlm.neuron.graph_hybrid import HybridGraphLinear
from graphlm.neuron.growable import GrowableEmbedding, GrowableLayerNorm, GrowableLinear
from graphlm.neuron.growth import add_attn_function_preserving, add_attn_smooth_start
from graphlm.neuron.hybrid_transformer import (
HybridGraphFFN,
HybridGraphTransformerBlock,
PlainTransformerBlock,
make_block,
)
from graphlm.neuron.rms_norm import RMSNorm

__all__ = [
"ChannelGraphLinear",
"GroupGraphLinear",
"GrowableEmbedding",
"GrowableLayerNorm",
"GrowableLinear",
"HybridGraphFFN",
"HybridGraphLinear",
"HybridGraphTransformerBlock",
"NeuronBlock",
"NeuronConfig",
"NeuronGrowingDecoder",
"PlainTransformerBlock",
"RMSNorm",
"SinusoidalAlpha",
"add_attn_function_preserving",
"add_attn_smooth_start",
"make_block",
]
Loading
Loading