元可塑性循环单元(Meta-RU)与 GRU 的对比研究与混合设计

实验环境 :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 做公平对比,最终回答两个问题:

  1. 原版 Meta-RU 能否作为可训练序列模型在实用任务上匹敌/超越 GRU?
  2. 能否"二者兼得"------同时获得 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 失败,按直觉施加两处修正:

  1. 门控注入项 h = (1−r)⊙h + r⊙(g+a)meta-gated
  2. 可学习 ρ (ρ 经 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. 结论

  1. 原版 Meta-RU 的"混沌边"机制是真实且可观测的(R 自适应、判别优势),但把内容与策略混在同一表达式里,导致记忆任务结构性失败;
  2. 两处直觉修正(门控注入、可学习 ρ)均为负优化------真正的瓶颈是候选态非线性,不是注入方式;
  3. 通过"繁殖项只改门、不改候选态"的分层设计,实现了二者兼得:Meta-GRU 在所有任务上不劣于 GRU,同时在混沌敏感任务上超越 GRU;
  4. 设计原则一句话:候选态管内容,门管策略,元可塑性只允许修改策略

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 结论

  1. 混合体在自然语言上同样成立metagru-update 语言建模困惑度 11.04(全场最低,优于 GRU 11.17) ,作者分类 86.0% 与 GRU(87.9%)统计等价;metagru-reset 两个任务均 ≈ GRU。
  2. 原版 meta-orig 在真实文本上暴露硬伤:作者分类仅 56.1%(随机 33%),远低于 GRU 87.9%------120 词窗口内的长程词法风格信号无法跨越,再次印证其结构性记忆缺陷;语言建模也垫底(11.70)。
  3. 与 §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 结论:混合体的优势是"分层的"

  1. 继承了 GRU 的梯度健康与记忆(§11.1/§11.2,所有记忆任务达标,LM 反超);
  2. 新增一个可独立调节的混沌敏感旋钮 R·h(1-h)→重置门:以零记忆代价换取判别类任务的表征可分离性(§11.3/§11.4);
  3. 净效应:六类任务上"不劣于 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.txtdata/dickens.txtdata/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()
相关推荐
Elastic 中国社区官方博客1 小时前
在 Elasticsearch 中回填时间序列数据:通过批量 API 加载数月的历史指标数据
大数据·运维·数据库·人工智能·elasticsearch·搜索引擎·全文检索
LlmCraft|大模型工程实践1 小时前
08 预训练语言模型:BERT 与 GPT
人工智能·深度学习·nlp
feasibility.1 小时前
ABot-World-0:当一块 RTX 5090 能编织无限世界——交互式世界模型的高德解法
人工智能·aigc·视频·地图·具身智能·英伟达·世界模型
leoZ2311 小时前
第 6 篇:SchemaForm 渲染器核心实现
前端·javascript·vue.js·人工智能·神经网络·机器学习·自然语言处理
故七月1 小时前
优胜劣汰·动态赋能——锦邻创享OPC社区的考核管理与退出机制
大数据·人工智能
枫叶林FYL1 小时前
【群体智能集群控制工程实践】第10章 无人舰队核心功能实现
大数据·人工智能·算法
昵称画2 小时前
POC验证怎么设计用例?不走过程的实操要点
大数据·数据库·人工智能·低代码·excel
2601_967212722 小时前
新一代电源轨道系统技术甄别维度与行业技术路线分析
大数据·网络·人工智能
Lucas_coding2 小时前
【Codex Remote】 Codex App通过SSH连接远程Linux开发环境
人工智能