Dense Chrono-Spike LM 学習とプロンプト検索
この文書は、現行実装の SpikingEvoTextLM / ChronoSpikeAttention / MetaSTDP / AEG / SNN-RAG を、コードと検証結果に基づいて整理したものです。
1. 実装の基本契約
現行の dense LM は、単純な論理回路の集まりではなく次の統合契約を満たす。
- token → spike train に変換する
TASEncoderDecoder - causal temporal attention を行う
ChronoSpikeAttention EvoLIF/Izhikevich/LIFのいずれかを選択可能な出力層AEGとMetaSTDPを optional に有効化可能- prompt-time retrieval を補助する
SNNRAGHybridと dense fallback
実運用上の安全設定は以下の通り。
EVOSPIKENET_FORCE_HARD_SPIKE_PARITY=trueEVOSPIKENET_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
SpikingEvoTextLMのEvoLIFpath ではLIFNeuronLayerを使い、resolve_evo_lif_scale_factor()によりスケールを制限するEVOSPIKENET_EVO_LIF_SCALE_FACTORは100.0を上限とする安全 cap を持つ
Izhikevich
IzhikevichNeuronLayerは float statevと回復変数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 を守りながら過去イベントの重みを減衰させる。
実装では tau が learnable_tau=True で学習可能、per_head_tau=True で head ごとに分離可能。time_steps と attention_axis に応じて、Q/K/V のシーケンス軸または time axis を処理する。
出力は output_lif / output_lif.init_leaky() と組み合わせて spike train に変換される。これは SNN と dense 論理を接続している箇所である。
4. SpikingEvoTextLM のデータフロー
SpikingEvoTextLM.forward() は以下の順で処理する。
- 入力 token を encoder で embedding / spike train 化する
AEGによる importance gating を任意適用する- transformer block で時系列情報を処理する
- spiking activation を time aggregate する
output_potential_sumを計算するmean_spike_activity,output_potential_sum,input_embeddingsを結合して readout に渡すdeep_logitsとdirect_logitsを混合して最終 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が defaultstable_convergeは AEG / Meta-STDP を off にするaggressive_convergenceは loss を早く下げるように reward と clipping を調整する- explicit env override が profile を上書きする
reward は loss => reward の変換で、raw, clipped_raw, ema_delta を利用できる。
という設計により、単純な loss の符号反転よりもノイズを抑えた adaptation に寄る。
6. データセット生成と長文日本語コーパスの流れ
学習データ処理の現行経路は次の通り。
get_training_corpus(args)_apply_rag_japanese_preprocess()_iter_preprocessed_hf_japanese_wikipedia_chunks()_tokenize_corpus_in_chunks()_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_embeddingとvector_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 の設計に近い。