|
26 | 26 | from veomni.distributed.torch_parallelize import build_parallelize_model |
27 | 27 | from veomni.models.auto import build_foundation_model |
28 | 28 | from veomni.optim import build_lr_scheduler, build_optimizer |
| 29 | +from veomni.utils.seqlen_pos_transform_utils import prepare_fa_kwargs_from_position_ids |
29 | 30 |
|
30 | 31 | import verl.utils.torch_functional as verl_F |
31 | 32 | from verl.trainer.config import CheckpointConfig |
@@ -102,7 +103,7 @@ def __init__( |
102 | 103 | dp_size=dp_size, |
103 | 104 | dp_replicate_size=data_parallel_replicate_size, |
104 | 105 | dp_shard_size=data_parallel_shard_size, |
105 | | - ep_size=self.engine_config.expert_parallel_size, |
| 106 | + extra_parallel_sizes=(self.engine_config.expert_parallel_size,), |
106 | 107 | ulysses_size=self.engine_config.ulysses_parallel_size, |
107 | 108 | dp_mode=self.data_parallel_mode, |
108 | 109 | ) |
@@ -214,7 +215,9 @@ def _build_model_optimizer(self): |
214 | 215 | enable_mixed_precision=self.engine_config.mixed_precision, |
215 | 216 | enable_gradient_checkpointing=self.model_config.enable_gradient_checkpointing, |
216 | 217 | enable_fsdp_offload=self.engine_config.enable_fsdp_offload, |
217 | | - basic_modules=module._no_split_modules + self.engine_config.basic_modules, |
| 218 | + basic_modules=list( |
| 219 | + set(getattr(module, "_no_split_modules", None) or []) | set(self.engine_config.basic_modules) |
| 220 | + ), |
218 | 221 | enable_reentrant=self.engine_config.enable_reentrant, |
219 | 222 | enable_forward_prefetch=self.engine_config.forward_prefetch, |
220 | 223 | ) |
@@ -575,19 +578,60 @@ def __call__(self, batch: Sequence[dict[str, "torch.Tensor"]]) -> dict[str, "tor |
575 | 578 | return batch |
576 | 579 |
|
577 | 580 |
|
| 581 | +def _prepare_veomni_flash_attention_kwargs(position_ids: torch.Tensor) -> dict[str, torch.Tensor | int]: |
| 582 | + """Normalize packed position_ids layout and derive varlen FlashAttention kwargs. |
| 583 | +
|
| 584 | + Supported formats for use_remove_padding=true: |
| 585 | + - 2D: (1, total_nnz) - standard packed format |
| 586 | + - 3D: (rope_dim, 1, total_nnz) - VeRL mRoPE packed format |
| 587 | + """ |
| 588 | + if position_ids.dim() == 2: |
| 589 | + # (1, total_nnz) - standard packed format |
| 590 | + fa_position_ids = position_ids |
| 591 | + elif position_ids.dim() == 3: |
| 592 | + # (rope_dim, 1, total_nnz) - VeRL mRoPE packed format |
| 593 | + if position_ids.shape[1] == 1: |
| 594 | + fa_position_ids = position_ids[0] |
| 595 | + else: |
| 596 | + raise ValueError( |
| 597 | + f"Unsupported 3D position_ids shape: {tuple(position_ids.shape)}, expected (rope_dim, 1, total_nnz)" |
| 598 | + ) |
| 599 | + else: |
| 600 | + raise ValueError( |
| 601 | + f"Unsupported position_ids rank: {position_ids.dim()}, " |
| 602 | + f"expected 2 (1, total_nnz) or 3 (rope_dim, 1, total_nnz)" |
| 603 | + ) |
| 604 | + |
| 605 | + (cu_seq_lens_q, cu_seq_lens_k), (max_length_q, max_length_k) = prepare_fa_kwargs_from_position_ids(fa_position_ids) |
| 606 | + return { |
| 607 | + "cu_seq_lens_q": cu_seq_lens_q, |
| 608 | + "cu_seq_lens_k": cu_seq_lens_k, |
| 609 | + "max_length_q": max_length_q, |
| 610 | + "max_length_k": max_length_k, |
| 611 | + } |
| 612 | + |
| 613 | + |
578 | 614 | @EngineRegistry.register(model_type="language_model", backend=["veomni"], device=["cuda", "npu"]) |
579 | 615 | class VeOmniEngineWithLMHead(VeOmniEngine, FSDPEngineWithLMHead): |
580 | 616 | def prepare_model_inputs(self, micro_batch: TensorDict): |
581 | 617 | # TODO: Cannot work properly for qwen_vl ulysses |
582 | 618 | model_inputs, output_args = super().prepare_model_inputs(micro_batch) |
583 | 619 | input_ids_rmpad = model_inputs["input_ids"] |
| 620 | + sp_enabled = parallel_state.get_parallel_state().sp_enabled |
| 621 | + sp_shard_collator = OmniSequenceShardCollator() if sp_enabled else None |
| 622 | + |
584 | 623 | if self.module.config.model_type in VL_TYPE2INDEX.keys(): |
585 | 624 | image_mask = input_ids_rmpad == VL_TYPE2INDEX[self.module.config.model_type]["IMAGE_INPUT_INDEX"] |
586 | 625 | video_mask = input_ids_rmpad == VL_TYPE2INDEX[self.module.config.model_type]["VIDEO_INPUT_INDEX"] |
587 | 626 | model_inputs.update({"image_mask": image_mask, "video_mask": video_mask}) |
588 | 627 |
|
589 | | - if parallel_state.get_parallel_state().sp_enabled: |
590 | | - omni_sequence_shard_collator = OmniSequenceShardCollator() |
591 | | - omni_sequence_shard_collator(model_inputs) |
| 628 | + if sp_enabled: |
| 629 | + sp_shard_collator(model_inputs) |
| 630 | + |
| 631 | + use_remove_padding = tu.get_non_tensor_data(data=micro_batch, key="use_remove_padding", default=True) |
| 632 | + if use_remove_padding and model_inputs.get("position_ids", None) is not None: |
| 633 | + model_inputs.update(_prepare_veomni_flash_attention_kwargs(model_inputs["position_ids"])) |
| 634 | + if sp_enabled: |
| 635 | + model_inputs["position_ids"] = sp_shard_collator.sp_slice(model_inputs["position_ids"], dim=-1) |
592 | 636 |
|
593 | 637 | return model_inputs, output_args |
0 commit comments