diff --git a/unirl/rollout/engine/fastvideo/engine.py b/unirl/rollout/engine/fastvideo/engine.py index 5fb1fedc1..24cdc24ab 100644 --- a/unirl/rollout/engine/fastvideo/engine.py +++ b/unirl/rollout/engine/fastvideo/engine.py @@ -181,6 +181,7 @@ def __init__( ) self._version = 0 + self._weight_update_calls = 0 self._generate_lock = threading.Lock() self._shutdown_lock = threading.Lock() self._shutdown_requested = False @@ -620,7 +621,7 @@ def update_weights_from_path(self, checkpoint_path: str, *, track_prefix: str = require(self._generator is not None, "fastvideo engine is offloaded/not initialized") self._generator.update_transformer_weights_from_path(checkpoint_path) self._last_weights_path = checkpoint_path - self._version += 1 + self._weight_update_calls += 1 logger.info("fastvideo transformer weights updated from %s", checkpoint_path) diff --git a/unirl/rollout/engine/sglang/engine.py b/unirl/rollout/engine/sglang/engine.py index 06557a160..a885aedef 100644 --- a/unirl/rollout/engine/sglang/engine.py +++ b/unirl/rollout/engine/sglang/engine.py @@ -152,6 +152,7 @@ def __init__( ) self._version = 0 + self._weight_update_calls = 0 def _prepare_generation(self, sample: Sample) -> Any: require( @@ -282,7 +283,7 @@ def update_weights_from_tensor( load_format=load_format, flush_cache=flush_cache, ) - self._version += 1 + self._weight_update_calls += 1 def init_weights_update_group( self, @@ -329,7 +330,7 @@ def update_weights_from_distributed( group_name=group_name, flush_cache=flush_cache, ) - self._version += 1 + self._weight_update_calls += 1 def destroy_weights_update_group( self, diff --git a/unirl/rollout/engine/sglang_diffusion/engine.py b/unirl/rollout/engine/sglang_diffusion/engine.py index a6a8c859e..2ab166119 100644 --- a/unirl/rollout/engine/sglang_diffusion/engine.py +++ b/unirl/rollout/engine/sglang_diffusion/engine.py @@ -98,6 +98,7 @@ def __init__( self.schedule_policy = self.adapter.schedule_policy() self._version = 0 + self._weight_update_calls = 0 self._generate_lock = threading.Lock() self._shutdown_lock = threading.Lock() self._shutdown_requested = False @@ -236,7 +237,7 @@ def update_weights_from_tensor( load_format=load_format, flush_cache=flush_cache, ) - self._version += 1 + self._weight_update_calls += 1 def init_weights_update_group( self, @@ -279,7 +280,7 @@ def update_weights_from_distributed( target_modules=target_modules, flush_cache=flush_cache, ) - self._version += 1 + self._weight_update_calls += 1 def destroy_weights_update_group( self, diff --git a/unirl/rollout/engine/vllm_omni/engine.py b/unirl/rollout/engine/vllm_omni/engine.py index 194b38139..bf19ceeb9 100644 --- a/unirl/rollout/engine/vllm_omni/engine.py +++ b/unirl/rollout/engine/vllm_omni/engine.py @@ -38,6 +38,7 @@ def __init__( ) -> None: self.cfg = config self._version = 0 + self._weight_update_calls = 0 self._generate_lock = threading.Lock() self._shutdown_lock = threading.Lock() self._shutdown_requested = False @@ -256,7 +257,7 @@ def update_weights_from_ipc( use_shm=use_shm, replica_rank=replica_rank, ) - self._version += 1 + self._weight_update_calls += 1 def init_weights_update_group( self, @@ -299,7 +300,7 @@ def update_weights_from_distributed( target_modules=target_modules, flush_cache=flush_cache, ) - self._version += 1 + self._weight_update_calls += 1 def destroy_weights_update_group( self, @@ -326,7 +327,7 @@ def update_weights_from_tensor( load_format=load_format, flush_cache=flush_cache, ) - self._version += 1 + self._weight_update_calls += 1 def set_lora_from_tensors( self,