Live2D

Note -「COLA」临界态驱动的炫酷训练

\[\mathscr{Lorain~wy~Lora~blea.} \newcommand{\DS}[0]{\displaystyle} % operators alias \newcommand{\opn}[1]{\operatorname{#1}} \newcommand{\card}[0]{\opn{card}} \newcommand{\lcm}[0]{\opn{lcm}} \newcommand{\char}[0]{\opn{char}} \newcommand{\Char}[0]{\opn{Char}} \newcommand{\Min}[0]{\opn{Min}} \newcommand{\rank}[0]{\opn{rank}} \newcommand{\Hom}[0]{\opn{Hom}} \newcommand{\End}[0]{\opn{End}} \newcommand{\im}[0]{\opn{im}} \newcommand{\tr}[0]{\opn{tr}} \newcommand{\diag}[0]{\opn{diag}} \newcommand{\coker}[0]{\opn{coker}} \newcommand{\id}[0]{\opn{id}} \newcommand{\sgn}[0]{\opn{sgn}} \newcommand{\Res}[0]{\opn{Res}} \newcommand{\Ad}[0]{\opn{Ad}} \newcommand{\ord}[0]{\opn{ord}} \newcommand{\Stab}[0]{\opn{Stab}} \newcommand{\conjeq}[0]{\sim_{\u{conj}}} \newcommand{\cent}[0]{\u{\degree C}} \newcommand{\Sym}[0]{\opn{Sym}} \newcommand{\Var}[0]{\opn{Var}} \newcommand{\wg}[0]{\wedge} \newcommand{\Wg}[0]{\bigwedge} \newcommand{\sq}[0]{\opn{\square}} % symbols alias \newcommand{\E}[0]{\exist} \newcommand{\A}[0]{\forall} \newcommand{\l}[0]{\left} \newcommand{\r}[0]{\right} \newcommand{\ox}[0]{\otimes} \newcommand{\lra}[0]{\leftrightarrow} \newcommand{\llra}[0]{\longleftrightarrow} \newcommand{\iso}[1]{\overset{\sim}{#1}} \newcommand{\eps}[0]{\varepsilon} \newcommand{\Ra}[0]{\Rightarrow} \newcommand{\Eq}[0]{\Leftrightarrow} \newcommand{\d}[0]{\mathrm{d}} \newcommand{\e}[0]{\mathrm{e}} \newcommand{\i}[0]{\mathrm{i}} \newcommand{\j}[0]{\mathrm{j}} \newcommand{\k}[0]{\mathrm{k}} \newcommand{\Ex}[0]{\mathbb{E}} \newcommand{\D}[0]{\mathbb{D}} \newcommand{\oo}[0]{\infty} \newcommand{\tto}[0]{\rightrightarrows} \newcommand{\mmap}[0]{\hookrightarrow} \newcommand{\emap}[0]{\twoheadrightarrow} \newcommand{\actl}[0]{\curvearrowright} \newcommand{\actr}[0]{\curvearrowleft} \newcommand{\nsubg}[0]{\triangleleft} \newcommand{\nsupg}[0]{\triangleright} \newcommand{\lin}[0]{\lim_{n\to\oo}} \newcommand{\linf}[0]{\liminf_{n\to\oo}} \newcommand{\lsup}[0]{\limsup_{n\to\oo}} \newcommand{\ser}[0]{\sum_{n=1}^\oo} \newcommand{\serz}[0]{\sum_{n=0}^\oo} \newcommand{\isoto}[0]{\overset\sim\to} \newcommand{\F}[0]{\mathbb F} \newcommand{\x}[0]{\times} \newcommand{\M}[0]{\mathbf{M}} \newcommand{\T}[0]{\intercal} \newcommand{\Co}[0]{\complement} \newcommand{\alp}[0]{\alpha} \newcommand{\lmd}[0]{\lambda} \newcommand{\mmid}[0]{\parallel} \newcommand{\loop}[0]{{\circlearrowleft}} \newcommand{\go}[0]{\triangleright} % symbols with parameters \newcommand{\der}[1]{\frac{\d}{\d #1}} \newcommand{\ul}[1]{\underline{#1}} \newcommand{\ol}[1]{\overline{#1}} \newcommand{\wt}[1]{\widetilde{#1}} \newcommand{\br}[1]{\l(#1\r)} \newcommand{\bk}[1]{\l[#1\r]} \newcommand{\ev}[1]{\l.#1\r|} \newcommand{\wh}[1]{\widehat{#1}} \newcommand{\eval}[1]{\l[\!\l[#1\r]\!\r]} \newcommand{\abs}[1]{\l|#1\r|} \newcommand{\bs}[1]{\boldsymbol{#1}} \newcommand{\dat}[1]{\bs{\mathrm{#1}}} \newcommand{\env}[2]{\begin{#1}#2\end{#1}} \newcommand{\ALI}[1]{\env{aligned}{#1}} \newcommand{\CAS}[1]{\env{cases}{#1}} \newcommand{\pmat}[1]{\env{pmatrix}{#1}} \newcommand{\algo}[1]{\begin{array}{r|l}#1\end{array}} \newcommand{\dary}[2]{\l|\begin{array}{#1}#2\end{array}\r|} \newcommand{\pary}[2]{\l(\begin{array}{#1}#2\end{array}\r)} \newcommand{\pblk}[4]{\l(\begin{array}{c|c}{#1}&{#2}\\\hline{#3}&{#4}\end{array}\r)} \newcommand{\u}[1]{\mathrm{#1}} \newcommand{\t}[1]{\text{#1}} \newcommand{\ts}[1]{\textsf{#1}} \newcommand{\tb}[1]{\textbf{#1}} \newcommand{\os}[2]{\overset{#1}{#2}} \newcommand{\lix}[1]{\lim_{x\to #1}} \newcommand{\ops}[1]{#1\cdots #1} \newcommand{\seq}[3]{{#1}_{#2}\ops,{#1}_{#3}} \newcommand{\dedu}[2]{\u{(#1)}\Ra\u{(#2)}} \newcommand{\prv}[3]{\DS{{\DS #1} \over {\DS #2}}~(#3)} \]

  又一个老生常谈的问题是, BPTT 的时空代价太过昂贵了, 我们就算获得了一个完美的类脑甚至仿脑的网络结构, 也难以通过朴素的 BPTT 为它训练真正合适的权重, 进而让它们无法发挥结构优势, 无法战胜传统模型.

  生物领域的实验揭示了自然界, 从神经元到种群行为, 广泛存在的临界态模式. COLA 的观察是, 如果将一个 RNN 处在近临界态, 那么神经元与自身的长程时间依赖和神经元与其他神经元的长程空间依赖就可能能够被神经元自身的状态所近似, 甚至可能被线性近似. 这样, 我们就可以去掉 BPTT 中昂贵的时间步展开过程, 在优化训练时空的同时让训练后模型更具稳定性.

  关于临界态行为, 一个 naive yet educational 的例子: https://www.bilibili.com/video/BV1Ep4y1D7rw.

RNN 的临界态

  以最大 Lyapunov 指数 \(\lmd_\max\approx 0\) 作为临界判准, 研究并尝试让 RNN 处于临界态.

  考虑一个朴素的 RNN 架构, 设 \(\dat u_t\in\R^D\) 为输入, \(\dat x_t\in\R^H\) 为预激活, \(\dat h_t\in\R^H\) 为隐状态, \(\phi\) 为非线性激活函数, 那么 RNN 的动力学行为是:

\[\ALI{ \dat x_t &= \dat W_{hh}\dat h_{t-1}+\dat W_{xh}\dat u_t+\dat b_h,\\ \dat h_t &= \phi(\dat x_t),\\ \dat o_t &= \dat W_{hy}\dat h_t+\dat b_y,\\ \dat p_t &= \t{softmax}(\dat o_t). } \]

\(\ell(\cdot,\cdot)\) 为特定向量损失函数, \(\theta\) 为所有可训练参数, \(\ol w_t\) 是关于时间的损失权重, 则我们可以一般地写出 \(T\) 步迭代的损失为

\[\mathcal L(\theta)=\sum_{t=1}^T\mathcal L_t=\sum_{t=1}^T\ol w_t\ell(\dat o_t,\dat y_t^\star). \]

其中

\[\frac{\part \mathcal L_t}{\part\dat x_t}=\frac{\part\dat h_t}{\part\dat x_t}\odot \dat W_{hy}^\T \frac{\part \mathcal L_t}{\part \dat o_t}. \]

(文中总是把 \(\frac{\part\dat h_t}{\part\dat x_t}\) 视为 \(\R^H\) 的向量, 虽然它应该是一个 \(\R^{H\x H}\) 的对角阵. 我们就按照文中的记号写吧.)

  COLA 首先用 Lyapunov 诊断把初始网络置于临界点的稳定侧 (\(\lmd_\max=0^-\)). 注意这个近临界状态只是一次性的初始化 scaffold, 我们不硬性约束整个训练过程中网络始终保持临界. 以 \(\dat h_t\) 为状态的动力系统的 Jacobi 矩阵为

\[\dat J_t:=\frac{\part\dat h_t}{\part \dat h_{t-1}}=\diag\br{\frac{\part\dat h_t}{\part\dat x_t}}\dat W_{hh}, \]

\(T\) 步离散演化下可估计系统的最大 Lyapunov 指数为

\[\lmd_\max\approx\frac{1}{T}\log\sigma_\max\br{\prod_{t=1}^T\dat J_t}. \]

其中 \(\sigma_\max\) 即表示线性变换的最大奇异值.

  回顾 Lyapunov 指数的相关结论: 沿数据驱动的状态轨迹, \(\lmd_\max<0\) 表示微小扰动平均收缩, \(\lmd_\max>0\) 表示对初值的混沌敏感性, \(\lmd_\max\approx0\) 则界定二者之间的临界边界.

  自然, 我们面临的第一个问题是: 如何把初始网络置于一个近似临界的状态呢? 我们尝试通过对 \(\dat W_{hh}\) 的线性增益来简单地完成这一目标. 首先归一化地设

\[\ol{\dat W}_{hh}:=\frac{\dat W_{hh}}{\|\dat W_{hh}\|_{\t{op}}},\quad \dat W_{hh}(g):=g\ol{\dat W}_{hh}. \]

其中 \(\|\cdot\|_{\t{op}}\) 是算子范数, 在 \(\mathcal L_2\)-范数的 \(\End(\R^H)\) 下也就是最大奇异值. 在稳定边界附近, 省略 QR 轨迹分析的数学过程, 论文经验地使用关于 \(\log g\) 的局部仿射近似:

\[\lmd_\max=\lmd_\max(g)\approx a+b\log g. \]

这样, 我们可以在假设 \(b\approx 1\) 的情况下在 \(g_1=1\) 处单点估计在 \(\lmd_\max(g)=0\) 时有

\[g_{\t{unit}}=\exp(-\lmd_{\max}(1)). \]

或者取 \(g_2>1\) 做两点估计

\[\hat b=\frac{\lmd_\max(g_2)-\lmd_\max(1)}{\log g_2},\quad g_{\t{two}}=\exp\br{-\frac{\lmd_\max(1)}{\hat b}}. \]

其中 \(g_2>1\) 应取为邻近的稳定侧探测点. 文中以两点估计作为默认方法.

临界态假设与 COLA 训练

  利用临界态性质消解隐状态梯度的时间和空间依赖, 优化 BPTT 训练.

  完成一次性的近临界初始化后, COLA 引入两个关于精确 BPTT 指导信号的经验近似. 回忆 BPTT 梯度

\[\dat\delta_t:=\frac{\part\mathcal L}{\dat x_t}=\frac{\part\mathcal L_t}{\part\dat x_t}+\frac{\part\dat h_t}{\part \dat x_t}\odot(\dat W_{hh}^\T\dat\delta_{t+1}). \]

我们正式引入临界态的假设:

  • 长程时间相关性: \(\delta_{t+1}^{(i)}\approx \alp_i\delta_t^{(i)}\), 标量系数 \(\alp_i\) 缓慢变化.
  • 长程空间相关性: \(\delta_t^{(j)}\approx \beta_{ji}\delta_t^{(i)}\), 标量系数 \(\beta_{ji}\) 缓慢变化.

其中第一条更准确地说是短时间窗口上的线性自回归连续性, 第二条表示同一步指导信号主要集中在少数空间模态上. 代回上式就有

\[\dat\delta_t\approx\frac{\part\mathcal L_t}{\part\dat x_t}+\dat A_t\dat\delta_t,\quad \dat A_t:=\diag\br{\frac{\part\dat h_t}{\part\dat x_t}}\dat W_{hh}^\T\diag(\dat\alp). \]

但解这个方程的计算代价比较高. 我们先逐项研究

\[\ALI{ \delta_t^{(i)}&\approx\frac{\part\mathcal L_t}{\part x_t^{(i)}}+\frac{\part h_t^{(i)}}{\part x_t^{(i)}}\sum_{j=1}^H W_{hh}^{(ji)}\alp_j\delta_t^{(j)}\\ &\approx\frac{\part\mathcal L_t}{\part x_t^{(i)}}+\frac{\part h_t^{(i)}}{\part x_t^{(i)}}\sum_{j=1}^H W_{hh}^{(ji)}\alp_j\beta_{ji}\delta_t^{(i)}\\ &=\frac{\part\mathcal L_t}{\part x_t^{(i)}}+\frac{\part h_t^{(i)}}{\part x_t^{(i)}}\cdot \delta_t^{(i)}\cdot \sum_{j=1}^H W_{hh}^{(ji)}\alp_j\beta_{ji}\\ &=:\frac{\part\mathcal L_t}{\part x_t^{(i)}}+\frac{\part h_t^{(i)}}{\part x_t^{(i)}}\cdot \delta_t^{(i)}\cdot \mu_i. } \]

\(\wt\delta_t^{(i)}\) 为上述方程解, 即

\[\wt\delta_t^{(i)}=\frac{\frac{\part \mathcal L_t}{\part x_t^{(i)}}}{1-\mu_i\frac{\part h_t^{(i)}}{\part x_t^{(i)}}}, \]

这给出我们对 \(\delta_t^{(i)}\) 的近似. 整理回向量形式就有

\[\wt{\dat\delta}_t=\frac{\part\mathcal L_t}{\part\dat x_t}\oslash \br{\bs 1-\dat\mu\odot \frac{\part\dat h_t}{\part\dat x_t}}. \]

  接下来的工作便是根据这个代理指导信号在线估计 \(\dat\alp\)\(\dat\mu\). 对于前者, 我们使用无截距线性自回归模型的最小二乘估计:

\[\alp_i^\star=\frac{\Ex\bk{\wt\delta_t^{(i)}\wt\delta_{t-1}^{(i)}}}{\Ex\bk{\br{\wt\delta_{t-1}^{(i)}}^2}} \]

利用其变化缓慢的假设, 在时间步上用 EMA 累计维护

\[\ALI{ S_i(t)&\gets\rho_\alp S_i(t-1)+(1-\rho_\alp)\wt\delta_t^{(i)}\wt\delta_{t-1}^{(i)},\\ Q_i(t)&\gets\rho_\alp Q_i(t-1)+(1-\rho_\alp)\br{\wt\delta_{t-1}^{(i)}}^2,\\ \alp_i(t)&\gets\t{clip}\br{\frac{S_i(t)}{Q_i(t)+\eps_\alp},\alp_\min,\alp_\max}. } \]

其中通常限制 \(\abs{\alp_i}<1\) 以保持 continuation model 稳定.

  对于后者, 单个时间步的闭式指导信号不能独立确定 \(\dat\mu\). 论文因此用

\[\dat{\wt \delta}_t\approx\dat \alp\odot\dat{\wt \delta}_{t-1} \]

作为判准: 如果 \(\dat{\wt\delta}_t\) 要近似真实 BPTT signal, 它也应该尽量复现假设中的短窗口时间结构. 代入近似形式 (这里就写为等号了) 有:

\[\frac{\part\mathcal L_t}{\part\dat x_t}\oslash \br{\bs 1-\dat\mu\odot \frac{\part\dat h_t}{\part\dat x_t}}=\dat\alp\odot \frac{\part\mathcal L_{t-1}}{\part\dat x_{t-1}}\oslash \br{\bs 1-\dat\mu\odot \frac{\part\dat h_{t-1}}{\part\dat x_{t-1}}}. \]

整理:

\[\Ra \frac{\part\mathcal L_t}{\part\dat x_t}-\frac{\part\mathcal L_t}{\part\dat x_t}\odot\frac{\part\dat h_{t-1}}{\part\dat x_{t-1}}\odot \dat\mu=\dat\alp\odot\frac{\part\mathcal L_{t-1}}{\part\dat x_{t-1}}-\dat\alp\odot\frac{\part\mathcal L_{t-1}}{\part\dat x_{t-1}}\odot\frac{\part\dat h_t}{\part\dat x_t}\odot\dat\mu. \]

那么:

\[\Ra \dat\mu\odot\dat A_t=\dat B_t;\\ \dat A_t:=\dat\alp\odot\frac{\part\mathcal L_{t-1}}{\part\dat x_{t-1}}\odot\frac{\part\dat h_t}{\part\dat x_t}-\frac{\part\mathcal L_t}{\part\dat x_t}\odot\frac{\part\dat h_{t-1}}{\part\dat x_{t-1}},\\ \dat B_t:=\dat\alp\odot\frac{\part\mathcal L_{t-1}}{\part\dat x_{t-1}}-\frac{\part\mathcal L_t}{\part\dat x_t}. \]

(注意 \(\dat A_t\)\(\dat B_t\) 是向量而非矩阵.) 对每个单元做 exponentially weighted least squares:

\[\mu_i(t)=\frac{\sum_{\tau\le t}\omega_\tau A_\tau^{(i)}B_\tau^{(i)}}{\sum_{\tau\le t}\omega_\tau\br{A_\tau^{(i)}}^2+\eps_\mu}, \]

并用两个 EMA 分别维护分子中的 \(A_t^{(i)}B_t^{(i)}\) 和分母中的 \(\br{A_t^{(i)}}^2\). 最后对结果做随当前激活导数变化的投影 \(\Pi\), 保证

\[\abs{\mu_i\frac{\part h_t^{(i)}}{\part x_t^{(i)}}}\leq 1-\eps_d, \]

并给分母设置一个小的数值下限.

  最后, hidden 和 recurrent 参数用 \(\wt{\dat\delta}_t\) 代替精确指导信号:

\[\ALI{ \Delta\dat W_{hh}&=-\eta\wt{\dat\delta}_t\dat h_{t-1}^\T,& \Delta\dat W_{xh}&=-\eta\wt{\dat\delta}_t\dat u_t^\T,& \Delta\dat b_h&=-\eta\wt{\dat\delta}_t. } \]

readout 参数只依赖当前时间步, 因而仍使用精确的梯度:

\[\Delta\dat W_{hy}=-\eta\frac{\part\mathcal L_t}{\part\dat o_t}\dat h_t^\T,\qquad \Delta\dat b_y=-\eta\frac{\part\mathcal L_t}{\part\dat o_t}. \]

文中 Algorithm 1 在当前步先用已有的 \(\dat\mu\) 计算 \(\wt{\dat\delta}_t\), 再更新 \(\dat\alp\) 和网络参数, 最后估计供下一时间步使用的新 \(\dat\mu\).

实验结果与讨论

  本节为 AIGC.

  设参数量 \(N_\theta=\Theta(H^2+HD+KH)\). COLA 和 BPTT 的渐近时间复杂度同为 \(O(TN_\theta)\), 但 COLA 不保存长度为 \(T\) 的激活轨迹, 因而 activation memory 为 \(O(1)\), 额外的 \(\dat\alp,\dat\mu\) 及其 EMA 状态只需要 \(O(H)\) 空间. 这不意味着小模型上的 wall-clock 必然快于 BPTT; 它的主要优势是严格在线、内存不随序列长度增长, 并且适合无法保存完整轨迹的流式任务.

  实验结论也需要按任务类型理解:

  • 在 Adding 和 Lorenz rollout 这类依赖长程信用分配与动力学稳定性的任务上, COLA 明显优于各个对照方法.
  • 在 Row-MNIST、Row-CIFAR10 和 UCI HAR 分类任务上, BPTT 或 TBPTT-10 通常仍然更强; COLA 是有竞争力的在线替代, 但并非普遍优于 BPTT.
  • 在字符语言模型上, COLA 接近较强的 TBPTT-1 或 SnAp-1, 但也不是所有任务上的最优方法.
  • ConvRNN 的静态图像 refinement 缺少真正演化的时间证据, 因而 scalar closure 的优势较弱. SNN 扩展依赖具体的 cell-task pairing; LSTM 有多条长期状态链, 论文只把它作为 empirical pilot, 而没有声称理论已经完整覆盖.

  因此, COLA 的核心 trade-off 是用 task-dependent 的梯度近似换取严格在线和常数 activation memory. 当 teaching signal 的短窗口时间延续性较强、空间上集中于少数主模态, 且 recurrence 主要沿一条耦合状态链传播时, scalar closure 最可信; 当这些条件明显不成立且允许完整离线反向传播时, BPTT 仍可能是更合适的选择.


Reference.

  • W. Wang, K. Gao, and G. Chen, “Global Credit Assignment via Dynamical Criticality”in 2026 International Conference on Machine Learning (ICML), 2026.
posted @ 2026-07-10 17:13  Rainybunny  阅读(92)  评论(3)    收藏  举报