deepseek-v4-pro 在TD3算法上的改进
deepseek-v4-pro 在TD3算法上的改进
借助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 万步 |
关键发现
- 早期一致:前 4 万步两者几乎重叠,因为 Critic 正在从随机探索中学习基础 Q 值
- 中期拉开:4–40 万步,HATD3 的 Trimmed Mean 开始体现优势,更准确的 TD target 带来的学习效率提升逐渐显现
- 后期加速:到达 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 间的多样性来自不同随机初始化即可,不需要给它们喂不同数据 |
更多推荐
所有评论(0)