已开启
添加dispatch_size参数校验 #907
Nurxat创建于 2月26日
添加dispatch_size参数校验 #907
已开启
共 1 个文件变更+78-0
| @@ -366,6 +366,84 @@ def validate_rl_args( | |||
| 366 | f"{int(rl_config.mini_batch_size/actor_data_parallel_size)})" | 366 | f"{int(rl_config.mini_batch_size/actor_data_parallel_size)})" |
| 367 | ) | 367 | ) |
| 368 | 368 | ||
| 369 | + # 若指定了自定义的dispatch_size,检查 dispatch_size 是否超过单个DP上生成的样本数量,避免出现DP上没有样本可用的情况 | ||
| 370 | + def _validate_dispatch_size(dispatch_size, global_batch_size, n_samples_per_prompt, data_parallel_size, component): | ||
| 371 | + if dispatch_size and dispatch_size > (global_batch_size * n_samples_per_prompt // data_parallel_size): | ||
| 372 | + raise ValueError( | ||
| 373 | + f"{component} {dispatch_size} cannot be greater than " | ||
| 374 | + f"the number of samples generated on each data parallel worker, which is " | ||
| 375 | + f"global_batch_size * n_samples_per_prompt // data_parallel_size = " | ||
| 376 | + f"{global_batch_size} * {n_samples_per_prompt} // {data_parallel_size} = " | ||
| 377 | + f"{global_batch_size * n_samples_per_prompt // data_parallel_size}" | ||
| 378 | + ) | ||
| 379 | + | ||
| 380 | + _validate_dispatch_size(rl_config.actor_rollout_dispatch_size, actor_config.global_batch_size, rl_config.n_samples_per_prompt, generate_config.data_parallel_size, "actor_rollout_dispatch_size") | ||
| 381 | + _validate_dispatch_size(rl_config.actor_logprob_dispatch_size, actor_config.global_batch_size, rl_config.n_samples_per_prompt, actor_data_parallel_size, "actor_logprob_dispatch_size") | ||
| 382 | + | ||
| 383 | + if ref_config: | ||
| 384 | + _validate_dispatch_size(rl_config.ref_dispatch_size, ref_config.global_batch_size, rl_config.n_samples_per_prompt, ref_data_parallel_size, "ref_dispatch_size") | ||
| 385 | + | ||
| 386 | + _validate_dispatch_size(rl_config.actor_update_dispatch_size, actor_config.global_batch_size, rl_config.n_samples_per_prompt, actor_data_parallel_size, "actor_update_dispatch_size") | ||
| 387 | + if rl_config.critic_resource: | ||
| 388 | + _validate_dispatch_size(rl_config.critic_update_dispatch_size, critic_config.global_batch_size, rl_config.n_samples_per_prompt, critic_data_parallel_size, "critic_update_dispatch_size") | ||
| 389 | + _validate_dispatch_size(rl_config.critic_value_dispatch_size, critic_config.global_batch_size, rl_config.n_samples_per_prompt, critic_data_parallel_size, "critic_value_dispatch_size") | ||
| 390 | + | ||
| 391 | + if rl_config.kl_dispatch_size and rl_config.kl_dispatch_size > (critic_config.global_batch_size * rl_config.n_samples_per_prompt): | ||
| 392 | + raise ValueError( | ||
| 393 | + f"kl_dispatch_size {rl_config.kl_dispatch_size} cannot be greater than " | ||
| 394 | + f"the number of samples generated on each data parallel worker, which is " | ||
| 395 | + f"global_batch_size * n_samples_per_prompt = " | ||
| 396 | + f"{critic_config.global_batch_size} * {rl_config.n_samples_per_prompt} = " | ||
| 397 | + f"{critic_config.global_batch_size * rl_config.n_samples_per_prompt}" | ||
| 398 | + ) | ||
| 399 | + | ||
| 400 | + if rl_config.reward_resource: | ||
| 401 | + reward_data_parallel_size = rl_config.reward_resource.num_npus // ( | ||
| 402 | + reward_config.tensor_model_parallel_size * | ||
| 403 | + reward_config.pipeline_model_parallel_size * | ||
| 404 | + reward_config.context_parallel_size | ||
| 405 | + ) | ||
| 406 | + _validate_dispatch_size(rl_config.reward_dispatch_size, reward_config.global_batch_size, rl_config.n_samples_per_prompt, reward_data_parallel_size, "reward_dispatch_size") | ||
| 407 | + | ||
| 408 | + if rl_config.dynamic_sampling_dispatch_size and rl_config.dynamic_sampling_dispatch_size > ( | ||
| 409 | + reward_config.global_batch_size * rl_config.n_samples_per_prompt // get_node_nums() | ||
| 410 | + ): | ||
| 411 | + raise ValueError( | ||
| 412 | + f"dynamic_sampling_dispatch_size {rl_config.dynamic_sampling_dispatch_size} cannot be greater than " | ||
| 413 | + f"the number of samples generated on each data parallel worker, which is " | ||
| 414 | + f"global_batch_size * n_samples_per_prompt // num_nodes_in_ray_cluster = " | ||
| 415 | + f"{reward_config.global_batch_size} * {rl_config.n_samples_per_prompt} // {get_node_nums()} = " | ||
| 416 | + f"{reward_config.global_batch_size * rl_config.n_samples_per_prompt // get_node_nums()}" | ||
| 417 | + ) | ||
| 418 | + | ||
| 419 | + if ( | ||
| 420 | + rl_config.reuse_image_embeds and rl_config.actor_image_embeds_dispatch_size | ||
| 421 | + and rl_config.actor_image_embeds_dispatch_size > ( | ||
| 422 | + actor_config.global_batch_size * rl_config.n_samples_per_prompt // actor_data_parallel_size | ||
| 423 | + ) | ||
| 424 | + ): | ||
| 425 | + raise ValueError( | ||
| 426 | + f"actor_image_embeds_dispatch_size {rl_config.actor_image_embeds_dispatch_size} cannot be greater than " | ||
| 427 | + f"the number of samples generated on each data parallel worker, which is " | ||
| 428 | + f"global_batch_size * n_samples_per_prompt // actor_data_parallel_size = " | ||
| 429 | + f"{actor_config.global_batch_size} * {rl_config.n_samples_per_prompt} // {actor_data_parallel_size} = " | ||
| 430 | + f"{actor_config.global_batch_size * rl_config.n_samples_per_prompt // actor_data_parallel_size}" | ||
| 431 | + ) | ||
| 432 | + | ||
| 433 | + if ( | ||
| 434 | + rl_config.reuse_image_embeds and rl_config.actor_image_embeds_dispatch_size | ||
| 435 | + and rl_config.actor_image_embeds_dispatch_size > ( | ||
| 436 | + vit_config.global_batch_size * rl_config.n_samples_per_prompt // vit_data_parallel_size | ||
| 437 | + ) | ||
| 438 | + ): | ||
| 439 | + raise ValueError( | ||
| 440 | + f"actor_image_embeds_dispatch_size {rl_config.actor_image_embeds_dispatch_size} cannot be greater than " | ||
| 441 | + f"the number of samples generated on each data parallel worker, which is " | ||
| 442 | + f"global_batch_size * n_samples_per_prompt // vit_data_parallel_size = " | ||
| 443 | + f"{vit_config.global_batch_size} * {rl_config.n_samples_per_prompt} // {vit_data_parallel_size} = " | ||
| 444 | + f"{vit_config.global_batch_size * rl_config.n_samples_per_prompt // vit_data_parallel_size}" | ||
| 445 | + ) | ||
| 446 | + | ||
| 369 | if rl_config.filter_groups_enable: | 447 | if rl_config.filter_groups_enable: |
| 370 | # 若开启dapo动态采样,update的gbs=filter_groups_train_batch_size | 448 | # 若开启dapo动态采样,update的gbs=filter_groups_train_batch_size |
| 371 | rl_config.actor_update_dispatch_size = ( | 449 | rl_config.actor_update_dispatch_size = ( |
这一行太长了,建议换行对齐参数,提高可读性。