deepseek-v4-pro 在TD3算法上的改进

TD3 官方仓库
HATD3 改进

借助AI学习TD3代码,好奇AI能否在 TD3 上进行改进, 在消耗 2 元后,deepseek-v4-pro 最终改进出 HATD3 算法,
并且自己写代码,设计实验验证,完成了整个闭环,在 RTX 4070 上运行了接近24h后得出实验结果,性能提升了一点点。
整个过程,几乎完全都是AI自己完成的,最令人震惊的是其自我纠正能力,官方 TD3 的代码是8前的仓库,许多 python 库早已迭代更新多次,AI 在终端中反复试错,最终将环境配置成功;在算法改进阶段,经历了几个版本,一开始提出了
不同的 critic 应该用不同的 batch 更新,个人比较认可的,这样能提高 critic 独立性,但最终运行结果发散了,原因是不满足贝尔曼公式的一致性。AI通过反复实验和矫正,将降低性能指标的改进的删除了,最终保留了能提高性能的小改进。
改进 1:Critic 集成 → Trimmed Mean(替代 min)

TD3 HATD3
Critic 数量 2 个(Q1, Q2 共享网络结构) 3 个独立 Critic
TD Target min(Q1, Q2) — 取最小值 Trimmed Mean — 去掉首尾取中位数

为什么更好? min(Q1, Q2) 过于保守:只要有一个 Critic 低估,整个 target 就被压低。Trimmed Mean 去掉最极端的高值和低值,取中位数,更鲁棒

改进 2:目标网络critic每步软更新

TD3 HATD3
更新频率 仅在 Actor 更新步(每 policy_freq=2 步) 每步都更新

以下是 deepseek-v4-pro 自我总结

HATD3 算法设计与实验报告

环境:HalfCheetah-v5 | 10 种子 | 100 万步 | PyTorch + Gymnasium


1. 算法概述

HATD3(Heterogeneous Adaptive TD3)是对 TD3 的改进算法,目标是在不依赖更大网络或更多计算资源的前提下,通过优化 Critic 集成机制来提升学习效率和最终性能。

最终版本(v3)在 HalfCheetah-v5 上取得了 +6.4% 的提升,且训练更稳定(标准差更小)。


2. 基线:TD3 是什么

组件 做法
Critic 2 个 Q 网络(Q1, Q2 共享网络结构)
TD Target min(Q1, Q2) — 取最小值防止过估计
目标网络更新 仅在 Actor 更新步更新(每 2 步一次)
Actor 更新 延迟更新,最大化 Q1(s, π(s))

3. 改进历程:v1 → v2 → v3

3.1 v1(全量改进版,~370 行)

尝试的改进:

改进 思路 结果
解耦采样 (Decoupled Sampling) 3 个 Critic 各自从 replay buffer 抽取不同的 batch ❌ 训练崩溃
多样性正则化 (Diversity Regularization) 惩罚 Critic 间 Q 值过于相似 ❌ 梯度链断裂
自适应噪声 (Adaptive Noise) Q 值方差大时加大探索噪声 ❌ 方向不明确
自适应更新频率 (Adaptive Frequency) Q 值方差大时更频繁更新 Actor ❌ 不稳定

最终性能:~30K 步时 reward = -138(TD3 同期 = 1884),完全发散。


3.2 v1 失败原因详细分析

Bug P0:TD Target 错配(最致命)
# v1 的逻辑:
# 1. 从 batch[0] 计算 TD target
target_Q = reward_0 + gamma * trimmed_mean(critic_targets(s'_0, a'_0))

# 2. 但用 batch[i] 的 (s,a) 去匹配 batch[0] 的 target
for i in range(3):
    loss_i = MSE(critic_i(s_i, a_i), target_Q)  # ← 错配!

后果:相当于用轨迹 A 的回报去训练轨迹 B 的 Critic,Bellman 方程完全失效,reward 暴跌。

Bug P1:多样性正则化的梯度陷阱
# v1 的 diversity loss
q_preds = torch.stack([critic_i(s, a).detach() for critic_i in critics])  # ← detach()
div_loss = -((q_preds - q_preds.mean()) ** 2).mean() + 1e-6
total_loss = critic_loss + lambda_div * div_loss

.detach() 切断了 Critic 输出到 div_loss 的梯度,div_loss 形同虚设。去掉 .detach() 后又引发新的错误:

RuntimeError: Trying to backward through the graph a second time

原因:第一个 Critic 的 backward() 释放了计算图,第二个 Critic 的 backward() 无法重用。加 retain_graph=True 后:

RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation

原因optimizer.step() 修改了参数(inplace),但 retain_graph=True 保留了旧的图,后续 backward 检测到参数版本冲突。

Bug P2:自适应噪声方向相反

v1 中 Q 值方差大(不确定时)反而增大噪声,正确的做法应该是不确定时减小噪声、保守行事。


3.3 v2(修复版,~350 行)

修复内容

  • 修正 TD target 错配问题
  • 尝试用 retain_graph=True 保留多样性正则化
  • 统一所有 Critic 的 backward 为单次

新问题retain_graph + optimizer.step() 的 inplace 冲突无法解决,多样性正则化被迫放弃。

最终性能:30K 步时 reward = -604(TD3 同期 = 1884),仍在发散。


3.4 v3(最终版,~191 行)

策略转变:既然复杂组件逐个失败,那就做减法,只保留核心改进。

最终保留的 3 个改进:

改进 做法 与 TD3 的区别
Critic 集成 + Trimmed Mean 3 个独立 Critic,去掉最高最低取中位数 TD3 用 2 个 Critic + min
目标网络每步更新 每步都软更新(不绑 Actor 步) TD3 仅在 Actor 更新步更新
共享 batch 所有 Critic 用同一批数据 v1 各不相同,v3 回归 TD3 一致

去掉的内容:解耦采样、多样性正则化、自适应噪声、自适应频率。


4. v3 最终架构

                    ┌──────────────────────┐
                    │   Replay Buffer       │
                    │   sample(batch)       │
                    └──────┬───────────────┘
                           │ (s, a, r, s', d)  ← 只采样一次,共享
                           ▼
              ┌────────────────────────┐
              │   TD Target 计算        │
              │   trim_mean(Q1,Q2,Q3)  │  ← 3 个 Target Critic
              │   target = r+γ·Q       │
              └──────┬─────────────────┘
                     │
    ┌────────────────┼────────────────┐
    ▼                ▼                ▼
Critic_0         Critic_1         Critic_2
    │                │                │
    ├──── loss_0 ────┴──── loss_1 ────┘
    │            │
    │     total_loss.backward()  ← 单次 backward
    │            │
    ▼            ▼
  每步 soft update 所有 target networks

关键设计决策:

  • 共享 batch:保证 TD target 与 (s, a) 精确匹配,多样性来自不同初始化而非数据错配
  • 单次 backward:避免多次 backward 引发的图释放和 inplace 冲突
  • Trimmed Mean:比 min 更鲁棒,不会因单个 Critic 的极端估计而严重压低 target
  • 每步更新 target t a u = 0.005 \\tau=0.005 tau=0.005 很小,不会导致目标剧烈变化,但能保证 Critic 始终用最新参考

5. 最终实验对比

实验设置

项目 设置
环境 HalfCheetah-v5
训练步数 1,000,000
种子数 0–9(共 10 个)
评估频率 每 5000 步评估 10 个 episode 取平均

性能对比

在这里插入图片描述

指标 TD3 HATD3 提升
最终(100万步) 10397 ± 755 11060 ± 616 +6.4%
峰值 10493 11193 +6.7%
标准差 755 616 -18% 更稳定
达 1000 所需步数 45,000 45,000 持平
达 3000 所需步数 70,000 70,000 持平
达 5000 所需步数 130,000 140,000 持平
达 8000 所需步数 385,000 345,000 快 4 万步
达 10000 所需步数 795,000 635,000 快 16 万步

关键发现

  1. 早期一致:前 4 万步两者几乎重叠,因为 Critic 正在从随机探索中学习基础 Q 值
  2. 中期拉开:4–40 万步,HATD3 的 Trimmed Mean 开始体现优势,更准确的 TD target 带来的学习效率提升逐渐显现
  3. 后期加速:到达 10000 奖励的时间比 TD3 快了 16 万步,说明集成 Critic 在训练收敛阶段帮助更大

6. 经验教训

教训 说明
不匹配的 TD target = 灾难 Bellman 方程的 reward/next_state 必须与当前 (s,a) 严格对应,任何错配都会导致发散
PyTorch autograd 的坑 多 Critic + 辅助 loss 场景下,backward() + retain_graph + optimizer.step() 三者极易产生 inplace 冲突
detach() 的双刃剑 本想"冻结"Q 值方差来稳定多样性正则化,但同时也切断了关键的梯度流,导致 loss 形同虚设
做减法比做加法更有效 v1 塞入 4 个改进全部失败,v3 只保留 2 个核心改动反而成功。简单的鲁棒改进 > 复杂的脆弱设计
集成多样性不靠数据错配 Critic 间的多样性来自不同随机初始化即可,不需要给它们喂不同数据
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐