Commit beeaa41d authored by jameskrw's avatar jameskrw
Browse files

update wandb log

parent 3ac67d0c
Loading
Loading
Loading
Loading
+9 −1
Original line number Diff line number Diff line
@@ -638,6 +638,7 @@ class QwenVLRolloutManger():
        outputs=[]
        scores=[]
        images=[]
        dones=[]
        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, self.envs[k].interface_config)
@@ -647,4 +648,11 @@ class QwenVLRolloutManger():
            outputs.append(ouput_rst['prompt'])
            scores.append(score)
            images.append(ouput_rst['image_data'])
        return inputs,outputs,scores,images
            dones.append(self.env_states[k]['done'])
        return {
            "inputs": inputs,
            "outputs": outputs,
            "scores": scores,
            "images": images,
            "dones": dones,
        }
+27 −5
Original line number Diff line number Diff line
@@ -715,6 +715,7 @@ class RayPPOTrainer(object):
        sample_outputs = []
        sample_scores = []
        sample_images = []
        sample_dones=[]
        
        if self.test_rollout_manager==None:
            self.test_rollout_manager = QwenVLRolloutManger(
@@ -749,19 +750,28 @@ class RayPPOTrainer(object):
            
            self.test_rollout_manager.reset(env_configs)
            self.test_rollout_manager.rollout_loop()
            inputs, outputs, scores,images = self.test_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)
            sample_images.extend(images)
            rst = self.test_rollout_manager.recording_to_log() # data source == inputs in our current setting, outputs=whole trjecotry
            sample_inputs.extend(rst['inputs'])
            sample_outputs.extend(rst['outputs'])
            sample_scores.extend(rst['scores'])
            sample_images.extend(rst['images'])
            sample_dones.extend(rst['dones'])
        
        self._maybe_log_val_generations_to_wandb(inputs=sample_inputs, outputs=sample_outputs, scores=sample_scores, images=sample_images)
        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)
        
        
        data_source_done = defaultdict(list)
        for data_source, dones in zip(sample_inputs, sample_dones):
            data_source_done[data_source].append(dones)
        for data_source, dones in data_source_done.items():
            metric_dict[f'val/test_done/{data_source}'] = np.mean(dones)
        print(f"[DEBUG] validation at global step {self.global_steps} ends")
        return metric_dict

@@ -1052,7 +1062,19 @@ class RayPPOTrainer(object):
                        # gen_batch_output = self.actor_rollout_wg.generate_sequences(gen_batch)
                        rollout_manager.rollout_loop()
                        final_gen_batch_output = rollout_manager.generate_batch_for_update()
                        rst= rollout_manager.recording_to_log() # data source == inputs in our current setting, outputs=whole 
                        
                        data_source_reward = defaultdict(list)
                        for data_source, scores in zip(rst["inputs"], rst["scores"]):
                            data_source_reward[data_source].append(scores)
                        for data_source, rewards in data_source_reward.items():
                            metrics[f'train/score/{data_source}'] = np.mean(rewards)

                        data_source_done = defaultdict(list)
                        for data_source, dones in zip(rst["inputs"], rst["dones"]):
                            data_source_done[data_source].append(dones)
                        for data_source, dones in data_source_done.items():
                            metrics[f'train/done/{data_source}'] = np.mean(dones)
                    print(f"[DEBUG] step {self.global_steps} rollout ends")
                    
                    # VAGEN: This is moved to before rollout because we don't use vllm n sample param for grpo due to multi-turn nature of our method