Skip to content

Commit 48b72b2

Browse files
tAnGjIa520claude
andcommitted
chore: update .gitignore and refine PPO training configurations
- Add Claude Code directories to .gitignore (.codex, codex_home/, .omc/, .claude/, scripts/, docs/) - Update CLAUDE.md with environment setup and PPO implementation details - Refine split training mode in train_unizero.py for better World Model/PPO separation - Enhance buffer tracking with latest_push_count and new_data_ratio in game_buffer.py - Improve PPO batch extraction and padding logic in game_buffer_unizero.py - Add loss_type parameter support in world_model.py for flexible training modes - Update PPO configurations with split training options and improved hyperparameters Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1 parent 2ce9597 commit 48b72b2

9 files changed

Lines changed: 361 additions & 66 deletions

File tree

.gitignore

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1450,4 +1450,12 @@ events.*
14501450
!/assets/pooltool/**
14511451
lzero/mcts/ctree/ctree_alphazero/pybind11
14521452

1453-
zoo/jericho/envs/z-machine-games-master
1453+
zoo/jericho/envs/z-machine-games-master
1454+
1455+
# Claude Code specific
1456+
.codex
1457+
codex_home/
1458+
.omc/
1459+
.claude/
1460+
scripts/
1461+
docs/

CLAUDE.md

Lines changed: 14 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -75,7 +75,10 @@ Replace MCTS planning with PPO policy gradient training while keeping UniZero's
7575

7676
**Policy:**
7777
- `lzero/policy/unizero.py` - `_forward_learn` unpacks PPO data (advantage, old_log_prob, return); `_forward_collect` supports pure policy mode (skips MCTS)
78-
- `lzero/policy/utils.py` - Added `ppo_error`, `ppo_policy_error`, `ppo_value_error`
78+
- `lzero/policy/utils.py` - Added `ppo_error`, `ppo_policy_error`, `ppo_value_error` (lines 716-865):
79+
- `ppo_policy_error`: Clipped surrogate loss with optional dual_clip, entropy bonus, KL/clipfrac monitoring
80+
- `ppo_value_error`: MSE loss with optional value clipping
81+
- `ppo_error`: Combined policy + value + entropy loss computation
7982

8083
**World Model:**
8184
- `lzero/model/unizero_world_models/world_model.py` - New `compute_loss_ppo()` method:
@@ -129,13 +132,21 @@ WorldModel.compute_loss_ppo()
129132
### Important Notes
130133

131134
- **old_log_prob storage:** Collector stores raw logits, not log probabilities
132-
- **GAE computation:** Happens in collector after episode completion via `_batch_compute_gae_for_pool()`
135+
- **GAE computation:** Happens in collector after episode completion via `_batch_compute_gae_for_pool()` using DI-engine's GAE implementation
136+
- **Return computation order:** Returns are computed BEFORE advantage normalization (fixed in commit b744cdac)
133137
- **Value normalization:** Supported in collector for stable training
134138
- **Batch structure:** current_batch has 11 elements (last 3 are advantage, old_log_prob, return for PPO)
139+
- **PPO loss components:** Policy (clipped surrogate) + Value (MSE with optional clipping) + Entropy bonus
135140

136141
## Known Issues
137142

138-
**Current debugging:** LunarLander PPO not converging on `lunarlander_disc_unizero_ppo_online_config`. Need to investigate PPO-related bugs.
143+
**Recent fixes (commits 2ce95971, b744cdac, d1415022):**
144+
- Fixed PPO value loss computation in world model
145+
- Fixed return computation order (must compute returns BEFORE advantage normalization)
146+
- Integrated DI-engine GAE functions in muzero_collector
147+
- Enhanced buffer/collector/world model for proper PPO data flow
148+
149+
**Current status:** LunarLander PPO implementation complete with proper GAE/PPO loss computation.
139150

140151
**Backup:** `/mnt/shared-storage-user/tangjia/unizero_ppo/LightZero-bak/`
141152

lzero/entry/train_unizero.py

Lines changed: 89 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -198,56 +198,104 @@ def train_unizero(
198198
data_sufficient = replay_buffer.get_num_of_game_segments() > batch_size
199199
else:
200200
data_sufficient = replay_buffer.get_num_of_transitions() > batch_size
201-
201+
202202
if not data_sufficient:
203203
logging.warning(
204204
f'Rank {rank}: The data in replay_buffer is not sufficient to sample a mini-batch: '
205205
f'batch_size: {batch_size}, replay_buffer: {replay_buffer}. Continue to collect now ....'
206206
)
207207
continue
208208

209-
# Execute multiple training rounds
210-
for i in range(update_per_collect):
211-
# ✅ 处理混合采样返回的元组 (train_data_new, train_data_old)
212-
sample_result = replay_buffer.sample(batch_size, policy)
213-
214-
if isinstance(sample_result, tuple):
215-
train_data_new, train_data_old = sample_result
216-
# 如果有新旧数据分离,先训练新数据
217-
if train_data_old is not None:
218-
# 训练新数据
219-
train_data_new.append(learner.train_iter)
220-
train_data_new.append(True) # is_old_data = False
221-
log_vars_new = learner.train(train_data_new, collector.envstep)
222-
223-
# 训练老数据
224-
train_data_old.append(learner.train_iter)
225-
train_data_old.append(False) # is_old_data = True
226-
log_vars_old = learner.train(train_data_old, collector.envstep)
227-
228-
# 更新优先级(如果使用)
229-
if cfg.policy.use_priority:
230-
replay_buffer.update_priority(train_data_new, log_vars_new[0]['value_priority_orig'])
231-
replay_buffer.update_priority(train_data_old, log_vars_old[0]['value_priority_orig'])
232-
else:
233-
# Fallback: 只有新数据
234-
train_data_new.append(learner.train_iter)
235-
log_vars = learner.train(train_data_new, collector.envstep)
236-
if cfg.policy.use_priority:
237-
replay_buffer.update_priority(train_data_new, log_vars[0]['value_priority_orig'])
238-
else:
239-
# 向后兼容:如果返回的不是元组,使用原始逻辑
240-
train_data = sample_result
209+
# ========== Split Training Mode (PriorZero-style) ==========
210+
if cfg.policy.get('split_ppo_wm_training', False):
211+
# Phase 1: World Model Training (use all data)
212+
wm_update_per_collect = cfg.policy.get('wm_update_per_collect', update_per_collect)
213+
logging.info(f"[Rank {rank}] [WM Training] Updates: {wm_update_per_collect}")
214+
215+
for i in range(wm_update_per_collect):
216+
train_data = replay_buffer.sample(batch_size, policy)
217+
if isinstance(train_data, tuple):
218+
train_data = train_data[0] # Use first element if tuple
241219
train_data.append(learner.train_iter)
242-
log_vars = learner.train(train_data, collector.envstep)
220+
train_data.append('world_model_only') # ✅ loss_type
221+
222+
log_vars_wm = learner.train(train_data, collector.envstep)
223+
243224
if cfg.policy.use_priority:
244-
replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig'])
245-
246-
if replay_buffer._cfg.reanalyze_ratio > 0 and i % 20 == 0:
247-
policy.recompute_pos_emb_diff_and_clear_cache()
248-
249-
if cfg.policy.use_wandb:
250-
policy.set_train_iter_env_step(learner.train_iter, collector.envstep)
225+
replay_buffer.update_priority(train_data, log_vars_wm[0]['value_priority_orig'])
226+
227+
# Mark position after WM training
228+
replay_buffer.mark_latest_transitions_consumed()
229+
logging.info(f"[Rank {rank}] [WM Training] Completed. Marked position: {replay_buffer.last_pos_in_transition}")
230+
231+
# Phase 2: PPO Training (use only new data)
232+
ppo_batch_size = cfg.policy.get('ppo_batch_size', batch_size)
233+
ppo_train_data = replay_buffer.fetch_latest_batch(ppo_batch_size, policy)
234+
235+
if ppo_train_data is not None:
236+
ppo_update_per_collect = cfg.policy.get('ppo_update_per_collect', update_per_collect)
237+
logging.info(f"[Rank {rank}] [PPO Training] Updates: {ppo_update_per_collect}")
238+
239+
for i in range(ppo_update_per_collect):
240+
ppo_train_data_copy = [ppo_train_data[0], ppo_train_data[1]] # Copy to avoid modifying original
241+
ppo_train_data_copy.append(learner.train_iter)
242+
ppo_train_data_copy.append('all') # ✅ loss_type: train both PPO and world model on new data
243+
244+
log_vars_ppo = learner.train(ppo_train_data_copy, collector.envstep)
245+
246+
# Mark position after PPO training
247+
replay_buffer.mark_latest_transitions_consumed()
248+
logging.info(f"[Rank {rank}] [PPO Training] Completed. Marked position: {replay_buffer.last_pos_in_transition}")
249+
250+
log_vars = {**log_vars_wm, **log_vars_ppo}
251+
else:
252+
logging.warning(f"[Rank {rank}] [PPO Training] No new data, skipping PPO update")
253+
log_vars = log_vars_wm
254+
255+
# ========== Original Training Mode ==========
256+
else:
257+
# Execute multiple training rounds
258+
for i in range(update_per_collect):
259+
# ✅ 处理混合采样返回的元组 (train_data_new, train_data_old)
260+
sample_result = replay_buffer.sample(batch_size, policy)
261+
262+
if isinstance(sample_result, tuple):
263+
train_data_new, train_data_old = sample_result
264+
# 如果有新旧数据分离,先训练新数据
265+
if train_data_old is not None:
266+
# 训练新数据
267+
train_data_new.append(learner.train_iter)
268+
train_data_new.append(True) # is_old_data = False
269+
log_vars_new = learner.train(train_data_new, collector.envstep)
270+
271+
# 训练老数据
272+
train_data_old.append(learner.train_iter)
273+
train_data_old.append(False) # is_old_data = True
274+
log_vars_old = learner.train(train_data_old, collector.envstep)
275+
276+
# 更新优先级(如果使用)
277+
if cfg.policy.use_priority:
278+
replay_buffer.update_priority(train_data_new, log_vars_new[0]['value_priority_orig'])
279+
replay_buffer.update_priority(train_data_old, log_vars_old[0]['value_priority_orig'])
280+
else:
281+
# Fallback: 只有新数据
282+
train_data_new.append(learner.train_iter)
283+
log_vars = learner.train(train_data_new, collector.envstep)
284+
if cfg.policy.use_priority:
285+
replay_buffer.update_priority(train_data_new, log_vars[0]['value_priority_orig'])
286+
else:
287+
# 向后兼容:如果返回的不是元组,使用原始逻辑
288+
train_data = sample_result
289+
train_data.append(learner.train_iter)
290+
log_vars = learner.train(train_data, collector.envstep)
291+
if cfg.policy.use_priority:
292+
replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig'])
293+
294+
if replay_buffer._cfg.reanalyze_ratio > 0 and i % 20 == 0:
295+
policy.recompute_pos_emb_diff_and_clear_cache()
296+
297+
if cfg.policy.use_wandb:
298+
policy.set_train_iter_env_step(learner.train_iter, collector.envstep)
251299

252300
# Clear replay buffer after training for online learning
253301
# if cfg.policy.get('online_learning', False):

lzero/mcts/buffer/game_buffer.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -70,11 +70,14 @@ def __init__(self, cfg: dict):
7070
self.num_of_collected_episodes = 0
7171
self.base_idx = 0
7272
self.clear_time = 0
73-
73+
7474
# ✅ 记录最新 push 的 segment 数量
7575
self.latest_push_count = 0
7676
self.new_data_ratio = self._cfg.get('new_data_ratio', 0.5) # 默认 50% 新数据
7777

78+
# ✅ PriorZero-style: 跟踪已消费的 transition 位置
79+
self.last_pos_in_transition = 0
80+
7881
@abstractmethod
7982
def sample(
8083
self, batch_size: int, policy: Union["MuZeroPolicy", "EfficientZeroPolicy", "SampledEfficientZeroPolicy", "GumbelMuZeroPolicy"]

0 commit comments

Comments
 (0)