fix: storage dones/rewards shape squeeze, NaN guard in AdaBoot
This commit is contained in:
@@ -181,8 +181,8 @@ class DreamWaQPPO(PPO):
|
|||||||
"""
|
"""
|
||||||
if not hasattr(self, '_adaboot_reward_buf'):
|
if not hasattr(self, '_adaboot_reward_buf'):
|
||||||
self._adaboot_reward_buf = []
|
self._adaboot_reward_buf = []
|
||||||
dones = storage.dones # (T, E) bool
|
dones = storage.dones.squeeze(-1) # (T, E, 1) → (T, E)
|
||||||
rewards = storage.rewards # (T, E) 每步奖励
|
rewards = storage.rewards.squeeze(-1) # (T, E, 1) → (T, E)
|
||||||
if dones is None or rewards is None:
|
if dones is None or rewards is None:
|
||||||
return
|
return
|
||||||
T, E = dones.shape
|
T, E = dones.shape
|
||||||
@@ -197,7 +197,8 @@ class DreamWaQPPO(PPO):
|
|||||||
if dones[t, e].item():
|
if dones[t, e].item():
|
||||||
# episode 结束:累计从 start 到 t 的奖励
|
# episode 结束:累计从 start 到 t 的奖励
|
||||||
ep_sum = rewards[start:t + 1, e].sum().item()
|
ep_sum = rewards[start:t + 1, e].sum().item()
|
||||||
ep_rewards.append(ep_sum)
|
if abs(ep_sum) < 1e9: # 过滤 NaN/Inf
|
||||||
|
ep_rewards.append(ep_sum)
|
||||||
start = t + 1
|
start = t + 1
|
||||||
|
|
||||||
if len(ep_rewards) < 8:
|
if len(ep_rewards) < 8:
|
||||||
@@ -213,12 +214,14 @@ class DreamWaQPPO(PPO):
|
|||||||
if len(self._adaboot_reward_buf) < 50:
|
if len(self._adaboot_reward_buf) < 50:
|
||||||
return
|
return
|
||||||
|
|
||||||
# 论文公式: CV = σ / μ
|
# 论文公式: CV = σ / μ(NaN 防护)
|
||||||
buf = self._adaboot_reward_buf
|
buf = self._adaboot_reward_buf
|
||||||
mean_r = sum(buf) / len(buf)
|
mean_r = sum(buf) / len(buf)
|
||||||
var_r = sum((r - mean_r) ** 2 for r in buf) / len(buf)
|
var_r = sum((r - mean_r) ** 2 for r in buf) / len(buf)
|
||||||
std_r = var_r ** 0.5
|
std_r = var_r ** 0.5
|
||||||
cv = std_r / (mean_r + 1e-6)
|
cv = std_r / (mean_r + 1e-6)
|
||||||
|
if not (0 <= cv < 1e6): # NaN 或 Inf → 保持当前 prob
|
||||||
|
return
|
||||||
# CV → bootstrap 概率(CV 高→不稳定→多用 GT;映射系数 5.0 可调)
|
# CV → bootstrap 概率(CV 高→不稳定→多用 GT;映射系数 5.0 可调)
|
||||||
self.actor._adaboot_prob = max(0.0, min(1.0, cv * 5.0))
|
self.actor._adaboot_prob = max(0.0, min(1.0, cv * 5.0))
|
||||||
|
|
||||||
|
|||||||
@@ -105,14 +105,14 @@ class rslrl:
|
|||||||
class Go1DreamWaQWalkRslrlPpo(RslrlCfg):
|
class Go1DreamWaQWalkRslrlPpo(RslrlCfg):
|
||||||
"""Go1 DreamWaQ walk — CENet VAE + 不对称特权观测。"""
|
"""Go1 DreamWaQ walk — CENet VAE + 不对称特权观测。"""
|
||||||
|
|
||||||
num_envs: int = 2048
|
num_envs: int = 4096 # 论文原值
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
runner = self.runner
|
runner = self.runner
|
||||||
|
|
||||||
# Runner 设置(严格对齐上游 LeggedRobotCfgPPO + Go1RoughCfgPPO)
|
# Runner 设置(严格对齐上游 LeggedRobotCfgPPO + Go1RoughCfgPPO)
|
||||||
runner.seed = 5 # 上游 seed=5
|
runner.seed = 5 # 上游 seed=5
|
||||||
runner.max_iterations = 5000 # 2048 envs 需要更多迭代补偿采样量
|
runner.max_iterations = 1000 # 论文原值
|
||||||
runner.num_steps_per_env = 24
|
runner.num_steps_per_env = 24
|
||||||
runner.experiment_name = "go1_dreamwaq_walk"
|
runner.experiment_name = "go1_dreamwaq_walk"
|
||||||
runner.save_interval = 50
|
runner.save_interval = 50
|
||||||
|
|||||||
Reference in New Issue
Block a user