Commit c2cf988b authored by jameskrw's avatar jameskrw
Browse files

updated rollout manager for service

parent 56033e73
Loading
Loading
Loading
Loading
+685 −0

File added.

Preview size limit exceeded, changes collapsed.

+4 −0
Original line number Diff line number Diff line
@@ -201,3 +201,7 @@ rollout_manager:
  use_gae_mask: True
  special_token_for_loss_mask: ['<|box_start|>', '<|box_end|>']
  truncation: ${data.truncation}
  base_url: http://localhost:5000
  use_service: False
  timeout: 60
  max_workers: 8
 No newline at end of file
+19 −2
Original line number Diff line number Diff line
@@ -41,7 +41,7 @@ from torch.utils.data import RandomSampler, SequentialSampler
from torchdata.stateful_dataloader import StatefulDataLoader

from vagen.mllm_agent.rollout import QwenVLRolloutManger

from vagen.mllm_agent.rollout_service import QwenVLRolloutMangerService
WorkerType = Type[Worker]


@@ -760,6 +760,15 @@ class RayPPOTrainer(object):
        # Lists to collect samples for the table
    
        if self.test_rollout_manager==None:
            if self.config.rollout_manager.get("use_service",False):
                self.test_rollout_manager =QwenVLRolloutMangerService(
                    actor_rollout_wg=self.actor_rollout_wg,
                    config=self.config.rollout_manager,
                    tokenizer=self.tokenizer,
                    processor=self.processor,
                    split="val",
                )
            else:
                self.test_rollout_manager =QwenVLRolloutManger(
                    actor_rollout_wg=self.actor_rollout_wg,
                    config=self.config.rollout_manager,
@@ -1021,7 +1030,15 @@ class RayPPOTrainer(object):
        # we start from step 1
        self.global_steps += 1


        if self.config.rollout_manager.get("use_service",False):
            rollout_manager = QwenVLRolloutMangerService(
                actor_rollout_wg=self.actor_rollout_wg,
                config=self.config.rollout_manager,
                tokenizer=self.tokenizer,
                processor=self.processor,
                split="train",
            )
        else:
            rollout_manager = QwenVLRolloutManger(
                actor_rollout_wg=self.actor_rollout_wg,
                config=self.config.rollout_manager,