@@ -63,11 +63,10 @@ def worldmodel() -> WorldModelTrainer.Config:
6363 optimizer = default_adamw (lr = 2e-4 , weight_decay = 1e-2 )
6464 optimizer .implementation = "fused_opt_states_bf16"
6565 local_world_size , world_size , num_nodes = _world_sizes ()
66- checkpoint_base_folder = _reporterv2_checkpoint_base_folder ()
66+ checkpoint_folder = _reporterv2_checkpoint_folder ()
6767
6868 return WorldModelTrainer .Config (
6969 hf_assets_path = "." ,
70- dump_folder = checkpoint_base_folder or "./outputs" ,
7170 loss = WorldModelLoss .Config (plan_loss_weight = 0.1 ),
7271 tokenizer = WorldModelTokenizer .Config (
7372 compressor_model = COMPRESSOR_MODEL ,
@@ -111,7 +110,7 @@ def worldmodel() -> WorldModelTrainer.Config:
111110 ),
112111 checkpoint = WorldModelTorchPackageCheckpointManager .Config (
113112 enable = True ,
114- folder = os . getenv ( "REPORTERV2_TRAINING_ID" ) or "checkpoint" ,
113+ folder = checkpoint_folder ,
115114 interval = validation_freq * 5 ,
116115 async_mode = "async" ,
117116 keep_latest_k = 0 ,
@@ -155,9 +154,10 @@ def _world_sizes() -> tuple[int, int, int]:
155154 return local_world_size , world_size , num_nodes
156155
157156
158- def _reporterv2_checkpoint_base_folder () -> str :
157+ def _reporterv2_checkpoint_folder () -> str :
159158 host = os .getenv ("REPORTERV2_HOST" )
160- return f"{ host .rstrip ('/' )} /checkpoint" if host else ""
159+ folder = os .getenv ("REPORTERV2_TRAINING_ID" ) or "checkpoint"
160+ return f"{ host .rstrip ('/' )} /checkpoint/{ folder } " if host else folder
161161
162162
163163def main () -> None :
0 commit comments