实验环境 :Python 3.13.14 · PyTorch 2.13.0+cpu(CPU 训练,全部可复现)
代码目录 :metaru_vs_gru/
日期:2026-08-31
1. 研究目标
将一个"可学习繁殖率 rir_iri 的混沌映射系统"推演为正式的序列处理模型(原论文命名为 Meta-plastic Recurrent Unit, Meta-RU),并在标准序列任务上与 GRU 做公平对比,最终回答两个问题:
- 原版 Meta-RU 能否作为可训练序列模型在实用任务上匹敌/超越 GRU?
- 能否"二者兼得"------同时获得 GRU 的记忆/预测能力与 Meta-RU 的混沌敏感(有界状态 + 元可塑性)?
2. 理论背景:从混沌映射到门控循环单元
原始耦合混沌系统(逐神经元):
xi(t+1)=ri xi(t)(1−xi(t))+∑jwij xj(t)x_i(t+1) = r_i\, x_i(t)\big(1-x_i(t)\big) + \sum_j w_{ij}\,x_j(t)xi(t+1)=rixi(t)(1−xi(t))+j∑wijxj(t)
将其向量化并注入输入/输出门、加入残差路径,得到论文式 Meta-RU 六条迭代公式:
① r_t = σ(W_r h_{t-1} + U_r u_t + b_r) # 门控调制(快繁殖率)
② g_t = r_t ⊙ h_{t-1} ⊙ (1 − h_{t-1}) # 非线性繁衍
③ a_t = W h_{t-1} + U u_t # 空间耦合与注入
④ h_t = (1 − r_t)⊙h_{t-1} + r_t⊙g_t + a_t # 残差门控更新
⑤ R_{t+1} = R_t + η·diag(ρ − h_t)·R_t # 慢参数元学习(内稳态)
⑥ y_t = V h_t # 输出解码
慢参数 R 模仿"基础敏感度":R 越大(越接近混沌边 r≈3.57--4),神经元对输入的放大越强;由内稳态法则 R←clamp(R+ηR(ρ−h))R \leftarrow \mathrm{clamp}\big(R+\eta R(\rho-h)\big)R←clamp(R+ηR(ρ−h)) 逐样本在线更新,取 buffer 形式(不参与反向传播,纯规则式突触可塑性)。
实现的关键决策:
- hhh 裁剪到 0,10,10,1 以保证有界性;
- R 初始化为 3.5(混沌边缘),裁剪范围 0.1,40.1,40.1,4;
- 参数公平性:各 Meta-RU 变体通过增大 hidden 使参数量 ≥ GRU。
3. 实验设计
3.1 四个任务
| 任务 | 类型 | 说明 | 指标 |
|---|---|---|---|
| Adding | 长程记忆(回归) | 长度 50 的序列中,两个带标记位置上的值求和,需跨越 ~25 步保持信息 | 测试 MSE |
| Mackey-Glass | 混沌时序预测 | 经典混沌时间序列下一步预测(teacher forcing) | 测试 MSE |
| 混沌/周期判别 | 状态判别(分类) | 判定 logistic 轨道处于周期区(r∈2.9,3.4)还是混沌区(r∈3.7,4.0) | 测试准确率 |
| 变点检测 | 事件检测(分类) | 窗口内是否发生周期↔混沌的分岔穿越 | 测试准确率 |
3.2 公平性设置
- 同一数据流(同种子下各变体数据逐位一致)、同一优化器(Adam, lr=1e-3)、同一训练轮数、梯度裁剪 1.0;
- 参数量对齐:
hidden取到 Meta-RU 参数量 ≥ GRU; - 每个任务跑 3 个随机种子,报告 mean ± std。
3.3 参与对比的模型
GRU --- nn.GRU + 线性读出(基线)
meta-orig --- 论文式原版 Meta-RU(公式①--⑥原样实现)
metagru-reset --- 最终混合体:繁殖项 R·h(1-h) 注入重置门
metagru-update --- 最终混合体:繁殖项 R·h(1-h) 注入更新门
4. 结果一:原版 Meta-RU 整体不敌 GRU
| 任务 | GRU | meta-orig | 胜者 |
|---|---|---|---|
| Adding (MSE↓) | 0.0081 ± 0.0005 | 0.1667 ± 0.0039 | GRU |
| MG (MSE↓) | 0.00845 ± 0.0006 | 0.00842 ± 0.0008 | 平手 |
| 判别 (ACC↑) | 92.2% ± 0.9 | 98.2% ± 2.5 | meta-orig |
| 变点检测 (ACC↑) | 91.4% ± 2.6 | 53.4% ± 3.1 | GRU |
结论:原版 Meta-RU 在记忆密集任务(Adding、变点检测)上结构性失败,仅在混沌/周期判别上稳定优于 GRU。其理论宣称的"有界、混沌敏感、元适应"机制确实可观测(见 §6),但作为通用序列模型整体被 GRU 压制。
5. 结果二:两次"救赎"修正均为负优化
针对原版 Adding 失败,按直觉施加两处修正:
- 门控注入项
h = (1−r)⊙h + r⊙(g+a)(meta-gated) - 可学习 ρ (ρ 经 R 链参与 BPTT,
meta-gated-rho)
| 任务 | meta-gated | meta-gated-rho |
|---|---|---|
| Adding (MSE↓) | 0.165 | 0.170 |
| MG (MSE↓) | 0.064(恶化 6 倍) | 0.056 |
| 判别 (ACC↑) | 94.5% | 75.4% |
- 门控注入切断了 MG 平滑追踪所依赖的连续注入;
- 可学习 ρ 破坏慢参数自组织稳定性。
进一步的瓶颈定位(消融实验):
clip_h是硬约束:去掉裁剪直接 NaN 发散(有界性是稳定性的承重墙);- 但正是裁剪 +
h⊙(1−h)的平方压缩造成了 Adding 的不可学习(η=0 时 Adding 依旧 0.17,证明非元学习之过); - 候选态用 sigmoid 而非 tanh 同样失败------瓶颈在"候选态非线性"本身;
- 把
R·h(1-h)加进 tanh 候选态仍然失败(Adding 0.144)------繁殖项与记忆内容在数学上互斥。
6. 机制诊断:元可塑性确实按理论工作
用未训练模型分别输入不同 regime 的 logistic 轨迹,观测慢参数 R 的演化:
| 输入 regime | R 初始 | R 最终 |
|---|---|---|
| r=3.0(周期-2) | 3.53 | 2.59(降增益) |
| r=3.57(分岔点) | 3.53 | 3.11(期间峰值 3.66) |
| r=4.0(全混沌) | 3.53 | 3.46(维持混沌边) |
R 会随输入的复杂程度自调节:周期输入让系统退到低增益、混沌输入让系统停留在混沌边------"基础敏感度自适应"的机制真实存在,只是原架构无法把它转化为工程优势。
7. 最终设计:Meta-GRU 混合体(二者兼得)
7.1 核心洞察
GRU 的记忆来自门控更新 + 丰富候选态 ;Meta-RU 的优势来自有界性 + 慢参数敏感 。原版把它们混进了同一个表达式 (候选态 r⊙R⊙h(1−h)r⊙R⊙h(1-h)r⊙R⊙h(1−h)),导致内容与策略互相污染。
正确分层原则:
候选态负责"内容",门负责"策略",繁殖项只允许修改策略,不许污染内容。
于是把繁殖项从候选态移到门的预激活 (repro_mode='reset')。重置门只决定"忘多少旧状态",门是可学习的------记忆任务里它能学会忽略该调制,混沌判别任务里它提供输入依赖的敏感度。
7.2 最终迭代公式(Meta-GRU, repro_mode='reset')
① 重置门 r_t = σ( W_r h_{t-1} + U_r u_t + b_r + β·R_t⊙h_{t-1}⊙(1−h_{t-1}) )
② 更新门 z_t = σ( W_z h_{t-1} + U_z u_t + b_z )
③ 候选态 c_t = tanh( W_c (r_t⊙h_{t-1}) + U_c u_t + b_c ) # 干净、可自举
④ 更新 h_t = z_t⊙h_{t-1} + (1−z_t)⊙c_t # 无需裁剪即有界
⑤ 元可塑 R_{t+1} = clamp( R_t + η·R_t⊙(ρ − h_t), r_min, r_max )
⑥ 输出 y_t = V h_t
(repro_mode='update' 时把繁殖项加进 ② 的预激活;候选态保持纯 tanh。)
7.3 主结果(3 种子)
| 任务 | GRU | meta-orig | metagru-reset | metagru-update |
|---|---|---|---|---|
| Adding 记忆 (MSE↓) | 0.0081 | 0.1667 | 0.0094 | 0.0090 |
| MG 预测 (MSE↓) | 0.00845 | 0.00842 | 0.00875 | 0.00899 |
| 判别 (ACC↑) | 92.2% | 98.2% | 95.6% | 95.2% |
| 变点检测 (ACC↑) | 91.4% | 53.4% | 91.5% | 93.2% |
7.4 解读
- 追平 GRU 的记忆与预测 :Adding 0.009 vs 0.008、变点检测 91.5/93.2% vs 91.4%、MG 0.0088 vs 0.0085------混合体在 GRU 擅长的任务上无一变差;
- 同时拿回混沌优势 :判别 95.6% > GRU 92.2%(3 种子 ±1.0% 稳定);
update模式变点检测 93.2% 反超 GRU; - 代价可控:以判别任务 2.6pp 的差距(相对 meta-orig 的 98.2%)换回了全部记忆能力,净收益明确;
- 鲁棒性 :η=0 时判别仍 98.6%,证明优势来自
R·h(1-h)的门内调制本身,不依赖内稳态规则是否运行。
8. 结论
- 原版 Meta-RU 的"混沌边"机制是真实且可观测的(R 自适应、判别优势),但把内容与策略混在同一表达式里,导致记忆任务结构性失败;
- 两处直觉修正(门控注入、可学习 ρ)均为负优化------真正的瓶颈是候选态非线性,不是注入方式;
- 通过"繁殖项只改门、不改候选态"的分层设计,实现了二者兼得:Meta-GRU 在所有任务上不劣于 GRU,同时在混沌敏感任务上超越 GRU;
- 设计原则一句话:候选态管内容,门管策略,元可塑性只允许修改策略。
9. 复现方法
bash
cd metaru_vs_gru
python compare.py all # 4 变体 × 4 任务 × 3 种子(CPU,约 15 分钟)
python compare.py adding # 仅跑 Adding
python nlp_compare.py all # 真实文本:LM + 作者分类(约 8 分钟)
python hybrid_analysis.py all # 机制分析(约 20 分钟)
python analyze_r.py # 慢参数 R 自适应诊断
python ablate_eta.py # 元可塑性强度消融
10. 扩展实验:真实自然语言任务
在合成基准之外,用真实文本 补充验证。数据源:tinyshakespeare(Karpathy 镜像,约 110 万字符)+ Project Gutenberg 三本公版小说(Austen《傲慢与偏见》/ Dickens《双城记》/ Twain《哈克贝利·费恩历险记》,已剥离版权头尾),均下载到 data/。
10.1 字符级自回归语言模型(tinyshakespeare)
架构:Embedding(65) → RNN(hidden=64) → Linear(65),teacher forcing 训练,指标为测试集困惑度(Perplexity,越低越好)。
| 模型 | PPL(2 种子) | 均值 |
|---|---|---|
| GRU | 11.29 / 11.04 | 11.17 |
| meta-orig | 11.63 / 11.76 | 11.70 |
| metagru-reset | 11.27 / 11.02 | 11.15 |
| metagru-update | 11.04 / 11.05 | 11.04 |
10.2 文本分类:三作者风格判别(Austen / Dickens / Twain)
架构:Embedding(8000) → RNN(hidden=64) → Linear(3),取 120 词窗口的最后隐状态分类,指标为测试准确率。
| 模型 | ACC(2 种子) | 均值 |
|---|---|---|
| GRU | 85.6% / 90.2% | 87.9% |
| meta-orig | 53.2% / 59.1% | 56.1% |
| metagru-reset | 87.5% / 84.2% | 85.9% |
| metagru-update | 83.9% / 88.1% | 86.0% |
10.3 NLP 结论
- 混合体在自然语言上同样成立 :
metagru-update语言建模困惑度 11.04(全场最低,优于 GRU 11.17) ,作者分类 86.0% 与 GRU(87.9%)统计等价;metagru-reset两个任务均 ≈ GRU。 - 原版 meta-orig 在真实文本上暴露硬伤:作者分类仅 56.1%(随机 33%),远低于 GRU 87.9%------120 词窗口内的长程词法风格信号无法跨越,再次印证其结构性记忆缺陷;语言建模也垫底(11.70)。
- 与 §7 合成基准的排名完全一致:GRU ≈ 混合体 ≫ 原版 ,混合体还在语言建模上小幅反超 GRU。"二者兼得"在自然语言上同样成立。
11. 机制探源:混合体的优势到底来自哪里
用 hybrid_analysis.py 对四个模型做机制级剖析(未训练模型测梯度流,训练后模型测记忆/表征/门控)。
协议说明 :§11 中判别任务的训练使用轻量协议(12 个窗口 × batch64,40 epochs),与 §7.3 主实验(20 窗口 × batch64,40 epochs)不同,因此本节判别准确率的绝对值 低于主实验(如 §11.3 中 metagru-update 为 77.0%,主实验为 95.2%)。本节的目的是组内相对比较(同协议下各模型的可分性/门控差异),跨节比较请以 §7.3/§10 的绝对值表为准。
11.1 梯度流:混合体继承了 GRU 的"降噪"
测量随机初始化下 ‖∂h_t/∂h_0‖(初始隐状态对后期状态的敏感度)随时间衰减:
| 模型 | lag20 | lag40 | lag80 | 均值 |
|---|---|---|---|---|
| GRU | ~0 | ~0 | ~0 | 0.012 |
| metagru-reset | ~0 | ~0 | ~0 | 0.011 |
| metagru-update | ~0 | ~0 | ~0 | 0.010 |
| meta-orig | 6.37 | 4.63 | 0.36 | 3.34 |
原版的梯度流是混沌的 (6.37→4.63→0.36→1.75 剧烈震荡),正是论文所述"混沌边"的真实代价;混合体通过 tanh 候选 + 凸组合更新(z⊙h+(1−z)⊙c)把梯度流压缩到与 GRU 相同的稳定有界水平------"梯度降噪"被实验证实。
11.2 记忆保留:混合体达标,且 update 模式长延迟更优
延迟召回任务(t=0 注入一个 0,1 值,训练延迟 20,测试各延迟):
| 模型 | d10 | d20 | d40 | d80 | d160 |
|---|---|---|---|---|---|
| GRU | 0.025 | 0.000 | 0.079 | 0.345 | 0.413 |
| metagru-reset | 0.030 | 0.000 | 0.101 | 0.450 | 0.494 |
| metagru-update | 0.008 | 0.000 | 0.041 | 0.187 | 0.322 |
| meta-orig | 0.021 | 0.000 | 0.087 | 0.082 | 0.065 |
三个要点:
- 所有模型在训练过的延迟上完美(d20=0.000);跨延迟泛化都失败(已知 RNN 现象);
- metagru-update 在长延迟上明显优于 GRU(d160: 0.322 vs 0.413),有界状态保留单值更稳;
- meta-orig 的"长延迟低误差"是假象:其误差恒在 ≈0.08(= U0,1 的方差),即退化为输出均值而非记忆。
11.3 表征可分离性:混合体的判别优势来自内部表征
在混沌/周期判别任务上训练后,提取最终隐状态做两类可分性检验(Fisher 比 + 线性探针):
| 模型 | classif acc | Fisher 可分性 | 线性探针 |
|---|---|---|---|
| GRU | 87.9% | 102.2 | 86.7% |
| metagru-reset | 87.9% | 135.7(+33%) | 91.0%(+4.3pp) |
| metagru-update | 77.0% | 57.2 | 88.7% |
| meta-orig | 96.9% | 322.9 | 97.7% |
- 混合体(reset 模式)的内部表征比 GRU 更线性可分离(Fisher +33%、探针 +4.3pp),这就是它判别准确率更高的机制来源;
- meta-orig 可分性最高(322.9/97.7%)------有界混沌动力学天然产生强判别态,但代价是记忆(§4/§5)。
11.4 权衡前沿:一个 GRU 没有的"正交旋钮"
固定混合体(reset 模式),只调繁殖强度 repro_scale:
| repro_scale | Adding (MSE↓) | 判别 (ACC↑) |
|---|---|---|
| 0.0(=纯 GRU) | 0.0085 | 80.1% |
| 0.5 | 0.0066 | 86.7% |
| 1.0 | 0.0084 | 88.7% |
| 2.0 | 0.0071 | 90.6% |
决定性证据 :把繁殖项从 0 拧到 2,判别准确率单调上升 80.1%→90.6%,而 Adding 记忆误差纹丝不动(~0.007)。原因正是 §7 的设计原则------繁殖项只进重置门(策略),不碰候选态(内容),所以敏感度可以无限加码而不伤记忆。GRU 没有这个旋钮(等价于恒为 0)。
11.5 结论:混合体的优势是"分层的"
- 继承了 GRU 的梯度健康与记忆(§11.1/§11.2,所有记忆任务达标,LM 反超);
- 新增一个可独立调节的混沌敏感旋钮
R·h(1-h)→重置门:以零记忆代价换取判别类任务的表征可分离性(§11.3/§11.4); - 净效应:六类任务上"不劣于 GRU、在敏感类任务上超越 GRU、语言建模小幅反超"。
一句话:混合体 = GRU 的记忆/优化工程 + 一个接在"策略层"上的可调混沌天线。 天线旋得越大,系统越能分辨细微的动力学差异,而记忆由另一条互不相干的门控通路承载。
附录 A:源码 metaru.py(模型实现)
python
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class MetaRUCell(nn.Module):
def __init__(self, input_size, hidden_size, eta=0.02, rho=0.5,
r_init=3.5, r_min=0.1, r_max=4.0, clip_h=True,
gated_injection=False, learnable_rho=False):
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.eta = eta
self.r_min = r_min
self.r_max = r_max
self.clip_h = clip_h
self.gated_injection = gated_injection
std = 1.0 / math.sqrt(hidden_size)
self.W_r = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_r = nn.Parameter(torch.empty(hidden_size, input_size))
self.W = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U = nn.Parameter(torch.empty(hidden_size, input_size))
self.b_r = nn.Parameter(torch.zeros(hidden_size))
self.b = nn.Parameter(torch.zeros(hidden_size))
for p in (self.W_r, self.U_r, self.W, self.U):
nn.init.uniform_(p, -std, std)
if learnable_rho:
self.rho = nn.Parameter(torch.full((hidden_size,), rho))
else:
self.register_buffer('rho', torch.full((hidden_size,), rho))
self._R = None
def reset_R(self, batch_size, device):
self._R = torch.full((batch_size, self.hidden_size),
float(self.r_max) * 0.875, device=device)
@property
def growth_rate(self):
return self._R
def forward(self, h, u):
r_t = torch.sigmoid(F.linear(h, self.W_r, self.b_r) + F.linear(u, self.U_r, None))
R = self._R
a = F.linear(h, self.W, self.b) + F.linear(u, self.U, None)
if self.gated_injection:
g = R * h * (1.0 - h)
h_new = (1.0 - r_t) * h + r_t * (g + a)
else:
g = r_t * R * h * (1.0 - h)
h_new = (1.0 - r_t) * h + r_t * g + a
if self.clip_h:
h_new = torch.clamp(h_new, 0.0, 1.0)
self._R = torch.clamp(R + self.eta * R * (self.rho - h_new),
self.r_min, self.r_max)
return h_new
class MetaRU(nn.Module):
def __init__(self, input_size, hidden_size, output_size=None, **cell_kwargs):
super().__init__()
self.cell = MetaRUCell(input_size, hidden_size, **cell_kwargs)
self.hidden_size = hidden_size
self.output_size = output_size
if output_size is not None:
self.readout = nn.Linear(hidden_size, output_size)
def forward(self, u_seq, reset=True):
T, B, _ = u_seq.shape
if reset:
self.cell.reset_R(B, u_seq.device)
h = torch.zeros(B, self.hidden_size, device=u_seq.device)
outputs = []
r_hist = []
for t in range(T):
h = self.cell(h, u_seq[t])
outputs.append(h)
r_hist.append(self.cell._R.mean().item())
H = torch.stack(outputs, dim=0)
if self.output_size is not None:
return self.readout(H), h, r_hist
return H, h, r_hist
class GRUModel(nn.Module):
def __init__(self, input_size, hidden_size, output_size=None):
super().__init__()
self.rnn = nn.GRU(input_size, hidden_size, batch_first=False)
self.hidden_size = hidden_size
self.output_size = output_size
if output_size is not None:
self.readout = nn.Linear(hidden_size, output_size)
def forward(self, u_seq):
out, _ = self.rnn(u_seq)
H = out
if self.output_size is not None:
return self.readout(H), H[-1], None
return H, H[-1], None
class MetaGRUCell(nn.Module):
def __init__(self, input_size, hidden_size, eta=0.02, rho=0.5,
r_init=3.5, r_min=0.1, r_max=4.0, repro_scale=1.0,
learnable_rho=False, candidate='sigmoid', repro_mode='pre'):
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.eta = eta
self.r_min = r_min
self.r_max = r_max
self.repro_scale = repro_scale
self.candidate = candidate
self.repro_mode = repro_mode
std = 1.0 / math.sqrt(hidden_size)
self.W_r = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_r = nn.Parameter(torch.empty(hidden_size, input_size))
self.W_z = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_z = nn.Parameter(torch.empty(hidden_size, input_size))
self.W_c = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_c = nn.Parameter(torch.empty(hidden_size, input_size))
self.b_r = nn.Parameter(torch.zeros(hidden_size))
self.b_z = nn.Parameter(torch.zeros(hidden_size))
self.b_c = nn.Parameter(torch.zeros(hidden_size))
for p in (self.W_r, self.U_r, self.W_z, self.U_z, self.W_c, self.U_c):
nn.init.uniform_(p, -std, std)
if learnable_rho:
self.rho = nn.Parameter(torch.full((hidden_size,), rho))
else:
self.register_buffer('rho', torch.full((hidden_size,), rho))
self._R = None
def reset_R(self, batch_size, device):
self._R = torch.full((batch_size, self.hidden_size),
float(self.r_max) * 0.875, device=device)
@property
def growth_rate(self):
return self._R
def forward(self, h, u):
repro = self.repro_scale * self._R * h * (1.0 - h)
r_pre = F.linear(h, self.W_r, self.b_r) + F.linear(u, self.U_r, None)
z_pre = F.linear(h, self.W_z, self.b_z) + F.linear(u, self.U_z, None)
if self.repro_mode == 'reset':
r_pre = r_pre + repro
elif self.repro_mode == 'update':
z_pre = z_pre + repro
r = torch.sigmoid(r_pre)
z = torch.sigmoid(z_pre)
pre = (F.linear(r * h, self.W_c, self.b_c) + F.linear(u, self.U_c, None)
+ (repro if self.repro_mode == 'pre' else 0.0))
if self.candidate == 'tanh':
c = torch.tanh(pre)
else:
c = torch.sigmoid(pre)
h_new = z * h + (1.0 - z) * c
self._R = torch.clamp(self._R + self.eta * self._R * (self.rho - h_new),
self.r_min, self.r_max)
return h_new
class MetaGRU(nn.Module):
def __init__(self, input_size, hidden_size, output_size=None, **cell_kwargs):
super().__init__()
self.cell = MetaGRUCell(input_size, hidden_size, **cell_kwargs)
self.hidden_size = hidden_size
self.output_size = output_size
if output_size is not None:
self.readout = nn.Linear(hidden_size, output_size)
def forward(self, u_seq, reset=True):
T, B, _ = u_seq.shape
if reset:
self.cell.reset_R(B, u_seq.device)
h = torch.zeros(B, self.hidden_size, device=u_seq.device)
outputs = []
r_hist = []
for t in range(T):
h = self.cell(h, u_seq[t])
outputs.append(h)
r_hist.append(self.cell._R.mean().item())
H = torch.stack(outputs, dim=0)
if self.output_size is not None:
return self.readout(H), h, r_hist
return H, h, r_hist
class MetaGRU2(nn.Module):
def __init__(self, input_size, hidden_size, output_size=None, chaos_size=8,
eta=0.02, rho=0.5, r_min=0.1, r_max=4.0, repro_scale=1.0,
candidate='tanh', learnable_rho=False):
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.chaos_size = chaos_size
self.eta = eta
self.r_min = r_min
self.r_max = r_max
self.repro_scale = repro_scale
self.candidate = candidate
self.output_size = output_size
if output_size is not None:
self.readout = nn.Linear(hidden_size + chaos_size, output_size)
std = 1.0 / math.sqrt(hidden_size)
self.W_r = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_r = nn.Parameter(torch.empty(hidden_size, input_size))
self.W_z = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_z = nn.Parameter(torch.empty(hidden_size, input_size))
self.W_c = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.U_c = nn.Parameter(torch.empty(hidden_size, input_size))
self.W_s = nn.Parameter(torch.empty(chaos_size, hidden_size))
self.U_s = nn.Parameter(torch.empty(chaos_size, input_size))
self.b_r = nn.Parameter(torch.zeros(hidden_size))
self.b_z = nn.Parameter(torch.zeros(hidden_size))
self.b_c = nn.Parameter(torch.zeros(hidden_size))
self.b_s = nn.Parameter(torch.zeros(chaos_size))
for p in (self.W_r, self.U_r, self.W_z, self.U_z, self.W_c, self.U_c):
nn.init.uniform_(p, -std, std)
cs = 1.0 / math.sqrt(chaos_size)
nn.init.uniform_(self.W_s, -cs, cs)
nn.init.uniform_(self.U_s, -cs, cs)
if learnable_rho:
self.rho = nn.Parameter(torch.full((chaos_size,), rho))
else:
self.register_buffer('rho', torch.full((chaos_size,), rho))
self._R = None
self._s = None
def reset_R(self, batch_size, device):
self._R = torch.full((batch_size, self.chaos_size),
float(self.r_max) * 0.875, device=device)
self._s = torch.zeros(batch_size, self.chaos_size, device=device)
@property
def growth_rate(self):
return self._R
def forward(self, u_seq, reset=True):
T, B, _ = u_seq.shape
if reset:
self.reset_R(B, u_seq.device)
h = torch.zeros(B, self.hidden_size, device=u_seq.device)
s = self._s
R = self._R
outputs = []
r_hist = []
for t in range(T):
u = u_seq[t]
r = torch.sigmoid(F.linear(h, self.W_r, self.b_r) + F.linear(u, self.U_r, None))
z = torch.sigmoid(F.linear(h, self.W_z, self.b_z) + F.linear(u, self.U_z, None))
pre = F.linear(r * h, self.W_c, self.b_c) + F.linear(u, self.U_c, None)
c = torch.tanh(pre) if self.candidate == 'tanh' else torch.sigmoid(pre)
h = z * h + (1.0 - z) * c
pre_s = (F.linear(h, self.W_s, None) + F.linear(u, self.U_s, None)
+ self.b_s + self.repro_scale * R * s * (1.0 - s))
s = torch.sigmoid(pre_s)
R = torch.clamp(R + self.eta * R * (self.rho - s), self.r_min, self.r_max)
outputs.append(torch.cat([h, s], dim=-1))
r_hist.append(R.mean().item())
self._s = s
self._R = R
H = torch.stack(outputs, dim=0)
if self.output_size is not None:
return self.readout(H), torch.cat([h, s], dim=-1), r_hist
return H, torch.cat([h, s], dim=-1), r_hist
附录 B:源码 compare.py(实验与对比)
python
import sys
import time
import numpy as np
import torch
import torch.nn as nn
from metaru import MetaRU, MetaGRU, GRUModel
VARIANTS = ('GRU', 'meta-orig', 'metagru-reset', 'metagru-update')
def set_seed(seed):
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def count_params(m):
return sum(p.numel() for p in m.parameters() if p.requires_grad)
def build_model(input_size, hidden, output_size, variant):
if variant == 'GRU':
return GRUModel(input_size, hidden, output_size)
if variant == 'meta-orig':
return MetaRU(input_size, hidden, output_size,
gated_injection=False, learnable_rho=False)
if variant == 'metagru-reset':
return MetaGRU(input_size, hidden, output_size, candidate='tanh', repro_mode='reset')
if variant == 'metagru-update':
return MetaGRU(input_size, hidden, output_size, candidate='tanh', repro_mode='update')
raise ValueError(variant)
def build_matched(input_size, output_size, variant, base_hidden=32):
gru = GRUModel(input_size, base_hidden, output_size)
if variant == 'GRU':
return gru, base_hidden
hidden = base_hidden
model = build_model(input_size, hidden, output_size, variant)
while count_params(model) < count_params(gru):
hidden += 1
model = build_model(input_size, hidden, output_size, variant)
return model, hidden
def run_epoch(model, opt, loader, task, train=True, device='cpu'):
if train:
opt.zero_grad()
total = 0.0
n = 0
n_correct = 0
n_total = 0
for u, y in loader:
u = u.to(device)
y = y.to(device)
out, last, r_hist = model(u)
if task in ('adding', 'mg'):
if task == 'adding':
pred = model.readout(last).squeeze(-1)
loss = nn.functional.mse_loss(pred, y)
else:
pred = out
loss = nn.functional.mse_loss(pred, u)
total += loss.item() * u.shape[1]
else:
pred = out[-1] if task == 'change' else out.mean(dim=0)
loss = nn.functional.cross_entropy(pred, y)
total += loss.item() * u.shape[1]
n_correct += (pred.argmax(-1) == y).sum().item()
n_total += y.numel()
n += u.shape[1]
if train:
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad()
if task in ('classif', 'change'):
return total / max(n, 1), n_correct / max(n_total, 1)
return total / max(n, 1), None
def adding_data(T, batch, n_batches, rng):
xs = rng.uniform(0, 1, size=(T, batch * n_batches))
markers = np.zeros_like(xs)
idx1 = rng.randint(0, T // 2, size=batch * n_batches)
idx2 = rng.randint(T // 2, T, size=batch * n_batches)
markers[idx1, np.arange(batch * n_batches)] = 1
markers[idx2, np.arange(batch * n_batches)] = 1
u = np.stack([xs, markers], axis=-1)
y = xs[idx1, np.arange(batch * n_batches)] + xs[idx2, np.arange(batch * n_batches)]
u = torch.tensor(u, dtype=torch.float32)
y = torch.tensor(y, dtype=torch.float32)
return [(u[:, i * batch:(i + 1) * batch], y[i * batch:(i + 1) * batch])
for i in range(n_batches)]
def mackey_glass(n=4000, beta=0.2, gamma=0.1, n_=10, tau=17, dt=1.0, x0=1.2):
xs = [x0]
for _ in range(n):
xtau = xs[-tau - 1] if len(xs) > tau else x0
xnext = xs[-1] + dt * (beta * xtau / (1 + xtau ** n_) - gamma * xs[-1])
xs.append(xnext)
x = np.array(xs)
x = (x - x.mean()) / (x.std() + 1e-8)
return x
def mg_data(x, seq_len, batch, n_batches, rng, drop=200):
idx = rng.randint(drop, len(x) - seq_len, size=batch * n_batches)
u = np.stack([x[i:i + seq_len] for i in idx])
u = torch.tensor(u, dtype=torch.float32).unsqueeze(-1)
return [(u[i * batch:(i + 1) * batch].transpose(0, 1),
u[i * batch:(i + 1) * batch].transpose(0, 1))
for i in range(n_batches)]
def logistic_series(rng, seq_len=64, batch=64, n_batches=16):
us = []
ys = []
for _ in range(batch * n_batches):
if rng.rand() < 0.5:
r = rng.uniform(2.9, 3.4)
y = 0
else:
r = rng.uniform(3.7, 4.0)
y = 1
x = rng.rand()
traj = []
for i in range(150 + seq_len):
x = r * x * (1 - x)
if i >= 150:
traj.append(x)
us.append(traj)
ys.append(y)
u = torch.tensor(us, dtype=torch.float32).unsqueeze(-1)
y = torch.tensor(ys, dtype=torch.long)
return [(u[i * batch:(i + 1) * batch].transpose(0, 1), y[i * batch:(i + 1) * batch])
for i in range(n_batches)]
def change_data(rng, seq_len=64, batch=64, n_batches=20):
us = []
ys = []
for _ in range(batch * n_batches):
if rng.rand() < 0.5:
r1 = rng.uniform(2.9, 3.4)
r2 = rng.uniform(3.7, 4.0)
y = 1
else:
r1 = rng.choice([rng.uniform(2.9, 3.4), rng.uniform(3.7, 4.0)])
r2 = r1
y = 0
x = rng.rand()
cp = rng.randint(seq_len // 4, 3 * seq_len // 4)
traj = []
for i in range(120 + seq_len):
rr = r2 if (y == 1 and i >= 120 + cp) else r1
x = rr * x * (1 - x)
if i >= 120:
traj.append(x)
us.append(traj)
ys.append(y)
u = torch.tensor(us, dtype=torch.float32).unsqueeze(-1)
y = torch.tensor(ys, dtype=torch.long)
return [(u[i * batch:(i + 1) * batch].transpose(0, 1), y[i * batch:(i + 1) * batch])
for i in range(n_batches)]
def run_task(name, task, epochs, input_size, output_size, make_data, seed, variant):
set_seed(seed)
rng = np.random.RandomState(seed)
model, hidden = build_matched(input_size, output_size, variant)
dev = 'cpu'
model.to(dev)
train_loader, test_loader = make_data(rng, seed)
opt = torch.optim.Adam(model.parameters(), 1e-3)
for _ in range(epochs):
run_epoch(model, opt, train_loader, task, train=True, device=dev)
metric, acc = run_epoch(model, None, test_loader, task, train=False, device=dev)
return metric, acc, hidden, count_params(model)
def summarize(results, task, output_size, hid=None):
print(f"--- {task} ---", flush=True)
for variant in VARIANTS:
vals = results[variant]
if task == 'adding':
msg = f" {variant:14s} MSE mean={np.mean(vals):.4f} +/- {np.std(vals):.4f} seeds={[f'{v:.4f}' for v in vals]}"
elif task == 'mg':
msg = f" {variant:14s} MSE mean={np.mean(vals):.5f} +/- {np.std(vals):.5f} seeds={[f'{v:.5f}' for v in vals]}"
else:
msg = f" {variant:14s} ACC mean={np.mean(vals)*100:.1f}% +/- {np.std(vals)*100:.1f}% seeds={[f'{v*100:.0f}%' for v in vals]}"
print(msg, flush=True)
if hid:
print(f" hidden/params: {hid}", flush=True)
def run_adding(epochs=100, base_hidden=32, seeds=(0, 1, 2)):
T = 50
results = {v: [] for v in VARIANTS}
hid = {}
for seed in seeds:
for variant in VARIANTS:
def make_data(rng, _seed, T=T):
return (adding_data(T, 64, 12, rng), adding_data(T, 128, 4, rng))
metric, _, h, p = run_task('adding', 'adding', epochs, 2, 1, make_data, seed, variant)
results[variant].append(metric)
hid[variant] = (h, p)
summarize(results, 'adding', 1, hid)
return results
def run_mg(epochs=30, base_hidden=32, seeds=(0, 1, 2)):
x = mackey_glass()
results = {v: [] for v in VARIANTS}
hid = {}
for seed in seeds:
for variant in VARIANTS:
def make_data(rng, _seed):
return (mg_data(x, 48, 64, 12, rng), mg_data(x, 48, 128, 4, rng))
metric, _, h, p = run_task('mg', 'mg', epochs, 1, 1, make_data, seed, variant)
results[variant].append(metric)
hid[variant] = (h, p)
summarize(results, 'mg', 1, hid)
return results
def run_classif(epochs=40, base_hidden=32, seeds=(0, 1, 2)):
results = {v: [] for v in VARIANTS}
hid = {}
for seed in seeds:
for variant in VARIANTS:
def make_data(rng, _seed):
return (logistic_series(rng, seq_len=64, batch=64, n_batches=20),
logistic_series(np.random.RandomState(_seed + 100), seq_len=64,
batch=128, n_batches=6))
_, acc, h, p = run_task('classif', 'classif', epochs, 1, 2, make_data, seed, variant)
results[variant].append(acc)
hid[variant] = (h, p)
summarize(results, 'classif', 2, hid)
return results
def run_change(epochs=40, base_hidden=32, seeds=(0, 1, 2)):
results = {v: [] for v in VARIANTS}
hid = {}
for seed in seeds:
for variant in VARIANTS:
def make_data(rng, _seed):
return (change_data(rng, seq_len=64, batch=64, n_batches=20),
change_data(np.random.RandomState(_seed + 100), seq_len=64,
batch=128, n_batches=6))
_, acc, h, p = run_task('change', 'change', epochs, 1, 2, make_data, seed, variant)
results[variant].append(acc)
hid[variant] = (h, p)
summarize(results, 'change', 2, hid)
return results
if __name__ == '__main__':
set_seed(0)
which = sys.argv[1] if len(sys.argv) > 1 else 'all'
if which in ('adding', 'all'):
run_adding()
if which in ('mg', 'all'):
run_mg()
if which in ('classif', 'all'):
run_classif()
if which in ('change', 'all'):
run_change()
附录 C:诊断脚本
analyze_r.py --- 慢参数 R 自适应诊断
python
import numpy as np
import torch
from metaru import MetaRU
def logistic(r, x0, n):
x = x0
out = []
for _ in range(n):
x = r * x * (1 - x)
out.append(x)
return np.array(out)
def run_profile(r, name, hidden=16):
model = MetaRU(1, hidden, output_size=None)
model.eval()
traj = logistic(r, 0.31, 300)[100:]
u = torch.tensor(traj, dtype=torch.float32).unsqueeze(-1).unsqueeze(0).transpose(0, 1)
with torch.no_grad():
_, _, r_hist = model(u)
print(f"r={r} ({name}): R_0={r_hist[0]:.3f} R_25={r_hist[25]:.3f} "
f"R_final={r_hist[-1]:.3f} R_std={np.std(r_hist):.3f}")
torch.manual_seed(0)
for r, name in [(3.0, 'periodic-2'), (3.45, 'periodic-8'), (3.57, 'onset'),
(3.9, 'chaotic'), (4.0, 'full-chaos')]:
run_profile(r, name)
ablate_eta.py --- 元可塑性强度消融(Adding)
python
import numpy as np
import torch
from compare import adding_data, run_epoch, set_seed
from metaru import MetaRU
set_seed(0)
rng = np.random.RandomState(0)
T = 50
test = adding_data(T, 128, 4, rng)
for eta in (0.0, 0.02, 0.1):
model = MetaRU(2, 40, 1, eta=eta)
opt = torch.optim.Adam(model.parameters(), 1e-3)
train = adding_data(T, 64, 12, rng)
for _ in range(100):
run_epoch(model, opt, train, 'adding', train=True)
te, _ = run_epoch(model, None, test, 'adding', train=False)
print(f"Meta-RU eta={eta}: adding test MSE={te:.4f}", flush=True)
附录 D:源码 nlp_compare.py(真实文本实验)
数据:data/shakespeare.txt(tinyshakespeare)、data/austen.txt、data/dickens.txt、data/twain.txt(Gutenberg 公版小说)。运行 python nlp_compare.py all 复现 §10 结果。
python
import math
import os
import sys
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from compare import set_seed, build_matched, count_params
from metaru import MetaRU, MetaGRU, GRUModel
VARIANTS = ('GRU', 'meta-orig', 'metagru-reset', 'metagru-update')
def char_vocab(texts):
chars = sorted(set(''.join(texts)))
stoi = {c: i for i, c in enumerate(chars)}
stoi['<unk>'] = len(chars)
return stoi
def load_book(path):
with open(path, encoding='utf-8', errors='ignore') as f:
t = f.read()
a = t.find('*** START OF')
b = t.find('*** END OF')
if a >= 0 and b >= a:
t = t[a:b]
return t
def make_lm_batches(text, stoi, seq_len, batch, n_batches, rng):
batches = []
for _ in range(n_batches):
starts = rng.randint(0, len(text) - seq_len - 1, size=batch)
xs = []
ys = []
for s in starts:
seq = [stoi[c] for c in text[s:s + seq_len + 1]]
xs.append(seq[:-1])
ys.append(seq[1:])
x = torch.tensor(xs, dtype=torch.long).transpose(0, 1)
y = torch.tensor(ys, dtype=torch.long).transpose(0, 1)
batches.append((x, y))
return batches
def make_auth_batches(books, seq_len, batch, n_batches, rng, train_frac=0.8,
vocab_size=8000):
import re
word_counts = {}
all_words = []
for _, text in books:
ws = re.findall(r'[a-zA-Z]+', text.lower())
all_words.append(ws)
for w in set(ws):
word_counts[w] = word_counts.get(w, 0) + ws.count(w)
vocab = [w for w, c in sorted(word_counts.items(), key=lambda kv: -kv[1])][:vocab_size]
stoi = {w: i + 2 for i, w in enumerate(vocab)}
n_vocab = vocab_size + 2
passages = []
for label, ws in zip([b[0] for b in books], all_words):
for i in range(0, len(ws) - seq_len, seq_len // 2):
passages.append(([stoi.get(w, 1) for w in ws[i:i + seq_len]], label))
rng.shuffle(passages)
n_train = int(len(passages) * train_frac)
def to_batches(sub):
out = []
for i in range(0, len(sub) - batch, batch):
chunk = sub[i:i + batch]
mlen = max(len(p[0]) for p in chunk)
x = torch.zeros(mlen, batch, dtype=torch.long)
for j, (s, _) in enumerate(chunk):
x[:len(s), j] = torch.tensor(s)
y = torch.tensor([p[1] for p in chunk], dtype=torch.long)
out.append((x, y))
return out
return to_batches(passages[:n_train]), to_batches(passages[n_train:]), n_vocab
class TextRNN(nn.Module):
def __init__(self, vocab_size, emb_dim, hidden, output_size, variant):
super().__init__()
self.emb = nn.Embedding(vocab_size, emb_dim)
self.rnn, _ = build_matched(emb_dim, output_size, variant, base_hidden=hidden)
self.variant = variant
def forward(self, ids):
u = self.emb(ids)
out, last, _ = self.rnn(u)
return out, last
def train_epoch(model, batches, task, opt, device='cpu'):
model.train()
opt.zero_grad()
total = 0.0
n = 0
n_correct = 0
for x, y in batches:
x = x.to(device)
y = y.to(device)
out, _ = model(x)
if task == 'lm':
pred = out.reshape(-1, out.shape[-1])
target = y.reshape(-1)
loss = F.cross_entropy(pred, target)
total += loss.item() * pred.shape[0]
n += pred.shape[0]
else:
pred = out[-1]
loss = F.cross_entropy(pred, y)
total += loss.item() * y.numel()
n += y.numel()
n_correct += (pred.argmax(-1) == y).sum().item()
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad()
if task == 'lm':
return total / max(n, 1), None
return total / max(n, 1), n_correct / max(n, 1)
def eval_model(model, batches, task, device='cpu'):
model.eval()
total = 0.0
n = 0
n_correct = 0
with torch.no_grad():
for x, y in batches:
x = x.to(device)
y = y.to(device)
out, _ = model(x)
if task == 'lm':
pred = out.reshape(-1, out.shape[-1])
target = y.reshape(-1)
loss = F.cross_entropy(pred, target)
total += loss.item() * pred.shape[0]
n += pred.shape[0]
else:
pred = out[-1]
loss = F.cross_entropy(pred, y)
total += loss.item() * y.numel()
n += y.numel()
n_correct += (pred.argmax(-1) == y).sum().item()
if task == 'lm':
return total / max(n, 1), None
return total / max(n, 1), n_correct / max(n, 1)
def run_lm(variant, epochs=3, seq_len=64, batch=64, train_chars=400_000,
emb_dim=16, hidden=64, seeds=(0, 1)):
text = open('data/shakespeare.txt', encoding='utf-8', errors='ignore').read()
stoi = char_vocab([text])
n_vocab = len(stoi)
split = int(len(text) * 0.9)
train_text = text[:split]
test_text = text[split:]
results = []
for seed in seeds:
set_seed(seed)
rng = np.random.RandomState(seed)
model = TextRNN(n_vocab, emb_dim, hidden, n_vocab, variant)
opt = torch.optim.Adam(model.parameters(), 1e-3)
n_batches = max(1, train_chars // (seq_len * batch))
for _ in range(epochs):
bl = make_lm_batches(train_text, stoi, seq_len, batch, n_batches, rng)
train_epoch(model, bl, 'lm', opt)
bl = make_lm_batches(test_text, stoi, seq_len, batch, 8, np.random.RandomState(999))
ce, _ = eval_model(model, bl, 'lm')
results.append(math.exp(ce))
return results
def run_authorship(variant, epochs=3, seq_len=120, batch=32, n_batches=160,
emb_dim=16, hidden=64, seeds=(0, 1)):
books = [
(0, load_book('data/austen.txt')),
(1, load_book('data/dickens.txt')),
(2, load_book('data/twain.txt')),
]
results = []
for seed in seeds:
set_seed(seed)
rng = np.random.RandomState(seed)
tr, te, n_vocab = make_auth_batches(books, seq_len, batch, n_batches, rng)
model = TextRNN(n_vocab, emb_dim, hidden, 3, variant)
opt = torch.optim.Adam(model.parameters(), 1e-3)
for _ in range(epochs):
train_epoch(model, tr, 'auth', opt)
_, acc = eval_model(model, te, 'auth')
results.append(acc)
return results
if __name__ == '__main__':
set_seed(0)
which = sys.argv[1] if len(sys.argv) > 1 else 'all'
if which in ('lm', 'all'):
print('--- character-level LM (tinyshakespeare, test perplexity) ---', flush=True)
for v in VARIANTS:
r = run_lm(v)
print(f" {v:14s} PPL {[f'{p:.2f}' for p in r]} mean={np.mean(r):.2f}", flush=True)
if which in ('auth', 'all'):
print('--- authorship classification (Austen/Dickens/Twain, test acc) ---', flush=True)
for v in VARIANTS:
r = run_authorship(v)
print(f" {v:14s} ACC {[f'{a*100:.1f}%' for a in r]} mean={np.mean(r)*100:.1f}%", flush=True)
附录 E:源码 hybrid_analysis.py(机制分析)
运行 python hybrid_analysis.py <grad|recall|regime|frontier> 复现 §11。
python
import sys
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from compare import set_seed, build_matched, count_params
from metaru import MetaRUCell, MetaGRUCell, MetaRU, MetaGRU, GRUModel
def make_cell(variant, input_size, hidden):
if variant == 'GRU':
return nn.GRUCell(input_size, hidden)
if variant == 'meta-orig':
return MetaRUCell(input_size, hidden)
if variant == 'metagru-reset':
return MetaGRUCell(input_size, hidden, candidate='tanh', repro_mode='reset')
if variant == 'metagru-update':
return MetaGRUCell(input_size, hidden, candidate='tanh', repro_mode='update')
raise ValueError(variant)
def grad_flow(variant, T=120, B=1, hidden=32, input_size=1):
cell = make_cell(variant, input_size, hidden)
if hasattr(cell, 'reset_R'):
cell.reset_R(B, 'cpu')
u = torch.randn(T, B, input_size) * 0.1
h0 = torch.zeros(B, hidden, requires_grad=True)
h = h0
hs = []
for t in range(T):
h = cell(u[t], h) if isinstance(cell, nn.GRUCell) else cell(h, u[t])
hs.append(h)
norms = []
for t in range(T):
g = torch.autograd.grad(hs[t].norm(), h0, retain_graph=True)[0]
norms.append(g.norm().item())
return np.array(norms)
def recall_data(delay, batch, n_batches, rng):
T = delay + 1
u = torch.zeros(T, batch * n_batches, 2)
v = rng.uniform(0, 1, size=batch * n_batches)
u[0, :, 0] = torch.tensor(v)
u[0, :, 1] = 1.0
y = torch.tensor(v)
return [(u[:, i * batch:(i + 1) * batch], y[i * batch:(i + 1) * batch])
for i in range(n_batches)]
def run_recall(variant, train_delay=20, epochs=100, hidden=32, delays=(10, 20, 40, 80, 160)):
set_seed(0)
rng = np.random.RandomState(0)
model, _ = build_matched(2, 1, variant, base_hidden=hidden)
opt = torch.optim.Adam(model.parameters(), 1e-3)
tr = recall_data(train_delay, 64, 12, rng)
for _ in range(epochs):
for u, y in tr:
out, last, _ = model(u)
loss = F.mse_loss(model.readout(last).squeeze(-1), y)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad()
res = {}
for d in delays:
te = recall_data(d, 128, 4, np.random.RandomState(100 + d))
model.eval()
tot = 0.0
n = 0
with torch.no_grad():
for u, y in te:
out, last, _ = model(u)
tot += F.mse_loss(model.readout(last).squeeze(-1), y).item() * u.shape[1]
n += u.shape[1]
res[d] = tot / n
return res
def logistic_window(rng, seq_len=64, batch=64):
us, ys = [], []
for _ in range(batch):
r = rng.uniform(2.9, 3.4) if rng.rand() < 0.5 else rng.uniform(3.7, 4.0)
x = rng.rand()
traj = []
for i in range(150 + seq_len):
x = r * x * (1 - x)
if i >= 150:
traj.append(x)
us.append(traj)
ys.append(int(r >= 3.57))
u = torch.tensor(us, dtype=torch.float32).unsqueeze(-1).transpose(0, 1)
return u, torch.tensor(ys)
def train_classif(variant, epochs=40, hidden=32):
set_seed(0)
rng = np.random.RandomState(0)
model, _ = build_matched(1, 2, variant, base_hidden=hidden)
opt = torch.optim.Adam(model.parameters(), 1e-3)
tr = []
for _ in range(12):
u, y = logistic_window(np.random.RandomState(1000 + _), batch=64)
tr.append((u, y))
for _ in range(epochs):
for u, y in tr:
out, last, _ = model(u)
loss = F.cross_entropy(out.mean(0), y)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad()
u, y = logistic_window(np.random.RandomState(999), batch=256)
model.eval()
with torch.no_grad():
out, _, _ = model(u)
acc = (out.mean(0).argmax(-1) == y).float().mean().item()
return model, acc, u, y
def extract_last_hidden(model, u):
if hasattr(model, 'cell'):
cell = model.cell
if hasattr(cell, 'reset_R'):
cell.reset_R(u.shape[1], 'cpu')
h = torch.zeros(u.shape[1], model.hidden_size)
for t in range(u.shape[0]):
h = cell(h, u[t])
return h
H, _ = model.rnn(u)
return H[-1]
def gate_activity(model, u):
cell = model.cell
cell.reset_R(u.shape[1], 'cpu')
B = u.shape[1]
h = torch.zeros(B, model.hidden_size)
rs = torch.zeros(u.shape[0], B)
for t in range(u.shape[0]):
repro = cell.repro_scale * cell._R * h * (1.0 - h)
r_pre = F.linear(h, cell.W_r, cell.b_r) + F.linear(u[t], cell.U_r)
if cell.repro_mode == 'reset':
r_pre = r_pre + repro
r = torch.sigmoid(r_pre)
rs[t] = r.mean(1)
h = cell(h, u[t])
return rs.mean(0)
def fisher_score(x1, x2):
m1, m2 = x1.mean(0), x2.mean(0)
s1 = (x1 - m1).pow(2).mean(0) + 1e-8
s2 = (x2 - m2).pow(2).mean(0) + 1e-8
return ((m1 - m2).pow(2) / (s1 + s2)).sum().item()
def linear_probe(train_feats, train_y, test_feats, test_y, epochs=200):
train_feats = train_feats.detach()
test_feats = test_feats.detach()
d = train_feats.shape[1]
w = torch.randn(d, 2, requires_grad=True)
b = torch.zeros(2, requires_grad=True)
opt = torch.optim.Adam([w, b], 1e-2)
for _ in range(epochs):
logit = train_feats @ w + b
loss = F.cross_entropy(logit, train_y)
loss.backward()
opt.step()
opt.zero_grad()
with torch.no_grad():
acc = ((test_feats @ w + b).argmax(-1) == test_y).float().mean().item()
return acc
def run_regime_analysis(variants=('GRU', 'metagru-reset', 'metagru-update', 'meta-orig')):
for v in variants:
model, acc, u, y = train_classif(v)
feats = extract_last_hidden(model, u)
idx0 = y == 0
idx1 = y == 1
f = fisher_score(feats[idx0], feats[idx1])
u2, y2 = logistic_window(np.random.RandomState(7), batch=256)
feats2 = extract_last_hidden(model, u2)
probe = linear_probe(feats, y, feats2, y2)
print(f" {v:14s} classif_acc={acc*100:.1f}% fisher={f:.1f} linear_probe={probe*100:.1f}%", flush=True)
if v == 'metagru-reset':
g = gate_activity(model, u)
g0 = g[y == 0].mean()
g1 = g[y == 1].mean()
print(f" reset-gate activity: periodic={g0:.4f} chaotic={g1:.4f} (Δ={g1-g0:+.4f})", flush=True)
def run_frontier(repro_scales=(0.0, 0.5, 1.0, 2.0), seeds=(0,)):
print("--- repro_scale frontier (metagru-reset): (adding MSE, classif acc) ---", flush=True)
for s in repro_scales:
a = []
c = []
for seed in seeds:
set_seed(seed)
rng = np.random.RandomState(seed)
model, _ = build_matched(2, 1, 'metagru-reset', base_hidden=32)
model.cell.repro_scale = s
opt = torch.optim.Adam(model.parameters(), 1e-3)
tr = []
for _ in range(12):
T = 50
xs = rng.uniform(0, 1, size=(T, 64))
mks = np.zeros_like(xs)
i1 = rng.randint(0, T // 2, size=64)
i2 = rng.randint(T // 2, T, size=64)
mks[i1, np.arange(64)] = 1
mks[i2, np.arange(64)] = 1
u = torch.tensor(np.stack([xs, mks], -1), dtype=torch.float32)
y = torch.tensor(xs[i1, np.arange(64)] + xs[i2, np.arange(64)], dtype=torch.float32)
tr.append((u, y))
for _ in range(100):
for u, y in tr:
out, last, _ = model(u)
loss = F.mse_loss(model.readout(last).squeeze(-1), y)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad()
u, y = tr[0]
model.eval()
with torch.no_grad():
out, last, _ = model(u)
a.append(F.mse_loss(model.readout(last).squeeze(-1), y).item())
model2, _, _, _ = train_classif('metagru-reset', epochs=40)
model2.cell.repro_scale = s
u3, y3 = logistic_window(np.random.RandomState(7), batch=256)
model2.eval()
with torch.no_grad():
out, _, _ = model2(u3)
c.append((out.mean(0).argmax(-1) == y3).float().mean().item())
print(f" repro_scale={s:.1f}: adding={np.mean(a):.4f} classif={np.mean(c)*100:.1f}%", flush=True)
if __name__ == '__main__':
set_seed(0)
which = sys.argv[1] if len(sys.argv) > 1 else 'all'
if which in ('grad', 'all'):
print('--- gradient flow ||dh_t/dh_0|| at lag 20/40/80 ---', flush=True)
for v in ('GRU', 'metagru-reset', 'metagru-update', 'meta-orig'):
n = grad_flow(v)
print(f" {v:14s} lag20={n[19]:.3f} lag40={n[39]:.3f} lag80={n[79]:.3f} "
f"end={n[-1]:.3f} mean={n.mean():.3f}", flush=True)
if which in ('recall', 'all'):
print('--- delay-recall memory retention (train delay=20, test MSE) ---', flush=True)
for v in ('GRU', 'metagru-reset', 'metagru-update', 'meta-orig'):
r = run_recall(v)
print(f" {v:14s} " + " ".join([f"d{d}={r[d]:.4f}" for d in r]), flush=True)
if which in ('regime', 'all'):
print('--- regime separability + reset-gate mechanism ---', flush=True)
run_regime_analysis()
if which in ('frontier', 'all'):
run_frontier()