提出新算法RHEL,用物理对称性突破长序列建模瓶颈
Learning long range dependencies through time reversal symmetry breaking
- 基于哈密顿系统时间反演破缺设计梯度计算方法
- 仅需3次前向传播即可训练超长序列模型(达5万步)
- 性能媲美传统BPTT,适合高能效长时序建模应用
深度状态空间模型(SSMs)重新激活基于物理的计算范式,使循环神经网络可自然嵌入动力系统。本文提出递归哈密顿回声学习(RHEL),该算法严格通过非耗散哈密顿系统的物理轨迹有限差分计算损失梯度。在机器学习层面,RHEL无需显式雅可比矩阵计算,且梯度估计无方差,仅需三次前向传播,与模型规模无关。我们首先在连续时间下引入RHEL,证明其与连续伴随状态法等价;为便于模拟,进一步提出离散版本,当应用于一类称为哈密顿递归单元(HRUs)的循环模块时,等价于随时间反向传播(BPTT)。此设定下,我们通过构建多层HRU层级结构——哈密顿状态空间模型(HSSMs),验证了RHEL的可扩展性。实验中,将RHEL用于训练具有线性和非线性动态的HSSMs,在多种时序任务上表现优异,涵盖中长序列分类与回归,最长序列长度达约50,000。结果表明,RHEL在所有模型和任务中均稳定匹配BPTT性能。本工作为设计可扩展、低能耗且具备自学习能力的物理驱动序列模型开辟了新路径。
原文摘要 · Abstract (English)
Deep State Space Models (SSMs) reignite physics-grounded compute paradigms, as RNNs could natively be embodied into dynamical systems. This calls for dedicated learning algorithms obeying to core physical principles, with efficient techniques to simulate these systems and guide their design. We propose Recurrent Hamiltonian Echo Learning (RHEL), an algorithm which provably computes loss gradients as finite differences of physical trajectories of non-dissipative, Hamiltonian systems. In ML terms, RHEL only requires three "forward passes" irrespective of model size, without explicit Jacobian computation, nor incurring any variance in the gradient estimation. Motivated by the physical realization of our algorithm, we first introduce RHEL in continuous time and demonstrate its formal equivalence with the continuous adjoint state method. To facilitate the simulation of Hamiltonian systems trained by RHEL, we propose a discrete-time version of RHEL which is equivalent to Backpropagation Through Time (BPTT) when applied to a class of recurrent modules which we call Hamiltonian Recurrent Units (HRUs). This setting allows us to demonstrate the scalability of RHEL by generalizing these results to hierarchies of HRUs, which we call Hamiltonian SSMs (HSSMs). We apply RHEL to train HSSMs with linear and nonlinear dynamics on a variety of time-series tasks ranging from mid-range to long-range classification and regression with sequence length reaching $\sim 50k$. We show that RHEL consistently matches the performance of BPTT across all models and tasks. This work opens new doors for the design of scalable, energy-efficient physical systems endowed with self-learning capabilities for sequence modelling.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。