コンテンツにスキップ

Sparse Evo-MemoryLM 詳細数理・メモリ設計

実装状態: 実装済み。対象は EvoSpikeNet-Core/evospikenet/sparse_event_memory.py
対象: SparseEventMemoryLMEventCSRLayerTiedFactorizedVocabularySparseEventMemoryBlockEventMemoryExpert
非対象: この文書は既存の密な SpikingEvoTextLM / ChronoSpikeAttention / SpikingFFN を置換する仕様ではない。選択時だけ使う別アーキテクチャである。

目次

  1. 必要性と設計目標
  2. 実装上の境界
  3. 記号
  4. 全体データフロー
  5. 因子化共有語彙の数理
  6. INT8 CSRシナプスの数理
  7. EvoLIF形式イベント状態の数理
  8. 専門家ルーティングと低ランクアダプタ
  9. 損失・Adam・局所可塑性
  10. 実メモリ配置と寿命
  11. 容量式と密モデルとの比較
  12. 学習・推論・生成の実行順序
  13. 実測構成例
  14. 制約、誤解しやすい点、運用指針

必要性と設計目標

密な言語モデルでVRAMが増える理由

密なTransformer系の1ブロックには、典型的にQ/K/V/出力射影とFFNがあり、隠れ次元を \(D\) とすると、主な重みは概ね

\[ 4D^2 + (D \cdot 4D + 4D \cdot D) = 12D^2 \]

に比例します。FP32のAdamで全重みを学習すると、最低でも重み・勾配・一次モーメント・二次モーメントが必要であり、概算は

\[ M_{\mathrm{Adam,min}} \approx 16P\ \mathrm{bytes} \]

です。ここで \(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_indicescol_indicesquantized_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 reconstructionMetaSTDPAEGは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コードブック

\[ Q \in \mathbb{Z}_8^{V \times R} \]

と、Adamで学習する共有基底

\[ E \in \mathbb{R}^{R \times D} \]

を持ちます。トークン \(w\) の入力表現は

\[ \mathbf{x}_w = Q_w E \in \mathbb{R}^{D} \]

です。入力埋め込みと出力投影で同じ \(Q,E\) を使うため、入力用と出力用の別々の \(V \times D\) 行列を持ちません。

出力射影

隠れ状態 \(\mathbf{h}\in\mathbb{R}^{D}\) から潜在語彙座標を

\[ \mathbf{z} = \frac{E\mathbf{h}}{\sqrt{R}} \in \mathbb{R}^{R} \]

として計算します。候補集合 \(\mathcal{C}\) のlogitは

\[ \ell_c = \mathbf{z}^{\mathsf{T}} Q_c, \quad c\in\mathcal{C} \]

です。

  • forward()では \(\mathcal{C}=\{0,\ldots,V-1\}\) として全語彙logitsを作ります。
  • sampled_cross_entropy()では、全正解語のunique集合とランダム負例のunique集合だけを候補にし、\(|\mathcal{C}|=C\) に抑えます。

候補分類の損失は

\[ \mathcal{L}_{\mathrm{sampled}}=-\frac{1}{BS}\sum_{b,t} \log\frac{\exp(\ell_{y_{b,t}})}{\sum_{c\in\mathcal{C}}\exp(\ell_c)} \]

です。これは候補集合上のcross entropyであり、全語彙softmaxと厳密に同じ目的関数ではありません。


INT8 CSRシナプスの数理

接続トポロジー

EventCSRLayer(D,D)は各出力行に \(F\) 個の接続を持ちます。

\[ F = \min(D,\max(1,\lceil\rho D\rceil)),\qquad N_{\mathrm{conn}}=D F \]

\(i\) の開始オフセット \(o_i\) と互いに素なstride \(s_i\) により、接続列は

\[ j_{i,p}=(o_i+s_i p)\bmod D, \quad p\in\{0,\ldots,F-1\} \]

として初期化されます。これはdense maskを作らずに、各行へ広がった固定近傍を与える実装です。

CSR値は

\[ w_{i,p}\in\{-127,\ldots,127\}\subset\mathbb{Z}_8 \]

として保存されます。入力 \(\mathbf{x}\) のシナプス出力は、実際にはfloatへキャストした値で

\[ [\mathcal{S}(\mathbf{x})]_i=\sum_{p=0}^{F-1}\mathrm{float}(w_{i,p})x_{j_{i,p}} \]

を計算します。INT8のまま積和演算する専用neuromorphic kernelではありません。

物理保持量

1 CSR層の保持バイト数は、実装のINT64 crow_indices、INT64 col_indices、INT8値から

\[ M_{\mathrm{CSR,layers}}=8(D+1)+8DF+DF=8(D+1)+9DF\ \mathrm{bytes} \]

です。これは一時float CSRとkernel workspaceを含みません。


EvoLIF形式イベント状態の数理

専門家 \(m\)、時刻 \(t\) の入力を \(\mathbf{x}_t\)、直前発火を \(\mathbf{s}_{t-1}\) とします。低ランクアダプタを

\[ \mathbf{a}_t=(\mathbf{x}_t A_{\downarrow})A_{\uparrow}, \quad A_{\downarrow}\in\mathbb{R}^{D\times A}, \quad A_{\uparrow}\in\mathbb{R}^{A\times D} \]

と定義します。入力・再帰CSR層の電流を合わせると

\[ \mathbf{i}_t=\mathcal{S}_{\mathrm{in}}(\mathbf{x}_t) +\mathcal{S}_{\mathrm{rec}}(\mathbf{s}_{t-1})+\mathbf{a}_t \]

です。

膜電位はINT16で、実装は次の固定小数点更新を行います。

\[ \widetilde{\mathbf{i}}_t=\mathrm{round}(\theta\,\mathrm{detach}(\mathbf{i}_t))\in\mathbb{Z}_{32} \]
\[ \mathbf{v}_t=\mathrm{clip}_{[-32768,32767]} \left( \left\lfloor\frac{\lambda\mathbf{v}_{t-1}}{256}\right\rfloor+ \widetilde{\mathbf{i}}_t \right)\in\mathbb{Z}_{16} \]
\[ \mathbf{s}_t=\mathbb{1}[\mathbf{v}_t\ge\theta], \qquad \mathbf{v}_t\leftarrow0\quad\text{for fired elements.} \]

順伝播で出力するイベントにはstraight-through surrogateを使います。

\[ \widetilde{\mathbf{s}}_t= \mathbf{s}_t+\sigma(\mathbf{i}_t)-\mathrm{stopgrad}(\sigma(\mathbf{i}_t)) \]

よって前向き値はhard spikeですが、逆伝播では \(\sigma(\mathbf{i}_t)\) を通る勾配がアダプタと入力側へ流れます。INT8値そのものはbufferであるためAdamの勾配を持ちません。

出力状態は

\[ \mathbf{h}_t=\mathrm{LayerNorm}(\mathbf{x}_t+\mathrm{Dropout}(\widetilde{\mathbf{s}}_t)) \]

です。


専門家ルーティングと低ランクアダプタ

入力系列 \(X\in\mathbb{R}^{B\times S\times D}\) の系列平均を

\[ \bar{\mathbf{x}}_b=\frac{1}{S}\sum_{t=1}^{S}\mathbf{x}_{b,t} \]

とし、ルータは

\[ \mathbf{r}_b=W_r\bar{\mathbf{x}}_b+\mathbf{b}_r\in\mathbb{R}^{M} \]

を計算します。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}})\) を使います。

ブロック出力は

\[ Y_b=\sum_{k=1}^{K}g_{b,k}\,f_{I_{b,k}}(X_b) \]

です。選択された専門家だけを実行するため、実行される専門家本体はおおむね \(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\) 出力の接続について、概念的な更新は

\[ \Delta w_{i,j}=\mathrm{round}\left(127\eta_h\,\bar{x}_j\bar{s}_i\right) \]
\[ w_{i,j}\leftarrow\mathrm{clip}_{[-127,127]}(w_{i,j}+\Delta w_{i,j}) \]

です。実装は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

全ブロック・全専門家の主シナプス接続数は

\[ N_{\mathrm{syn}}=2LMDF \]

です。前の係数2は各専門家がinput CSRとrecurrent CSRを1つずつ持つためです。CSR永続保持量は

\[ M_{\mathrm{syn,persistent}}=2LM\left(8(D+1)+9DF\right)\ \mathrm{bytes} \]

です。

学習可能parameter数の主項は

\[ P_{\mathrm{trainable}} =RD+L\left[(DM+M)+M(2DA+2D)\right] \]

です。これは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()で構成した例です。

\[ V=32768,\ D=2048,\ L=12,\ R=128,\ \rho=0.005,\ M=2,\ K=1,\ A=16 \]

このとき \(F=\lceil0.005\times2048\rceil=11\)、主CSR接続数は

\[ N_{\mathrm{syn}}=2\times12\times2\times2048\times11=1,081,344 \]

です。実行確認では次を得ています。

指標
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()を測定してください。


制約、誤解しやすい点、運用指針

制約

  1. PyTorch CSRはbeta: 利用可能なGPU kernel、autograd挙動、性能はPyTorch/CUDA版に依存します。
  2. 局所則は全体誤差の勾配ではない: 主シナプスの更新はCEから直接最適化されません。
  3. Top-kは負荷均衡を持たない: 現実装にはexpert load-balancing loss、capacity factor、expert offloadはありません。
  4. 生成cacheなし: 長い生成ではprompt再評価が重なります。
  5. 全語彙forwardは大きい: 推論で\((B,S,V)\)logitsを明示生成します。
  6. 重みの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"で管理する。

関連文書