Sparse Evo-MemoryLM 詳細数理・メモリ設計
実装状態: 実装済み。対象は
EvoSpikeNet-Core/evospikenet/sparse_event_memory.py。
対象:SparseEventMemoryLM、EventCSRLayer、TiedFactorizedVocabulary、SparseEventMemoryBlock、EventMemoryExpert。
非対象: この文書は既存の密なSpikingEvoTextLM/ChronoSpikeAttention/SpikingFFNを置換する仕様ではない。選択時だけ使う別アーキテクチャである。
目次
- 必要性と設計目標
- 実装上の境界
- 記号
- 全体データフロー
- 因子化共有語彙の数理
- INT8 CSRシナプスの数理
- EvoLIF形式イベント状態の数理
- 専門家ルーティングと低ランクアダプタ
- 損失・Adam・局所可塑性
- 実メモリ配置と寿命
- 容量式と密モデルとの比較
- 学習・推論・生成の実行順序
- 実測構成例
- 制約、誤解しやすい点、運用指針
必要性と設計目標
密な言語モデルでVRAMが増える理由
密なTransformer系の1ブロックには、典型的にQ/K/V/出力射影とFFNがあり、隠れ次元を \(D\) とすると、主な重みは概ね
に比例します。FP32のAdamで全重みを学習すると、最低でも重み・勾配・一次モーメント・二次モーメントが必要であり、概算は
です。ここで \(P\) は学習対象パラメータ数です。これはアクティベーション、CUDAワークスペース、入力/出力logitsを含まない下限です。
EvoLIFやsoftmax-freeのChronoSpikeAttentionを用いても、密なQ/K/V/FFN重みとそれらのAdam状態が残る限り、この \(D^2\) の主項は消えません。
Sparse Evo-MemoryLMが採る分離
本モデルは容量を次の3領域に分離します。
| 領域 | 役割 | 保持形式 | 学習方式 |
|---|---|---|---|
| 主記憶 | 入力・再帰シナプス | 固定トポロジーCSR、INT8値 | 局所Hebbian更新 |
| 小規模可塑部 | 語彙基底、ルータ、アダプタ、LayerNorm | 浮動小数nn.Parameter |
Adam |
| 動的状態 | 膜電位、直前スパイク、活動統計 | INT16 / 浮動小数一時テンソル | 系列ごとに作成・解放 |
目的は「総パラメータ数を増やしてもメモリがゼロになる」ことではありません。大きい主シナプス行列をAdam対象から外し、GPUに必要なAdam状態を小さくすることです。
実装上の境界
実装の実態を以下に固定します。
- CSRの
crow_indices・col_indices・quantized_valuesはbufferであり、nn.Parameterではありません。 quantized_valuesはINT8で保持されますが、forward()ではsparse.mmのため一時的にfloat CSRへ変換されます。- 接続トポロジーは初期化時に作られ、学習中に接続の追加・削除・再配線は行いません。
- ルータは系列平均を入力としてTop-k専門家を選びます。Top-kのインデックス選択そのものは離散的であり、選択境界に対する微分可能なルーティングではありません。
EventMemoryExpertは時刻方向に反復します。ただし現実装は出力状態をPython listに蓄積して最後にstackするため、出力シーケンスとそのautograd情報は保持されます。「膜電位だけで系列全体の活性化メモリが一定」という実装ではありません。generate()は毎トークンで現在までの全promptを再評価します。KV cacheや再帰状態キャッシュは現実装にありません。--ssl-task reconstruction、MetaSTDP、AEGはSparse Evo-MemoryLMの学習分岐では使用しません。
記号
| 記号 | 意味 |
|---|---|
| \(V\) | 語彙サイズ |
| \(D\) | d_model、イベント状態次元 |
| \(R\) | factor_rank、共有語彙因子rank |
| \(L\) | num_transformer_blocks、疎イベントブロック数 |
| \(M\) | num_experts、ブロックごとの専門家数 |
| \(K\) | router_top_k、系列ごとに実行する専門家数 |
| \(\rho\) | connectivity、CSR行あたりの接続密度 |
| \(F\) | CSR行あたりfan-in。\(F=\min(D,\max(1,\lceil\rho D\rceil))\) |
| \(A\) | adapter_rank |
| \(B\) | バッチサイズ |
| \(S\) | 系列長 |
| \(C\) | sampled softmax候補語彙数 |
| \(\theta\) | event_threshold |
| \(\lambda\) | event_leak、実装では256を分母とする整数係数 |
| \(\eta_h\) | plasticity_learning_rate |
全体データフロー
flowchart TD
A[Token IDs] --> B[INT8 codebook lookup]
B --> C[Shared factor basis]
C --> D[Sequence router]
D --> E[Top-k sparse event experts]
E --> F[INT8 CSR input and recurrent synapses]
F --> G[INT16 membrane and spike state]
G --> H[LayerNorm output state]
H --> I[Sampled vocabulary projection]
I --> J[Cross-entropy]
J --> K[Adam on small trainable tensors]
K --> L[Local Hebbian update of INT8 CSR values]
classDef fixed fill:#e8f3ff,stroke:#1976d2,color:#102a43
classDef trainable fill:#e8f5e9,stroke:#2e7d32,color:#102a43
classDef state fill:#fff3e0,stroke:#ef6c00,color:#102a43
class B,F fixed
class C,D,H,I,K trainable
class G,L state
各ブロックは同じ入力系列を受け、選択された専門家の出力をルータ重みで加算します。ブロック出力が次のブロック入力となります。
因子化共有語彙の数理
保存形式
語彙層は、固定INT8コードブック
と、Adamで学習する共有基底
を持ちます。トークン \(w\) の入力表現は
です。入力埋め込みと出力投影で同じ \(Q,E\) を使うため、入力用と出力用の別々の \(V \times D\) 行列を持ちません。
出力射影
隠れ状態 \(\mathbf{h}\in\mathbb{R}^{D}\) から潜在語彙座標を
として計算します。候補集合 \(\mathcal{C}\) のlogitは
です。
forward()では \(\mathcal{C}=\{0,\ldots,V-1\}\) として全語彙logitsを作ります。sampled_cross_entropy()では、全正解語のunique集合とランダム負例のunique集合だけを候補にし、\(|\mathcal{C}|=C\) に抑えます。
候補分類の損失は
です。これは候補集合上のcross entropyであり、全語彙softmaxと厳密に同じ目的関数ではありません。
INT8 CSRシナプスの数理
接続トポロジー
各EventCSRLayer(D,D)は各出力行に \(F\) 個の接続を持ちます。
行 \(i\) の開始オフセット \(o_i\) と互いに素なstride \(s_i\) により、接続列は
として初期化されます。これはdense maskを作らずに、各行へ広がった固定近傍を与える実装です。
CSR値は
として保存されます。入力 \(\mathbf{x}\) のシナプス出力は、実際にはfloatへキャストした値で
を計算します。INT8のまま積和演算する専用neuromorphic kernelではありません。
物理保持量
1 CSR層の保持バイト数は、実装のINT64 crow_indices、INT64 col_indices、INT8値から
です。これは一時float CSRとkernel workspaceを含みません。
EvoLIF形式イベント状態の数理
専門家 \(m\)、時刻 \(t\) の入力を \(\mathbf{x}_t\)、直前発火を \(\mathbf{s}_{t-1}\) とします。低ランクアダプタを
と定義します。入力・再帰CSR層の電流を合わせると
です。
膜電位はINT16で、実装は次の固定小数点更新を行います。
順伝播で出力するイベントにはstraight-through surrogateを使います。
よって前向き値はhard spikeですが、逆伝播では \(\sigma(\mathbf{i}_t)\) を通る勾配がアダプタと入力側へ流れます。INT8値そのものはbufferであるためAdamの勾配を持ちません。
出力状態は
です。
専門家ルーティングと低ランクアダプタ
入力系列 \(X\in\mathbb{R}^{B\times S\times D}\) の系列平均を
とし、ルータは
を計算します。topkで選択した専門家添字を \(I_b\) とします。
- \(K=1\) のとき、選択専門家のゲートは \(g_{b,1}=\sigma(r_{b,I_b})\) です。単一要素softmaxが常に1となりルータ勾配が消えることを避けるためです。
- \(K>1\) のとき、選択logit上で \(g_{b,k}=\mathrm{softmax}(r_{b,I_{b,k}})\) を使います。
ブロック出力は
です。選択された専門家だけを実行するため、実行される専門家本体はおおむね \(K\) に比例します。ただし全専門家のCSR bufferとadapter parameterはモデルに常駐します。
損失・Adam・局所可塑性
Adamの対象
Adamに渡されるのは以下だけです。
- 共有基底 \(E\)
- 各ブロックのルータ \(W_r,\mathbf{b}_r\)
- 各専門家の \(A_{\downarrow},A_{\uparrow}\)
- LayerNormのscaleとbias
CSR indices、INT8値、語彙codebook、膜電位はnn.Parameterではありません。
局所Hebbian更新
各forwardは、入力、再帰入力、spikeの系列・バッチ平均を記録します。第 \(j\) 入力と第 \(i\) 出力の接続について、概念的な更新は
です。実装はCSR接続だけに対してindex_selectでこの式を評価します。各記録に対し入力CSRと再帰CSRの2更新を行い、処理後に活動記録をclearします。
この更新はCE損失の厳密な勾配降下ではありません。主シナプスの教師信号は局所相関だけであり、収束・精度・忘却は実データで評価する必要があります。
実メモリ配置と寿命
永続配置
| オブジェクト | dtype | 概算サイズ | 寿命 |
|---|---|---|---|
| CSR値 | INT8 | \(2LMDF\) bytes | モデル寿命 |
| CSR列index | INT64 | \(16LMDF\) bytes | モデル寿命 |
| CSR行ptr | INT64 | \(16LM(D+1)\) bytes | モデル寿命 |
| 語彙codebook | INT8 | \(VR\) bytes | モデル寿命 |
| 共有基底 | fp32標準 | \(4RD\) bytes | モデル寿命、Adam対象 |
| router | fp32標準 | \(4L(DM+M)\) bytes | モデル寿命、Adam対象 |
| adapters | fp32標準 | \(8LMDA\) bytes | モデル寿命、Adam対象 |
| LayerNorm | fp32標準 | \(8LMD\) bytes | モデル寿命、Adam対象 |
memory_report()は、上表のうち学習可能parameter数、CSR接続数・CSR保持量、語彙codebook量を返します。dtypeは利用者が後から変えられるため、trainable parameterのバイト数そのものは報告しません。
1 forward中の一時配置
| オブジェクト | 代表shape | 主なdtype | 備考 |
|---|---|---|---|
| token code lookup | \((B,S,R)\) | INT8 | index_select出力 |
| code cast / embedding | \((B,S,R)\) / \((B,S,D)\) | fp32標準 | 基底dtypeに追従 |
| membrane | \((B,D)\) | INT16 | 各専門家forwardで新規作成 |
| previous spikes | \((B,D)\) | 入力dtype | 実装上は浮動小数 |
| synaptic current / event | \((B,D)\) | 浮動小数 | 時刻ごとに計算 |
| output list + stack | \(S\)個の\((B,D)\) | 浮動小数+autograd | 系列長に比例 |
| float CSR値 | \(DF\) | 浮動小数 | quantized_values.to(dtype)で作成 |
| plasticity record | 3個の\((D)\) | fp32 | 実行専門家・forwardごとに蓄積 |
| sampled logits | \((BS,C)\) | 浮動小数 | 全語彙ではなく候補数に比例 |
CSR値はINT8で保持していても、現行PyTorch実装のforwardはfloat値のCSRテンソルを一時生成します。またspikeはインデックス列ではなく\((B,D)\)のdense tensorです。従って「計算中の全テンソルが疎・整数」という意味ではありません。
容量式と密モデルとの比較
Sparse Evo-MemoryLM
全ブロック・全専門家の主シナプス接続数は
です。前の係数2は各専門家がinput CSRとrecurrent CSRを1つずつ持つためです。CSR永続保持量は
です。
学習可能parameter数の主項は
です。これはbiasなしadapterとLayerNormのweight/biasを含む現実装に対応します。FP32 Adamの概算下限は \(16P_{\mathrm{trainable}}\) bytesです。
密なブロックとの対比
| 項目 | 密ChronoSpike + FFNの主項 | Sparse Evo-MemoryLMの主項 |
|---|---|---|
| ブロック重み | \(12D^2\) | CSR接続 \(2MDF\) + adapters \(2MDA\) |
| 主シナプス値 | 通常fp32 parameter | INT8 buffer |
| 主シナプスAdam | 必要 | 不要 |
| 語彙 | 典型的に入力/出力で \(O(VD)\) | INT8 \(VR\) + 共有基底 \(RD\) |
| 学習出力 | 全語彙CEなら \(O(BSV)\) | sampled CEなら \(O(BSC)\) |
接続率が小さく \(F\ll D\)、かつ \(A,R\ll D\) なら、密な \(D^2\) parameterとそのAdam状態を大きく減らせます。一方、CSR indexがINT64であるため、極端な低接続率以外では値1 byteだけを比較してはいけません。
学習・推論・生成の実行順序
学習
sequenceDiagram
participant T as Trainer
participant V as TiedFactorizedVocabulary
participant B as SparseEventMemoryBlock
participant X as EventMemoryExpert
participant O as Adam
participant H as Local Hebbian updater
T->>V: embed(input token IDs)
V-->>B: hidden states
loop each block
B->>B: mean pool and Top-k route
B->>X: selected sequences only
loop each token position
X->>X: CSR input + CSR recurrence + adapter
X->>X: INT16 membrane / hard spike / surrogate
end
X-->>B: normalized sequence states
end
B-->>T: final hidden states
T->>V: sampled candidate projection and CE
T->>O: backward + optimizer.step()
T->>H: apply_local_plasticity()
apply_local_plasticity()はoptimizer.step()後に呼びます。例外、NaN検出、勾配stepのskipなどでAdam更新が成功しない場合は、古い活動記録を使わないためにclear_local_plasticity()を呼びます。
推論と生成
forward(token_ids)は全語彙logitsを返します。generate()は温度と任意のtop-kで次トークンをsampleしますが、各生成時刻にprompt全体をforward()へ渡します。長い生成での計算量/一時メモリを抑えるKV cache・膜状態キャッシュは未実装です。
実測構成例
次は実装済みのmemory_report()で構成した例です。
このとき \(F=\lceil0.005\times2048\rceil=11\)、主CSR接続数は
です。実行確認では次を得ています。
| 指標 | 値 |
|---|---|
| Adam対象parameter | 1,982,488 |
| 固定INT8 CSR接続 | 1,081,344 |
| CSR永続保持量 | 約10.03 MiB |
| 語彙INT8コード保持量 | 4.00 MiB |
これはモデル本体の静的報告です。候補logits、autograd、float CSR変換、optimizer、CUDA allocator、tokenizerは別途必要です。RTX 2070 8GBなどで利用する場合も、実際のbatch/sequence長でtorch.cuda.max_memory_allocated()を測定してください。
制約、誤解しやすい点、運用指針
制約
- PyTorch CSRはbeta: 利用可能なGPU kernel、autograd挙動、性能はPyTorch/CUDA版に依存します。
- 局所則は全体誤差の勾配ではない: 主シナプスの更新はCEから直接最適化されません。
- Top-kは負荷均衡を持たない: 現実装にはexpert load-balancing loss、capacity factor、expert offloadはありません。
- 生成cacheなし: 長い生成ではprompt再評価が重なります。
- 全語彙forwardは大きい: 推論で\((B,S,V)\)logitsを明示生成します。
- 重みのINT8保持とINT8演算は異なる: forward中のCSR値はfloatへキャストされます。
推奨運用
- 学習は
sampled_cross_entropy()を使い、--sampled-negativesを測定しながら設定する。 - まず小さい \(D,L,M\) と短い \(S\) でloss、発火率、
LocalSynapseUpdates、ピークVRAMを確認する。 --sparse-connectivityを上げる前にCSR indexの保持量とsparse kernel性能を測定する。- checkpoint再開時は
architecture: sparse_event_memory、語彙サイズ、\(D,R,L,M,K,A,\rho\) が一致することを確認する。 - Docker SDKへ保存する際は、model/config/memory reportに加えてtokenizer archiveを同じ
llm_type="SparseEventMemoryLM"で管理する。