已开启
添加dispatch_size参数校验 #907
添加dispatch_size参数校验 #907
已开启
Nurxat创建于 2月26日
1 个文件变更+78-0
Mmindspeed_rl/config_cls/validate_config.py+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")
叶蕴瑶
叶蕴瑶叶蕴瑶6月4日

这一行太长了,建议换行对齐参数,提高可读性。

likedislike
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_size448 # 若开启dapo动态采样,update的gbs=filter_groups_train_batch_size
371 rl_config.actor_update_dispatch_size = (449 rl_config.actor_update_dispatch_size = (