基于 Neural CDE 的混合连续时间策略(HCT)框架解析:从理论定义到 NDP 实现 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载导读本文以hct/readme/index.md为骨架结合其配套文档 HCT_ReadMe.ipynb 与仓库源码系统讲解 Google Research 开源包hctHybrid Continuous Time Policies混合连续时间策略的核心技术如何把低频图像观测 高频状态观测统一建模为从观测元组到控制函数的因果映射并通过NDPNeural Dynamic Policies架构以参数化二阶 ODE 生成连续控制轨迹。读完本文你将掌握 HCT 的形式化定义、离散化记号M/N/插值时间 τ、NDP 三大模块与 DMP 微分方程、三种计算控制流的方式以及基于增强 ODE 积分损失的模仿学习训练流程并能在 hct 仓库中定位到对应的 Flax/JAX 实现。背景Multiscale Sensor Fusion 与 IROS 2022 论文hct/readme/index.md明确指出本包包含实现Hybrid Continuous Time Policies框架所需的全部代码该框架来自论文Multiscale Sensor Fusion and Continuous Control with Neural CDEs发表于 IROS 2022。仓库的核心动机是解决机器人连续控制中的多尺度传感器融合问题视觉图像传感器通常以较低频率采样而关节编码器、IMU 等本体感受传感器可以以高得多的频率采样。传统做法要么丢弃高频信息要么把两者强制对齐到同一时间网格HCT 则把策略输出本身建模为一个时间连续的控制函数让低频视觉观测与高频状态观测在函数空间中自然融合。仓库布局如下与本文主题一一对应hct/readme/HCT_ReadMe.ipynb完整的记号说明、NDP 数学定义与三种 flow 计算方式的文档型 Colab其 markdown 注释部分即本包的核心设计文档hct/ndp/ndp_model.pyNDP 的 Flax 模块实现编码器、解码器、ODE 积分hct/ndp/ndp_utils.py基于 Flax Optax 的训练工具hct/common共享的 ResNet/MLP 构件、工具函数与类型标注hct/examples/NDP_HCT_Aloha.ipynb在 ALOHA 机器人任务上的使用示例hct/tests/test_ndp.py对三种 flow 计算方式与单步训练的验证测试。HCT 策略的形式化定义设图像观测 $s_t \in \mathbb{R}^{H \times W \times C}$ 在每个离散时刻 $t \in \mathbb{N}$ 到达而在相邻离散时刻 $t$ 与 $t1$ 之间我们可以从其他感知模态获得连续时间实际可放宽为更高频率的状态观测记为函数 $x_t(\cdot): \tau \in [0,T] \rightarrow \mathbb{R}^n$。变量 $\tau$ 称为插值时间interpolation time用于索引两次图像观测之间的连续时间$\tau 0$ 对应离散时刻 $t$$\tau T$ 对应离散时刻 $t1$。HCT 策略是一个从观测元组 $o_t : (s_t, x_t(\cdot))$ 到控制函数 $u_t(\cdot): \tau \in [0,T] \rightarrow U$ 的泛函映射其中 $U$ 是控制空间为记号简便设 $U \mathbb{R}^m$。在 MDP 记号下从 $o_t$ 映射出的动作 $a_t$ 即控制函数 $u_t(\cdot)$与分层策略hierarchical policies形成自然类比。由于策略必须能够利用实时到达的观测生成动作该映射还必须满足因果性causal——生成 $u_t(\tau)$ 时只能使用 $\tau$ 之前含已经观测到的信息。这一动作即函数的视角正是全文的关键架构设计围绕函数空间展开而非离散动作向量。离散时间测量与插值记号函数化表示有利于设计架构但实际系统中观测与动作都以固定频率收发因此文档引入如下记号在区间 $\tau \in [0, T)$ 内输出 $M 1$ 个等间隔动作位于 $\tau_0 0, \tau_1 \frac{T}{M}, \ldots, \tau_{M-1} \frac{(M-1)T}{M}$在 $\tau T$ 处归零恰好与下一张图像 $s_{t1}$ 到达对齐。因此控制频率是图像观测频率的 $M$ 倍。记 $\mathbf{u}t : {u_t(\tau_0), \ldots, u_t(\tau{M-1})}$ 为区间内的全部动作集合。假设信号 $x_t(\cdot)$ 的观测频率是控制频率的 $N \geq 1$ 倍。定义 $\mathbf{x}t^{i}$ 为子区间 $[\tau_{i-1}, \tau_i]$$i 1, \ldots, M-1$内的 $N1$ 个等间隔观测 $$\mathbf{x}t^{i} {x_t(\tau{i-1}),; x_t(\tau{i-1}\tfrac{T}{MN}),; \ldots,; x_t(\tau_i)}$$为便于批处理额外定义 $\mathbf{x}t^0$ 为 $u{t-1}(\tau_{M-1})$ 与 $u_t(\tau_0)$ 之间即上一区间的末尾到本区间起点的 $N1$ 个高频状态观测。下图仓库中的 hct/readme/InFuser.png由 HCT_ReadMe.ipynb 引用直观展示了 $s_t$、$x_t(\cdot)$、$u_t(\cdot)$ 与插值时间 $\tau$ 之间的时序结构关系由于因果性约束架构必须满足以下函数关系$$\begin{eqnarray} (s_t, \mathbf{x}_t^0) \rightarrow u_t(0) \ (s_t, \mathbf{x}_t^{0}, \ldots, \mathbf{x}_t^{j}) \rightarrow u_t(\tau_j), \quad j 1, \ldots, M-1 \end{eqnarray}$$即第一个动作只依赖当前图像与跨区间高频观测后续每个动作可以额外利用到 $\tau_j$ 之前到达的高频观测块。这为下文 NDP 架构的开环生成整段控制函数提供了因果性注脚——只要控制函数的生成不依赖 $\tau$ 之后的观测即可。NDPNeural Dynamic PoliciesHCT 框架下的一个具体实现是NDP它改编自 Neural Dynamic Policies 论文。其关键特点在于只使用观测 $s_t$ 与 $x_t(0)$通过求解一个参数化二阶 ODE开环open-loop生成整个区间 $\tau \in [0, T]$ 的控制函数 $u_t(\cdot)$。注意 NDP 中固定 $T 1$因此 $M$ 个动作输出在 $\tau \in {0, \frac{1}{M}, \ldots, \frac{M-1}{M}}$。DMP动态运动基元NDP 的控制函数由Dynamic Movement PrimitiveDMP定义其核心是一组参数化耦合二阶 ODE$$\begin{eqnarray} \dfrac{d^2 u_t(\tau)}{d \tau^2} : \alpha_u \left(\beta (g_t - u_t(\tau)) - \dfrac{d u_t(\tau)}{d \tau}\right) f_t(\phi_t(\tau)), \quad \tau \in [0, 1) \ \dfrac{d \phi_t(\tau)}{d \tau} : -\alpha_\phi \phi_t(\tau), \quad \tau \in [0, 1) \end{eqnarray}$$其中 $\alpha_u, \alpha_\phi, \beta \in \mathbb{R}$ 为正的常数超参数$\phi_t$ 是初值为 $\phi_t(0) 1$ 的相位函数随时间指数衰减。驱动力函数 $f_t$ 的形式为$$f_t(\phi) \dfrac{\phi}{\sum_{k1}^K \psi_k(\phi)} (W_t \psi(\phi)) \circ (g_t - u_t(0))$$其中 $\psi_k$ 是高斯径向基函数$\psi(\phi) : (\psi_1(\phi), \ldots, \psi_K(\phi))$$K$ 为基函数个数符号 $\circ$ 表示逐元素乘。整个方程组的可学习参数是目标向量$g_t \in \mathbb{R}^m$ 与权重矩阵$W_t \in \mathbb{R}^{m \times K}$。直觉上方程的第一项构成一个弹簧-阻尼系统把 $u_t(\cdot)$ 从初始条件牵引向目标 $g_t$类似 PD 控制项而相位衰减的 $f_t(\phi)$ 在早期提供大的形状调制、随时间衰减从而灵活塑造轨迹形状的同时保证收敛到目标。三大模块NDP 架构由三个模块组成见 hct/readme/HCT_ReadMe.ipynb编码器Encoder把观测 $(s_t, x_t(0))$ 映射为图像嵌入 $z_{s_t}$ 与 DMP 参数 ${g_t, W_t}$——因此得名neural dynamicpolicies解码器Decoder把 $(z_{s_t}, x_t(0))$ 映射为 DMP-ODE 的初始条件 $\left(u_t(0), \frac{d u_t(0)}{d\tau}\right)$DMP-ODE即上文的耦合常微分方程组。常量 ${\alpha_u, \alpha_\phi, \beta, K}$ 作为超参数保留。源码视角Flax 模块与默认超参数NDP 模型在 hct/ndp/ndp_model.py 中实现为 Flax Linen 模块NDP其中编码器与解码器分别对应NDPEncoder与NDPDecoder类。NDPEncoder图像嵌入与 DMP 参数预测class NDPEncoder(nn.Module): Encoder module for NDP. image, hf_obs -- (zs, g, W) action_dim: int zs_dim: int 64 # 图像嵌入维度 zs_width: int 128 # 图像编码网络宽度 num_basis_fncs: int 4 # RBF 基函数个数 K activation nn.relu其前向流程为图像先经 ResNetBatchNorm一个 5 次 stride-2 下采样、带 BatchNorm 的 ResNet 卷积编码器最终输出embed_dim维向量得到 $z_s$将激活后的 $z_s$ 与高频观测hf_obs拼接后分别经goal_mapMLP[2*action_dim, action_dim]预测目标向量 $g_t$经weights_mapMLP[num_weights, num_weights]其中num_weights K * action_dim预测展平后的权重矩阵 $W_t$。NDPDecoder初始条件预测class NDPDecoder(nn.Module): zs, hf_obs -- u(0), u_dot(0) action_dim: int zo_dim: int 32解码器把拼接后的嵌入再经一个三层 MLPfusion_map融合为 $z_o$随后把 $z_o$ 与hf_obs再次拼接经out_map宽度依次为2*out_dim*4, out_dim*2, out_dim其中out_dim 2 * action_dim输出 $2 \times action_dim$ 维的初始条件 $(u(0), \dot{u}(0))$。NDP 主模块ODE 组装与超参数NDP主模块聚合以上两部分并组装 DMP-ODE。其默认超参数可在实例化时覆盖参数默认值含义action_dim必填动作空间维度 $m$num_actions必填两次观测之间的动作数 $M$zs_dim/zs_width64/128图像嵌入维度 / 图像编码器宽度zo_dim32解码器宽度num_basis_fncs4高斯 RBF 基函数个数 $K$alpha_p1.0相位方程衰减系数 $\alpha_\phi$alpha_u10.0二阶 ODE 增益 $\alpha_u$beta2.5弹簧-阻尼比例系数 $\beta$ode_solverdiffrax.Tsit5()数值积分器ode_solver_dt1e-2积分步长adjointdiffrax.RecursiveCheckpointAdjoint()可微求解的反向伴随方法setup中完成的关键组装hct/ndp/ndp_model.pyRBF 中心与带宽rbf_centers exp(-alpha_p * linspace(0, 1, K))rbf_h K / rbf_centers并按 $\psi \exp(-rbf_h \cdot (p - rbf_centers)^2)$ 计算基函数值forcing functionndp_forcer实现 $f_t(\phi) \frac{1}{\sum_k \psi_k} (W_t \psi) \cdot \phi \cdot (g_t - u_0)$与文档公式一致积分网格step_delta 1 / num_actions并断言ode_solver_dt step_deltasample_times arange(num_actions) * step_delta即 $M$ 个动作的输出时刻真值插值interp_control用jnp.interp对真值动作样本做线性插值供训练损失在任意 $\tau$ 处取值。ODE 右侧_ndp_ode标记为nn.nowrap把状态 $(\mathbf{u}, \dot{\mathbf{u}}, p)$ 拆分为三部分返回$\ddot{u} \alpha_u (\beta (g - u) - \dot{u}) f_t(p)$$\dot{p} -\alpha_p \cdot p$_aug_ode则在 NDP-ODE 之外附加代价 ODE用于模仿学习损失见下文。三种方式计算控制流Flow文档HCT_ReadMe.ipynb给出三种求解 NDP、获得 $t$ 与 $t1$ 之间控制动作的方式测试 hct/tests/test_ndp.py 对三者的一致性做了验证断言u_pred与逐步计算的每个动作np.allclose。方式一批量一次算完默认前向model.apply(params, batch_images, batch_hf_obs)params模块参数batch_images图像张量批次$s_t$batch_hf_obs高频观测批次$x_t(0)$。输出中每个样本形状为 $M \times m$。这对应NDP.__call__内部调用compute_ndp_flow(images, hf_obs, self.sample_times)即在sample_times上求值并输出全部 $M$ 个动作。方式二批量求解稠密时间网格model.apply(params, batch_images, batch_hf_obs, pred_times, methodndp_model.compute_ndp_flow)在任意稠密时间向量pred_times要求从 0 开始且最大值 ≤ 1见compute_ndp_flow中的assert jnp.max(pred_times) 1.上求解 $u_t(\cdot)$输出形状为len(pred_times) × m。底层通过diffrax.diffeqsolve结合SaveAt(tspred_times)一次积分完成适合可视化、曲线重建或作为子模块嵌入更大网络。方式三单样本逐步作为策略在线执行# 提取 NDP 模型的逐步函数 re_init, step_fwd model.step_functions # 给定新 (image, hf_obs)计算初始动作与 NDP 参数 ndp_state, ndp_args re_init(params, image, hf_obs) # u(0) ndp_state[:model.action_dim] # 逐步计算 u(tau_1)...u(tau_{M-1}) tau 0. for i in range(1, M): ndp_state, tau step_fwd(params, ndp_state, tau, ndp_args) # u(tau_i) ndp_state[:model.action_dim]其中re_init内部通过utils.unbatch_flax_fn为单样本补批次轴调用encode/decode得到初始状态 $(u(0), \dot u(0), \phi(0)1.0)$ 与 ODE 参数ndp_args (weights, goal, u0)step_fwd则用diffrax在[tau, tau step_delta]上推进一步见 hct/ndp/ndp_model.py 的step_functions属性与_step_ndp。这种方式与每 $\frac{T}{M}$ 周期执行一个动作的真实部署节奏完全一致。基于模仿学习的训练把损失写成积分设 $\hat{\mathbf{u}}_t {\hat{u}_t(0), \ldots, \hat{u}_t(\frac{M-1}{M})}$ 为两次离散时刻 $t$ 与 $t1$ 之间观测到的真值动作序列$\hat{u}t(\cdot)$ 为其线性插值函数NDP 模型在观测条件下生成的控制函数记为 $u{t,\theta}(\cdot)$$\theta$ 为全部可学习参数。模仿损失定义为积分$$I_t(\theta) : \int_{0}^{\frac{M-1}{M}} l(\hat{u}t(\tau), u{t,\theta}(\tau))\ d\tau,$$其中 $l(\cdot, \cdot) \mapsto \mathbb{R}$ 是度量观测动作与预测动作差异的损失函数例如测试 hct/tests/test_ndp.py 中使用的平方误差jnp.sum(jnp.square(u_true - u_pred))。关键洞察该积分等价于如下辅助 ODE 的终值$$\dfrac{d J(\tau)}{d\tau} l(\hat{u}t(\tau), u{t,\theta}(\tau)), \quad \tau \in [0, \tfrac{M-1}{M}], \quad J(0) 0.$$因此计算损失时只需求解一组增强 ODE——同时包含 NDP-ODE 与上述代价 ODE对应源码中的_aug_ode。给定批量观测batch_images、batch_hf_obs与真值动作batch_true_actions每个样本形状 $M \times m$一次前向即同时得到预测动作与逐样本损失batch_pred_actions, batch_losses model.apply( params, batch_images, batch_hf_obs, batch_true_actions, methodndp_model.compute_augmented_flow)batch_losses对 batch 求平均后即为可自动微分的最终损失可无缝接入任意训练管线。训练管线Flax Optax本库使用 Flax 与 Optax 构建训练状态见 hct/ndp/ndp_utils.pycreate_ndp_train_state(model, key, learning_rate, weight_decay, batch_images, batch_hf_obs)用model.init初始化参数构造utils.TrainStateBN在标准TrainState基础上携带batch_stats以支持 BatchNorm并把apply_fn绑定到compute_augmented_flowmake_optax_adam(learning_rate, weight_decay)weight_decay 0时使用optax.adamw否则使用optax.adamoptimize_ndp(state, images, hf_obs, u_true)以jax.pmap按axis_namebatch做数据并行loss_fn内调用带mutable[batch_stats]的 apply 更新 BatchNorm 统计量jax.value_and_grad求梯度pmean跨设备聚合损失/梯度/统计量最后apply_gradients更新参数。数据侧hct/common/utils.py 提供BatchManager随机置换的批管理器next_pmapped_batch按设备数切分批次、split_across_devices、数据集归一化normalize/compute_norm_stats其中图像被强制归一化为零均值、标准差 255以及save_model/restore_model检查点工具。类型别名集中在 hct/common/typing.py如LossFunction Callable[[chex.Array, chex.Array], jnp.float_]。安装与运行仓库以 src-layout 组织安装依赖声明于 hct/setup.pyabsl-py、chex、jax、flax、diffraxODE 求解核心、optax、dm-haiku、numpy、jupyter、jupyter_http_over_ws与matplotlib。仓库根目录提供一键脚本 hct/run.sh其执行流程为# 1) 重组为 src-layoutcommon/、ndp/、__init__.py 移入 src/hct/ mkdir src mkdir src/hct mv common src/hct/ mv ndp src/hct/ mv __init__.py src/hct/ # 2) 创建虚拟环境并安装 python3 -m venv hct-env source hct-env/bin/activate pip install -e . # 3) 运行冒烟测试 python tests/test_ndp.py安装后可直接用 hct/tests/test_ndp.py 验证框架可用性该测试构造一个小型 NDPaction_dim2, num_actions4损失取平方误差依次验证1批量前向输出形状为(2, 4, 2)且与增强 flow 的预测一致2稠密轨迹pred_times输出形状为(2, 50, 2)3step_functions逐步计算的每个动作与批量结果逐点一致4create_ndp_train_stateoptimize_ndp完成一步训练且损失与参数保持有限。端到端示例可参考 hct/examples/NDP_HCT_Aloha.ipynbALOHA 机器人任务。小结hct 包把多尺度传感器融合 连续时间控制落到一个可复现的 JAX/Flax 实现上形式上定义 HCT 策略为满足因果性的泛函映射实践上通过 NDP 的编码器-解码器-DMP-ODE 三段式结构把图像与高频状态统一压缩成二阶 ODE 的参数与初值并用增强 ODE 终值把模仿损失变成可微分的积分。无论是想复现论文实验、在自定义机器人环境上做多频率策略学习还是研究神经 ODE 与策略网络的结合都可以从 hct/ndp/ndp_model.py、hct/readme/HCT_ReadMe.ipynb 与 hct/examples/NDP_HCT_Aloha.ipynb 出发深入。如使用本代码库请引用文首所述的 IROS 2022 论文。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐Backtrader多时间框架混合策略开发实战Backtrader多时间框架混合策略开发实战 概述 在量化交易中多时间框架分析是一种常见且强大的技术分析方法。Backtrader作为一款功能强大的Pyth金融科技数据分析pyalgotrade多时间框架策略从分钟到日线的完整实现想要在量化交易中实现更精准的信号判断和风险控制pyalgotrade多时间框架策略就是你的终极解决方案 这个强大的Python算法交易库让不同时间周期的金融科技gh_mirrors/we/WebServer定时器实现基于小根堆的连接超时管理策略gh_mirrors/we/WebServer定时器实现基于小根堆的连接超时管理策略 连接超时管理的技术痛点与解决方案 在高并发网络编程中TCP连接的超时管后端网络上一篇fre:ac音频转换器完全上手指南6步走完CD抓轨、批量转码与音乐库整理下一篇3步搞定Kubernetes文件系统备份Velero PV数据保护实战指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考