コンテンツにスキップ

Dense Chrono-Spike LM 学習とプロンプト検索

この文書は、現行実装の SpikingEvoTextLM / ChronoSpikeAttention / MetaSTDP / AEG / SNN-RAG を、コードと検証結果に基づいて整理したものです。

1. 実装の基本契約

現行の dense LM は、単純な論理回路の集まりではなく次の統合契約を満たす。

  • token → spike train に変換する TASEncoderDecoder
  • causal temporal attention を行う ChronoSpikeAttention
  • EvoLIF / Izhikevich / LIF のいずれかを選択可能な出力層
  • AEGMetaSTDP を optional に有効化可能
  • prompt-time retrieval を補助する SNNRAGHybrid と dense fallback

実運用上の安全設定は以下の通り。

  • EVOSPIKENET_FORCE_HARD_SPIKE_PARITY=true
  • EVOSPIKENET_TRAIN_CONTINUOUS_RELAXATION=false
  • これ以外の mode は明示的に opt-in される場合のみ有効化される

2. ニューロン構造とタイプ

LIF

  • LIFNeuronLayer は integer-based の membrane potential / threshold / leak を持つ
  • 推論時は hard spike を出す
  • 勾配経路は evo_lif_spike_with_ste() により identity 近似を保つ

EvoLIF

  • SpikingEvoTextLMEvoLIF path では LIFNeuronLayer を使い、resolve_evo_lif_scale_factor() によりスケールを制限する
  • EVOSPIKENET_EVO_LIF_SCALE_FACTOR100.0 を上限とする安全 cap を持つ

Izhikevich

  • IzhikevichNeuronLayer は float state v と回復変数 u を持つ
  • しきい値を超えたときに spike し、reset を行う
  • izhikevich_spike_with_ste() は hard spike を維持しつつ、タスク勾配を次層に通す

学習時の補助ロジック

  • spike_activation_with_residual_gradient()spikes + residual_scale * tanh(current) を使う
  • 短時間の spike がゼロに落ちるときでも downstream gradient を消失させない

3. ChronoSpikeAttention の現行数理

ChronoSpikeAttention は causal mask を守りながら過去イベントの重みを減衰させる。

\[ M(t, t') = \exp\left(-\max(0, t - t') / \tau\right) \]

実装では taulearnable_tau=True で学習可能、per_head_tau=True で head ごとに分離可能。time_stepsattention_axis に応じて、Q/K/V のシーケンス軸または time axis を処理する。

出力は output_lif / output_lif.init_leaky() と組み合わせて spike train に変換される。これは SNN と dense 論理を接続している箇所である。

4. SpikingEvoTextLM のデータフロー

SpikingEvoTextLM.forward() は以下の順で処理する。

  1. 入力 token を encoder で embedding / spike train 化する
  2. AEG による importance gating を任意適用する
  3. transformer block で時系列情報を処理する
  4. spiking activation を time aggregate する
  5. output_potential_sum を計算する
  6. mean_spike_activity, output_potential_sum, input_embeddings を結合して readout に渡す
  7. deep_logitsdirect_logits を混合して最終 logits を返す

最終の混合は

\[ \text{logits} = \sigma(\alpha) \cdot \text{direct\_logits} + (1 - \sigma(\alpha)) \cdot \text{deep\_logits} \]

であり、\(\alpha = \text{sigmoid}(\text{readout\_direct\_logit\_mix})\) として学習される。

5. 学習: convergence profile と reward signal

学習用の convergence profile は examples/train_spiking_evospikenet_lm.py_resolve_convergence_profile() にある。

  • stable_baseline が default
  • stable_converge は AEG / Meta-STDP を off にする
  • aggressive_convergence は loss を早く下げるように reward と clipping を調整する
  • explicit env override が profile を上書きする

reward は loss => reward の変換で、raw, clipped_raw, ema_delta を利用できる。

\[ \text{reward} = -\Delta \text{EMA}(\text{loss}) \]

という設計により、単純な loss の符号反転よりもノイズを抑えた adaptation に寄る。

6. データセット生成と長文日本語コーパスの流れ

学習データ処理の現行経路は次の通り。

  1. get_training_corpus(args)
  2. _apply_rag_japanese_preprocess()
  3. _iter_preprocessed_hf_japanese_wikipedia_chunks()
  4. _tokenize_corpus_in_chunks()
  5. _build_next_token_dataset()

核心は、chunk で長い corpus を切り分け、tokenize して shift-by-one データセットを作ることだ。これにより、GPU メモリ全体を使い切ることなく大きな日本語 corpus を扱える。

7. プロンプト時検索と dense fallback

evospikenet/snn_rag.py では prompt-time retrieval を SNNRAGHybrid が担う。

  • query は spike encoding と dense embedding の両方を持つ
  • ChronoSpikeAttention の時間的な相関を使って spike-based scoring を行う
  • KnowledgeGraphIntegrator がオンデバイスな lexical graph を作る
  • メイン検索が低信頼の場合、dense_vector_embeddingvector_similarity の fallback を行う

この fallback は「スパイク検出が弱いときに dense retriever を補助する」設計であり、spike_fallback_threshold によって閾値制御される。実際に、unknown prompt や low-confidence query 時には dense fallback が有効になる。

8. まとめ

現行の実装で最も重要なのは、

  • hard-spike parity を default に保ちつつ
  • 収束に必要な continuous / residual / direct readout を局所的に許容する
  • 収束 profile と reward mode を stable_baseline で安定化しつつ環境変数 override を許す
  • prompt-time retrieval と dense fallback をコア LM 実装と切り離して補助する

という三層構造である。

これは、単純な「生物学的に真っ当な SNN」ではなく、実学習で安定して動く production 向けの spiking LM の設計に近い。