Skip to content

Commit 1f96d92

Browse files
committed
[veomni] feat: bump veomni to v0.1.8
- Fix parallel_state init parameter (ep_size -> extra_parallel_sizes) and use set type for basic_modules - Automatically rewrites `flash_attention_2/3/4` to VeOmni SP-aware variants - Add _prepare_veomni_flash_attention_kwargs to precompute cu_seq_lens for packed sequences - Slice position_ids when sp_enabled
1 parent b9d71f9 commit 1f96d92

7 files changed

Lines changed: 76 additions & 10 deletions

File tree

‎.github/workflows/e2e_ppo_trainer_veomni_vllm.yml‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -117,7 +117,8 @@ jobs:
117117
run: |
118118
pip3 install -r requirements-test.txt
119119
pip3 install --no-deps -e .
120-
pip3 install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.4
120+
pip3 install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.8 --ignore-requires-python --no-deps
121+
pip3 install transformers==4.57.3
121122
- name: Prepare GSM8K dataset
122123
run: |
123124
ray stop --force

‎.github/workflows/e2e_sft_llm.yml‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,8 @@ jobs:
105105
pip3 install peft
106106
pip3 install -r requirements-test.txt
107107
pip3 install --no-deps -e .
108-
pip3 install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.4
108+
pip3 install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.8 --ignore-requires-python --no-deps
109+
pip3 install transformers==4.57.3
109110
- name: Prepare gsm8k dataset
110111
run: |
111112
ray stop --force

‎.github/workflows/e2e_sft_llm_ascend.yml‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -96,7 +96,8 @@ jobs:
9696
- name: Install the current repository
9797
run: |
9898
pip install --no-deps -e .
99-
pip install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.4
99+
pip install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.8 --ignore-requires-python --no-deps
100+
pip install transformers==4.57.3
100101
pip install pandas==2.3.3
101102
pip uninstall -y mbridge
102103
pip install git+https://github.com/ISEEKYAN/mbridge.git@89eb10

‎.github/workflows/e2e_sft_vlm.yml‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,8 @@ jobs:
105105
pip3 install peft
106106
pip3 install -r requirements-test.txt
107107
pip3 install --no-deps -e .
108-
pip3 install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.4
108+
pip3 install git+https://github.com/ByteDance-Seed/VeOmni.git@v0.1.8 --ignore-requires-python --no-deps
109+
pip3 install transformers==4.57.3
109110
- name: Prepare pokemon-gpt4o-captions dataset
110111
run: |
111112
ray stop --force

‎tests/special_e2e/sft/test_sft_engine_all.sh‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ echo "run with sp2 fsdp_size2 num_gpus8 fsdp_strategy fsdp2"
2525
BACKEND=fsdp SP_SIZE=2 FSDP_SIZE=2 NUM_GPUS=8 FSDP_STRATEGY=fsdp2 bash tests/special_e2e/sft/run_sft_engine.sh
2626

2727
# test with veomni
28-
echo "run with sp2 fsdp_size4 num_gpus8 fsdp_strategy fsdp2"
28+
echo "run with sp2 fsdp_size4 num_gpus8 fsdp_strategy fsdp2 backend veomni"
2929
BACKEND=veomni SP_SIZE=2 FSDP_SIZE=4 NUM_GPUS=8 FSDP_STRATEGY=fsdp2 bash tests/special_e2e/sft/run_sft_engine.sh
3030

3131

‎verl/workers/config/engine.py‎

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,8 @@
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
1414

15+
import logging
16+
import os
1517
import warnings
1618
from dataclasses import dataclass, field
1719
from typing import Any, Callable, Literal, Optional
@@ -37,6 +39,10 @@
3739
]
3840

3941

42+
logger = logging.getLogger(__name__)
43+
logger.setLevel(os.getenv("VERL_LOGGING_LEVEL", "INFO"))
44+
45+
4046
# TODO: rename to RouterReplayConfig after removing the legacy implementation
4147
@dataclass
4248
class EngineRouterReplayConfig(BaseConfig):
@@ -311,6 +317,8 @@ class VeOmniEngineConfig(EngineConfig):
311317
312318
"""
313319

320+
_mutable_fields = EngineConfig._mutable_fields | {"attn_implementation"}
321+
314322
wrap_policy: dict[str, Any] = field(default_factory=dict)
315323
offload_policy: bool = False
316324
reshard_after_forward: bool = True
@@ -342,6 +350,16 @@ def __post_init__(self):
342350
super().__post_init__()
343351
assert self.strategy in ["veomni"], f"strategy {self.strategy} not supported"
344352

353+
replacements = {
354+
"flash_attention_2": "veomni_flash_attention_2_with_sp",
355+
"flash_attention_3": "veomni_flash_attention_3_with_sp",
356+
"flash_attention_4": "veomni_flash_attention_4_with_sp",
357+
}
358+
if self.attn_implementation in replacements:
359+
new_impl = replacements[self.attn_implementation]
360+
logger.info(f"Replacing attn_implementation from '{self.attn_implementation}' to '{new_impl}'")
361+
self.attn_implementation = new_impl
362+
345363

346364
@dataclass
347365
class TorchtitanEngineConfig(EngineConfig):

‎verl/workers/engine/veomni/transformer_impl.py‎

Lines changed: 49 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from veomni.distributed.torch_parallelize import build_parallelize_model
2727
from veomni.models.auto import build_foundation_model
2828
from veomni.optim import build_lr_scheduler, build_optimizer
29+
from veomni.utils.seqlen_pos_transform_utils import prepare_fa_kwargs_from_position_ids
2930

3031
import verl.utils.torch_functional as verl_F
3132
from verl.trainer.config import CheckpointConfig
@@ -102,7 +103,7 @@ def __init__(
102103
dp_size=dp_size,
103104
dp_replicate_size=data_parallel_replicate_size,
104105
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,),
106107
ulysses_size=self.engine_config.ulysses_parallel_size,
107108
dp_mode=self.data_parallel_mode,
108109
)
@@ -214,7 +215,9 @@ def _build_model_optimizer(self):
214215
enable_mixed_precision=self.engine_config.mixed_precision,
215216
enable_gradient_checkpointing=self.model_config.enable_gradient_checkpointing,
216217
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+
),
218221
enable_reentrant=self.engine_config.enable_reentrant,
219222
enable_forward_prefetch=self.engine_config.forward_prefetch,
220223
)
@@ -575,19 +578,60 @@ def __call__(self, batch: Sequence[dict[str, "torch.Tensor"]]) -> dict[str, "tor
575578
return batch
576579

577580

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+
578614
@EngineRegistry.register(model_type="language_model", backend=["veomni"], device=["cuda", "npu"])
579615
class VeOmniEngineWithLMHead(VeOmniEngine, FSDPEngineWithLMHead):
580616
def prepare_model_inputs(self, micro_batch: TensorDict):
581617
# TODO: Cannot work properly for qwen_vl ulysses
582618
model_inputs, output_args = super().prepare_model_inputs(micro_batch)
583619
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+
584623
if self.module.config.model_type in VL_TYPE2INDEX.keys():
585624
image_mask = input_ids_rmpad == VL_TYPE2INDEX[self.module.config.model_type]["IMAGE_INPUT_INDEX"]
586625
video_mask = input_ids_rmpad == VL_TYPE2INDEX[self.module.config.model_type]["VIDEO_INPUT_INDEX"]
587626
model_inputs.update({"image_mask": image_mask, "video_mask": video_mask})
588627

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)
592636

593637
return model_inputs, output_args

0 commit comments

Comments
 (0)