Commit 91a2c3bf authored by jameskrw's avatar jameskrw
Browse files

minor test

parent 8b54afc6
Loading
Loading
Loading
Loading
+5 −2
Original line number Diff line number Diff line
@@ -46,7 +46,7 @@ python3 -m vagen.trainer.main_ppo \
    actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=1 \
    actor_rollout_ref.rollout.tensor_model_parallel_size=2 \
    actor_rollout_ref.rollout.name=vllm \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.4 \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.2 \
    actor_rollout_ref.rollout.enable_chunked_prefill=False \
    actor_rollout_ref.rollout.enforce_eager=False \
    actor_rollout_ref.rollout.free_cache_engine=False \
@@ -67,7 +67,7 @@ python3 -m vagen.trainer.main_ppo \
    trainer.logger=['console','wandb'] \
    trainer.project_name='vagen' \
    trainer.experiment_name='mask_gae_mask_loss_turnwise_reward_bi_level' \
    trainer.n_gpus_per_node=8 \
    trainer.n_gpus_per_node=4 \
    trainer.nnodes=1 \
    trainer.save_freq=70 \
    trainer.test_freq=20 \
@@ -80,4 +80,7 @@ python3 -m vagen.trainer.main_ppo \
    trainer.val_before_train=True \
    trainer.val_generations_to_log_to_wandb=8 \
    rollout_manager.n_trajectory=1 \
    data.truncation=error \
    rollout_manager.truncation=error \

    2>&1 | tee mask_gae_mask_loss_turnwise_reward_bi_level.log
+2 −2
Original line number Diff line number Diff line
@@ -142,7 +142,7 @@ class QwenVLRolloutManger():
            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"
            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()
@@ -156,7 +156,7 @@ class QwenVLRolloutManger():
                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
            assert new_input_ids[b][seq_len-right_pad_tokens]==pad_token_id
            
            # shift right to fill the wholes
            holes_to_fill=1