Spaces:
Configuration error
Configuration error
Upload README.md with huggingface_hub
Browse files
README.md
CHANGED
|
@@ -1,17 +1,128 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SRCN v2 — Spiking Recurrent Columnar Network for Chinese Language Modeling
|
| 2 |
+
|
| 3 |
+
**396M 参数脉冲神经网络**,在 4×RTX 3090 上的字符级中文语言模型。
|
| 4 |
+
|
| 5 |
+
## 项目概述
|
| 6 |
+
|
| 7 |
+
SRCN(Spiking Recurrent Columnar Network)是一种全脉冲循环神经网络,使用 LIF 神经元模型 + 替代梯度(Fast Sigmoid Surrogate)进行训练。v2 版本针对中文语言建模任务进行了架构优化和训练稳定性改进。
|
| 8 |
+
|
| 9 |
+
### 关键指标 (当前状态)
|
| 10 |
+
|
| 11 |
+
| 指标 | 值 |
|
| 12 |
+
|------|-----|
|
| 13 |
+
| 参数 | 396M |
|
| 14 |
+
| 架构 | 160 列 × 384 神经元/列 × 8 伙伴连接 |
|
| 15 |
+
| Motor 神经元 | 33,408 |
|
| 16 |
+
| 词表 | 8,455 字符 |
|
| 17 |
+
| 批次 | B=128 (有效 B=512) |
|
| 18 |
+
| 吞吐 | ~3,770 tok/s (4×3090) |
|
| 19 |
+
| 显存 | ~20.8 GB 峰值 |
|
| 20 |
+
| 当前 Loss | **4.21** (困惑度 ~68) |
|
| 21 |
+
| 随机基线 | 9.04 (困惑度 ~8,455) |
|
| 22 |
+
|
| 23 |
+
## 架构
|
| 24 |
+
|
| 25 |
+
```
|
| 26 |
+
TemporalPhaseEncoder (80Hz sin相位编码)
|
| 27 |
+
│ I_inj (B, C, M)
|
| 28 |
+
▼
|
| 29 |
+
SRCNLayer (脉冲循环层, 189M参数)
|
| 30 |
+
│ NMDA (α=0.98) + AMPA (α=0.667) 突触
|
| 31 |
+
│ I_nmda / I_ampa / V 均有 clamp 安全阀
|
| 32 |
+
│ FastSigmoidSurrogate 替代梯度
|
| 33 |
+
│ W_raw: 每列 8 个伙伴列,bmm 计算突触电流
|
| 34 |
+
▼
|
| 35 |
+
Motor Readout (33K 脉冲 → LayerNorm → 4096 → ReLU → 8455)
|
| 36 |
+
│ 2层MLP head,无weight decay
|
| 37 |
+
▼
|
| 38 |
+
CrossEntropy Loss
|
| 39 |
+
```
|
| 40 |
+
|
| 41 |
+
### 关键设计决策
|
| 42 |
+
|
| 43 |
+
1. **低 SR (9-12%) 保持稀疏性**:每步仅 ~4,000/33,400 个神经元发放,能效比高
|
| 44 |
+
2. **V_th 自适应平衡**:ε=5e-5,目标 a_target=0.015,max=5.0
|
| 45 |
+
3. **W_raw 受控增长**:低 LR (5e-5) + 高 weight decay (5e-4)
|
| 46 |
+
4. **Encoder 保护梯度**:LR=3e-4,防止 W_raw 通过复发压制编码器信号
|
| 47 |
+
5. **MLP head 无 weight decay**:防止 CrossEntropy 类不平衡 (8455:1) 萎缩权重
|
| 48 |
+
|
| 49 |
+
## 训练稳定性 — 已解决的关键问题
|
| 50 |
+
|
| 51 |
+
### NaN 崩溃(已修)
|
| 52 |
+
- **根因**: NMDA (α=0.98) 50× 放大 + FP16 bmm 溢出
|
| 53 |
+
- **修复**: I_syn clamp ±500, I_nmda clamp 1000, V clamp ±100
|
| 54 |
+
- **辅助**: NaN handler 自动重置状态 + V_th_persist 防污染
|
| 55 |
+
|
| 56 |
+
### W_raw ↔ Encoder 学习失衡(已修)
|
| 57 |
+
- **问题**: W_raw 34× 增长,Encoder 0× 学习 → 复发主导 → 梯度消失
|
| 58 |
+
- **修复**: W_raw LR 5e-5 + wd 5e-4, Encoder LR 3e-4
|
| 59 |
+
|
| 60 |
+
### Motor 输出信息瓶颈(已修)
|
| 61 |
+
- **问题**: 原始 13,440 motor 神经元 × 13% SR = 1,747 bit/step
|
| 62 |
+
- **修复**: motor_ratio 从 22% → 54%,33,408 motor 神经元
|
| 63 |
+
- **配合**: LayerNorm + 2层 MLP head (→4096→8455)
|
| 64 |
+
|
| 65 |
+
## 生成示例 (Loss 4.28)
|
| 66 |
+
|
| 67 |
+
```
|
| 68 |
+
'小明和小红一' → '小明和小红一共有多少个苹果,他想' ← 完美数学题
|
| 69 |
+
'学习计算机要' → '学习计算机要求求出小明手上有10' ← 训练数据模式
|
| 70 |
+
'地球是太阳系' → '地球是太阳系统。 "我们可以得到' ← 差一字
|
| 71 |
+
'今天天气' → '今天天气,但它是在生活中的经' ← 语义连贯
|
| 72 |
+
'这个问题很' → '这个问题很好,大家都有很多的钱' ← 通顺中文
|
| 73 |
+
'我喜欢吃' → '我喜欢吃的东西,我们可以用来' ← 合理续写
|
| 74 |
+
'他昨天去了' → '他昨天去了,我们可以用除法法。' ← 数学推理
|
| 75 |
+
```
|
| 76 |
+
|
| 77 |
+
模型已学会中文句法 + 常识关联 + 数学题模式,远超随机基线。
|
| 78 |
+
|
| 79 |
+
## 文件说明
|
| 80 |
+
|
| 81 |
+
| 文件 | 用途 |
|
| 82 |
+
|------|------|
|
| 83 |
+
| `train_multi.py` | 主训练脚本 (DDP, B=128, 30min checkpoint) |
|
| 84 |
+
| `srcn_model.py` | 模型定义 (Encoder → SRCNLayer → MLP head) |
|
| 85 |
+
| `dataset.py` | 字符级 tokenizer + PackedChineseDataset |
|
| 86 |
+
| `watch.py` | 实时监视器 (Loss/SR/VRAM 趋势) |
|
| 87 |
+
| `run.sh` | 启动脚本 |
|
| 88 |
+
| `audit.py` / `diagnose_grad.py` / `full_diag.py` | 诊断工具 |
|
| 89 |
+
| `annotated_corpus.jsonl` | 训练语料 (需自行准备,~124MB) |
|
| 90 |
+
| `packed_dataset_340m.pkl` | 打包数据集 (需自行生成) |
|
| 91 |
+
| `vocab_tokenizer_v3.pkl` | Tokenizer 词表 |
|
| 92 |
+
|
| 93 |
+
## 运行
|
| 94 |
+
|
| 95 |
+
```bash
|
| 96 |
+
# 准备数据
|
| 97 |
+
python3 dataset.py # 生成 packed_dataset_340m.pkl 和 vocab_tokenizer_v3.pkl
|
| 98 |
+
|
| 99 |
+
# 启动训练 (4 GPU, B=128)
|
| 100 |
+
SRCN_B=128 torchrun --nproc_per_node=4 train_multi.py
|
| 101 |
+
|
| 102 |
+
# 监视
|
| 103 |
+
python3 watch.py
|
| 104 |
+
```
|
| 105 |
+
|
| 106 |
+
```bash
|
| 107 |
+
# 自定义配置
|
| 108 |
+
SRCN_B=96 SRCN_C=160 SRCN_M=384 SRCN_K=8 torchrun --nproc_per_node=4 train_multi.py
|
| 109 |
+
```
|
| 110 |
+
|
| 111 |
+
## 训练历史
|
| 112 |
+
|
| 113 |
+
| 阶段 | Loss | 关键改动 |
|
| 114 |
+
|------|------|---------|
|
| 115 |
+
| v1 初始 | 6.1-6.3 | gain=3→13, V_th=2.0, 梯度断裂修复 |
|
| 116 |
+
| v2 稳定 | 5.6-5.9 | Clamp 安全阀, 移除 empty_cache, DDP bucket 修复 |
|
| 117 |
+
| v3 扩容 | 4.45 | Motor 33K, MLP head, LayerNorm, wd=0 |
|
| 118 |
+
| v4 释放 | **4.21** | a_target 0.015→0.10, V_th 不再压制信息 |
|
| 119 |
+
|
| 120 |
+
## 硬件需求
|
| 121 |
+
|
| 122 |
+
- 4× NVIDIA RTX 3090 (24GB)
|
| 123 |
+
- B=128 需要 ~21GB 显存
|
| 124 |
+
- B=96 约需 ~18GB, B=64 约需 ~14GB
|
| 125 |
+
|
| 126 |
+
## License
|
| 127 |
+
|
| 128 |
+
MIT
|