Commit 4d7c8b6e authored by jameskrw's avatar jameskrw
Browse files

refactored rollout manger, add validation

parent 5cdb8022
Loading
Loading
Loading
Loading
+67 −45
Original line number Diff line number Diff line
@@ -40,8 +40,8 @@ class QwenVLRolloutManger():
        self.actor_rollout_wg = actor_rollout_wg
        self.verbose = verbose
        self.truncation = truncation
        self.recorder= None # defaultdict(list)
        self.envs = None # dict
        self.recorder= None # defaultdict(list) env_id:record
        self.envs = None # dict env_id:EnvInterface
        self.env_states = None # dict
        self.batch_idx_to_env_id = None # dict

@@ -101,6 +101,8 @@ class QwenVLRolloutManger():
    def _compute_loss_mask(self, input_ids, attention_mask):
        """
        Compute loss mask for the input ids and attention mask
        We only do loss for the tokens in input_ids that are wrapped by special tokens (by defualt they're <|box_start|> and <|box_end|>)
        
        
        Args:
            input_ids: (batch_size, seq_len)
@@ -111,7 +113,7 @@ class QwenVLRolloutManger():
            attention_mask: (batch_size, seq_len)
            loss_mask: (batch_size, seq_len)
        
        - There will be different stratgy to handel special tokens in the list
        - There will be different stratgy to handel special tokens in the input_ids
        - 1. remove them, in this case we need to fill the hole by adding pad in the right and shift the sequence left
        - 2. keep them, attention mask will be 0 for them
        - 3. Replace them with pad token
@@ -287,29 +289,12 @@ class QwenVLRolloutManger():
        self.recorder[env_id].append(record_entry)


    def _generate_input_item(
            self, 
    def _single_recording_to_prompt(self,
                            recording: List[Dict], 
                            step: int, 
                            window_size: int = None,
        ):
        """
        Given a recording, generate the input for MLLM
                            last_question: bool = False,):
        
        Args:
            recording: List of dictionaries containing recorded environment interactions
            step: Current step to generate input for
            window_size: Number of past steps to include in the context
        
        Returns:
            Dictionary containing properly formatted inputs for the MLLM
            - prompts: task instruction
            - responses: responses generated from prompts
            - input_ids, attention_mask, position_ids: prompts and responses generated from prompts
            - position_ids: 
                - position_ids for prompts: rope
                - rest postion_ids: refer to vllm_rollout_spmd.py to check how to compute
        """
        assert step >= 0
        start_step = max(0, step - window_size) if window_size is not None else 0
        end_step = step
@@ -326,6 +311,7 @@ class QwenVLRolloutManger():
                llm_raw_response = record['info']['llm_raw_response']
                filtered_llm_raw_response = self._handle_special_tokens(llm_raw_response, compute_loss_mask=False)
                chat.append({"role": "assistant", "content": filtered_llm_raw_response})
            if i<len(history)-1 or last_question:
                chat.append({"role": "user", "content": record['text_template']})

        image_data=[]
@@ -334,9 +320,39 @@ class QwenVLRolloutManger():
                for img in record['image_data']:
                    image_data.append(img)
            
        prompt_with_chat_template = self.tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=False)
        return {
            "prompt": prompt_with_chat_template,
            "image_data": image_data,
        }
        
    def _generate_input_item(
            self, 
            recording: List[Dict], 
            step: int, 
            window_size: int = None,
        ):
        """
        Given a recording, generate the input for MLLM
        
        Args:
            recording: List of dictionaries containing recorded environment interactions
            step: Current step to generate input for
            window_size: Number of past steps to include in the context
        
        Returns:
            Dictionary containing properly formatted inputs for the MLLM
            - prompts: task instruction
            - responses: responses generated from prompts
            - input_ids, attention_mask, position_ids: prompts and responses generated from prompts
            - position_ids: 
                - position_ids for prompts: rope
                - rest postion_ids: refer to vllm_rollout_spmd.py to check how to compute
        """
        rst=self._single_recording_to_prompt(recording, step, window_size,last_question=True)
        prompt_with_chat_template=rst['prompt']
        image_data=rst['image_data']        
        has_images = len(image_data) > 0        
        prompt_with_chat_template = self.tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=False)

        row_dict = {}
        if has_images:  # expand image token
@@ -384,33 +400,17 @@ class QwenVLRolloutManger():
                - rest postion_ids: refer to vllm_rollout_spmd.py to check how to compute

        """
        assert step > 0 # we must update with at least one response
        start_step = max(0, step - window_size) if window_size is not None else 0
        end_step = step
        assert len(recording) >= end_step+1, 'History length is not enough'
        history = recording[start_step: end_step + 1]



        # handle prompt, prompt=pad_token since we now have everything in response and compute a loss mask for them
        prompt_with_chat_template=self.tokenizer.pad_token 
        
        # handle response
        response_chat = []
        env_id = history[0]['env_id']
        response_chat.append({"role": "system", "content": self.envs[env_id].get_task_instruction()})
        for i, record in enumerate(history):
            if i>0:
                llm_raw_response = record['info']['llm_raw_response']
                filtered_llm_raw_response = self._handle_special_tokens(llm_raw_response, compute_loss_mask=True)
                response_chat.append({"role": "assistant", "content": filtered_llm_raw_response})
            if i<len(history)-1:
                response_chat.append({"role": "user", "content": record['text_template']})
        response_with_chat_template = self.tokenizer.apply_chat_template(response_chat, add_generation_prompt=False, tokenize=False)
        image_data=[]
        for record in history[:-1]: # do not collect last step's obseravation data since we don't use it for update
            if 'image_data' in record:
                for img in record['image_data']:
                    image_data.append(img)
        response_rst=self._single_recording_to_prompt(recording, step, window_size, last_question=False)
        response_with_chat_template=response_rst['prompt']
        image_data=response_rst['image_data']
       
        has_images = len(image_data) > 0
        row_dict = {}
        if has_images:  # expand image token
@@ -571,10 +571,32 @@ class QwenVLRolloutManger():
            row_dict = self._generate_input_final_item(
                recording=self.recorder[env_id],
                step=self.env_states[env_id]['step'],
                window_size=self.config.window_size,
                window_size=None,
            )
            row_dict['reward_model'] = {"style": "given", "ground_truth": {"reward": self.envs[env_id].get_traj_reward()}}
            batch_list.append(row_dict)
        batch_dict = collate_fn(batch_list)
        batch = DataProto.from_single_dict(batch_dict)
        return batch
    
    
    def recording_to_log(self):
        """
        Get the recording of all environments
        
        Returns:
            Dictionary containing the recording of all environments
        """
        inputs=[]
        outputs=[]
        scores=[]
        for k,v in self.recorder.items():
            step=self.env_states[k]['step']
            input_str=self.envs[k].name_repr()+self.envs[k].config_repr(self.envs[k].env_config)
            ouput_rst=self._single_recording_to_prompt(v, step, window_size=None, last_question=False)
            output_str=ouput_rst['prompt']
            score=self.envs[k].get_traj_reward()
            inputs.append(input_str)
            outputs.append(output_str)
            scores.append(score)
        return inputs,outputs,scores
+18 −11
Original line number Diff line number Diff line
%% Cell type:code id: tags:

``` python
from verl.utils import hf_tokenizer, hf_processor
import torch
```

%% Cell type:code id: tags:

``` python
model_name = "Qwen/Qwen2.5-VL-3B-Instruct"
processor = hf_processor(model_name)
tokenizer = hf_tokenizer(model_name)
```

%% Output

    Using a slow image processor as `use_fast` is unset and a slow processor was saved with this model. `use_fast=True` will be the default behavior in v4.48, even if the model was saved with a slow processor. This will result in minor differences in outputs. You'll still be able to use a slow processor with `use_fast=False`.

%% Cell type:code id: tags:

``` python
```

%% Cell type:code id: tags:

``` python
import verl.utils.torch_functional as verl_F
import torch
from verl.utils.model import compute_position_id_with_mask
response_with_chat_template='Abc def'
prompt_with_chat_template='ascas'
input_ids_response, attention_mask_response = verl_F.tokenize_and_postprocess_data(prompt=response_with_chat_template,
                                                                         tokenizer=tokenizer,
                                                                         max_length=10,
                                                                         pad_token_id=tokenizer.pad_token_id,
                                                                         left_pad=False,
                                                                         truncation='error')
input_ids_prompt, attention_mask_prompt = verl_F.tokenize_and_postprocess_data(prompt=prompt_with_chat_template,
                                                                         tokenizer=tokenizer,
                                                                         max_length=10,
                                                                         pad_token_id=tokenizer.pad_token_id,
                                                                         left_pad=True,
                                                                         truncation='error')
attention_mask_prompt=torch.zeros_like(input_ids_prompt) # All prompt will be masked



input_ids_prompt=input_ids_prompt[0]
attention_mask_prompt=attention_mask_prompt[0]
input_ids_response=input_ids_response[0]
attention_mask_response=attention_mask_response[0]
loss_mask_prompt = torch.zeros_like(attention_mask_prompt)


input_ids = torch.cat([input_ids_prompt, input_ids_response], dim=-1)
attention_mask = torch.cat([attention_mask_prompt, attention_mask_response], dim=-1)



position_ids_prompt = compute_position_id_with_mask(attention_mask_prompt)

response_length = input_ids_response.shape[0]
delta_position_id = torch.arange(1, response_length + 1, device=position_ids_prompt.device)

```

%% Cell type:code id: tags:

``` python
delta_position_id
```

%% Output

    tensor([ 1,  2,  3,  4,  5,  6,  7,  8,  9, 10])

%% Cell type:code id: tags:

``` python
position_ids_prompt[-1]=10
```

%% Cell type:code id: tags:

``` python
position_ids_response = position_ids_prompt[-1:] + delta_position_id
```

%% Cell type:code id: tags:

``` python
position_ids_response
```

%% Output

    tensor([11, 12, 13, 14, 15, 16, 17, 18, 19, 20])

%% Cell type:code id: tags:

``` python
input_ids.shape
```

%% Output

    torch.Size([1, 1024])
    torch.Size([20])

%% Cell type:code id: tags:

``` python
print(attention_mask[:10])
```

%% Output

    tensor([0, 0, 0, 0, 0, 0, 0, 0, 0, 0])

%% Cell type:code id: tags:

``` python
rst=tokenizer.batch_decode([input_ids], skip_special_tokens=True)
```

%% Cell type:code id: tags:

``` python
new_input_ids, new_attention_mask, new_loss_mask=loss_mask(input_ids, attention_mask)
print(tokenizer.decode(new_input_ids[0][:30]))
print(new_attention_mask[0][:30])
print(new_loss_mask[0][:30])
rst
```

%% Output

    DEBUG:input_ids.shape:torch.Size([1, 1024])
    I love you<|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|><|endoftext|>
    tensor([1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
            0, 0, 0, 0, 0, 0])
    tensor([1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
            0, 0, 0, 0, 0, 0])
    ['ascasAbc def']
+128 −56
Original line number Diff line number Diff line
@@ -24,7 +24,7 @@ from enum import Enum
from pprint import pprint
from typing import Type, Dict
from copy import deepcopy

from collections import defaultdict
import numpy as np
from codetiming import Timer
from omegaconf import OmegaConf, open_dict
@@ -408,6 +408,9 @@ class RayPPOTrainer(object):

        self._validate_config()
        self._create_dataloader()
        self.test_rollout_config=None
        self.test_rollout_manager=None
        

    def _validate_config(self):
        config = self.config
@@ -609,89 +612,158 @@ class RayPPOTrainer(object):
        wandb.log({"val/generations": new_table}, step=self.global_steps)
        self.validation_table = new_table

    
    # Original _validate
    # def _validate(self):
    #     reward_tensor_lst = []
    #     data_source_lst = []

    #     # Lists to collect samples for the table
    #     sample_inputs = []
    #     sample_outputs = []
    #     sample_scores = []

    #     for test_data in self.val_dataloader:
    #         test_batch = DataProto.from_single_dict(test_data)

    #         # we only do validation on rule-based rm
    #         if self.config.reward_model.enable and test_batch[0].non_tensor_batch['reward_model']['style'] == 'model':
    #             return {}

    #         # Store original inputs
    #         input_ids = test_batch.batch['input_ids']
    #         input_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in input_ids]
    #         sample_inputs.extend(input_texts)

    #         if 'multi_modal_inputs' in test_batch.non_tensor_batch.keys():
    #             test_gen_batch = test_batch.pop(
    #                 batch_keys=['input_ids', 'attention_mask', 'position_ids'],
    #                 non_tensor_batch_keys=['raw_prompt_ids', 'multi_modal_data', 'multi_modal_inputs'],
    #             )
    #         else:
    #             test_gen_batch = test_batch.pop(
    #                 batch_keys=['input_ids', 'attention_mask', 'position_ids'],
    #                 non_tensor_batch_keys=['raw_prompt_ids'],
    #             )

    #         test_gen_batch.meta_info = {
    #             'eos_token_id': self.tokenizer.eos_token_id,
    #             'pad_token_id': self.tokenizer.pad_token_id,
    #             'recompute_log_prob': False,
    #             'do_sample': False,
    #             'validate': True,
    #         }

    #         # pad to be divisible by dp_size
    #         test_gen_batch_padded, pad_size = pad_dataproto_to_divisor(test_gen_batch, self.actor_rollout_wg.world_size)
    #         test_output_gen_batch_padded = self.actor_rollout_wg.generate_sequences(test_gen_batch_padded)
    #         # unpad
    #         test_output_gen_batch = unpad_dataproto(test_output_gen_batch_padded, pad_size=pad_size)
    #         print('validation generation end')

    #         # Store generated outputs
    #         output_ids = test_output_gen_batch.batch['responses']
    #         output_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in output_ids]
    #         sample_outputs.extend(output_texts)

    #         test_batch = test_batch.union(test_output_gen_batch)

    #         # evaluate using reward_function
    #         reward_tensor = self.val_reward_fn(test_batch)

    #         # Store scores
    #         scores = reward_tensor.sum(-1).cpu().tolist()
    #         sample_scores.extend(scores)

    #         reward_tensor_lst.append(reward_tensor)
    #         data_source_lst.append(test_batch.non_tensor_batch.get('data_source', ['unknown'] * reward_tensor.shape[0]))

    #     self._maybe_log_val_generations_to_wandb(inputs=sample_inputs, outputs=sample_outputs, scores=sample_scores)

    #     reward_tensor = torch.cat(reward_tensor_lst, dim=0).sum(-1).cpu()  # (batch_size,)
    #     data_sources = np.concatenate(data_source_lst, axis=0)

    #     # evaluate test_score based on data source
    #     data_source_reward = {}
    #     for i in range(reward_tensor.shape[0]):
    #         data_source = data_sources[i]
    #         if data_source not in data_source_reward:
    #             data_source_reward[data_source] = []
    #         data_source_reward[data_source].append(reward_tensor[i].item())

    #     metric_dict = {}
    #     for data_source, rewards in data_source_reward.items():
    #         metric_dict[f'val/test_score/{data_source}'] = np.mean(rewards)

    #     return metric_dict
    
    # Agentic Setting _validate
    def _validate(self):
        reward_tensor_lst = []
        data_source_lst = []

        # Lists to collect samples for the table
        sample_inputs = []
        sample_outputs = []
        sample_scores = []

        for test_data in self.val_dataloader:
            test_batch = DataProto.from_single_dict(test_data)
        if self.test_rollout_config==None:
            self.test_rollout_config = QwenVLRolloutConifg(
                max_trajectories_length=self.config.data.max_trajectories_length,
                max_turns=self.config.max_turns,
                n_gpu_per_node=self.config.trainer.n_gpus_per_node,
            )
        
            # we only do validation on rule-based rm
            if self.config.reward_model.enable and test_batch[0].non_tensor_batch['reward_model']['style'] == 'model':
                return {}
        if self.test_rollout_manager==None:
            self.test_rollout_manager = QwenVLRolloutManger(
                actor_rollout_wg=self.actor_rollout_wg,
                tokenizer=self.tokenizer,
                config=self.test_rollout_config,
                processor=self.processor,
                verbose=self.config.trainer.get('verbose', False),
            )
        
            # Store original inputs
            input_ids = test_batch.batch['input_ids']
            input_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in input_ids]
            sample_inputs.extend(input_texts)
        for batch_dict in self.val_dataloader:
            
            if 'multi_modal_inputs' in test_batch.non_tensor_batch.keys():
                test_gen_batch = test_batch.pop(
            batch: DataProto = DataProto.from_single_dict(batch_dict)
            # pop these keys so it will not cause error when rollout
            if 'multi_modal_inputs' in batch.non_tensor_batch.keys():
                batch.pop(
                    batch_keys=['input_ids', 'attention_mask', 'position_ids'],
                    non_tensor_batch_keys=['raw_prompt_ids', 'multi_modal_data', 'multi_modal_inputs'],
                )
            else:
                test_gen_batch = test_batch.pop(
                batch.pop(
                    batch_keys=['input_ids', 'attention_mask', 'position_ids'],
                    non_tensor_batch_keys=['raw_prompt_ids'],
                )

            test_gen_batch.meta_info = {
                'eos_token_id': self.tokenizer.eos_token_id,
                'pad_token_id': self.tokenizer.pad_token_id,
                'recompute_log_prob': False,
                'do_sample': False,
                'validate': True,
            }
            env_configs = [
                    EnvConfig(env_name=batch.non_tensor_batch['extra_info'][i]['env_name'],
                              env_config=batch.non_tensor_batch['extra_info'][i]['env_config'],
                              seed=batch.non_tensor_batch['extra_info'][i]['seed'])
                    for i in range(len(batch))
                ]
            
            # pad to be divisible by dp_size
            test_gen_batch_padded, pad_size = pad_dataproto_to_divisor(test_gen_batch, self.actor_rollout_wg.world_size)
            test_output_gen_batch_padded = self.actor_rollout_wg.generate_sequences(test_gen_batch_padded)
            # unpad
            test_output_gen_batch = unpad_dataproto(test_output_gen_batch_padded, pad_size=pad_size)
            self.rollout_manager.reset(env_configs)
            print('validation generation start')
            self.rollout_manager.rollout_loop()
            print('validation generation end')

            # Store generated outputs
            output_ids = test_output_gen_batch.batch['responses']
            output_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in output_ids]
            sample_outputs.extend(output_texts)

            test_batch = test_batch.union(test_output_gen_batch)

            # evaluate using reward_function
            reward_tensor = self.val_reward_fn(test_batch)

            # Store scores
            scores = reward_tensor.sum(-1).cpu().tolist()
            inputs, outputs, scores = self.rollout_manager.recording_to_log() # data source == inputs in our current setting, outputs=whole trjecotry
            sample_inputs.extend(inputs)
            sample_outputs.extend(outputs)
            sample_scores.extend(scores)
        
            reward_tensor_lst.append(reward_tensor)
            data_source_lst.append(test_batch.non_tensor_batch.get('data_source', ['unknown'] * reward_tensor.shape[0]))

        self._maybe_log_val_generations_to_wandb(inputs=sample_inputs, outputs=sample_outputs, scores=sample_scores)

        reward_tensor = torch.cat(reward_tensor_lst, dim=0).sum(-1).cpu()  # (batch_size,)
        data_sources = np.concatenate(data_source_lst, axis=0)

        # evaluate test_score based on data source
        data_source_reward = {}
        for i in range(reward_tensor.shape[0]):
            data_source = data_sources[i]
            if data_source not in data_source_reward:
                data_source_reward[data_source] = []
            data_source_reward[data_source].append(reward_tensor[i].item())

        metric_dict = {}
        data_source_reward = defaultdict(list)
        for data_source, scores in zip(sample_inputs, sample_scores):
            data_source_reward[data_source].append(scores)
        for data_source, rewards in data_source_reward.items():
            metric_dict[f'val/test_score/{data_source}'] = np.mean(rewards)

        return metric_dict

            
            

    def init_workers(self):
        """Init resource pool and worker group"""
        self.resource_pool_manager.create_resource_pool()