diff --git a/docs/figures/neuron/phase14/attention_adj.png b/docs/figures/neuron/phase14/attention_adj.png new file mode 100644 index 0000000..d613e24 Binary files /dev/null and b/docs/figures/neuron/phase14/attention_adj.png differ diff --git a/docs/figures/neuron/phase14/loss_curves.png b/docs/figures/neuron/phase14/loss_curves.png new file mode 100644 index 0000000..c8e8d40 Binary files /dev/null and b/docs/figures/neuron/phase14/loss_curves.png differ diff --git a/notebooks/02-function-level/12-phase13-hybrid-transformer.ipynb b/notebooks/02-function-level/12-phase13-hybrid-transformer.ipynb index 3d731db..fab0d65 100644 --- a/notebooks/02-function-level/12-phase13-hybrid-transformer.ipynb +++ b/notebooks/02-function-level/12-phase13-hybrid-transformer.ipynb @@ -144,49 +144,7 @@ "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'}\")" - ] + "source": "fig, ax = plt.subplots(1, 1, figsize=(10, 5))\ncolors = {\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}\nwindow = 50 # rolling mean for smoothing\n\nfor arch in ARCHS:\n losses_per_seed = [results[(arch, s)][\"losses\"] for s in SEEDS]\n # rolling mean — slice 시작점 +1 시프트로 window 와 divisor 일치 (CodeRabbit #3304780219)\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 # 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\nax.set_xlabel(\"step\")\nax.set_ylabel(f\"loss (rolling mean w={window})\")\nax.set_title(\"Phase 13 — Transformer FFN: 4 arch loss curves (mean ± σ over 2 seeds)\")\nax.legend(loc=\"upper right\")\nax.grid(alpha=0.3)\nplt.tight_layout()\n\nout_dir = Path(\"../../runs/notebook-neuron-phase13\")\nout_dir.mkdir(parents=True, exist_ok=True)\nfig.savefig(out_dir / \"loss_curves.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(f\"saved: {out_dir / 'loss_curves.png'}\")" }, { "cell_type": "markdown", diff --git a/notebooks/02-function-level/13-phase14-graph-attention.ipynb b/notebooks/02-function-level/13-phase14-graph-attention.ipynb new file mode 100644 index 0000000..3fa3362 --- /dev/null +++ b/notebooks/02-function-level/13-phase14-graph-attention.ipynb @@ -0,0 +1,336 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "# 13-phase14-graph-attention\n", + "\n", + "**neuron Phase 14** — attention 의 qkv / out 까지 `HybridGraphLinear` 로 교체 (**block 전체가 graph**). Phase 13 은 FFN 만 graph 였고, Phase 14 에서 paradigm 의 *full graph block* 단계 진입.\n", + "\n", + "핵심 가설:\n", + "1. **function preservation (attention 포함)** — full graph block (adj=full/full) ≈ plain block?\n", + "2. **attention graph 가 학습에 도움?** — Phase 13 (FFN-only) vs Phase 14 (full graph) 동일 arch 비교에서 final_loss 개선?\n", + "3. **dual scale-corrected 우위 재현** — Phase 12/13 의 `hybrid_around_one_around_one` 우위가 full graph 에서도 유지?\n", + "4. **파라미터 증가 대비 효과** — full graph 는 attention adj 도 추가되어 파라미터 ↑. ROI?\n", + "\n", + "설계: 4 arch × 2 seed × {Phase 13 (FFN-only), Phase 14 (full graph)} = 16 run.\n", + "- arch: plain / hybrid_full_full / hybrid_full_around_one / hybrid_around_one_around_one\n", + "- plain 은 use_full_graph 무관 (둘 다 PlainTransformerBlock)\n", + "- → 실제 비교: 4 arch × 2 seed × 2 mode - 2 (plain 중복) = 14 unique run\n", + "\n", + "데이터: TinyShakespeare (char-LM, block_size=64)\n", + "시드: [42, 123]\n", + "작성일: 2026-05-26\n", + "연관: Issue [#69](https://github.com/EinSofINTEREST/GraphLM/issues/69) / Phase 13 baseline PR [#68](https://github.com/EinSofINTEREST/GraphLM/pull/68)" + ] + }, + { + "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 13 과 동일 hyperparameter (공정 비교)\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", + "SEEDS = [42, 123]\n", + "ARCHS = [\n", + " \"plain\",\n", + " \"hybrid_full_full\",\n", + " \"hybrid_full_around_one\",\n", + " \"hybrid_around_one_around_one\",\n", + "]\n", + "MODES = [(\"phase13_ffn_only\", False), (\"phase14_full_graph\", True)]\n", + "\n", + "# arch x mode 별 파라미터 수 비교\n", + "print(\"\\n== Parameter count by arch × mode ==\")\n", + "for arch in ARCHS:\n", + " for mode_name, use_full in MODES:\n", + " if arch == \"plain\" and use_full:\n", + " continue # plain 은 use_full 무관 (PlainTransformerBlock 동일)\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=use_full,\n", + " )\n", + " print(f\" {arch:32s} {mode_name:25s} params = {count_parameters(m):,}\")" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. Sweep 실행 (4 arch × 2 seed × 2 mode = 14 unique run)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "results = {}\n", + "for arch in ARCHS:\n", + " for mode_name, use_full in MODES:\n", + " if arch == \"plain\" and use_full:\n", + " continue\n", + " for seed in SEEDS:\n", + " key = (arch, mode_name, seed)\n", + " print(f\"\\n== arch={arch} 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=FFN_DIM,\n", + " n_layers=N_LAYERS,\n", + " group_size=GROUP_SIZE,\n", + " arch=arch,\n", + " use_full_graph=use_full,\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} {'mode':25s} {'seed':>6s} {'final_loss':>12s} {'perplexity':>12s}\")\n", + "print(\"-\" * 95)\n", + "for (arch, mode_name, seed), out in results.items():\n", + " fl = out[\"final_loss\"]\n", + " print(f\"{arch:32s} {mode_name:25s} {seed:>6d} {fl:>12.4f} {safe_perplexity(fl):>12.2f}\")\n", + "\n", + "# arch × mode 별 평균\n", + "print(\"\\n== Arch × mode summary (mean ± σ across seeds) ==\")\n", + "summary = {}\n", + "for arch in ARCHS:\n", + " for mode_name, use_full in MODES:\n", + " if arch == \"plain\" and use_full:\n", + " continue\n", + " vals = [results[(arch, mode_name, s)][\"final_loss\"] for s in SEEDS]\n", + " mean = statistics.mean(vals)\n", + " std = statistics.stdev(vals) if len(vals) > 1 else 0.0\n", + " summary[(arch, mode_name)] = (mean, std)\n", + " print(\n", + " f\" {arch:32s} {mode_name:25s} {mean:.4f} ± {std:.4f} (perplexity ≈ {safe_perplexity(mean):.2f})\"\n", + " )\n", + "\n", + "# 자동 verdict\n", + "print(\"\\n== Verdict ==\")\n", + "plain_loss = summary[(\"plain\", \"phase13_ffn_only\")][0]\n", + "ff_p14 = summary[(\"hybrid_full_full\", \"phase14_full_graph\")][0]\n", + "aa_p13 = summary[(\"hybrid_around_one_around_one\", \"phase13_ffn_only\")][0]\n", + "aa_p14 = summary[(\"hybrid_around_one_around_one\", \"phase14_full_graph\")][0]\n", + "\n", + "# 1. function preservation 확장 — full graph 도 plain 근방?\n", + "diff_ff_p14 = abs(ff_p14 - plain_loss)\n", + "verdict_1 = \"PASS\" if diff_ff_p14 < 0.15 else \"FAIL\"\n", + "print(\n", + " f\"1. full graph function preservation: |hybrid_full_full(p14) - plain| = {diff_ff_p14:.4f} [{verdict_1}]\"\n", + ")\n", + "\n", + "# 2. attention graph 효과 — 동일 arch 에서 Phase 14 ≤ Phase 13 + 0.05?\n", + "diff_aa = aa_p14 - aa_p13\n", + "verdict_2 = \"PASS\" if diff_aa <= 0.05 else \"FAIL\"\n", + "print(f\"2. attention graph not hurting (aa): Phase14 - Phase13 = {diff_aa:+.4f} [{verdict_2}]\")\n", + "\n", + "# 3. 모두 finite\n", + "all_finite = all(math.isfinite(out[\"final_loss\"]) for out in results.values())\n", + "verdict_3 = \"PASS\" if all_finite else \"FAIL\"\n", + "print(f\"3. all-finite stability (RMSNorm + graph attention): {all_finite} [{verdict_3}]\")\n", + "\n", + "# 4. dual scale-corrected 우위 in full graph — aa_p14 가 ff_p14 보다 ≤?\n", + "diff_aa_vs_ff = aa_p14 - ff_p14\n", + "verdict_4 = \"PASS\" if diff_aa_vs_ff <= 0.05 else \"FAIL\"\n", + "print(\n", + " f\"4. around_one×around_one ≤ full_full + 0.05 (p14): diff = {diff_aa_vs_ff:+.4f} [{verdict_4}]\"\n", + ")" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## 4. Loss curve 시각화 — Phase 13 (실선) vs Phase 14 (점선)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "10", + "metadata": {}, + "outputs": [], + "source": "fig, ax = plt.subplots(1, 1, figsize=(12, 6))\ncolors = {\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}\nwindow = 50\n\nfor arch in ARCHS:\n for mode_name, use_full in MODES:\n if arch == \"plain\" and use_full:\n continue\n losses_per_seed = [results[(arch, mode_name, s)][\"losses\"] for s in SEEDS]\n # rolling mean — slice 시작점 +1 시프트로 window 와 divisor 일치 (CodeRabbit #3304780219)\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 color = colors[arch]\n # Phase 13 = solid, Phase 14 = dashed\n linestyle = \"-\" if mode_name == \"phase13_ffn_only\" else \"--\"\n label = f\"{arch} ({mode_name.replace('_', ' ')})\"\n ax.plot(steps, mean, label=label, color=color, linewidth=1.5, linestyle=linestyle)\n ax.fill_between(steps, mean - std, mean + std, color=color, alpha=0.10)\n\nax.set_xlabel(\"step\")\nax.set_ylabel(f\"loss (rolling mean w={window})\")\nax.set_title(\n \"Phase 14 — full graph block: 4 arch × {Phase 13 FFN-only, Phase 14 full graph} (mean ± σ over 2 seeds)\"\n)\nax.legend(loc=\"upper right\", fontsize=8)\nax.grid(alpha=0.3)\nplt.tight_layout()\n\nout_dir = Path(\"../../runs/notebook-neuron-phase14\")\nout_dir.mkdir(parents=True, exist_ok=True)\nfig.savefig(out_dir / \"loss_curves.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(f\"saved: {out_dir / 'loss_curves.png'}\")" + }, + { + "cell_type": "markdown", + "id": "11", + "metadata": {}, + "source": [ + "## 5. attention adj 시각화 (Phase 14, hybrid_around_one_around_one 만)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "12", + "metadata": {}, + "outputs": [], + "source": [ + "# Phase 14 의 attention qkv / out 의 학습된 adj_outer 모양 확인 (block 0, seed 42)\n", + "snap = results[(\"hybrid_around_one_around_one\", \"phase14_full_graph\", 42)][\"final_adj\"][0]\n", + "print(f\"snap keys (Phase 14 full graph): {sorted(snap.keys())}\")\n", + "\n", + "fig, axes = plt.subplots(2, 2, figsize=(10, 8))\n", + "for ax, layer_name in zip(axes.flat, [\"qkv\", \"out\", \"fc1\", \"fc2\"], strict=True):\n", + " outer = snap[layer_name][\"outer\"].numpy()\n", + " im = ax.imshow(outer, cmap=\"RdBu_r\", vmin=-2, vmax=2)\n", + " ax.set_title(f\"{layer_name} — adj_outer (block 0)\")\n", + " ax.set_xlabel(\"G_in\")\n", + " ax.set_ylabel(\"G_out\")\n", + " plt.colorbar(im, ax=ax, fraction=0.046)\n", + "\n", + "plt.tight_layout()\n", + "fig.savefig(out_dir / \"attention_adj.png\", dpi=150, bbox_inches=\"tight\")\n", + "plt.show()\n", + "print(f\"saved: {out_dir / 'attention_adj.png'}\")" + ] + }, + { + "cell_type": "markdown", + "id": "13", + "metadata": {}, + "source": [ + "## 6. 결론 / 다음 단계\n", + "\n", + "(셀 출력 보고 사용자가 채울 영역)\n", + "\n", + "- full graph block 도 function preservation 성립?\n", + "- attention graph 가 학습에 유의미한 효과?\n", + "- qkv vs out 의 adj 학습 패턴 차이?\n", + "\n", + "**Phase 15 후보**:\n", + "- sparsity-driven prune — adj magnitude < threshold edge 영구 제거 → dead channel 발생 (DST 계열, training-time dynamic parameter count 의 진정한 첫 단계)\n", + "- Net2Net / LiGO 식 grow — 학습 중 channel / group 추가 (function preservation 유지)" + ] + } + ], + "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 +} \ No newline at end of file diff --git a/src/graphlm/neuron/__init__.py b/src/graphlm/neuron/__init__.py index 84d0c0b..5151b8f 100644 --- a/src/graphlm/neuron/__init__.py +++ b/src/graphlm/neuron/__init__.py @@ -11,25 +11,30 @@ NeuronGrowingDecoder, SinusoidalAlpha, ) +from graphlm.neuron.graph_attention import HybridGraphCausalSelfAttention from graphlm.neuron.graph_channel import ChannelGraphLinear from graphlm.neuron.graph_group import GroupGraphLinear 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 ( + FullGraphTransformerBlock, HybridGraphFFN, HybridGraphTransformerBlock, PlainTransformerBlock, make_block, + make_full_block, ) from graphlm.neuron.rms_norm import RMSNorm __all__ = [ "ChannelGraphLinear", + "FullGraphTransformerBlock", "GroupGraphLinear", "GrowableEmbedding", "GrowableLayerNorm", "GrowableLinear", + "HybridGraphCausalSelfAttention", "HybridGraphFFN", "HybridGraphLinear", "HybridGraphTransformerBlock", @@ -42,4 +47,5 @@ "add_attn_function_preserving", "add_attn_smooth_start", "make_block", + "make_full_block", ] diff --git a/src/graphlm/neuron/graph_attention.py b/src/graphlm/neuron/graph_attention.py new file mode 100644 index 0000000..d7be7ce --- /dev/null +++ b/src/graphlm/neuron/graph_attention.py @@ -0,0 +1,87 @@ +"""Phase 14 — HybridGraphCausalSelfAttention (qkv + out 을 HybridGraphLinear 로). + +Phase 13 은 FFN 만 graph 였고 attention 은 표준 ``nn.Linear``. Phase 14 는 **qkv / out 도 +HybridGraphLinear** 로 교체하여 block 전체가 graph 가 되는 단계. + +설계: +- ``qkv``: hidden_dim → 3·hidden_dim (rectangular — identity outer 미지원) +- ``out``: hidden_dim → hidden_dim (square — identity 이론상 가능하나 통일성 위해 미사용) +- sdpa (scaled dot-product attention) 은 그대로 standard +- function preservation: adj_outer=full + adj_inner=full + 같은 W → standard CausalSelfAttention forward 동치 + +0-init 거부 + magnitude rule 은 underlying ``HybridGraphLinear`` 가 상속. +""" + +from __future__ import annotations + +import torch.nn.functional as F +from torch import Tensor, nn + +from graphlm.neuron.graph_hybrid import AdjInnerInit, AdjOuterInit, HybridGraphLinear + + +class HybridGraphCausalSelfAttention(nn.Module): + """Causal multi-head self-attention with HybridGraphLinear qkv + out. + + Args: + hidden_dim: Transformer hidden size (n_heads 의 배수). + n_heads: number of attention heads. + group_size: HybridGraphLinear 의 block size. hidden_dim 과 3·hidden_dim 모두 배수여야 함. + adj_outer_init: ``"full"`` / ``"uniform_around_one"`` (identity 는 qkv rectangular 라 미지원). + adj_inner_init: ``"full"`` / ``"uniform_around_one"``. + dropout: attention dropout (sdpa 의 dropout_p). + + Forward: + x: ``(B, T, hidden_dim)`` → y: same shape + function preserving when both adj = full and W = same as standard nn.Linear init. + """ + + def __init__( + self, + hidden_dim: int, + n_heads: int, + group_size: int, + *, + adj_outer_init: AdjOuterInit = "full", + adj_inner_init: AdjInnerInit = "full", + dropout: float = 0.0, + ): + super().__init__() + if hidden_dim % n_heads != 0: + raise ValueError(f"hidden_dim {hidden_dim} not divisible by n_heads {n_heads}") + if adj_outer_init == "identity": + raise ValueError( + "HybridGraphCausalSelfAttention 은 adj_outer_init='identity' 미지원 — " + "qkv 는 hidden_dim → 3·hidden_dim 으로 rectangular 라 정방 identity 정의 불가. " + "'full' 또는 'uniform_around_one' 사용." + ) + self.n_heads = n_heads + self.head_dim = hidden_dim // n_heads + self.qkv = HybridGraphLinear( + hidden_dim, + 3 * hidden_dim, + group_size=group_size, + adj_outer_init=adj_outer_init, + adj_inner_init=adj_inner_init, + bias=False, + ) + self.out = HybridGraphLinear( + hidden_dim, + hidden_dim, + group_size=group_size, + adj_outer_init=adj_outer_init, + adj_inner_init=adj_inner_init, + bias=False, + ) + self.dropout = dropout + + def forward(self, x: Tensor) -> Tensor: + B, T, C = x.shape + qkv = self.qkv(x).reshape(B, T, 3, self.n_heads, self.head_dim) + qkv = qkv.permute(2, 0, 3, 1, 4) + q, k, v = qkv[0], qkv[1], qkv[2] + out = F.scaled_dot_product_attention( + q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0 + ) + out = out.transpose(1, 2).reshape(B, T, C) + return self.out(out) diff --git a/src/graphlm/neuron/hybrid_transformer.py b/src/graphlm/neuron/hybrid_transformer.py index 4e77b4a..4f29b12 100644 --- a/src/graphlm/neuron/hybrid_transformer.py +++ b/src/graphlm/neuron/hybrid_transformer.py @@ -23,6 +23,7 @@ from torch import Tensor, nn from graphlm.neuron.backbone import CausalSelfAttention +from graphlm.neuron.graph_attention import HybridGraphCausalSelfAttention from graphlm.neuron.graph_hybrid import AdjInnerInit, AdjOuterInit, HybridGraphLinear from graphlm.neuron.rms_norm import RMSNorm @@ -165,6 +166,53 @@ def forward(self, x: Tensor) -> Tensor: return x + self.ffn(self.rms2(x)) +class FullGraphTransformerBlock(nn.Module): + """Phase 14 — attention + FFN 둘 다 HybridGraphLinear 인 full graph block. + + Phase 13 의 ``HybridGraphTransformerBlock`` (FFN-only graph) 와 동일 forward 구조이되, + ``CausalSelfAttention`` (표준 nn.Linear) 대신 ``HybridGraphCausalSelfAttention`` + (qkv + out 도 HybridGraphLinear) 사용. **block 내 모든 Linear 가 graph**. + + adj_outer_init 의 ``"identity"`` 는 qkv (rectangular) + FFN fc1 (rectangular) 둘 다 + rejection 하므로 사용 불가. + """ + + def __init__( + self, + hidden_dim: int, + n_heads: int, + ffn_dim: int, + group_size: int, + *, + adj_outer_init: AdjOuterInit = "full", + adj_inner_init: AdjInnerInit = "full", + dropout: float = 0.0, + ): + super().__init__() + self.rms1 = RMSNorm(hidden_dim) + self.attn = HybridGraphCausalSelfAttention( + hidden_dim, + n_heads, + group_size=group_size, + adj_outer_init=adj_outer_init, + adj_inner_init=adj_inner_init, + dropout=dropout, + ) + self.rms2 = RMSNorm(hidden_dim) + self.ffn = HybridGraphFFN( + hidden_dim, + ffn_dim, + group_size=group_size, + adj_outer_init=adj_outer_init, + adj_inner_init=adj_inner_init, + dropout=dropout, + ) + + def forward(self, x: Tensor) -> Tensor: + x = x + self.attn(self.rms1(x)) + return x + self.ffn(self.rms2(x)) + + def make_block( arch: Arch, hidden_dim: int, @@ -207,3 +255,51 @@ def make_block( dropout=dropout, ) raise ValueError(f"unknown arch: {arch}") + + +def make_full_block( + arch: Arch, + hidden_dim: int, + n_heads: int, + ffn_dim: int, + group_size: int, + dropout: float = 0.0, +) -> nn.Module: + """4 가지 arch 중 하나로 Phase 14 **full graph** Transformer block 생성. + + Phase 13 ``make_block`` 과 동일 arch literal 이되, hybrid_* 는 attention 도 graph 화. + plain 은 Phase 13 와 동일 (PlainTransformerBlock). + """ + if arch == "plain": + return PlainTransformerBlock(hidden_dim, n_heads, ffn_dim, dropout=dropout) + if arch == "hybrid_full_full": + return FullGraphTransformerBlock( + hidden_dim, + n_heads, + ffn_dim, + group_size=group_size, + adj_outer_init="full", + adj_inner_init="full", + dropout=dropout, + ) + if arch == "hybrid_full_around_one": + return FullGraphTransformerBlock( + hidden_dim, + n_heads, + ffn_dim, + group_size=group_size, + adj_outer_init="full", + adj_inner_init="uniform_around_one", + dropout=dropout, + ) + if arch == "hybrid_around_one_around_one": + return FullGraphTransformerBlock( + hidden_dim, + n_heads, + ffn_dim, + group_size=group_size, + adj_outer_init="uniform_around_one", + adj_inner_init="uniform_around_one", + dropout=dropout, + ) + raise ValueError(f"unknown arch: {arch}") diff --git a/src/graphlm/neuron/hybrid_transformer_demo.py b/src/graphlm/neuron/hybrid_transformer_demo.py index 4d3343f..8e95892 100644 --- a/src/graphlm/neuron/hybrid_transformer_demo.py +++ b/src/graphlm/neuron/hybrid_transformer_demo.py @@ -23,8 +23,10 @@ from graphlm.data.tinyshakespeare import TinyShakespeareDataset, iter_random_batches from graphlm.neuron.hybrid_transformer import ( Arch, + FullGraphTransformerBlock, HybridGraphTransformerBlock, make_block, + make_full_block, ) from graphlm.neuron.rms_norm import RMSNorm from graphlm.utils import set_seed @@ -49,6 +51,9 @@ class HybridTransformerTrainConfig: group_size: int arch: Arch dropout: float = 0.0 + # Phase 14: True → make_full_block (attention + FFN 둘 다 graph), + # False → make_block (FFN-only graph, Phase 13 default) + use_full_graph: bool = False # train block_size: int = 64 @@ -80,15 +85,18 @@ def __init__( arch: Arch, group_size: int, dropout: float = 0.0, + use_full_graph: bool = False, ): super().__init__() self.arch = arch + self.use_full_graph = use_full_graph self.max_seq_len = max_seq_len self.tok_emb = nn.Embedding(vocab_size, hidden_dim) self.pos_emb = nn.Embedding(max_seq_len, hidden_dim) + block_factory = make_full_block if use_full_graph else make_block self.blocks = nn.ModuleList( [ - make_block( + block_factory( arch, hidden_dim=hidden_dim, n_heads=n_heads, @@ -114,31 +122,41 @@ def forward(self, x: Tensor) -> Tensor: return self.lm_head(h) -def _block_iter(model: HybridGraphTransformerLM) -> Iterator[HybridGraphTransformerBlock]: - """모델의 hybrid block 만 yield (plain 은 skip).""" +def _block_iter( + model: HybridGraphTransformerLM, +) -> Iterator[HybridGraphTransformerBlock | FullGraphTransformerBlock]: + """모델의 hybrid block 만 yield (plain 은 skip). Phase 13 + Phase 14 둘 다 지원.""" for blk in model.blocks: - if isinstance(blk, HybridGraphTransformerBlock): + if isinstance(blk, HybridGraphTransformerBlock | FullGraphTransformerBlock): yield blk +def _snapshot_layer(layer) -> dict[str, Tensor]: + """HybridGraphLinear 한 layer 의 outer/inner snapshot.""" + return { + "outer": layer.adj_outer.detach().cpu().clone(), + "inner": layer.adj_inner.detach().cpu().clone(), + } + + def _snapshot_adj(model: HybridGraphTransformerLM) -> list[dict[str, dict[str, Tensor]]] | None: - """hybrid arch 인 경우 각 block 의 FFN adj snapshot (Phase 12 demo 와 동일 hierarchy).""" + """hybrid arch 인 경우 각 block 의 adj snapshot. + + - Phase 13 (FFN-only graph): ``{"fc1", "fc2"}`` + - Phase 14 (full graph): ``{"qkv", "out", "fc1", "fc2"}`` — attention adj 까지 포함 + """ if model.arch == "plain": return None snapshots: list[dict[str, dict[str, Tensor]]] = [] for blk in _block_iter(model): - snapshots.append( - { - "fc1": { - "outer": blk.ffn.fc1.adj_outer.detach().cpu().clone(), - "inner": blk.ffn.fc1.adj_inner.detach().cpu().clone(), - }, - "fc2": { - "outer": blk.ffn.fc2.adj_outer.detach().cpu().clone(), - "inner": blk.ffn.fc2.adj_inner.detach().cpu().clone(), - }, - } - ) + snap: dict[str, dict[str, Tensor]] = { + "fc1": _snapshot_layer(blk.ffn.fc1), + "fc2": _snapshot_layer(blk.ffn.fc2), + } + if isinstance(blk, FullGraphTransformerBlock): + snap["qkv"] = _snapshot_layer(blk.attn.qkv) + snap["out"] = _snapshot_layer(blk.attn.out) + snapshots.append(snap) return snapshots @@ -159,6 +177,7 @@ def train_hybrid_transformer_lm(config: HybridTransformerTrainConfig) -> dict: arch=config.arch, group_size=config.group_size, dropout=config.dropout, + use_full_graph=config.use_full_graph, ).to(config.device) data_iter = iter_random_batches( config.dataset, batch_size=config.batch_size, block_size=config.block_size, seed=config.seed diff --git a/tests/neuron/test_graph_attention.py b/tests/neuron/test_graph_attention.py new file mode 100644 index 0000000..c17ef32 --- /dev/null +++ b/tests/neuron/test_graph_attention.py @@ -0,0 +1,122 @@ +"""Tests for graphlm.neuron.graph_attention — Phase 14 graph attention.""" + +from __future__ import annotations + +import pytest +import torch + +from graphlm.neuron.backbone import CausalSelfAttention +from graphlm.neuron.graph_attention import HybridGraphCausalSelfAttention +from graphlm.neuron.graph_hybrid import HybridGraphLinear + + +def test_shape(): + attn = HybridGraphCausalSelfAttention(hidden_dim=32, n_heads=4, group_size=8) + x = torch.randn(2, 16, 32) + assert attn(x).shape == (2, 16, 32) + + +def test_qkv_out_are_hybrid_graph_linear(): + attn = HybridGraphCausalSelfAttention(hidden_dim=32, n_heads=4, group_size=8) + assert isinstance(attn.qkv, HybridGraphLinear) + assert isinstance(attn.out, HybridGraphLinear) + # qkv: 32 → 96, out: 32 → 32 + assert attn.qkv.in_features == 32 + assert attn.qkv.out_features == 96 + assert attn.out.in_features == 32 + assert attn.out.out_features == 32 + + +def test_identity_outer_rejected(): + """qkv 가 rectangular 라 identity outer 미지원.""" + with pytest.raises(ValueError, match="rectangular"): + HybridGraphCausalSelfAttention( + hidden_dim=32, n_heads=4, group_size=8, adj_outer_init="identity" + ) + + +def test_heads_not_divisible_raises(): + with pytest.raises(ValueError, match="not divisible by n_heads"): + HybridGraphCausalSelfAttention(hidden_dim=33, n_heads=4, group_size=11) + + +def test_function_preservation_full_full_matches_standard(): + """adj=full/full + 같은 W → standard CausalSelfAttention 와 forward 동치 (atol=1e-5).""" + torch.manual_seed(0) + hidden, n_heads, k = 32, 4, 8 + hg_attn = HybridGraphCausalSelfAttention(hidden, n_heads, group_size=k) + std_attn = CausalSelfAttention(hidden, n_heads, dropout=0.0) + + _copy_hybrid_to_plain(hg_attn.qkv, std_attn.qkv) + _copy_hybrid_to_plain(hg_attn.out, std_attn.out) + + hg_attn.eval() + std_attn.eval() + x = torch.randn(2, 8, hidden) + with torch.no_grad(): + y_hg = hg_attn(x) + y_std = std_attn(x) + assert torch.allclose(y_hg, y_std, atol=1e-5), ( + f"attention forward 차이: max |diff| = {(y_hg - y_std).abs().max().item()}" + ) + + +@pytest.mark.parametrize( + "outer,inner", + [ + ("full", "full"), + ("full", "uniform_around_one"), + ("uniform_around_one", "uniform_around_one"), + ], +) +def test_gradient_flows_all_params(outer, inner): + attn = HybridGraphCausalSelfAttention( + hidden_dim=16, + n_heads=4, + group_size=4, + adj_outer_init=outer, + adj_inner_init=inner, + ) + x = torch.randn(2, 4, 16, requires_grad=False) + attn(x).sum().backward() + null_grad = [n for n, p in attn.named_parameters() if p.grad is None] + assert not null_grad, f"grad 없는 파라미터: {null_grad}" + + +@pytest.mark.parametrize("bad", ["zero", "zeros"]) +def test_zero_init_outer_rejected(bad): + """0-init outer 거부 — underlying HybridGraphLinear 로부터 상속.""" + with pytest.raises(ValueError, match="vanishing"): + HybridGraphCausalSelfAttention( + hidden_dim=32, + n_heads=4, + group_size=8, + adj_outer_init=bad, # type: ignore[arg-type] + ) + + +@pytest.mark.parametrize("bad", ["zero", "zeros"]) +def test_zero_init_inner_rejected(bad): + """0-init inner 거부.""" + with pytest.raises(ValueError, match="vanishing"): + HybridGraphCausalSelfAttention( + hidden_dim=32, + n_heads=4, + group_size=8, + adj_inner_init=bad, # type: ignore[arg-type] + ) + + +# ── helpers ── + + +def _copy_hybrid_to_plain(hg, plain): + """HybridGraphLinear 의 block weight → nn.Linear 표준 weight 형식으로 복사.""" + in_f, out_f = hg.in_features, hg.out_features + k = hg.group_size + W_std = torch.zeros(out_f, in_f) + for go in range(hg.n_groups_out): + for gi in range(hg.n_groups_in): + W_std[go * k : (go + 1) * k, gi * k : (gi + 1) * k] = hg.weight[go, gi].T + with torch.no_grad(): + plain.weight.copy_(W_std) diff --git a/tests/neuron/test_hybrid_transformer.py b/tests/neuron/test_hybrid_transformer.py index f366a52..27f942b 100644 --- a/tests/neuron/test_hybrid_transformer.py +++ b/tests/neuron/test_hybrid_transformer.py @@ -7,11 +7,13 @@ from torch import nn from graphlm.neuron.hybrid_transformer import ( + FullGraphTransformerBlock, HybridGraphFFN, HybridGraphTransformerBlock, PlainFFN, PlainTransformerBlock, make_block, + make_full_block, ) # ── HybridGraphFFN ─────────────────────────────────────────── @@ -189,6 +191,88 @@ def test_block_dropout_propagates_to_ffn(block_cls, kwargs): assert block.ffn.dropout.p == 0.5, "FFN 의 dropout p 가 block dropout 과 불일치" +# ── FullGraphTransformerBlock (Phase 14) ───────────────────── + + +def test_full_block_shape(): + block = FullGraphTransformerBlock(hidden_dim=32, n_heads=4, ffn_dim=64, group_size=8) + x = torch.randn(2, 16, 32) + assert block(x).shape == (2, 16, 32) + + +def test_full_block_function_preservation_against_plain(): + """full graph block (adj=full/full) + plain block 의 동일 weight 로 forward 동일.""" + torch.manual_seed(0) + hidden, n_heads, ffn_d, k = 16, 4, 32, 4 + full = FullGraphTransformerBlock(hidden, n_heads, ffn_d, group_size=k) + plain = PlainTransformerBlock(hidden, n_heads, ffn_d) + + plain.rms1.load_state_dict(full.rms1.state_dict()) + plain.rms2.load_state_dict(full.rms2.state_dict()) + # attention: HybridGraphLinear qkv / out → standard nn.Linear 로 복사 + _copy_hybrid_to_plain(full.attn.qkv, plain.attn.qkv) + _copy_hybrid_to_plain(full.attn.out, plain.attn.out) + # FFN: 동일 패턴 + _copy_hybrid_to_plain(full.ffn.fc1, plain.ffn.fc1) + _copy_hybrid_to_plain(full.ffn.fc2, plain.ffn.fc2) + + full.eval() + plain.eval() + x = torch.randn(2, 8, hidden) + with torch.no_grad(): + y_full = full(x) + y_plain = plain(x) + assert torch.allclose(y_full, y_plain, atol=1e-5), ( + f"full graph block forward 차이: max |diff| = {(y_full - y_plain).abs().max().item()}" + ) + + +def test_full_block_gradient_flows_all_params(): + block = FullGraphTransformerBlock( + hidden_dim=16, + n_heads=4, + ffn_dim=32, + group_size=4, + adj_outer_init="uniform_around_one", + adj_inner_init="uniform_around_one", + ) + x = torch.randn(2, 4, 16) + block(x).sum().backward() + null_grad = [n for n, p in block.named_parameters() if p.grad is None] + assert not null_grad, f"grad 없는 파라미터: {null_grad}" + + +@pytest.mark.parametrize( + "arch", + ["plain", "hybrid_full_full", "hybrid_full_around_one", "hybrid_around_one_around_one"], +) +def test_make_full_block_all_archs_forward(arch): + block = make_full_block(arch, hidden_dim=16, n_heads=4, ffn_dim=32, group_size=4) + x = torch.randn(2, 8, 16) + assert block(x).shape == (2, 8, 16) + + +def test_make_full_block_plain_uses_nn_linear(): + """make_full_block 의 plain 은 make_block 의 plain 과 동일 (PlainTransformerBlock).""" + block = make_full_block("plain", 16, 4, 32, 4) + assert isinstance(block, PlainTransformerBlock) + + +def test_make_full_block_hybrid_uses_full_graph(): + """make_full_block 의 hybrid_* 는 FullGraphTransformerBlock (attention 도 graph).""" + block = make_full_block("hybrid_full_full", 16, 4, 32, 4) + assert isinstance(block, FullGraphTransformerBlock) + # attention 도 HybridGraphLinear 인지 확인 + from graphlm.neuron.graph_attention import HybridGraphCausalSelfAttention + + assert isinstance(block.attn, HybridGraphCausalSelfAttention) + + +def test_make_full_block_unknown_raises(): + with pytest.raises(ValueError, match="unknown arch"): + make_full_block("bogus", 16, 4, 32, 4) # type: ignore[arg-type] + + # ── helpers ──────────────────────────────────────────────────