@@ -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):
0 commit comments