Commit 9cbfc79f authored by jameskrw's avatar jameskrw
Browse files

added multi turn reward, move rollout config to ppo_trainer

parent ea6204ed
Loading
Loading
Loading
Loading

vagen/mllm_agent/image.png

deleted100644 → 0
−1.67 MiB
Loading image diff...

vagen/mllm_agent/image2.png

deleted100644 → 0
−641 KiB
Loading image diff...
+12 −38
Original line number Diff line number Diff line
@@ -17,30 +17,18 @@ import vagen.env
from vagen.env.register import REGISTERED_ENVS
from vagen.env.base import EnvConfig,IMAGE_PLACEHOLDER

@dataclass
class QwenVLRolloutConifg:
    window_size: int = 5
    max_trajectory_length: int = 3072
    max_turns: int = 5
    n_gpu_per_node: int = 1 # used for multigpu batch balancing
    sptk_for_loss_mask: List[str] = field(init=False, default_factory=lambda: ['<|box_start|>', '<|box_end|>'])
    end_turn_token: str = field(init=False, default='<|im_end|>') # specially for Qwen2.5
    
class QwenVLRolloutManger():
    def __init__(self,
                 actor_rollout_wg,
                 config,
                 tokenizer: PreTrainedTokenizer,
                 config: QwenVLRolloutConifg,
                 processor: Optional[ProcessorMixin] = None,
                 verbose: bool = False,
                 truncation='error',
                 ):
        self.tokenizer = tokenizer
        self.processor = processor
        self.config = config
        self.actor_rollout_wg = actor_rollout_wg
        self.verbose = verbose
        self.truncation = truncation
        self.recorder= None # defaultdict(list) env_id:record
        self.envs = None # dict env_id:EnvInterface
        self.env_states = None # dict
@@ -55,8 +43,8 @@ class QwenVLRolloutManger():
        llm_raw_response = re.sub(r'<image>', '', llm_raw_response)
        if prep_for_loss_mask:
            # filtering special tokens for llm_raw_response, then adding them to the beginning and end of the response for loss mask computation
            sptk_b = self.config.sptk_for_loss_mask[0]
            sptk_e = self.config.sptk_for_loss_mask[1]
            sptk_b = self.config.special_token_for_loss_mask[0]
            sptk_e = self.config.special_token_for_loss_mask[1]
            llm_raw_response = re.sub(sptk_e, '', llm_raw_response)
            llm_raw_response = re.sub(sptk_b, '', llm_raw_response)
            llm_raw_response = sptk_b + llm_raw_response + sptk_e
@@ -134,8 +122,8 @@ class QwenVLRolloutManger():
        """
        
        # Get token IDs for special tokens and pad token
        sptk_b = self.tokenizer.convert_tokens_to_ids(self.config.sptk_for_loss_mask[0])
        sptk_e = self.tokenizer.convert_tokens_to_ids(self.config.sptk_for_loss_mask[1])
        sptk_b = self.tokenizer.convert_tokens_to_ids(self.config.special_token_for_loss_mask[0])
        sptk_e = self.tokenizer.convert_tokens_to_ids(self.config.special_token_for_loss_mask[1])
        pad_token_id = self.tokenizer.pad_token_id

        batch_size = input_ids.shape[0]
@@ -458,13 +446,13 @@ class QwenVLRolloutManger():
                                                                         max_length=self.config.max_trajectory_length-1, # -1 for the prompt padding token
                                                                         pad_token_id=self.tokenizer.pad_token_id,
                                                                         left_pad=False,
                                                                         truncation=self.truncation)
                                                                         truncation=self.config.truncation)
        input_ids_prompt, attention_mask_prompt = verl_F.tokenize_and_postprocess_data(prompt=prompt_with_chat_template,
                                                                         tokenizer=self.tokenizer,
                                                                         max_length=1,
                                                                         pad_token_id=self.tokenizer.pad_token_id,
                                                                         left_pad=True,
                                                                         truncation=self.truncation)
                                                                         truncation=self.config.truncation)
        attention_mask_prompt=torch.zeros_like(input_ids_prompt) # All prompt will be masked
        
        
@@ -505,22 +493,23 @@ class QwenVLRolloutManger():
            delta_position_id = torch.arange(1, response_length + 1, device=position_ids_prompt.device)
            position_ids_response = position_ids_prompt[-1:] + delta_position_id
        
        if self.config.use_multi_turn_reward:
            reward_positions = torch.nonzero(token_level_reward_mask).squeeze(-1)
            multi_turn_token_level_reward = torch.zeros_like(token_level_reward_mask, dtype=torch.float)
            assert len(reward_positions) == len(rewards), "Number of rewards does not match number of reward positions"
            for idx,reward in enumerate(rewards):
                multi_turn_token_level_reward[reward_positions[idx]] = reward
            
            row_dict["multi_turn_token_level_reward"] = multi_turn_token_level_reward # (seq_len,) 
        if self.config.use_loss_mask:
            row_dict['loss_mask'] = loss_mask
        position_ids = torch.cat([position_ids_prompt, position_ids_response], dim=-1)
        row_dict['prompts'] = input_ids_prompt
        row_dict['responses'] = input_ids_response
        row_dict['input_ids'] = input_ids
        row_dict['attention_mask'] = attention_mask
        row_dict['position_ids'] = position_ids
        row_dict['loss_mask'] = loss_mask
        index = row_dict.get("extra_info", {}).get("index", 0)
        row_dict["index"] = index
        row_dict["multi_turn_token_level_reward"] = multi_turn_token_level_reward # (seq_len) later need to convert to (response_len)
        return row_dict

    @torch.no_grad()
@@ -592,14 +581,7 @@ class QwenVLRolloutManger():
            gen_batch.non_tensor_batch['raw_prompt_ids'] = raw_prompt_ids_array
            
            output_batch = self.actor_rollout_wg.generate_sequences(gen_batch)
            ##DEBUG
            # print(f"[DEBUG] rollout turn {step}")
            # print(f"[DEBUG] rollout output_batch.non_tensor_batch.keys(): {output_batch.non_tensor_batch.keys()}")
            # print(f"[DEBUG] rollout output_batch.batch.keys(): {output_batch.batch.keys()}")
            # print(f"[DEBUG] rollout output_batch.batch['input_ids'].shape: {output_batch.batch['input_ids'].shape}")
            # print(f"[DEBUG] rollout output_batch.batch['attention_mask'].shape: {output_batch.batch['attention_mask'].shape}")
            # print(f"[DEBUG] rollout output_batch.batch['position_ids'].shape: {output_batch.batch['position_ids'].shape}")
            # print(f"[DEBUG] --------------------------------------------")
            
            
            
            responses_str = self.tokenizer.batch_decode(
@@ -632,14 +614,6 @@ class QwenVLRolloutManger():
            batch_list.append(row_dict)
        batch_dict = collate_fn(batch_list)
        batch = DataProto.from_single_dict(batch_dict)
        ##DEBUG
        # print(f"[DEBUG] final trajectory")
        # print(f"[DEBUG] rollout batch.non_tensor_batch.keys(): {batch.non_tensor_batch.keys()}")
        # print(f"[DEBUG] rollout batch.batch.keys(): {batch.batch.keys()}")
        # print(f"[DEBUG] rollout batch.batch['input_ids'].shape: {batch.batch['input_ids'].shape}")
        # print(f"[DEBUG] rollout batch.batch['attention_mask'].shape: {batch.batch['attention_mask'].shape}")
        # print(f"[DEBUG] rollout batch.batch['loss_mask'].shape: {batch.batch['loss_mask'].shape}")
        # print(f"[DEBUG] --------------------------------------------")
        return batch
    
    @torch.no_grad()
+0 −134
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([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
rst
```

%% Output

    ['ascasAbc def']
+0 −191
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
text_template = "The quick brown fox jumps over the lazy dog.<image1asdasdqwa>, <image2>, <image3>."
import re
image_keys=re.findall(r'<image[a-zA-Z0-9]*>', text_template)
print(image_keys)
```

%% Output

    ['<image1asdasdqwa>', '<image2>', '<image3>']

%% 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
tokenizer.pad_token
```

%% Output

    '<|endoftext|>'

%% Cell type:code id: tags:

``` python
tokenizer.decode(tokenizer.pad_token_id)
```

%% Output

    '<|endoftext|>'

%% Cell type:code id: tags:

``` python
def loss_mask(input_ids, attention_mask):
    sptk_b = tokenizer.convert_tokens_to_ids('<|box_start|>')
    sptk_e = tokenizer.convert_tokens_to_ids('<|box_end|>')
    pad_token_id = tokenizer.pad_token_id

    print(f"DEBUG:input_ids.shape:{input_ids.shape}")
    batch_size = input_ids.shape[0]
    seq_len = input_ids.shape[1]

    # Initialize output tensors with same shape as inputs
    new_input_ids = input_ids.clone()
    new_attention_mask = attention_mask.clone()
    loss_mask = torch.zeros_like(input_ids)
    new_loss_mask = torch.zeros_like(input_ids)
    # Process each example in the batch
    for b in range(batch_size):
        # Count right padding tokens using attention mask
        right_pad_tokens = (new_input_ids[b] == pad_token_id).sum().item()

        # Assert that initial padding tokens have attention mask of 0
        assert torch.all(attention_mask[b, -right_pad_tokens:] == 0), "right padding tokens must have attention mask of 0"

        # Find special token indices
        sptk_b_indices = (input_ids[b] == sptk_b).nonzero().flatten()
        sptk_e_indices = (input_ids[b] == sptk_e).nonzero().flatten()

        # Create a mask for tokens that should compute loss
        hole_pos=[] # initialize holes position list with last padding token position
        for start_pos, end_pos in zip(sptk_b_indices, sptk_e_indices):
            loss_mask[b][start_pos+1:end_pos] = 1
            hole_pos.append(start_pos.item())
            hole_pos.append(end_pos.item())
        hole_pos.append(seq_len-right_pad_tokens)
        assert new_input_ids[b][seq_len-right_pad_tokens]==pad_token_id

        # shift right to fill the wholes
        holes_to_fill=1
        for i in range(0,len(hole_pos)-1):
            start_pos = hole_pos[i]
            end_pos = hole_pos[i+1]
            new_loss_mask[b,start_pos+1-holes_to_fill:end_pos-holes_to_fill]=loss_mask[b,start_pos+1:end_pos]
            new_input_ids[b,start_pos+1-holes_to_fill:end_pos-holes_to_fill]=input_ids[b,start_pos+1:end_pos]
            new_attention_mask[b,start_pos+1-holes_to_fill:end_pos-holes_to_fill]=attention_mask[b,start_pos+1:end_pos]
            holes_to_fill+=1

        valid_tokens = seq_len-right_pad_tokens-len(hole_pos)+1 # the number of non-special tokens and non-padding tokens
        new_loss_mask[b][valid_tokens:]=0
        new_input_ids[b][valid_tokens:]=pad_token_id
        new_attention_mask[b][valid_tokens:]=0

    return new_input_ids, new_attention_mask, new_loss_mask
```

%% Cell type:code id: tags:

``` python
import verl.utils.torch_functional as verl_F
import torch
prompt_with_chat_template=''
input_ids,attention_mask=verl_F.tokenize_and_postprocess_data(prompt=prompt_with_chat_template,
                                        tokenizer=tokenizer,
                                        max_length=1024,
                                        pad_token_id=tokenizer.pad_token_id,
                                        left_pad=False,
                                        truncation="error",
                                        )
```

%% Cell type:code id: tags:

``` python
input_ids.shape
```

%% Output

    torch.Size([1, 1024])

%% Cell type:code id: tags:

``` python
print(tokenizer.decode(input_ids[0]))
print(attention_mask[:,:10])
```

%% Output

    ---------------------------------------------------------------------------
    TypeError                                 Traceback (most recent call last)
Cell     In[25], line 1
    ----> 1 print(tokenizer.decode(input_ids[0]))
          2 print(attention_mask[:,:10])
File     ~/miniconda3/envs/vagen/lib/python3.11/site-packages/transformers/tokenization_utils_base.py:3860, in PreTrainedTokenizerBase.decode(self, token_ids, skip_special_tokens, clean_up_tokenization_spaces, **kwargs)
       3857 # Convert inputs to python lists
       3858 token_ids = to_py_obj(token_ids)
    -> 3860 return self._decode(
       3861     token_ids=token_ids,
       3862     skip_special_tokens=skip_special_tokens,
       3863     clean_up_tokenization_spaces=clean_up_tokenization_spaces,
       3864     **kwargs,
       3865 )
File     ~/miniconda3/envs/vagen/lib/python3.11/site-packages/transformers/tokenization_utils_fast.py:668, in PreTrainedTokenizerFast._decode(self, token_ids, skip_special_tokens, clean_up_tokenization_spaces, **kwargs)
        666 if isinstance(token_ids, int):
        667     token_ids = [token_ids]
    --> 668 text = self._tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens)
        670 clean_up_tokenization_spaces = (
        671     clean_up_tokenization_spaces
        672     if clean_up_tokenization_spaces is not None
        673     else self.clean_up_tokenization_spaces
        674 )
        675 if clean_up_tokenization_spaces:
    TypeError: argument 'ids': 'float' object cannot be interpreted as an integer

%% Cell type:code id: tags:

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

%% 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])
```

%% 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])
Loading