diff --git a/examples/multimodal_dev/models/base.py b/examples/multimodal_dev/models/base.py index 6161e0b3ae0..d0597f540ef 100644 --- a/examples/multimodal_dev/models/base.py +++ b/examples/multimodal_dev/models/base.py @@ -45,10 +45,8 @@ def _cp_split_tensor(tensor, seq_dim, cp_size, cp_rank): class _NoCPGroup: - """Dummy size-1 process group used to bypass MRoPE's BSHD-style - zigzag of pre-computed THD freqs (Megatron-Core gap: - ``MultimodalRotaryEmbedding.forward`` lacks the ``not packed_seq`` - skip that plain ``RotaryEmbedding`` has). + """Dummy size-1 process group used to bypass BSHD-style CP slicing + for THD MRoPE call sites that do not pass ``packed_seq=True``. """ def size(self): diff --git a/examples/multimodal_dev/models/qwen35_vl/configuration.py b/examples/multimodal_dev/models/qwen35_vl/configuration.py index 81f73148314..8f65fe25a02 100644 --- a/examples/multimodal_dev/models/qwen35_vl/configuration.py +++ b/examples/multimodal_dev/models/qwen35_vl/configuration.py @@ -105,6 +105,13 @@ def get_qwen35_vl_vision_config( if num_layers_override is not None: num_layers = num_layers_override + vision_head_dim = vcfg["kv_channels"] + assert vision_head_dim % 4 == 0, ( + "Qwen3.5-VL vision RoPE expects the per-head dimension to split " + f"evenly across row/column frequencies, got {vision_head_dim}" + ) + vision_rope_axis_dim = vision_head_dim // 4 + return TransformerConfig( num_layers=num_layers, hidden_size=vcfg["hidden_size"], @@ -120,6 +127,12 @@ def get_qwen35_vl_vision_config( bias_activation_fusion=False, apply_query_key_layer_scaling=False, apply_rope_fusion=False, + # Vision RoPE is 2D row/column RoPE. Represent it as sectioned raw + # mRoPE with a zero temporal section so the fused mRoPE dispatcher can + # reuse the same Triton kernel when rope fusion is enabled. + mrope_section=[0, vision_rope_axis_dim, vision_rope_axis_dim], + mrope_interleaved=False, + rotary_interleaved=False, bf16=False, ) diff --git a/examples/multimodal_dev/models/qwen35_vl/specs.py b/examples/multimodal_dev/models/qwen35_vl/specs.py index 22fb4e616eb..eac6d543a04 100644 --- a/examples/multimodal_dev/models/qwen35_vl/specs.py +++ b/examples/multimodal_dev/models/qwen35_vl/specs.py @@ -17,6 +17,7 @@ from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_block import TransformerBlockSubmodules from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.utils import nvtx_range_pop, nvtx_range_push def _apply_rope_fp32(t, freqs, config, cu_seqlens=None, mscale=1.0, cp_group=None): @@ -25,35 +26,55 @@ def _apply_rope_fp32(t, freqs, config, cu_seqlens=None, mscale=1.0, cp_group=Non Mirrors ``Qwen3VLSelfAttention.apply_rotary_pos_emb_absolute`` in Megatron-Bridge with ``apply_rotary_pos_emb_in_fp32=True``. """ - from megatron.core import parallel_state - from megatron.core.models.common.embeddings.rope_utils import ( - _apply_rotary_pos_emb_bshd, - _apply_rotary_pos_emb_thd, - ) + from megatron.core.models.common.embeddings import rope_utils + from megatron.core.models.common.embeddings.rope_utils import apply_rotary_pos_emb orig_dtype = t.dtype - t_fp32 = t.float() - - if cu_seqlens is None: - out = _apply_rotary_pos_emb_bshd( - t_fp32, - freqs, - rotary_interleaved=config.rotary_interleaved, - multi_latent_attention=getattr(config, 'multi_latent_attention', False), - mscale=mscale, - ) - else: - if cp_group is None: - cp_group = parallel_state.get_context_parallel_group() - out = _apply_rotary_pos_emb_thd( - t_fp32, + if ( + cu_seqlens is not None + and getattr(config, "apply_rope_fusion", False) + and getattr(config, "mrope_section", None) is not None + and getattr(config, "rotary_interleaved", False) is False + and getattr(config, "multi_latent_attention", False) is False + and mscale == 1.0 + and t.dim() == 3 + and freqs.dim() == 4 + and freqs.shape[0] == 3 + and cp_group is not None + and rope_utils.fused_apply_mrope_thd is not None + and rope_utils.get_fused_mrope_thd_unavailable_reason is not None + ): + unavailable_reason = rope_utils.get_fused_mrope_thd_unavailable_reason( + t, cu_seqlens, freqs, rotary_interleaved=config.rotary_interleaved, - multi_latent_attention=getattr(config, 'multi_latent_attention', False), - mscale=mscale, - cp_group=cp_group, + cp_size=cp_group.size(), + cp_rank=cp_group.rank(), ) + if unavailable_reason is None: + return rope_utils.fused_apply_mrope_thd( + t, + cu_seqlens, + freqs, + config.mrope_section, + interleaved_mrope=config.mrope_interleaved, + rotary_interleaved=config.rotary_interleaved, + cp_size=cp_group.size(), + cp_rank=cp_group.rank(), + fp32_compute=True, + ) + + t_fp32 = t.float() + out = apply_rotary_pos_emb( + t_fp32, + freqs, + config=config, + cu_seqlens=cu_seqlens, + mscale=mscale, + cp_group=cp_group, + mla_rotary_interleaved=getattr(config, 'multi_latent_attention', False), + ) return out.to(orig_dtype) @@ -65,9 +86,19 @@ def _apply_rope_fp32_no_cp(t, freqs, config, cu_seqlens=None, mscale=1.0, cp_gro incorrectly split the vision seqlens. This wrapper substitutes a trivial group so the vision RoPE sees the full packed sequence. """ - return _apply_rope_fp32( - t, freqs, config, cu_seqlens, mscale, cp_group=_NO_CP_GROUP, - ) + range_name = "qwen35_vl.vision_encoder.rope_apply" + nvtx_range_push(range_name) + try: + return _apply_rope_fp32( + t, + freqs, + config, + cu_seqlens, + mscale, + cp_group=_NO_CP_GROUP, + ) + finally: + nvtx_range_pop(range_name) class Qwen35VLVisionSelfAttention(SelfAttention): diff --git a/examples/multimodal_dev/models/qwen35_vl/vision_encoder.py b/examples/multimodal_dev/models/qwen35_vl/vision_encoder.py index 8e8a6146a7f..1dc00221141 100644 --- a/examples/multimodal_dev/models/qwen35_vl/vision_encoder.py +++ b/examples/multimodal_dev/models/qwen35_vl/vision_encoder.py @@ -26,15 +26,10 @@ import torch.nn.functional as F from torch import Tensor -from megatron.core.models.common.vision_module.vision_module import ( - VisionModule, -) -from megatron.core.packed_seq_params import PackedSeqParams -from megatron.core.tensor_parallel.layers import ( - ColumnParallelLinear, - RowParallelLinear, -) from megatron.core.extensions.transformer_engine import TENorm +from megatron.core.models.common.vision_module.vision_module import VisionModule +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_block import TransformerBlock @@ -322,9 +317,7 @@ def __init__( # --- Transformer blocks --- if transformer_layer_spec is None: - from examples.multimodal_dev.models.qwen35_vl.specs import ( - get_qwen35_vl_vision_spec, - ) + from examples.multimodal_dev.models.qwen35_vl.specs import get_qwen35_vl_vision_spec transformer_layer_spec = get_qwen35_vl_vision_spec() self.decoder = TransformerBlock( @@ -455,7 +448,9 @@ def _compute_rotary_pos_emb(self, grid_thw: Tensor) -> Tensor: grid_thw: ``[num_images, 3]`` (T, H, W) per image. Returns: - ``[total_patches, head_dim // 2]`` raw RoPE frequencies. + Raw sectioned frequencies ``[3, 1, total_patches, head_dim // 2]`` + when ``config.mrope_section`` is set. Otherwise returns the legacy + ``[total_patches, head_dim // 2]`` row/column frequency tensor. """ merge = self.spatial_merge_size grid_thw_list = grid_thw.tolist() @@ -512,7 +507,27 @@ def _compute_rotary_pos_emb(self, grid_thw: Tensor) -> Tensor: embeddings = freq_table[pos_ids] embeddings = embeddings.flatten(1) - return embeddings + + mrope_section = getattr(self.config, "mrope_section", None) + if mrope_section is None: + return embeddings + + sec_t, sec_h, sec_w = (int(section) for section in mrope_section) + if sec_t != 0 or sec_h + sec_w != embeddings.shape[-1]: + raise ValueError( + "Qwen3.5-VL vision RoPE expects mrope_section " + f"[0, row_dim, col_dim] summing to {embeddings.shape[-1]}, " + f"got {mrope_section}" + ) + + raw_freqs = embeddings.new_zeros( + 3, 1, embeddings.shape[0], embeddings.shape[1], + ) + raw_freqs[1, 0, :, :sec_h] = embeddings[:, :sec_h] + raw_freqs[2, 0, :, sec_h : sec_h + sec_w] = embeddings[ + :, sec_h : sec_h + sec_w + ] + return raw_freqs # --------------------------------------------------------------- # PackedSeqParams for variable-length attention @@ -575,8 +590,9 @@ def forward( # 3. 2D Vision RoPE rot_freqs = self._compute_rotary_pos_emb(grid_thw) - emb = torch.cat((rot_freqs, rot_freqs), dim=-1) - rot_freqs_expanded = emb.unsqueeze(1).unsqueeze(1) + if getattr(self.config, "mrope_section", None) is None: + emb = torch.cat((rot_freqs, rot_freqs), dim=-1) + rot_freqs = emb.unsqueeze(1).unsqueeze(1) # 4. Transformer blocks with PackedSeqParams packed_seq_params = self._build_packed_seq_params(grid_thw) @@ -584,7 +600,7 @@ def forward( hidden_states = self.decoder( hidden_states=hidden_states, attention_mask=None, - rotary_pos_emb=rot_freqs_expanded, + rotary_pos_emb=rot_freqs, packed_seq_params=packed_seq_params, ) hidden_states = hidden_states.squeeze(1) diff --git a/examples/multimodal_dev/pretrain_multimodal.py b/examples/multimodal_dev/pretrain_multimodal.py index 9031339f360..3792f05f9f1 100644 --- a/examples/multimodal_dev/pretrain_multimodal.py +++ b/examples/multimodal_dev/pretrain_multimodal.py @@ -78,6 +78,7 @@ def model_provider( ) vision_config.bf16 = language_config.bf16 vision_config.fp16 = language_config.fp16 + vision_config.apply_rope_fusion = language_config.apply_rope_fusion if getattr(args, "recompute_vision", False): vision_config.recompute_granularity = "full" diff --git a/examples/multimodal_dev/scripts/run_qwen35_vl.sh b/examples/multimodal_dev/scripts/run_qwen35_vl.sh index 80a2a94671a..44c1fb5e2a5 100755 --- a/examples/multimodal_dev/scripts/run_qwen35_vl.sh +++ b/examples/multimodal_dev/scripts/run_qwen35_vl.sh @@ -12,8 +12,16 @@ # TP, EP, PP: parallelism sizes # MBS, GBS: micro/global batch sizes # NUM_LAYERS, NUM_EXPERTS: override for proxy testing +# MTP_NUM_LAYERS: number of MTP layers (default: 1, set 0 to disable) +# LINEAR_ATTENTION_FREQ: every Nth decoder layer uses standard attention (default: 4; set 1 to force all standard attention) +# DATASET_PROVIDER: cord_v2 (default) or mock +# TOKENIZER_TYPE: HuggingFaceTokenizer (default) or NullTokenizer +# NO_ROPE_FUSION: set to 1 to pass --no-rope-fusion for baseline profiling +# SAVE_CHECKPOINTS: set to 0 to skip checkpoint saves in short profiling runs # LAUNCHER: torchrun (default) or python +# TORCHRUN_PYTHON: Python executable for LAUNCHER=torchrun (default: python) # PROFILE: set to 1 to enable Nsight Systems profiling (default: 0) +# NVTX_RANGES: set to 1 to emit Megatron custom NVTX ranges when PROFILE=1 (default: 1) # PROFILE_STEP_START/PROFILE_STEP_END: profiled iteration window (default: 4-5) # example script: @@ -34,10 +42,13 @@ else NUM_NODES=${NNODES:-1} fi PROFILE=${PROFILE:-0} +NVTX_RANGES=${NVTX_RANGES:-1} PROFILE_STEP_START=${PROFILE_STEP_START:-4} PROFILE_STEP_END=${PROFILE_STEP_END:-5} PROFILE_RANKS=${PROFILE_RANKS:-0} LAUNCHER=${LAUNCHER:-torchrun} +TORCHRUN_PYTHON=${TORCHRUN_PYTHON:-python} +NO_ROPE_FUSION=${NO_ROPE_FUSION:-0} MODEL_VARIANT=${MODEL_VARIANT:-proxy} VISION_NUM_LAYERS=${VISION_NUM_LAYERS:-} @@ -45,6 +56,8 @@ VISION_NUM_LAYERS=${VISION_NUM_LAYERS:-} # Batch sizes MBS=${MBS:-2} GBS=${GBS:-16} +MTP_NUM_LAYERS=${MTP_NUM_LAYERS:-1} +LINEAR_ATTENTION_FREQ=${LINEAR_ATTENTION_FREQ:-4} # Parallelism TP=${TP:-1} @@ -186,6 +199,9 @@ USE_PACKED_SEQUENCE=${USE_PACKED_SEQUENCE:-0} if [ "$USE_PACKED_SEQUENCE" -eq 1 ]; then EXP_NAME+="_thd" fi +if [ "$NO_ROPE_FUSION" -eq 1 ]; then + EXP_NAME+="_no_rope_fusion" +fi MEGATRON_LM_PATH="${MEGATRON_LM_PATH:-$(cd "$(dirname "$0")/../../.." && pwd)}" ROOT_DIR="${ROOT_DIR:-${MEGATRON_LM_PATH}/local/}" @@ -241,13 +257,17 @@ TRAINING_ARGS=( --enable-experimental --manual-gc --manual-gc-interval 50 - --mtp-num-layers 1 - --mtp-loss-scaling-factor 0.1 --sft --use-flash-attn # --attention-backend flash --calculate-per-token-loss ) +if [ "$MTP_NUM_LAYERS" -gt 0 ]; then + TRAINING_ARGS+=( + --mtp-num-layers "$MTP_NUM_LAYERS" + --mtp-loss-scaling-factor 0.1 + ) +fi PROFILE_ARGS=() NSYS_CMD=() @@ -258,6 +278,9 @@ if [ "$PROFILE" = "1" ]; then --profile-step-end "$PROFILE_STEP_END" --profile-ranks "$PROFILE_RANKS" ) + if [ "$NVTX_RANGES" -eq 1 ]; then + PROFILE_ARGS+=( --nvtx-ranges ) + fi NSYS_OUTPUT_DIR="${CHECKPOINT_STORE_PATH}/nsys" mkdir -p "$NSYS_OUTPUT_DIR" @@ -274,12 +297,11 @@ if [ "$PROFILE" = "1" ]; then fi # --- Logging & Checkpointing --- +SAVE_CHECKPOINTS=${SAVE_CHECKPOINTS:-1} SAVE_INTERVAL=${SAVE_INTERVAL:-500} EVAL_AND_LOGGING_ARGS=( --log-interval 1 - --save-interval "$SAVE_INTERVAL" --eval-interval 500 - --save "$CHECKPOINT_STORE_PATH" --eval-iters 10 --tensorboard-dir "$TENSORBOARD_LOGS_PATH" --wandb-project "$WANDB_PROJECT" @@ -289,27 +311,44 @@ EVAL_AND_LOGGING_ARGS=( --log-timers-to-tensorboard --log-params-norm ) +if [ "$SAVE_CHECKPOINTS" -eq 1 ]; then + EVAL_AND_LOGGING_ARGS+=( + --save-interval "$SAVE_INTERVAL" + --save "$CHECKPOINT_STORE_PATH" + ) +fi # --- Tokenizer --- TOKENIZER_MODEL=${TOKENIZER_MODEL:-Qwen/Qwen3.5-397B-A17B} +TOKENIZER_TYPE=${TOKENIZER_TYPE:-HuggingFaceTokenizer} +VOCAB_SIZE=${VOCAB_SIZE:-248320} TOKENIZER_ARGS=( - --tokenizer-type HuggingFaceTokenizer - --tokenizer-model "$TOKENIZER_MODEL" + --tokenizer-type "$TOKENIZER_TYPE" ) +if [ "$TOKENIZER_TYPE" = "NullTokenizer" ]; then + TOKENIZER_ARGS+=( --vocab-size "$VOCAB_SIZE" ) +else + TOKENIZER_ARGS+=( --tokenizer-model "$TOKENIZER_MODEL" ) +fi # --- Multimodal-specific --- +DATASET_PROVIDER=${DATASET_PROVIDER:-cord_v2} +HF_PROCESSOR_PATH=${HF_PROCESSOR_PATH-Qwen/Qwen3.5-397B-A17B} +IMAGE_SEQ_LENGTH=${IMAGE_SEQ_LENGTH:-256} MULTIMODAL_ARGS=( --model-arch qwen35_vl --model-variant "$MODEL_VARIANT" - --dataset-provider cord_v2 - --hf-processor-path Qwen/Qwen3.5-397B-A17B + --dataset-provider "$DATASET_PROVIDER" --use-vanilla-collate-fn --image-token-id 248056 --image-size 224 --total-seq-length "$SEQ_LEN" - --image-seq-length 256 + --image-seq-length "$IMAGE_SEQ_LENGTH" --vision-num-layers "$VISION_NUM_LAYERS" ) +if [ -n "$HF_PROCESSOR_PATH" ]; then + MULTIMODAL_ARGS+=( --hf-processor-path "$HF_PROCESSOR_PATH" ) +fi if [ "$USE_PACKED_SEQUENCE" -eq 1 ]; then MULTIMODAL_ARGS+=( --use-packed-sequence ) @@ -341,7 +380,7 @@ GPT_MODEL_ARGS=( --attention-dropout 0.0 --hidden-dropout 0.0 --experimental-attention-variant gated_delta_net - --linear-attention-freq 4 + --linear-attention-freq "$LINEAR_ATTENTION_FREQ" --linear-conv-kernel-dim 4 --linear-key-head-dim 128 --linear-value-head-dim 128 @@ -350,6 +389,9 @@ GPT_MODEL_ARGS=( --make-vocab-size-divisible-by 485 --moe-router-force-load-balancing ) +if [ "$NO_ROPE_FUSION" -eq 1 ]; then + GPT_MODEL_ARGS+=( --no-rope-fusion ) +fi # --- Tied / untied embeddings --- # 0.8B, 2B, 4B use tied embeddings; all other variants untie them. @@ -457,9 +499,18 @@ echo " GPUs per node: $GPUS_PER_NODE" echo " Num nodes: $NUM_NODES" echo " TP=$TP EP=$EP PP=$PP CP=$CP" echo " MBS=$MBS GBS=$GBS" +echo " MTP layers: $MTP_NUM_LAYERS" +echo " Linear attn freq: $LINEAR_ATTENTION_FREQ" echo " Launcher: $LAUNCHER" +if [ "$LAUNCHER" = "torchrun" ]; then + echo " Torchrun py: $TORCHRUN_PYTHON" +fi echo " FSDP: $USE_FSDP" echo " PROFILE: $PROFILE" +echo " RoPE fusion: $([ "$NO_ROPE_FUSION" -eq 1 ] && echo off || echo on)" +echo " Dataset: $DATASET_PROVIDER" +echo " Tokenizer: $TOKENIZER_TYPE" +echo " Checkpoints: $([ "$SAVE_CHECKPOINTS" -eq 1 ] && echo on || echo off)" if [ -n "$CKPT_LOAD" ]; then echo " CKPT_LOAD: $CKPT_LOAD" echo " CKPT_FORMAT: ${CKPT_FORMAT:-auto}" @@ -468,13 +519,18 @@ fi if [ "$PROFILE" = "1" ]; then echo " Profile steps: ${PROFILE_STEP_START}-${PROFILE_STEP_END}" echo " Profile ranks: $PROFILE_RANKS" + echo " NVTX ranges: $([ "$NVTX_RANGES" -eq 1 ] && echo on || echo off)" fi echo "================================================================" if [ "$LAUNCHER" = "python" ]; then LAUNCH_CMD=( python $MEGATRON_LM_PATH/examples/multimodal_dev/pretrain_multimodal.py ) elif [ "$LAUNCHER" = "torchrun" ]; then - LAUNCH_CMD=( torchrun "${DISTRIBUTED_ARGS[@]}" $MEGATRON_LM_PATH/examples/multimodal_dev/pretrain_multimodal.py ) + LAUNCH_CMD=( + "$TORCHRUN_PYTHON" -m torch.distributed.run + "${DISTRIBUTED_ARGS[@]}" + $MEGATRON_LM_PATH/examples/multimodal_dev/pretrain_multimodal.py + ) else echo "Unsupported LAUNCHER=$LAUNCHER (expected torchrun or python)" >&2 exit 1 diff --git a/examples/multimodal_dev/tests/test_vision_rope_fusion.py b/examples/multimodal_dev/tests/test_vision_rope_fusion.py new file mode 100644 index 00000000000..602115af1fb --- /dev/null +++ b/examples/multimodal_dev/tests/test_vision_rope_fusion.py @@ -0,0 +1,182 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Tests for Qwen3.5-VL vision RoPE fusion dispatch.""" + +import os +import sys +from types import SimpleNamespace + +import pytest +import torch + +_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../..")) +if _REPO_ROOT not in sys.path: + sys.path.insert(0, _REPO_ROOT) + +import megatron.core.models.common.embeddings.rope_utils as rope_utils +from examples.multimodal_dev.models.qwen35_vl.configuration import get_qwen35_vl_vision_config +from examples.multimodal_dev.models.qwen35_vl.specs import _apply_rope_fp32_no_cp +from examples.multimodal_dev.models.qwen35_vl.vision_encoder import Qwen35VLVisionEncoder +from megatron.core.fusions.fused_mrope import is_fused_mrope_available, mrope_freqs_to_rotary_emb + + +class _FakeVisionRotaryEmbedding: + def __init__(self, axis_dim): + self.axis_dim = axis_dim + + def __call__(self, seqlen, device=None): + device = torch.device("cpu") if device is None else device + positions = torch.arange(seqlen, device=device, dtype=torch.float32)[:, None] + dims = torch.arange(self.axis_dim, device=device, dtype=torch.float32)[None, :] + return positions * 0.125 + dims * 0.01 + + +def test_vision_config_sets_2d_rope_as_sectioned_raw_mrope(): + config = get_qwen35_vl_vision_config(variant="0.8b") + + assert config.kv_channels == 64 + assert config.mrope_section == [0, 16, 16] + assert config.mrope_interleaved is False + assert config.rotary_interleaved is False + + +def test_vision_raw_mrope_freqs_match_legacy_materialized_rope(): + grid_thw = torch.tensor([[1, 4, 4], [2, 2, 2]], dtype=torch.long) + legacy_encoder = SimpleNamespace( + spatial_merge_size=2, + rot_pos_emb=_FakeVisionRotaryEmbedding(axis_dim=16), + config=SimpleNamespace(mrope_section=None), + ) + fused_encoder = SimpleNamespace( + spatial_merge_size=2, + rot_pos_emb=_FakeVisionRotaryEmbedding(axis_dim=16), + config=SimpleNamespace(mrope_section=[0, 16, 16]), + ) + + legacy_freqs = Qwen35VLVisionEncoder._compute_rotary_pos_emb(legacy_encoder, grid_thw) + raw_freqs = Qwen35VLVisionEncoder._compute_rotary_pos_emb(fused_encoder, grid_thw) + + expected = torch.cat((legacy_freqs, legacy_freqs), dim=-1).unsqueeze(1).unsqueeze(1) + converted = mrope_freqs_to_rotary_emb( + raw_freqs, + [0, 16, 16], + interleaved_mrope=False, + rotary_interleaved=False, + ) + + assert raw_freqs.shape == (3, 1, legacy_freqs.shape[0], legacy_freqs.shape[1]) + torch.testing.assert_close(converted, expected) + + +def test_vision_fp32_wrapper_dispatches_raw_freqs_to_fused_mrope_thd(monkeypatch): + calls = {} + + def fake_fused_apply_mrope_thd( + t, + cu_seqlens, + freqs, + mrope_section, + interleaved_mrope=False, + rotary_interleaved=False, + cp_size=1, + cp_rank=0, + fp32_compute=False, + ): + calls["t_shape"] = tuple(t.shape) + calls["t_dtype"] = t.dtype + calls["cu_seqlens"] = cu_seqlens.tolist() + calls["freqs_shape"] = tuple(freqs.shape) + calls["mrope_section"] = list(mrope_section) + calls["interleaved_mrope"] = interleaved_mrope + calls["rotary_interleaved"] = rotary_interleaved + calls["cp_size"] = cp_size + calls["cp_rank"] = cp_rank + calls["fp32_compute"] = fp32_compute + return t + 1.0 + + monkeypatch.setattr(rope_utils, "fused_apply_mrope_thd", fake_fused_apply_mrope_thd) + monkeypatch.setattr(rope_utils, "get_fused_mrope_thd_unavailable_reason", lambda *args, **kwargs: None) + + config = SimpleNamespace( + apply_rope_fusion=True, + mrope_section=[0, 2, 2], + mrope_interleaved=False, + rotary_interleaved=False, + multi_latent_attention=False, + ) + t = torch.zeros(6, 2, 8, dtype=torch.bfloat16) + freqs = torch.zeros(3, 1, 6, 4, dtype=torch.float32) + cu_seqlens = torch.tensor([0, 3, 6], dtype=torch.int32) + + out = _apply_rope_fp32_no_cp(t, freqs, config, cu_seqlens=cu_seqlens) + + assert out.dtype == torch.bfloat16 + torch.testing.assert_close(out, torch.ones_like(out)) + assert calls == { + "t_shape": (6, 2, 8), + "t_dtype": torch.bfloat16, + "cu_seqlens": [0, 3, 6], + "freqs_shape": (3, 1, 6, 4), + "mrope_section": [0, 2, 2], + "interleaved_mrope": False, + "rotary_interleaved": False, + "cp_size": 1, + "cp_rank": 0, + "fp32_compute": True, + } + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_vision_fused_rope_matches_unfused_forward_backward_cuda(): + generator = torch.Generator(device="cuda").manual_seed(1234) + total_tokens = 64 + num_heads = 2 + head_dim = 72 + half_rotary_dim = head_dim // 2 + section = [0, half_rotary_dim // 2, half_rotary_dim // 2] + cu_seqlens = torch.tensor([0, total_tokens], dtype=torch.int32, device="cuda") + + t_ref = torch.randn( + total_tokens, + num_heads, + head_dim, + device="cuda", + dtype=torch.bfloat16, + generator=generator, + requires_grad=True, + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + freqs = torch.randn( + 3, + 1, + total_tokens, + half_rotary_dim, + device="cuda", + dtype=torch.float32, + generator=generator, + ) + + ref_config = SimpleNamespace( + apply_rope_fusion=False, + mrope_section=section, + mrope_interleaved=False, + rotary_interleaved=False, + multi_latent_attention=False, + ) + fused_config = SimpleNamespace( + apply_rope_fusion=True, + mrope_section=section, + mrope_interleaved=False, + rotary_interleaved=False, + multi_latent_attention=False, + ) + + ref = _apply_rope_fp32_no_cp(t_ref, freqs, ref_config, cu_seqlens=cu_seqlens) + out = _apply_rope_fp32_no_cp(t_fused, freqs, fused_config, cu_seqlens=cu_seqlens) + torch.testing.assert_close(ref.float(), out.float(), rtol=2.0e-2, atol=5.0e-2) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), rtol=2.0e-2, atol=5.0e-2) diff --git a/megatron/core/fusions/fused_mla_yarn_rope_apply.py b/megatron/core/fusions/fused_mla_yarn_rope_apply.py index 6eed7581d03..4319af230ff 100644 --- a/megatron/core/fusions/fused_mla_yarn_rope_apply.py +++ b/megatron/core/fusions/fused_mla_yarn_rope_apply.py @@ -41,12 +41,17 @@ def _get_thd_token_idx(cu_seqlens, pid_m, seq_num, cp_rank, cp_size): last_cum_seqlen = cur_cum_seqlen seq_idx += 1 if cp_size > 1: - if token_idx < this_seq_len // 2: - token_idx = token_idx + cp_rank * this_seq_len // 2 + first_cp_seg = (this_seq_len + 1) // 2 + second_cp_seg = this_seq_len // 2 + if token_idx < first_cp_seg: + token_idx = token_idx + cp_rank * first_cp_seg else: - token_idx = (token_idx - this_seq_len // 2) + ( - 2 * cp_size - cp_rank - 1 - ) * this_seq_len // 2 + token_idx = ( + token_idx + - first_cp_seg + + cp_size * first_cp_seg + + (cp_size - cp_rank - 1) * second_cp_seg + ) return token_idx diff --git a/megatron/core/fusions/fused_mrope.py b/megatron/core/fusions/fused_mrope.py new file mode 100644 index 00000000000..6ebad4df933 --- /dev/null +++ b/megatron/core/fusions/fused_mrope.py @@ -0,0 +1,871 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Triton fused multimodal RoPE apply. + +The fused path consumes the raw three-axis mRoPE frequencies with shape +``[3, batch, seq, rotary_dim / 2]`` and applies the rotation directly to a BSHD +tensor. It supports both Qwen2-VL section-based mRoPE and Qwen3.5-VL +stride-3 interleaved mRoPE layouts. +""" + +from __future__ import annotations + +from typing import List, Optional +from unittest.mock import MagicMock + +import torch + +from megatron.core.utils import null_decorator + +try: + import triton + import triton.language as tl + + HAVE_TRITON = True +except ImportError: + HAVE_TRITON = False + +if not HAVE_TRITON: + triton = MagicMock() + triton.jit = null_decorator + tl = MagicMock() + + +def _smallest_power_of_2_at_least(x: int) -> int: + block = 1 + while block < x: + block *= 2 + return block + + +def _expected_interleaved_mrope_section(half_rotary_dim: int) -> tuple[int, int, int]: + return ((half_rotary_dim + 2) // 3, (half_rotary_dim + 1) // 3, half_rotary_dim // 3) + + +def _validate_mrope_section( + mrope_section: List[int], half_rotary_dim: int, interleaved_mrope: bool +) -> tuple[int, int, int]: + assert len(mrope_section) == 3, f"mrope_section must have length 3, got {mrope_section}" + + sec_t, sec_h, sec_w = (int(section) for section in mrope_section) + assert ( + min(sec_t, sec_h, sec_w) >= 0 + ), f"mrope_section values must be non-negative, got {mrope_section}" + assert half_rotary_dim > 0, "raw mRoPE rotary dim must be greater than 0" + assert ( + sec_t + sec_h + sec_w == half_rotary_dim + ), f"mrope_section {mrope_section} must sum to rotary_dim / 2 = {half_rotary_dim}" + if interleaved_mrope: + expected = _expected_interleaved_mrope_section(half_rotary_dim) + assert (sec_t, sec_h, sec_w) == expected, ( + f"interleaved mRoPE with rotary_dim / 2 = {half_rotary_dim} requires " + f"mrope_section {list(expected)}, got {mrope_section}" + ) + return sec_t, sec_h, sec_w + + +def _validate_mrope_inputs( + t: torch.Tensor, freqs: torch.Tensor, mrope_section: List[int], interleaved_mrope: bool +) -> tuple[int, int, int, int, int, int, int, int]: + assert t.dim() == 4, f"t must have shape [seq, batch, heads, head_dim], got {t.shape}" + assert freqs.dim() == 4, ( + "raw mRoPE freqs must have shape [3, batch, seq, rotary_dim / 2], " f"got {freqs.shape}" + ) + + seq, batch, heads, head_dim = t.shape + axes, freq_batch, freq_seq, half_rotary_dim = freqs.shape + assert axes == 3, f"raw mRoPE freqs first dimension must be 3, got {axes}" + assert ( + freq_batch == batch and freq_seq == seq + ), f"freqs shape {tuple(freqs.shape)} is incompatible with t shape {tuple(t.shape)}" + + sec_t, sec_h, sec_w = _validate_mrope_section(mrope_section, half_rotary_dim, interleaved_mrope) + + rotary_dim = half_rotary_dim * 2 + assert ( + rotary_dim <= head_dim + ), f"raw mRoPE rotary dim {rotary_dim} exceeds input head dim {head_dim}" + return seq, batch, heads, head_dim, half_rotary_dim, sec_t, sec_h, sec_w + + +def _validate_mrope_thd_inputs( + t: torch.Tensor, + cu_seqlens: torch.Tensor, + freqs: torch.Tensor, + mrope_section: List[int], + interleaved_mrope: bool, + cp_size: int, +) -> tuple[int, int, int, int, int, int, int]: + assert t.dim() == 3, f"t must have shape [tokens, heads, head_dim], got {t.shape}" + assert freqs.dim() == 4, ( + "raw mRoPE freqs must have shape [3, 1, total_seqlen, rotary_dim / 2], " + f"got {freqs.shape}" + ) + assert cu_seqlens.dim() == 1, f"cu_seqlens must be 1D, got {cu_seqlens.shape}" + + tokens, heads, head_dim = t.shape + axes, freq_batch, freq_seq, half_rotary_dim = freqs.shape + assert axes == 3, f"raw mRoPE freqs first dimension must be 3, got {axes}" + assert freq_batch == 1, ( + "raw mRoPE THD freqs must have singleton batch dimension, " f"got {freqs.shape}" + ) + assert freq_seq == tokens * cp_size, ( + "raw mRoPE THD freqs sequence length must match local tokens times cp_size, " + f"got freqs.shape[2]={freq_seq}, tokens={tokens}, cp_size={cp_size}" + ) + + sec_t, sec_h, sec_w = _validate_mrope_section(mrope_section, half_rotary_dim, interleaved_mrope) + rotary_dim = half_rotary_dim * 2 + assert ( + rotary_dim <= head_dim + ), f"raw mRoPE rotary dim {rotary_dim} exceeds input head dim {head_dim}" + return tokens, heads, head_dim, half_rotary_dim, sec_t, sec_h, sec_w + + +def get_fused_mrope_unavailable_reason( + t: Optional[torch.Tensor] = None, + freqs: Optional[torch.Tensor] = None, + rotary_interleaved: bool = False, +) -> Optional[str]: + """Return why fused mRoPE cannot run, or None when it is launchable.""" + if not HAVE_TRITON: + return "Triton is not available" + if rotary_interleaved: + return "rotary_interleaved=True is not supported" + if t is None or freqs is None: + return None + if not t.is_cuda or not freqs.is_cuda: + return "Triton fused mRoPE requires CUDA tensors" + if t.device != freqs.device: + return ( + "Triton fused mRoPE requires t and freqs on the same device, " + f"got {t.device} and {freqs.device}" + ) + if freqs.dtype != torch.float32: + return f"raw mRoPE freqs must be float32, got {freqs.dtype}" + if t.dtype not in (torch.float16, torch.bfloat16, torch.float32): + return f"input dtype {t.dtype} is not supported" + if t.stride(-1) != 1: + return f"input head dimension must be contiguous, got stride {t.stride()}" + try: + capability = torch.cuda.get_device_capability(t.device) + except RuntimeError as exc: + return f"could not query CUDA device capability: {exc}" + if capability < (7, 0): + return f"requires CUDA compute capability >= 7.0, got {capability[0]}.{capability[1]}" + if t.dtype == torch.bfloat16 and capability < (8, 0): + return ( + "requires CUDA compute capability >= 8.0 for bfloat16 inputs, " + f"got {capability[0]}.{capability[1]}" + ) + return None + + +def get_fused_mrope_thd_unavailable_reason( + t: Optional[torch.Tensor] = None, + cu_seqlens: Optional[torch.Tensor] = None, + freqs: Optional[torch.Tensor] = None, + rotary_interleaved: bool = False, + cp_size: int = 1, + cp_rank: int = 0, +) -> Optional[str]: + """Return why fused THD mRoPE cannot run, or None when it is launchable.""" + if not HAVE_TRITON: + return "Triton is not available" + if rotary_interleaved: + return "rotary_interleaved=True is not supported" + if cp_size < 1: + return f"cp_size must be positive, got {cp_size}" + if cp_rank < 0 or cp_rank >= cp_size: + return f"cp_rank must be in [0, {cp_size}), got {cp_rank}" + if t is None or cu_seqlens is None or freqs is None: + return None + if t.dim() != 3: + return ( + f"THD fused mRoPE expects t with shape [tokens, heads, head_dim], got {tuple(t.shape)}" + ) + if freqs.dim() != 4: + return ( + "raw mRoPE THD freqs must have shape [3, 1, total_seqlen, rotary_dim / 2], " + f"got {tuple(freqs.shape)}" + ) + if cu_seqlens.dim() != 1: + return f"cu_seqlens must be 1D, got {tuple(cu_seqlens.shape)}" + if not t.is_cuda or not freqs.is_cuda or not cu_seqlens.is_cuda: + return "Triton fused THD mRoPE requires CUDA tensors" + if t.device != freqs.device or t.device != cu_seqlens.device: + return ( + "Triton fused THD mRoPE requires t, freqs, and cu_seqlens on the same device, " + f"got {t.device}, {freqs.device}, and {cu_seqlens.device}" + ) + if freqs.dtype != torch.float32: + return f"raw mRoPE freqs must be float32, got {freqs.dtype}" + if t.dtype not in (torch.float16, torch.bfloat16, torch.float32): + return f"input dtype {t.dtype} is not supported" + if cu_seqlens.dtype not in (torch.int32, torch.int64): + return f"cu_seqlens dtype {cu_seqlens.dtype} is not supported" + if t.stride(-1) != 1: + return f"input head dimension must be contiguous, got stride {t.stride()}" + if freqs.shape[0] != 3 or freqs.shape[1] != 1: + return ( + "raw mRoPE THD freqs must have shape [3, 1, total_seqlen, rotary_dim / 2], " + f"got {tuple(freqs.shape)}" + ) + if cp_size > 1 and freqs.shape[2] % cp_size != 0: + return ( + "raw mRoPE THD freqs sequence length must be divisible by context parallel size, " + f"got freqs.shape[2]={freqs.shape[2]}, cp_size={cp_size}" + ) + if cp_size > 1: + # Guard: each packed sub-sequence length must satisfy seqlen % cp_size == 0. + seq_bounds = cu_seqlens.tolist() + for seq_start, seq_end in zip(seq_bounds[:-1], seq_bounds[1:]): + if (seq_end - seq_start) % cp_size != 0: + return ( + "each packed THD sub-sequence length must be divisible by context " + f"parallel size, got sub-sequence length {seq_end - seq_start} " + f"with cp_size={cp_size}" + ) + if freqs.shape[2] != t.shape[0] * cp_size: + return ( + "raw mRoPE THD freqs sequence length must match local tokens times cp_size, " + f"got freqs.shape[2]={freqs.shape[2]}, tokens={t.shape[0]}, cp_size={cp_size}" + ) + try: + capability = torch.cuda.get_device_capability(t.device) + except RuntimeError as exc: + return f"could not query CUDA device capability: {exc}" + if capability < (7, 0): + return f"requires CUDA compute capability >= 7.0, got {capability[0]}.{capability[1]}" + if t.dtype == torch.bfloat16 and capability < (8, 0): + return ( + "requires CUDA compute capability >= 8.0 for bfloat16 inputs, " + f"got {capability[0]}.{capability[1]}" + ) + return None + + +def can_launch_fused_mrope( + t: Optional[torch.Tensor] = None, + freqs: Optional[torch.Tensor] = None, + rotary_interleaved: bool = False, +) -> bool: + """Return whether the Triton fused mRoPE kernel can be launched.""" + return get_fused_mrope_unavailable_reason(t, freqs, rotary_interleaved) is None + + +def can_launch_fused_mrope_thd( + t: Optional[torch.Tensor] = None, + cu_seqlens: Optional[torch.Tensor] = None, + freqs: Optional[torch.Tensor] = None, + rotary_interleaved: bool = False, + cp_size: int = 1, + cp_rank: int = 0, +) -> bool: + """Return whether the Triton fused THD mRoPE kernel can be launched.""" + return ( + get_fused_mrope_thd_unavailable_reason( + t, + cu_seqlens, + freqs, + rotary_interleaved=rotary_interleaved, + cp_size=cp_size, + cp_rank=cp_rank, + ) + is None + ) + + +def mrope_freqs_to_rotary_emb( + freqs: torch.Tensor, + mrope_section: List[int], + interleaved_mrope: bool = False, + rotary_interleaved: bool = False, +) -> torch.Tensor: + """Convert raw mRoPE freqs to the unfused RoPE embedding layout. + + Args: + freqs: Raw mRoPE frequencies with shape ``[3, batch, seq, rotary_dim / 2]``. + mrope_section: Temporal, height, and width channel sections. + interleaved_mrope: Use Qwen3.5-VL stride-3 T/H/W layout when True. Use + Qwen2-VL section layout when False. + rotary_interleaved: Use adjacent-pair RoPE layout when True. This is + available for reference conversion; fused Triton currently supports + split-half layout only. + + Returns: + Tensor with shape ``[seq, batch, 1, rotary_dim]``. + """ + assert freqs.dim() == 4, ( + "raw mRoPE freqs must have shape [3, batch, seq, rotary_dim / 2], " f"got {freqs.shape}" + ) + assert freqs.size(0) == 3, f"raw mRoPE freqs first dimension must be 3, got {freqs.size(0)}" + assert len(mrope_section) == 3, f"mrope_section must have length 3, got {mrope_section}" + + half_rotary_dim = freqs.size(-1) + sec_t, sec_h, sec_w = _validate_mrope_section(mrope_section, half_rotary_dim, interleaved_mrope) + + if interleaved_mrope: + freqs_out = freqs[0].clone() + for dim_idx, offset in enumerate((1, 2), start=1): + length = int(mrope_section[dim_idx]) * 3 + idx = slice(offset, length, 3) + freqs_out[..., idx] = freqs[dim_idx, ..., idx] + if rotary_interleaved: + batch = freqs_out.shape[0] + emb = torch.stack( + (freqs_out.reshape(batch, -1, 1), freqs_out.reshape(batch, -1, 1)), dim=-1 + ) + emb = emb.view(batch, freqs_out.shape[1], -1) + else: + emb = torch.cat((freqs_out, freqs_out), dim=-1) + else: + if rotary_interleaved: + batch = freqs.shape[1] + emb = torch.stack( + (freqs.reshape(3, batch, -1, 1), freqs.reshape(3, batch, -1, 1)), dim=-1 + ).view(3, batch, freqs.shape[2], -1) + mrope_section_doubled = list(mrope_section) * 2 + emb = torch.cat( + [chunk[i % 3] for i, chunk in enumerate(emb.split(mrope_section_doubled, dim=-1))], + dim=-1, + ) + else: + freqs_out = torch.empty_like(freqs[0]) + freqs_out[..., :sec_t] = freqs[0, ..., :sec_t] + freqs_out[..., sec_t : sec_t + sec_h] = freqs[1, ..., sec_t : sec_t + sec_h] + freqs_out[..., sec_t + sec_h :] = freqs[2, ..., sec_t + sec_h :] + emb = torch.cat((freqs_out, freqs_out), dim=-1) + return emb[..., None, :].transpose(0, 1).contiguous() + + +@triton.jit +def _mrope_axis( + k, + SEC_T: tl.constexpr, + SEC_H: tl.constexpr, + SEC_W: tl.constexpr, + INTERLEAVED_MROPE: tl.constexpr, +): + if INTERLEAVED_MROPE: + rem = k % 3 + section_idx = k // 3 + is_h = (rem == 1) & (section_idx < SEC_H) + is_w = (rem == 2) & (section_idx < SEC_W) + return tl.where(is_h, 1, tl.where(is_w, 2, 0)) + + is_h = (k >= SEC_T) & (k < (SEC_T + SEC_H)) + is_w = k >= (SEC_T + SEC_H) + return tl.where(is_h, 1, tl.where(is_w, 2, 0)) + + +@triton.jit +def _fused_mrope_kernel( + T, + FREQS, + OUT, + t_s_seq, + t_s_batch, + t_s_head, + t_s_dim, + f_s_axis, + f_s_batch, + f_s_seq, + f_s_dim, + o_s_seq, + o_s_batch, + o_s_head, + o_s_dim, + HEAD_DIM: tl.constexpr, + HALF_ROTARY_DIM: tl.constexpr, + PASS_DIM: tl.constexpr, + SEC_T: tl.constexpr, + SEC_H: tl.constexpr, + SEC_W: tl.constexpr, + INTERLEAVED_MROPE: tl.constexpr, + ROTARY_INTERLEAVED: tl.constexpr, + INVERSE: tl.constexpr, + BLOCK_HALF: tl.constexpr, + BLOCK_PASS: tl.constexpr, +): + seq_idx = tl.program_id(0) + batch_idx = tl.program_id(1) + head_idx = tl.program_id(2) + + k = tl.arange(0, BLOCK_HALF) + mask = k < HALF_ROTARY_DIM + + axis = _mrope_axis(k, SEC_T, SEC_H, SEC_W, INTERLEAVED_MROPE) + + freqs_offset = axis * f_s_axis + batch_idx * f_s_batch + seq_idx * f_s_seq + k * f_s_dim + freqs = tl.load(FREQS + freqs_offset, mask=mask, other=0.0) + # Match PyTorch pointwise dtype semantics: cast cos/sin before the multiply. + cos_v = tl.cos(freqs).to(OUT.dtype.element_ty) + sin_v = tl.sin(freqs).to(OUT.dtype.element_ty) + if INVERSE: + sin_v = -sin_v + + t_base = T + seq_idx * t_s_seq + batch_idx * t_s_batch + head_idx * t_s_head + out_base = OUT + seq_idx * o_s_seq + batch_idx * o_s_batch + head_idx * o_s_head + + if ROTARY_INTERLEAVED: + lo_offset = (2 * k) * t_s_dim + hi_offset = (2 * k + 1) * t_s_dim + out_lo_offset = (2 * k) * o_s_dim + out_hi_offset = (2 * k + 1) * o_s_dim + else: + lo_offset = k * t_s_dim + hi_offset = (k + HALF_ROTARY_DIM) * t_s_dim + out_lo_offset = k * o_s_dim + out_hi_offset = (k + HALF_ROTARY_DIM) * o_s_dim + + t_lo = tl.load(t_base + lo_offset, mask=mask, other=0.0).to(OUT.dtype.element_ty) + t_hi = tl.load(t_base + hi_offset, mask=mask, other=0.0).to(OUT.dtype.element_ty) + + lo_cos = (t_lo * cos_v).to(OUT.dtype.element_ty) + hi_sin = (t_hi * sin_v).to(OUT.dtype.element_ty) + hi_cos = (t_hi * cos_v).to(OUT.dtype.element_ty) + lo_sin = (t_lo * sin_v).to(OUT.dtype.element_ty) + + out_lo = (lo_cos - hi_sin).to(OUT.dtype.element_ty) + out_hi = (hi_cos + lo_sin).to(OUT.dtype.element_ty) + + tl.store(out_base + out_lo_offset, out_lo, mask=mask) + tl.store(out_base + out_hi_offset, out_hi, mask=mask) + + if PASS_DIM > 0: + pass_idx = tl.arange(0, BLOCK_PASS) + pass_mask = pass_idx < PASS_DIM + src_dim = 2 * HALF_ROTARY_DIM + pass_idx + pass_values = tl.load(t_base + src_dim * t_s_dim, mask=pass_mask, other=0.0) + tl.store(out_base + src_dim * o_s_dim, pass_values, mask=pass_mask) + + +@triton.jit +def _fused_mrope_thd_kernel( + T, + CU_SEQLENS, + FREQS, + OUT, + t_s_token, + t_s_head, + t_s_dim, + cu_s_idx, + f_s_axis, + f_s_seq, + f_s_dim, + o_s_token, + o_s_head, + o_s_dim, + NUM_SEQS, + HEAD_DIM: tl.constexpr, + HALF_ROTARY_DIM: tl.constexpr, + PASS_DIM: tl.constexpr, + SEC_T: tl.constexpr, + SEC_H: tl.constexpr, + SEC_W: tl.constexpr, + INTERLEAVED_MROPE: tl.constexpr, + ROTARY_INTERLEAVED: tl.constexpr, + INVERSE: tl.constexpr, + CP_SIZE: tl.constexpr, + CP_RANK: tl.constexpr, + FP32_COMPUTE: tl.constexpr, + BLOCK_HALF: tl.constexpr, + BLOCK_PASS: tl.constexpr, +): + token_idx = tl.program_id(0) + head_idx = tl.program_id(1) + + freq_seq_idx = token_idx + seq_i = 0 + while seq_i < NUM_SEQS: + global_start = tl.load(CU_SEQLENS + seq_i * cu_s_idx) + global_end = tl.load(CU_SEQLENS + (seq_i + 1) * cu_s_idx) + local_start = global_start // CP_SIZE + local_end = global_end // CP_SIZE + in_seq = (token_idx >= local_start) & (token_idx < local_end) + local_offset = token_idx - local_start + + if CP_SIZE > 1: + local_seq_len = local_end - local_start + first_cp_seg = (local_seq_len + 1) // 2 + second_cp_seg = local_seq_len // 2 + first_freq_idx = global_start + CP_RANK * first_cp_seg + local_offset + second_freq_idx = ( + global_end - (CP_RANK + 1) * second_cp_seg + (local_offset - first_cp_seg) + ) + seq_freq_idx = tl.where(local_offset < first_cp_seg, first_freq_idx, second_freq_idx) + else: + seq_freq_idx = global_start + local_offset + + freq_seq_idx = tl.where(in_seq, seq_freq_idx, freq_seq_idx) + seq_i += 1 + + k = tl.arange(0, BLOCK_HALF) + mask = k < HALF_ROTARY_DIM + axis = _mrope_axis(k, SEC_T, SEC_H, SEC_W, INTERLEAVED_MROPE) + + freqs_offset = axis * f_s_axis + freq_seq_idx * f_s_seq + k * f_s_dim + freqs = tl.load(FREQS + freqs_offset, mask=mask, other=0.0) + if FP32_COMPUTE: + cos_v = tl.cos(freqs) + sin_v = tl.sin(freqs) + else: + cos_v = tl.cos(freqs).to(OUT.dtype.element_ty) + sin_v = tl.sin(freqs).to(OUT.dtype.element_ty) + if INVERSE: + sin_v = -sin_v + + t_base = T + token_idx * t_s_token + head_idx * t_s_head + out_base = OUT + token_idx * o_s_token + head_idx * o_s_head + + if ROTARY_INTERLEAVED: + lo_offset = (2 * k) * t_s_dim + hi_offset = (2 * k + 1) * t_s_dim + out_lo_offset = (2 * k) * o_s_dim + out_hi_offset = (2 * k + 1) * o_s_dim + else: + lo_offset = k * t_s_dim + hi_offset = (k + HALF_ROTARY_DIM) * t_s_dim + out_lo_offset = k * o_s_dim + out_hi_offset = (k + HALF_ROTARY_DIM) * o_s_dim + + if FP32_COMPUTE: + t_lo = tl.load(t_base + lo_offset, mask=mask, other=0.0).to(tl.float32) + t_hi = tl.load(t_base + hi_offset, mask=mask, other=0.0).to(tl.float32) + else: + t_lo = tl.load(t_base + lo_offset, mask=mask, other=0.0).to(OUT.dtype.element_ty) + t_hi = tl.load(t_base + hi_offset, mask=mask, other=0.0).to(OUT.dtype.element_ty) + + if FP32_COMPUTE: + lo_cos = t_lo * cos_v + hi_sin = t_hi * sin_v + hi_cos = t_hi * cos_v + lo_sin = t_lo * sin_v + else: + lo_cos = (t_lo * cos_v).to(OUT.dtype.element_ty) + hi_sin = (t_hi * sin_v).to(OUT.dtype.element_ty) + hi_cos = (t_hi * cos_v).to(OUT.dtype.element_ty) + lo_sin = (t_lo * sin_v).to(OUT.dtype.element_ty) + + if FP32_COMPUTE: + out_lo = lo_cos - hi_sin + out_hi = hi_cos + lo_sin + else: + out_lo = (lo_cos - hi_sin).to(OUT.dtype.element_ty) + out_hi = (hi_cos + lo_sin).to(OUT.dtype.element_ty) + + tl.store(out_base + out_lo_offset, out_lo, mask=mask) + tl.store(out_base + out_hi_offset, out_hi, mask=mask) + + if PASS_DIM > 0: + pass_idx = tl.arange(0, BLOCK_PASS) + pass_mask = pass_idx < PASS_DIM + src_dim = 2 * HALF_ROTARY_DIM + pass_idx + pass_values = tl.load(t_base + src_dim * t_s_dim, mask=pass_mask, other=0.0) + tl.store(out_base + src_dim * o_s_dim, pass_values, mask=pass_mask) + + +def _launch_fused_mrope( + t: torch.Tensor, + freqs: torch.Tensor, + mrope_section: List[int], + interleaved_mrope: bool, + rotary_interleaved: bool, + inverse: bool, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + unavailable_reason = get_fused_mrope_unavailable_reason(t, freqs, rotary_interleaved) + assert unavailable_reason is None, unavailable_reason + + seq, batch, heads, head_dim, half_rotary_dim, sec_t, sec_h, sec_w = _validate_mrope_inputs( + t, freqs, mrope_section, interleaved_mrope + ) + + if out is None: + out = torch.empty_like(t) + else: + assert out.shape == t.shape and out.dtype == t.dtype + assert ( + out.stride(-1) == 1 + ), f"fused mRoPE requires output contiguous head dimension, got {out.stride()}" + + block_half = _smallest_power_of_2_at_least(half_rotary_dim) + pass_dim = head_dim - (2 * half_rotary_dim) + block_pass = _smallest_power_of_2_at_least(max(pass_dim, 1)) + + grid = (seq, batch, heads) + _fused_mrope_kernel[grid]( + t, + freqs, + out, + t.stride(0), + t.stride(1), + t.stride(2), + t.stride(3), + freqs.stride(0), + freqs.stride(1), + freqs.stride(2), + freqs.stride(3), + out.stride(0), + out.stride(1), + out.stride(2), + out.stride(3), + HEAD_DIM=head_dim, + HALF_ROTARY_DIM=half_rotary_dim, + PASS_DIM=pass_dim, + SEC_T=sec_t, + SEC_H=sec_h, + SEC_W=sec_w, + INTERLEAVED_MROPE=interleaved_mrope, + ROTARY_INTERLEAVED=rotary_interleaved, + INVERSE=inverse, + BLOCK_HALF=block_half, + BLOCK_PASS=block_pass, + num_warps=4, + ) + return out + + +def _launch_fused_mrope_thd( + t: torch.Tensor, + cu_seqlens: torch.Tensor, + freqs: torch.Tensor, + mrope_section: List[int], + interleaved_mrope: bool, + rotary_interleaved: bool, + inverse: bool, + cp_size: int, + cp_rank: int, + fp32_compute: bool = False, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + unavailable_reason = get_fused_mrope_thd_unavailable_reason( + t, + cu_seqlens, + freqs, + rotary_interleaved=rotary_interleaved, + cp_size=cp_size, + cp_rank=cp_rank, + ) + assert unavailable_reason is None, unavailable_reason + + tokens, heads, head_dim, half_rotary_dim, sec_t, sec_h, sec_w = _validate_mrope_thd_inputs( + t, cu_seqlens, freqs, mrope_section, interleaved_mrope, cp_size + ) + + if out is None: + out = torch.empty_like(t) + else: + assert out.shape == t.shape and out.dtype == t.dtype + assert ( + out.stride(-1) == 1 + ), f"fused THD mRoPE requires output contiguous head dimension, got {out.stride()}" + + block_half = _smallest_power_of_2_at_least(half_rotary_dim) + pass_dim = head_dim - (2 * half_rotary_dim) + block_pass = _smallest_power_of_2_at_least(max(pass_dim, 1)) + num_seqs = cu_seqlens.numel() - 1 + + grid = (tokens, heads) + _fused_mrope_thd_kernel[grid]( + t, + cu_seqlens, + freqs, + out, + t.stride(0), + t.stride(1), + t.stride(2), + cu_seqlens.stride(0), + freqs.stride(0), + freqs.stride(2), + freqs.stride(3), + out.stride(0), + out.stride(1), + out.stride(2), + num_seqs, + HEAD_DIM=head_dim, + HALF_ROTARY_DIM=half_rotary_dim, + PASS_DIM=pass_dim, + SEC_T=sec_t, + SEC_H=sec_h, + SEC_W=sec_w, + INTERLEAVED_MROPE=interleaved_mrope, + ROTARY_INTERLEAVED=rotary_interleaved, + INVERSE=inverse, + CP_SIZE=cp_size, + CP_RANK=cp_rank, + FP32_COMPUTE=fp32_compute, + BLOCK_HALF=block_half, + BLOCK_PASS=block_pass, + num_warps=4, + ) + return out + + +class _FusedMRoPE(torch.autograd.Function): + """Autograd wrapper for fused mRoPE. + + The raw frequency table is generated from position IDs and inverse frequencies, + so gradients are only propagated to the rotated tensor. + """ + + @staticmethod + def forward(ctx, t, freqs, mrope_section, interleaved_mrope, rotary_interleaved): + assert not freqs.requires_grad, "fused mRoPE expects non-gradient raw frequency tensors" + ctx.mrope_section = tuple(int(section) for section in mrope_section) + ctx.interleaved_mrope = bool(interleaved_mrope) + ctx.rotary_interleaved = bool(rotary_interleaved) + ctx.save_for_backward(freqs) + return _launch_fused_mrope( + t, + freqs, + ctx.mrope_section, + ctx.interleaved_mrope, + ctx.rotary_interleaved, + inverse=False, + ) + + @staticmethod + def backward(ctx, grad_output): + (freqs,) = ctx.saved_tensors + grad_input = _launch_fused_mrope( + grad_output.contiguous(), + freqs, + ctx.mrope_section, + ctx.interleaved_mrope, + ctx.rotary_interleaved, + inverse=True, + ) + return grad_input, None, None, None, None + + +class _FusedMRoPETHD(torch.autograd.Function): + """Autograd wrapper for fused THD mRoPE.""" + + @staticmethod + def forward( + ctx, + t, + cu_seqlens, + freqs, + mrope_section, + interleaved_mrope, + rotary_interleaved, + cp_size, + cp_rank, + fp32_compute, + ): + assert not freqs.requires_grad, "fused THD mRoPE expects non-gradient raw frequency tensors" + ctx.mrope_section = tuple(int(section) for section in mrope_section) + ctx.interleaved_mrope = bool(interleaved_mrope) + ctx.rotary_interleaved = bool(rotary_interleaved) + ctx.cp_size = int(cp_size) + ctx.cp_rank = int(cp_rank) + ctx.fp32_compute = bool(fp32_compute) + ctx.save_for_backward(cu_seqlens, freqs) + return _launch_fused_mrope_thd( + t, + cu_seqlens, + freqs, + ctx.mrope_section, + ctx.interleaved_mrope, + ctx.rotary_interleaved, + inverse=False, + cp_size=ctx.cp_size, + cp_rank=ctx.cp_rank, + fp32_compute=ctx.fp32_compute, + ) + + @staticmethod + def backward(ctx, grad_output): + cu_seqlens, freqs = ctx.saved_tensors + grad_input = _launch_fused_mrope_thd( + grad_output.contiguous(), + cu_seqlens, + freqs, + ctx.mrope_section, + ctx.interleaved_mrope, + ctx.rotary_interleaved, + inverse=True, + cp_size=ctx.cp_size, + cp_rank=ctx.cp_rank, + fp32_compute=ctx.fp32_compute, + ) + return grad_input, None, None, None, None, None, None, None, None + + +def fused_apply_mrope( + t: torch.Tensor, + freqs: torch.Tensor, + mrope_section: List[int], + interleaved_mrope: bool = False, + rotary_interleaved: bool = False, +) -> torch.Tensor: + """Apply multimodal RoPE with a fused Triton kernel. + + Args: + t: Input tensor with shape ``[seq, batch, heads, head_dim]``. + freqs: Raw mRoPE frequencies with shape ``[3, batch, seq, rotary_dim / 2]``. + mrope_section: Temporal, height, and width channel sections. + interleaved_mrope: Use Qwen3.5-VL stride-3 T/H/W layout when True. Use + Qwen2-VL section layout when False. + rotary_interleaved: Must be False. The integrated fused mRoPE path + currently supports split-half RoPE layout. + + Returns: + Rotated tensor with the same shape and dtype as ``t``. + """ + return _FusedMRoPE.apply(t, freqs, mrope_section, interleaved_mrope, rotary_interleaved) + + +def fused_apply_mrope_thd( + t: torch.Tensor, + cu_seqlens: torch.Tensor, + freqs: torch.Tensor, + mrope_section: List[int], + interleaved_mrope: bool = False, + rotary_interleaved: bool = False, + cp_size: int = 1, + cp_rank: int = 0, + fp32_compute: bool = False, +) -> torch.Tensor: + """Apply multimodal RoPE to THD-packed tensors with a fused Triton kernel. + + Args: + t: Input tensor with shape ``[total_tokens, heads, head_dim]``. + cu_seqlens: Global cumulative sequence lengths for the packed batch. + freqs: Raw mRoPE frequencies with shape ``[3, 1, total_seqlen, rotary_dim / 2]``. + mrope_section: Temporal, height, and width channel sections. + interleaved_mrope: Use Qwen3.5-VL stride-3 T/H/W layout when True. + rotary_interleaved: Must be False. + cp_size: Context parallel world size for THD token mapping. + cp_rank: Context parallel rank for THD token mapping. + fp32_compute: Apply the rotary math in fp32 and cast directly to output dtype. + + Returns: + Rotated tensor with the same shape and dtype as ``t``. + """ + return _FusedMRoPETHD.apply( + t, + cu_seqlens, + freqs, + mrope_section, + interleaved_mrope, + rotary_interleaved, + cp_size, + cp_rank, + fp32_compute, + ) + + +def is_fused_mrope_available() -> bool: + """Return whether the Triton mRoPE fusion can be used on this host. + + This does not check tensor device, dtype, stride, or CUDA capability. Use + ``can_launch_fused_mrope`` or ``get_fused_mrope_unavailable_reason`` with + tensors before dispatching to the fused kernel. + """ + if not torch.cuda.is_available(): + return False + return can_launch_fused_mrope() diff --git a/megatron/core/models/common/embeddings/rope_utils.py b/megatron/core/models/common/embeddings/rope_utils.py index c97f738771b..0c89cdbf084 100644 --- a/megatron/core/models/common/embeddings/rope_utils.py +++ b/megatron/core/models/common/embeddings/rope_utils.py @@ -16,6 +16,7 @@ from megatron.core import parallel_state logger = logging.getLogger(__name__) +_ROPE_FUSION_FALLBACK_WARNINGS: set[str] = set() try: from megatron.core.extensions.transformer_engine import fused_apply_rotary_pos_emb @@ -29,6 +30,26 @@ fused_apply_rotary_pos_emb_thd = None +try: + from megatron.core.fusions.fused_mrope import ( + can_launch_fused_mrope_thd, + fused_apply_mrope, + fused_apply_mrope_thd, + get_fused_mrope_thd_unavailable_reason, + get_fused_mrope_unavailable_reason, + is_fused_mrope_available, + mrope_freqs_to_rotary_emb, + ) +except ImportError: + can_launch_fused_mrope_thd = None + fused_apply_mrope = None + fused_apply_mrope_thd = None + get_fused_mrope_thd_unavailable_reason = None + get_fused_mrope_unavailable_reason = None + is_fused_mrope_available = None + mrope_freqs_to_rotary_emb = None + + try: from flash_attn.layers.rotary import apply_rotary_emb as apply_rotary_emb_flash except ImportError: @@ -41,10 +62,101 @@ 'apply_rotary_pos_emb_with_cos_sin', 'fused_apply_rotary_pos_emb', 'fused_apply_rotary_pos_emb_thd', + 'can_launch_fused_mrope_thd', + 'fused_apply_mrope', + 'fused_apply_mrope_thd', + 'get_fused_mrope_thd_unavailable_reason', + 'get_fused_mrope_unavailable_reason', + 'is_fused_mrope_available', + 'mrope_freqs_to_rotary_emb', 'get_pos_emb_on_this_cp_rank', ] +def _is_raw_mrope_freqs(t: Tensor, freqs: Tensor, config: TransformerConfig) -> bool: + """Return whether freqs is the raw 3-axis mRoPE tensor for fused apply.""" + if config.mrope_section is None or freqs.dim() != 4 or freqs.shape[0] != 3: + return False + if sum(config.mrope_section) != freqs.shape[-1] or freqs.shape[-1] * 2 > t.shape[-1]: + return False + if t.dim() == 4: + return freqs.shape[1] == t.shape[1] and freqs.shape[2] == t.shape[0] + if t.dim() == 3: + return freqs.shape[1] == 1 + return False + + +def _is_raw_mrope_freqs_thd( + t: Tensor, freqs: Tensor, cu_seqlens: Tensor, config: TransformerConfig, cp_size: int +) -> bool: + """Return whether freqs is raw mRoPE for THD layout, or fail on raw-like bad shapes.""" + if config.mrope_section is None or freqs.dim() != 4 or freqs.shape[0] != 3: + return False + if t.dim() != 3: + raise ValueError( + f"raw mRoPE THD expects t with shape [tokens, heads, head_dim], got {tuple(t.shape)}" + ) + if sum(config.mrope_section) != freqs.shape[-1] or freqs.shape[-1] * 2 > t.shape[-1]: + return False + + if freqs.shape[1] != 1: + raise ValueError( + "raw mRoPE THD freqs must have singleton batch dimension with shape " + f"[3, 1, total_seqlen, rotary_dim / 2], got {tuple(freqs.shape)}" + ) + if cp_size > 1 and freqs.shape[2] % cp_size != 0: + raise ValueError( + "raw mRoPE THD freqs sequence length must be divisible by context parallel size, " + f"got freqs.shape[2]={freqs.shape[2]}, cp_size={cp_size}" + ) + expected_total_seqlen = t.shape[0] * cp_size + if freqs.shape[2] != expected_total_seqlen: + raise ValueError( + "raw mRoPE THD freqs sequence length must match local tokens times cp_size, " + f"got freqs.shape[2]={freqs.shape[2]}, tokens={t.shape[0]}, cp_size={cp_size}" + ) + if cu_seqlens.dim() != 1: + raise ValueError(f"raw mRoPE THD cu_seqlens must be 1D, got {tuple(cu_seqlens.shape)}") + return True + + +def _raw_mrope_freqs_to_emb(freqs: Tensor, config: TransformerConfig) -> Tensor: + assert mrope_freqs_to_rotary_emb is not None, "mRoPE frequency conversion is unavailable." + return mrope_freqs_to_rotary_emb( + freqs, + config.mrope_section, + interleaved_mrope=config.mrope_interleaved, + rotary_interleaved=config.rotary_interleaved, + ) + + +def _warn_rope_fusion_fallback_once(key: str, message: str) -> None: + if key in _ROPE_FUSION_FALLBACK_WARNINGS: + return + _ROPE_FUSION_FALLBACK_WARNINGS.add(key) + warnings.warn(message, stacklevel=2) + + +def _fused_mrope_unavailable_warning_key(reason: str, thd: bool = False) -> str: + prefix = "triton-mrope-thd-unavailable" if thd else "triton-mrope-unavailable" + reason_lower = reason.lower() + if "triton is not available" in reason_lower: + category = "import" + elif "cuda tensors" in reason_lower or "same device" in reason_lower: + category = "device" + elif "dtype" in reason_lower or "float32" in reason_lower: + category = "dtype" + elif "stride" in reason_lower or "contiguous" in reason_lower: + category = "stride" + elif "capability" in reason_lower: + category = "capability" + elif "rotary_interleaved" in reason_lower: + category = "rotary-interleaved" + else: + category = "other" + return f"{prefix}-{category}" + + def get_pos_emb_on_this_cp_rank( pos_emb: Tensor, seq_dim: int, cp_group: torch.distributed.ProcessGroup ) -> Tensor: @@ -181,20 +293,21 @@ def _get_thd_freqs_on_this_cp_rank( compatibility. """ if cp_size > 1: - cp_seg = x.size(0) // 2 + first_cp_seg = (x.size(0) + 1) // 2 + second_cp_seg = x.size(0) // 2 full_seqlen = cp_size * x.size(0) # Apply offset to both forward and backward segments for context parallelism - # offset=0: traditional behavior, freqs[0:cp_seg] and freqs[...] - # offset>0: exact mapping, freqs[offset+0:offset+cp_seg] and freqs[offset+...] + # offset=0: traditional behavior, freqs[0:first_cp_seg] and freqs[...] + # offset>0: exact mapping, freqs[offset+0:offset+first_cp_seg] and freqs[offset+...] return torch.cat( [ - freqs[offset + cp_rank * cp_seg : offset + (cp_rank + 1) * cp_seg], + freqs[offset + cp_rank * first_cp_seg : offset + (cp_rank + 1) * first_cp_seg], freqs[ offset + full_seqlen - - (cp_rank + 1) * cp_seg : offset + - (cp_rank + 1) * second_cp_seg : offset + full_seqlen - - cp_rank * cp_seg + - cp_rank * second_cp_seg ], ] ) @@ -205,6 +318,84 @@ def _get_thd_freqs_on_this_cp_rank( return freqs[offset : offset + x.size(0)] +def _get_thd_raw_mrope_freqs_on_this_cp_rank( + cp_rank: int, cp_size: int, x: Tensor, freqs: Tensor, offset: int = 0 +) -> Tensor: + """Get raw mRoPE frequency slices for this CP rank in THD layout.""" + if cp_size > 1: + first_cp_seg = (x.size(0) + 1) // 2 + second_cp_seg = x.size(0) // 2 + full_seqlen = cp_size * x.size(0) + return torch.cat( + [ + freqs[ + :, :, offset + cp_rank * first_cp_seg : offset + (cp_rank + 1) * first_cp_seg + ], + freqs[ + :, + :, + offset + + full_seqlen + - (cp_rank + 1) * second_cp_seg : offset + + full_seqlen + - cp_rank * second_cp_seg, + ], + ], + dim=2, + ) + else: + return freqs[:, :, offset : offset + x.size(0)] + + +def _get_thd_cp_splits(cu_seqlens: Tensor, cp_size: int) -> tuple[list[int], list[int]]: + """Return global sequence offsets and per-rank sequence lengths for THD CP fallback.""" + cu_seqlens_list = cu_seqlens.tolist() + local_seqlens = [] + for seq_start, seq_end in zip(cu_seqlens_list[:-1], cu_seqlens_list[1:]): + seq_len = seq_end - seq_start + if cp_size > 1 and seq_len % cp_size != 0: + raise ValueError( + "THD sequence lengths must be divisible by context parallel size, " + f"got sequence length {seq_len}, cp_size={cp_size}" + ) + local_seqlens.append(seq_len // cp_size) + return cu_seqlens_list, local_seqlens + + +def _pack_thd_raw_mrope_freqs( + t: Tensor, + cu_seqlens: Tensor, + freqs: Tensor, + cp_group: torch.distributed.ProcessGroup, + total_seqlen: Optional[int] = None, +) -> Tensor: + """Pack raw mRoPE freqs into the same local token order as THD tensor ``t``.""" + cp_size = cp_group.size() + cp_rank = cp_group.rank() + cu_seqlens_list, seqlens = _get_thd_cp_splits(cu_seqlens, cp_size) + sequence_splits = torch.split(t, seqlens) + if total_seqlen is None: + total_seqlen = cu_seqlens_list[-1] + assert freqs.size(2) == total_seqlen, ( + f"raw mRoPE THD freqs sequence length {freqs.size(2)} must match " + f"cu_seqlens[-1] = {total_seqlen}" + ) + + freq_slices = [] + for i, x in enumerate(sequence_splits): + seq_start_offset = cu_seqlens_list[i] + freq_slices.append( + _get_thd_raw_mrope_freqs_on_this_cp_rank(cp_rank, cp_size, x, freqs, seq_start_offset) + ) + + packed_freqs = torch.cat(freq_slices, dim=2) + assert packed_freqs.shape[2] == t.shape[0], ( + f"packed raw mRoPE freqs sequence length {packed_freqs.shape[2]} " + f"does not match THD tensor length {t.shape[0]}" + ) + return packed_freqs.contiguous() + + def _apply_rotary_pos_emb_thd( t: Tensor, cu_seqlens: Tensor, @@ -240,21 +431,21 @@ def _apply_rotary_pos_emb_thd( raise ValueError("cp_group must be provided for THD format RoPE") cp_size = cp_group.size() cp_rank = cp_group.rank() - seqlens = ((cu_seqlens[1:] - cu_seqlens[:-1]) // cp_size).tolist() + cu_seqlens_list, seqlens = _get_thd_cp_splits(cu_seqlens, cp_size) # Handle two different frequency tensor formats: - # 1. If freqs.size(0) == cu_seqlens[-1]: freqs contains all positions across all sequences + # 1. If freqs.size(0) == cu_seqlens_list[-1]: freqs contains all positions across all sequences # -> Use offset-based mapping for exact positional correspondence # 2. Otherwise: freqs contains only max sequence length positions # -> Use traditional mapping without offsets (map first :seqlen part) - if freqs.dim() >= 1 and freqs.size(0) == cu_seqlens[-1]: + if freqs.dim() >= 1 and freqs.size(0) == cu_seqlens_list[-1]: # CASE 1: Exact mapping with offsets # Build packed freqs in one pass, then apply once to the whole packed tensor sequence_splits = torch.split(t, seqlens) freq_slices = [] for i, x in enumerate(sequence_splits): # cu_seqlens[i] is the starting offset of this sequence in the original batch - seq_start_offset = cu_seqlens[i].item() + seq_start_offset = cu_seqlens_list[i] freq_slices.append( _get_thd_freqs_on_this_cp_rank(cp_rank, cp_size, x, freqs, seq_start_offset) ) @@ -311,40 +502,254 @@ def apply_rotary_pos_emb( if cp_group is None: cp_group = parallel_state.get_context_parallel_group() + is_raw_mrope_freqs = ( + _is_raw_mrope_freqs(t, freqs, config) + if cu_seqlens is None + else _is_raw_mrope_freqs_thd(t, freqs, cu_seqlens, config, cp_group.size()) + ) + if config.apply_rope_fusion: if cu_seqlens is None: + force_unfused_mrope = False + if is_raw_mrope_freqs: + unavailable_reason = None + can_try_fused_mrope = ( + fused_apply_mrope is not None + and get_fused_mrope_unavailable_reason is not None + and not mla_rotary_interleaved + and not inverse + and mscale == 1.0 + ) + if can_try_fused_mrope: + unavailable_reason = get_fused_mrope_unavailable_reason( + t, freqs, config.rotary_interleaved + ) + use_fused_mrope = can_try_fused_mrope and unavailable_reason is None + if use_fused_mrope: + return fused_apply_mrope( + t, + freqs, + config.mrope_section, + interleaved_mrope=config.mrope_interleaved, + rotary_interleaved=config.rotary_interleaved, + ) + + if unavailable_reason is not None: + _warn_rope_fusion_fallback_once( + _fused_mrope_unavailable_warning_key(unavailable_reason), + f"Triton fused mRoPE is unavailable: {unavailable_reason}. " + "Using unfused implementation.", + ) + force_unfused_mrope = True + unavailable_is_rotary_interleaved = ( + unavailable_reason is not None + and "rotary_interleaved" in unavailable_reason.lower() + ) + if mscale != 1.0: + _warn_rope_fusion_fallback_once( + "triton-mrope-mscale", + f"mscale={mscale} is not supported by Triton fused mRoPE. " + "Using unfused implementation.", + ) + force_unfused_mrope = True + if mla_rotary_interleaved: + _warn_rope_fusion_fallback_once( + "triton-mrope-mla-rotary-interleaved", + "Triton fused mRoPE does not support MLA-style interleaving in RoPE. " + "Using unfused implementation.", + ) + force_unfused_mrope = True + if inverse: + _warn_rope_fusion_fallback_once( + "triton-mrope-inverse", + "inverse RoPE is not supported by Triton fused mRoPE. " + "Using unfused implementation.", + ) + force_unfused_mrope = True + if config.rotary_interleaved and not unavailable_is_rotary_interleaved: + _warn_rope_fusion_fallback_once( + "triton-mrope-rotary-interleaved", + "Triton fused mRoPE currently supports rotary_interleaved=False. " + "Using unfused implementation.", + ) + force_unfused_mrope = True + freqs = _raw_mrope_freqs_to_emb(freqs, config) + is_raw_mrope_freqs = False + if force_unfused_mrope: + return _apply_rotary_pos_emb_bshd( + t, + freqs, + rotary_interleaved=config.rotary_interleaved, + mla_rotary_interleaved=mla_rotary_interleaved, + mscale=mscale, + inverse=inverse, + mla_output_remove_interleaving=mla_output_remove_interleaving, + ) + # NOTE: TE backends do not support mRoPE in bshd format when bs > 1. use_unfused = False if config.mrope_section is not None and freqs.shape[1] > 1: # TODO: Add a check in TransformerConfig and remove this unfused implementation. - warnings.warn( - "apply_rope_fusion does not support mRoPE in bshd format when bs > 1. " - "Please set apply_rope_fusion to false. This will become an error in v0.16." + _warn_rope_fusion_fallback_once( + "te-mrope-bshd-batch", + "Transformer Engine fused RoPE does not support mRoPE in bshd format when " + "bs > 1 without raw mRoPE freqs. Using unfused implementation.", ) use_unfused = True if mscale != 1.0: - warnings.warn( + _warn_rope_fusion_fallback_once( + "te-rope-mscale", f"mscale={mscale} is not supported by TE's fused RoPE. " - "Using unfused implementation." + "Using unfused implementation.", ) use_unfused = True if mla_rotary_interleaved: - warnings.warn( - "apply_rope_fusion does not support MLA-style interleaving in RoPE." - "Using unfused implementation." + _warn_rope_fusion_fallback_once( + "te-rope-mla-rotary-interleaved", + "apply_rope_fusion does not support MLA-style interleaving in RoPE. " + "Using unfused implementation.", ) use_unfused = True if inverse: - warnings.warn( + _warn_rope_fusion_fallback_once( + "te-rope-inverse", "inverse RoPE is not supported by TE's fused RoPE. " - "Using unfused implementation." + "Using unfused implementation.", + ) + use_unfused = True + if fused_apply_rotary_pos_emb is None: + _warn_rope_fusion_fallback_once( + "te-rope-unavailable", + "Transformer Engine fused RoPE is unavailable. Using unfused implementation.", ) use_unfused = True if not use_unfused: - assert fused_apply_rotary_pos_emb is not None, "apply_rope_fusion is not available." return fused_apply_rotary_pos_emb(t, freqs, interleaved=config.rotary_interleaved) else: - assert fused_apply_rotary_pos_emb_thd is not None, "apply_rope_fusion is not available." + if is_raw_mrope_freqs: + use_fused_mrope_thd = ( + fused_apply_mrope_thd is not None + and can_launch_fused_mrope_thd is not None + and get_fused_mrope_thd_unavailable_reason is not None + and mscale == 1.0 + and not mla_rotary_interleaved + and not inverse + and not config.rotary_interleaved + ) + if use_fused_mrope_thd: + unavailable_reason = get_fused_mrope_thd_unavailable_reason( + t, + cu_seqlens, + freqs, + rotary_interleaved=config.rotary_interleaved, + cp_size=cp_group.size(), + cp_rank=cp_group.rank(), + ) + if unavailable_reason is None: + return fused_apply_mrope_thd( + t, + cu_seqlens, + freqs, + config.mrope_section, + interleaved_mrope=config.mrope_interleaved, + rotary_interleaved=config.rotary_interleaved, + cp_size=cp_group.size(), + cp_rank=cp_group.rank(), + ) + _warn_rope_fusion_fallback_once( + _fused_mrope_unavailable_warning_key(unavailable_reason, thd=True), + f"Triton fused mRoPE for THD layout is unavailable: " + f"{unavailable_reason}. Using unfused implementation.", + ) + else: + has_unsupported_option = False + if mscale != 1.0: + _warn_rope_fusion_fallback_once( + "triton-mrope-thd-mscale", + f"mscale={mscale} is not supported by Triton fused mRoPE for THD " + "layout. Using unfused implementation.", + ) + has_unsupported_option = True + if mla_rotary_interleaved: + _warn_rope_fusion_fallback_once( + "triton-mrope-thd-mla-rotary-interleaved", + "Triton fused mRoPE for THD layout does not support MLA-style " + "interleaving in RoPE. Using unfused implementation.", + ) + has_unsupported_option = True + if inverse: + _warn_rope_fusion_fallback_once( + "triton-mrope-thd-inverse", + "inverse RoPE is not supported by Triton fused mRoPE for THD layout. " + "Using unfused implementation.", + ) + has_unsupported_option = True + if config.rotary_interleaved: + _warn_rope_fusion_fallback_once( + "triton-mrope-thd-rotary-interleaved", + "Triton fused mRoPE for THD layout currently supports " + "rotary_interleaved=False. Using unfused implementation.", + ) + has_unsupported_option = True + if not has_unsupported_option: + _warn_rope_fusion_fallback_once( + "triton-mrope-thd-unavailable", + "Triton fused mRoPE for THD layout is unavailable. " + "Using unfused implementation.", + ) + freqs = _raw_mrope_freqs_to_emb(freqs, config) + return _apply_rotary_pos_emb_thd( + t, + cu_seqlens, + freqs, + rotary_interleaved=config.rotary_interleaved, + mla_rotary_interleaved=mla_rotary_interleaved, + mscale=mscale, + cp_group=cp_group, + inverse=inverse, + mla_output_remove_interleaving=mla_output_remove_interleaving, + ) + use_unfused_thd = False + if mscale != 1.0: + _warn_rope_fusion_fallback_once( + "te-rope-thd-mscale", + f"mscale={mscale} is not supported by TE's fused RoPE for THD layout. " + "Using unfused implementation.", + ) + use_unfused_thd = True + if mla_rotary_interleaved: + _warn_rope_fusion_fallback_once( + "te-rope-thd-mla-rotary-interleaved", + "TE fused RoPE for THD layout does not support MLA-style interleaving " + "in RoPE. Using unfused implementation.", + ) + use_unfused_thd = True + if inverse: + _warn_rope_fusion_fallback_once( + "te-rope-thd-inverse", + "inverse RoPE is not supported by TE's fused RoPE for THD layout. " + "Using unfused implementation.", + ) + use_unfused_thd = True + if fused_apply_rotary_pos_emb_thd is None: + _warn_rope_fusion_fallback_once( + "te-rope-thd-unavailable", + "Transformer Engine fused RoPE for THD layout is unavailable. " + "Using unfused implementation.", + ) + use_unfused_thd = True + if use_unfused_thd: + return _apply_rotary_pos_emb_thd( + t, + cu_seqlens, + freqs, + rotary_interleaved=config.rotary_interleaved, + mla_rotary_interleaved=mla_rotary_interleaved, + mscale=mscale, + cp_group=cp_group, + inverse=inverse, + mla_output_remove_interleaving=mla_output_remove_interleaving, + ) return fused_apply_rotary_pos_emb_thd( t, cu_seqlens, @@ -354,6 +759,9 @@ def apply_rotary_pos_emb( interleaved=config.rotary_interleaved, ) # use unfused implementation + if is_raw_mrope_freqs: + freqs = _raw_mrope_freqs_to_emb(freqs, config) + if cu_seqlens is None: return _apply_rotary_pos_emb_bshd( t, diff --git a/megatron/core/models/common/embeddings/rotary_pos_embedding.py b/megatron/core/models/common/embeddings/rotary_pos_embedding.py index 804bdb7c537..e056ffeb9be 100644 --- a/megatron/core/models/common/embeddings/rotary_pos_embedding.py +++ b/megatron/core/models/common/embeddings/rotary_pos_embedding.py @@ -341,6 +341,8 @@ def forward( position_ids: torch.Tensor, mrope_section: List[int], cp_group: Optional[torch.distributed.ProcessGroup] = None, + return_raw_freqs: bool = False, + packed_seq: bool = False, ) -> Tensor: """Forward pass of multimodal RoPE embedding. @@ -350,9 +352,14 @@ def forward( height and width in rope calculation. cp_group (torch.distributed.ProcessGroup, optional): Context parallel group. Defaults to None. + return_raw_freqs (bool, optional): If True, return the raw per-axis frequencies with + shape [3, batchsize, seqlens, dim / 2] for fused mRoPE application. + packed_seq (bool, optional): Whether the sequence uses THD packing. Packed sequences + keep full position frequencies because THD RoPE applies CP partitioning later. Returns: - Tensor: Embeddings after applying RoPE. + Tensor: Embeddings after applying RoPE, or raw per-axis frequencies when + return_raw_freqs is True. """ seq = position_ids.to(device=self.inv_freq.device, dtype=self.inv_freq.dtype) @@ -366,6 +373,13 @@ def forward( # shape (3, bs, seq_length, dim) freqs = (inv_freq_expanded @ seq_expanded).transpose(2, 3) + if cp_group is None: + cp_group = self.cp_group + if return_raw_freqs: + if cp_group is not None and cp_group.size() > 1 and not packed_seq: + freqs = get_pos_emb_on_this_cp_rank(freqs, 2, cp_group) + return freqs.contiguous() + # first part even vector components, second part odd vector components, # 2 * dim in dimension size if self.interleaved_mrope: @@ -376,9 +390,9 @@ def forward( emb = torch.cat((freqs, freqs), dim=-1) # shape (bs, seq_length, 2 * dim) else: bs = freqs.shape[0] - emb = torch.stack((freqs.view(bs, -1, 1), freqs.view(bs, -1, 1)), dim=-1).view( - bs, freqs.shape[1], -1 - ) + emb = torch.stack( + (freqs.reshape(bs, -1, 1), freqs.reshape(bs, -1, 1)), dim=-1 + ).view(bs, freqs.shape[1], -1) else: # Original section-based layout (Qwen2-VL style). if not self.rotary_interleaved: @@ -386,8 +400,8 @@ def forward( else: bs = freqs.shape[1] emb = torch.stack( - (freqs.view(3, bs, -1, 1), freqs.view(3, bs, -1, 1)), dim=-1 - ).view(3, bs, freqs.shape[0], -1) + (freqs.reshape(3, bs, -1, 1), freqs.reshape(3, bs, -1, 1)), dim=-1 + ).view(3, bs, freqs.shape[2], -1) # generate freqs with mrope_section: cycle T/H/W per section chunk mrope_section_doubled = list(mrope_section) * 2 emb = torch.cat( @@ -396,9 +410,7 @@ def forward( # shape (seq_length, bs, 1, 2 * dim) emb = emb[..., None, :].transpose(0, 1).contiguous() - if cp_group is None: - cp_group = self.cp_group - if cp_group is not None and cp_group.size() > 1: + if cp_group is not None and cp_group.size() > 1 and not packed_seq: # slice rotary_pos_emb along sequence dimension and select the parition of the current # CP rank emb = get_pos_emb_on_this_cp_rank(emb, 0, cp_group) diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index 7bce9d96d2c..7adcc03ecd1 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -147,6 +147,7 @@ def __init__( self.mtp_process = mtp_block_spec is not None and mtp_on_this_rank( self.config, ignore_virtual=False, vp_stage=vp_stage ) + self._fused_mrope_available = False self.fuse_linear_cross_entropy = ( self.config.cross_entropy_loss_fusion @@ -209,6 +210,13 @@ def __init__( assert ( self.mrope_section is not None ), "mrope require mrope_section setting, but we got None from TransformerConfig" + if self.config.apply_rope_fusion and not self.config.rotary_interleaved: + try: + from megatron.core.fusions.fused_mrope import is_fused_mrope_available + + self._fused_mrope_available = is_fused_mrope_available() + except ImportError: + self._fused_mrope_available = False # Cache for RoPE tensors which do not change between iterations. self.rotary_pos_emb_cache = {} @@ -402,10 +410,25 @@ def _preprocess( ) elif self.position_embedding_type == 'mrope' and not self.config.multi_latent_attention: if self.training or not self.config.flash_decode: + packed_seq = packed_seq_params is not None and packed_seq_params.qkv_format == 'thd' + use_fused_mrope = False + use_raw_mrope_freqs = ( + self.config.apply_rope_fusion and not self.config.rotary_interleaved + ) + if self.config.fused_single_qkv_rope: + use_raw_mrope_freqs = False + # Inference indexes rotary_pos_emb as seq-major materialized embeddings. + # Raw mRoPE freqs are axis-major and are only safe for the normal decoder path. + if in_inference_mode: + use_raw_mrope_freqs = False + if use_raw_mrope_freqs: + use_fused_mrope = self._fused_mrope_available rotary_pos_emb = self.rotary_pos_emb( position_ids, self.mrope_section, cp_group=packed_seq_params.cp_group if packed_seq_params is not None else None, + return_raw_freqs=use_fused_mrope, + packed_seq=packed_seq, ) else: # Flash decoding uses precomputed cos and sin for RoPE diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index 31e06a84b48..e9eb8122d57 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -449,6 +449,7 @@ def _build_per_layer_rotary_pos_emb(self, rotary_base: float) -> None: rotary_interleaved=self.config.rotary_interleaved, seq_len_interpolation_factor=seq_len_interpolation_factor, rotary_base=rotary_base, + interleaved_mrope=self.config.mrope_interleaved, ) self.mrope_section = self.config.mrope_section assert ( diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index c53c024a80a..c6eca1c5bcf 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -2232,7 +2232,18 @@ def __post_init__(self): "It is experimental and may change in future versions." ) else: - if self.rotary_interleaved: + fused_mrope_available = False + # Triton fused mRoPE supports split-half RoPE only. Keep rotary_interleaved + # configs on the TE validation path so the TE >= 2.3 check still applies. + if self.mrope_section is not None and not self.rotary_interleaved: + try: + from megatron.core.fusions.fused_mrope import is_fused_mrope_available + + fused_mrope_available = is_fused_mrope_available() + except ImportError: + fused_mrope_available = False + + if self.rotary_interleaved and not fused_mrope_available: if not is_te_min_version("2.3.0"): raise ValueError( "rotary_interleaved does not work with apply_rope_fusion for " @@ -2244,9 +2255,14 @@ def __post_init__(self): fused_apply_rotary_pos_emb_thd, ) - if fused_apply_rotary_pos_emb is None and fused_apply_rotary_pos_emb_thd is None: + if ( + fused_apply_rotary_pos_emb is None + and fused_apply_rotary_pos_emb_thd is None + and not fused_mrope_available + ): raise ValueError( - "apply_rope_fusion is not available. Please install TE >= 1.4." + "apply_rope_fusion is not available. Please install TE >= 1.4 " + "or Triton for fused mRoPE." ) if self.fused_single_qkv_rope: diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 853973b92cd..9678a0b6e46 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -1630,7 +1630,7 @@ def validate_args(args, defaults={}): # Legacy RoPE arguments if args.use_rotary_position_embeddings: args.position_embedding_type = 'rope' - if args.position_embedding_type != 'rope': + if args.position_embedding_type not in ('rope', 'mrope'): args.apply_rope_fusion = False # Would just need to add 'NoPE' as a position_embedding_type to support this, but for now diff --git a/tests/unit_tests/fusions/test_fused_mrope.py b/tests/unit_tests/fusions/test_fused_mrope.py new file mode 100644 index 00000000000..b033f6b9bed --- /dev/null +++ b/tests/unit_tests/fusions/test_fused_mrope.py @@ -0,0 +1,1311 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +import warnings +from types import SimpleNamespace + +import pytest +import torch + +import megatron.core.models.common.embeddings.rope_utils as rope_utils +from megatron.core import parallel_state +from megatron.core.fusions.fused_mrope import ( + fused_apply_mrope, + fused_apply_mrope_thd, + get_fused_mrope_thd_unavailable_reason, + get_fused_mrope_unavailable_reason, + is_fused_mrope_available, + mrope_freqs_to_rotary_emb, +) +from megatron.core.models.common.embeddings import apply_rotary_pos_emb +from megatron.core.models.common.embeddings.rope_utils import ( + _ROPE_FUSION_FALLBACK_WARNINGS, + _apply_rotary_pos_emb_bshd, + _apply_rotary_pos_emb_thd, +) +from megatron.core.models.common.embeddings.rotary_pos_embedding import MultimodalRotaryEmbedding +from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.transformer.transformer_config import TransformerConfig +from tests.unit_tests.test_utilities import Utils + + +class FakeCPGroup: + def __init__(self, size=1, rank=0): + self._size = size + self._rank = rank + + def size(self): + return self._size + + def rank(self): + return self._rank + + +class FakeDynamicInferenceContext: + def is_dynamic_batching(self): + return True + + def is_static_batching(self): + return False + + +class FakeStaticInferenceContext: + def is_dynamic_batching(self): + return False + + def is_static_batching(self): + return True + + +@pytest.fixture(autouse=True) +def clear_rope_fusion_fallback_warnings(): + _ROPE_FUSION_FALLBACK_WARNINGS.clear() + yield + _ROPE_FUSION_FALLBACK_WARNINGS.clear() + + +def _dtype_tols(dtype): + if dtype == torch.bfloat16: + return dict(rtol=2.0e-2, atol=5.0e-2) + if dtype == torch.float16: + return dict(rtol=3.0e-3, atol=1.0e-2) + return dict(rtol=1.0e-6, atol=1.0e-6) + + +def _make_inputs( + dtype=torch.bfloat16, + requires_grad=False, + head_dim=20, + rotary_dim=16, + mrope_section=None, + interleaved_mrope=False, + batch=2, +): + seq = 32 + heads = 3 + if mrope_section is None: + mrope_section = [3, 3, 2] if interleaved_mrope else [2, 3, 3] + + generator = torch.Generator(device="cuda").manual_seed(1234) + t = torch.randn( + seq, + batch, + heads, + head_dim, + dtype=dtype, + device="cuda", + generator=generator, + requires_grad=requires_grad, + ) + freqs = torch.randn( + 3, batch, seq, rotary_dim // 2, dtype=torch.float32, device="cuda", generator=generator + ) + return t, freqs, mrope_section + + +def _make_position_ids(seq, batch): + base = torch.arange(seq, device="cuda", dtype=torch.long) + batch_offsets = torch.arange(batch, device="cuda", dtype=torch.long) + return ( + torch.stack((base, base * 2 + 3, base * 3 + 5), dim=0)[:, None, :] + + batch_offsets[None, :, None] + ).contiguous() + + +def _make_thd_inputs( + dtype=torch.bfloat16, + requires_grad=False, + interleaved_mrope=False, + cp_size=1, + padded_seq_lens=(12, 16), + head_dim=20, + rotary_dim=16, + mrope_section=None, +): + total_seq = sum(padded_seq_lens) + local_seq = total_seq // cp_size + heads = 3 + if mrope_section is None: + mrope_section = [3, 3, 2] if interleaved_mrope else [2, 3, 3] + + generator = torch.Generator(device="cuda").manual_seed(5678) + t = torch.randn( + local_seq, + heads, + head_dim, + dtype=dtype, + device="cuda", + generator=generator, + requires_grad=requires_grad, + ) + freqs = torch.randn( + 3, 1, total_seq, rotary_dim // 2, dtype=torch.float32, device="cuda", generator=generator + ) + cu_seqlens = torch.tensor([0, padded_seq_lens[0], total_seq], dtype=torch.int32, device="cuda") + return t, freqs, cu_seqlens, mrope_section + + +def _make_mrope_config( + num_attention_heads, mrope_section, interleaved_mrope=False, rotary_interleaved=False +): + return TransformerConfig( + num_attention_heads=num_attention_heads, + num_layers=1, + apply_rope_fusion=True, + mrope_section=mrope_section, + mrope_interleaved=interleaved_mrope, + rotary_interleaved=rotary_interleaved, + ) + + +def _fallback_warnings(recorded_warnings): + return [ + warning + for warning in recorded_warnings + if issubclass(warning.category, UserWarning) + and "Using unfused implementation" in str(warning.message) + ] + + +def _thd_cp_freq_indices(cu_seqlens_cpu, cp_size, cp_rank): + indices = [] + for global_start, global_end in zip(cu_seqlens_cpu[:-1], cu_seqlens_cpu[1:]): + local_seq_len = (global_end - global_start) // cp_size + first_cp_seg = (local_seq_len + 1) // 2 + second_cp_seg = local_seq_len // 2 + indices.extend( + range( + global_start + cp_rank * first_cp_seg, global_start + (cp_rank + 1) * first_cp_seg + ) + ) + indices.extend( + range(global_end - (cp_rank + 1) * second_cp_seg, global_end - cp_rank * second_cp_seg) + ) + return indices + + +def _assert_thd_cp_freq_index_coverage(cu_seqlens_cpu, cp_size): + expected = [] + actual = [] + for global_start, global_end in zip(cu_seqlens_cpu[:-1], cu_seqlens_cpu[1:]): + expected.extend(range(global_start, global_end)) + for cp_rank in range(cp_size): + actual.extend(_thd_cp_freq_indices(cu_seqlens_cpu, cp_size, cp_rank)) + assert sorted(actual) == expected + assert len(set(actual)) == len(actual) + + +@pytest.mark.parametrize("use_packed_seq", [False, True]) +def test_gpt_mrope_eval_requests_raw_freqs_when_fusion_available(use_packed_seq): + captured_kwargs = {} + + def fake_rotary_pos_emb(*args, **kwargs): + captured_kwargs.update(kwargs) + return "raw-mrope-freqs" + + model = SimpleNamespace( + training=False, + pre_process=False, + mtp_process=False, + position_embedding_type="mrope", + config=SimpleNamespace( + multi_latent_attention=False, + flash_decode=False, + apply_rope_fusion=True, + rotary_interleaved=False, + cuda_graph_impl=None, + fused_single_qkv_rope=False, + ), + rotary_pos_emb=fake_rotary_pos_emb, + mrope_section=[2, 3, 3], + _fused_mrope_available=True, + ) + packed_seq_params = ( + SimpleNamespace(qkv_format="thd", cp_group=FakeCPGroup()) if use_packed_seq else None + ) + + output = GPTModel._preprocess( + model, + input_ids=torch.zeros(1, 4, dtype=torch.long), + position_ids=torch.zeros(3, 1, 4, dtype=torch.long), + decoder_input=torch.zeros(4, 1, 12), + packed_seq_params=packed_seq_params, + ) + + assert output[1] == "raw-mrope-freqs" + assert captured_kwargs["return_raw_freqs"] is True + assert captured_kwargs["packed_seq"] is use_packed_seq + + +def test_gpt_mrope_eval_keeps_materialized_freqs_with_fused_single_qkv_rope(): + captured_kwargs = {} + + def fake_rotary_pos_emb(*args, **kwargs): + captured_kwargs.update(kwargs) + return "materialized-mrope-freqs" + + model = SimpleNamespace( + training=False, + pre_process=False, + mtp_process=False, + position_embedding_type="mrope", + config=SimpleNamespace( + multi_latent_attention=False, + flash_decode=False, + apply_rope_fusion=True, + rotary_interleaved=False, + cuda_graph_impl=None, + fused_single_qkv_rope=True, + ), + rotary_pos_emb=fake_rotary_pos_emb, + mrope_section=[2, 3, 3], + _fused_mrope_available=True, + ) + + output = GPTModel._preprocess( + model, + input_ids=torch.zeros(1, 4, dtype=torch.long), + position_ids=torch.zeros(3, 1, 4, dtype=torch.long), + decoder_input=torch.zeros(4, 1, 12), + ) + + assert output[1] == "materialized-mrope-freqs" + assert captured_kwargs["return_raw_freqs"] is False + + +def test_gpt_mrope_dynamic_inference_keeps_materialized_freqs(): + captured_kwargs = {} + + def fake_rotary_pos_emb(*args, **kwargs): + captured_kwargs.update(kwargs) + return "materialized-mrope-freqs" + + model = SimpleNamespace( + training=False, + pre_process=False, + mtp_process=False, + position_embedding_type="mrope", + config=SimpleNamespace( + multi_latent_attention=False, + flash_decode=False, + apply_rope_fusion=True, + rotary_interleaved=False, + cuda_graph_impl=None, + fused_single_qkv_rope=False, + ), + rotary_pos_emb=fake_rotary_pos_emb, + mrope_section=[2, 3, 3], + _fused_mrope_available=True, + ) + + output = GPTModel._preprocess( + model, + input_ids=torch.zeros(1, 4, dtype=torch.long), + position_ids=torch.zeros(3, 1, 4, dtype=torch.long), + decoder_input=torch.zeros(4, 1, 12), + inference_context=FakeDynamicInferenceContext(), + ) + + assert output[1] == "materialized-mrope-freqs" + assert captured_kwargs["return_raw_freqs"] is False + + +def test_gpt_mrope_static_inference_keeps_materialized_freqs(): + captured_kwargs = {} + + def fake_rotary_pos_emb(*args, **kwargs): + captured_kwargs.update(kwargs) + return "materialized-mrope-freqs" + + model = SimpleNamespace( + training=False, + pre_process=False, + mtp_process=False, + position_embedding_type="mrope", + config=SimpleNamespace( + multi_latent_attention=False, + flash_decode=False, + apply_rope_fusion=True, + rotary_interleaved=False, + cuda_graph_impl=None, + fused_single_qkv_rope=False, + ), + rotary_pos_emb=fake_rotary_pos_emb, + mrope_section=[2, 3, 3], + _fused_mrope_available=True, + ) + + output = GPTModel._preprocess( + model, + input_ids=torch.zeros(1, 4, dtype=torch.long), + position_ids=torch.zeros(3, 1, 4, dtype=torch.long), + decoder_input=torch.zeros(4, 1, 12), + inference_context=FakeStaticInferenceContext(), + ) + + assert output[1] == "materialized-mrope-freqs" + assert captured_kwargs["return_raw_freqs"] is False + + +def test_is_fused_mrope_available_requires_cuda(monkeypatch): + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + + assert not is_fused_mrope_available() + + +def test_transformer_config_rejects_fused_mrope_without_cuda_or_te(monkeypatch): + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb", None) + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb_thd", None) + + with pytest.raises(ValueError, match="apply_rope_fusion is not available"): + TransformerConfig( + num_attention_heads=1, num_layers=1, apply_rope_fusion=True, mrope_section=[1, 1, 1] + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +@pytest.mark.parametrize("head_dim", [16, 20]) +def test_fused_mrope_matches_unfused_forward_backward(interleaved_mrope, head_dim): + t_ref, freqs, mrope_section = _make_inputs( + requires_grad=True, head_dim=head_dim, interleaved_mrope=interleaved_mrope + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_bshd(t_ref, emb, rotary_interleaved=False) + out = fused_apply_mrope( + t_fused, freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +def test_apply_rotary_pos_emb_bshd_eval_uses_triton_without_te(interleaved_mrope, monkeypatch): + t, freqs, mrope_section = _make_inputs(interleaved_mrope=interleaved_mrope, batch=1) + config = _make_mrope_config(t.shape[2], mrope_section, interleaved_mrope) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False) + + fused_calls = 0 + orig_fused_apply_mrope = rope_utils.fused_apply_mrope + + def wrapped_fused_apply_mrope(*args, **kwargs): + nonlocal fused_calls + fused_calls += 1 + return orig_fused_apply_mrope(*args, **kwargs) + + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb", None) + monkeypatch.setattr(rope_utils, "fused_apply_mrope", wrapped_fused_apply_mrope) + with torch.no_grad(), warnings.catch_warnings(record=True) as recorded_warnings: + warnings.simplefilter("always") + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup()) + + assert fused_calls == 1 + assert not _fallback_warnings(recorded_warnings) + assert not out.requires_grad + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize( + "fallback_kwargs, warning_match", + [ + ({"mscale": 1.25}, "mscale=1.25 is not supported by Triton fused mRoPE"), + ({"inverse": True}, "inverse RoPE is not supported by Triton fused mRoPE"), + ], +) +def test_apply_rotary_pos_emb_raw_mrope_fallbacks_match_unfused(fallback_kwargs, warning_match): + t, freqs, mrope_section = _make_inputs() + config = TransformerConfig( + num_attention_heads=t.shape[2], + num_layers=1, + apply_rope_fusion=True, + mrope_section=mrope_section, + ) + + with pytest.warns(UserWarning, match=warning_match): + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup(), **fallback_kwargs) + + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False, **fallback_kwargs) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + with warnings.catch_warnings(record=True) as repeated_warnings: + warnings.simplefilter("always") + out_again = apply_rotary_pos_emb( + t, freqs, config, cp_group=FakeCPGroup(), **fallback_kwargs + ) + assert not repeated_warnings + torch.testing.assert_close(ref.float(), out_again.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize( + "fallback_kwargs, config_kwargs, expected_warning_key, warning_text", + [ + ( + {"mscale": 1.25}, + {}, + "triton-mrope-mscale", + "mscale=1.25 is not supported by Triton fused mRoPE", + ), + ( + {"inverse": True}, + {}, + "triton-mrope-inverse", + "inverse RoPE is not supported by Triton fused mRoPE", + ), + ( + {"mla_rotary_interleaved": True}, + {}, + "triton-mrope-mla-rotary-interleaved", + "does not support MLA-style interleaving", + ), + ( + {}, + {"rotary_interleaved": True}, + "triton-mrope-unavailable-rotary-interleaved", + "rotary_interleaved=True is not supported", + ), + ], +) +def test_apply_rotary_pos_emb_raw_mrope_fallback_emits_single_warning( + fallback_kwargs, config_kwargs, expected_warning_key, warning_text +): + t, freqs, mrope_section = _make_inputs() + config = _make_mrope_config(t.shape[2], mrope_section, **config_kwargs) + + with warnings.catch_warnings(record=True) as recorded_warnings: + warnings.simplefilter("always") + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup(), **fallback_kwargs) + + fallback_warnings = _fallback_warnings(recorded_warnings) + assert len(fallback_warnings) == 1 + assert warning_text in str(fallback_warnings[0].message) + assert _ROPE_FUSION_FALLBACK_WARNINGS == {expected_warning_key} + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, rotary_interleaved=config.rotary_interleaved + ) + ref = _apply_rotary_pos_emb_bshd( + t, emb, rotary_interleaved=config.rotary_interleaved, **fallback_kwargs + ) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +def test_interleaved_mrope_rejects_inconsistent_sections(): + freqs = torch.randn(3, 2, 8, 8, dtype=torch.float32) + + with pytest.raises(AssertionError, match="interleaved mRoPE"): + mrope_freqs_to_rotary_emb(freqs, [2, 3, 3], interleaved_mrope=True) + + +def test_raw_mrope_cpu_falls_back_to_unfused(): + t = torch.randn(8, 1, 3, 20, dtype=torch.float32) + freqs = torch.randn(3, 1, 8, 8, dtype=torch.float32) + mrope_section = [2, 3, 3] + config = SimpleNamespace( + apply_rope_fusion=True, + mrope_section=mrope_section, + mrope_interleaved=False, + rotary_interleaved=False, + ) + + unavailable_reason = get_fused_mrope_unavailable_reason(t, freqs) + assert unavailable_reason is not None + with pytest.warns( + UserWarning, match="(CUDA tensors|Triton is not available).*Using unfused implementation" + ): + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup()) + assert _ROPE_FUSION_FALLBACK_WARNINGS in ( + {"triton-mrope-unavailable-device"}, + {"triton-mrope-unavailable-import"}, + ) + + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False) + torch.testing.assert_close(ref, out) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_raw_mrope_unsupported_dtype_falls_back_to_unfused(): + t, freqs, mrope_section = _make_inputs(dtype=torch.float64) + config = TransformerConfig( + num_attention_heads=t.shape[2], + num_layers=1, + apply_rope_fusion=True, + mrope_section=mrope_section, + ) + + assert "dtype" in get_fused_mrope_unavailable_reason(t, freqs) + with pytest.warns(UserWarning, match="dtype.*Using unfused implementation"): + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup()) + assert _ROPE_FUSION_FALLBACK_WARNINGS == {"triton-mrope-unavailable-dtype"} + + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False) + torch.testing.assert_close(ref, out) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +def test_apply_rotary_pos_emb_dispatches_raw_mrope(interleaved_mrope): + t, freqs, mrope_section = _make_inputs(interleaved_mrope=interleaved_mrope) + config = TransformerConfig( + num_attention_heads=t.shape[2], + num_layers=1, + apply_rope_fusion=True, + mrope_section=mrope_section, + mrope_interleaved=interleaved_mrope, + ) + + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup()) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_raw_mrope_unsupported_freq_dtype_warning_key_is_dtype(): + t, freqs, mrope_section = _make_inputs() + freqs = freqs.to(torch.float16) + config = TransformerConfig( + num_attention_heads=t.shape[2], + num_layers=1, + apply_rope_fusion=True, + mrope_section=mrope_section, + ) + + assert "float32" in get_fused_mrope_unavailable_reason(t, freqs) + with pytest.warns(UserWarning, match="float32.*Using unfused implementation"): + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup()) + assert _ROPE_FUSION_FALLBACK_WARNINGS == {"triton-mrope-unavailable-dtype"} + + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +def test_apply_rotary_pos_emb_raw_mrope_checks_triton_availability_once(monkeypatch): + t = torch.randn(4, 1, 2, 8, dtype=torch.float32) + freqs = torch.randn(3, 1, 4, 4, dtype=torch.float32) + config = SimpleNamespace( + apply_rope_fusion=True, + mrope_section=[1, 1, 2], + mrope_interleaved=False, + rotary_interleaved=False, + ) + + calls = 0 + + def fake_unavailable_reason(*args, **kwargs): + nonlocal calls + calls += 1 + return None + + monkeypatch.setattr(rope_utils, "get_fused_mrope_unavailable_reason", fake_unavailable_reason) + monkeypatch.setattr(rope_utils, "fused_apply_mrope", lambda *args, **kwargs: t + 1) + + out = apply_rotary_pos_emb(t, freqs, config, cp_group=FakeCPGroup()) + + assert calls == 1 + torch.testing.assert_close(out, t + 1) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.parametrize("layout", ["bshd", "thd"]) +def test_materialized_mrope_falls_back_without_te_fused_rope(monkeypatch, layout): + if layout == "bshd": + t, freqs, mrope_section = _make_inputs(batch=1) + cu_seqlens = None + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb", None) + else: + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb_thd", None) + + config = _make_mrope_config(t.shape[-2], mrope_section) + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + + with pytest.warns(UserWarning, match="Transformer Engine fused RoPE.*unavailable"): + out = apply_rotary_pos_emb(t, emb, config, cu_seqlens, cp_group=FakeCPGroup()) + + if layout == "bshd": + ref = _apply_rotary_pos_emb_bshd(t, emb, rotary_interleaved=False) + else: + ref = _apply_rotary_pos_emb_thd(t, cu_seqlens, emb, cp_group=FakeCPGroup()) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.parametrize( + "fallback_kwargs, expected_warning_key, warning_text", + [ + ( + {"mscale": 1.25}, + "te-rope-thd-mscale", + "mscale=1.25 is not supported by TE's fused RoPE for THD layout", + ), + ( + {"inverse": True}, + "te-rope-thd-inverse", + "inverse RoPE is not supported by TE's fused RoPE for THD layout", + ), + ( + {"mla_rotary_interleaved": True}, + "te-rope-thd-mla-rotary-interleaved", + "does not support MLA-style interleaving", + ), + ], +) +def test_materialized_thd_mrope_option_fallbacks_do_not_call_te( + monkeypatch, fallback_kwargs, expected_warning_key, warning_text +): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + config = _make_mrope_config(t.shape[1], mrope_section) + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + + def unexpected_te_thd_call(*args, **kwargs): + raise AssertionError("TE THD fused RoPE should not be called") + + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb_thd", unexpected_te_thd_call) + with warnings.catch_warnings(record=True) as recorded_warnings: + warnings.simplefilter("always") + out = apply_rotary_pos_emb( + t, emb, config, cu_seqlens, cp_group=FakeCPGroup(), **fallback_kwargs + ) + + fallback_warnings = _fallback_warnings(recorded_warnings) + assert len(fallback_warnings) == 1 + assert warning_text in str(fallback_warnings[0].message) + assert _ROPE_FUSION_FALLBACK_WARNINGS == {expected_warning_key} + + ref = _apply_rotary_pos_emb_thd(t, cu_seqlens, emb, cp_group=FakeCPGroup(), **fallback_kwargs) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +@pytest.mark.parametrize("cp_size, cp_rank", [(1, 0), (2, 0), (2, 1)]) +def test_fused_mrope_thd_matches_unfused_forward_backward( + interleaved_mrope, cp_size, cp_rank, monkeypatch +): + t_ref, freqs, cu_seqlens, mrope_section = _make_thd_inputs( + requires_grad=True, interleaved_mrope=interleaved_mrope, cp_size=cp_size + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + cp_group = FakeCPGroup(size=cp_size, rank=cp_rank) + config = TransformerConfig( + num_attention_heads=t_ref.shape[1], + num_layers=1, + context_parallel_size=cp_size, + apply_rope_fusion=True, + mrope_section=mrope_section, + mrope_interleaved=interleaved_mrope, + ) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_thd(t_ref, cu_seqlens, emb, cp_group=cp_group) + + fused_calls = 0 + orig_fused_apply_mrope_thd = rope_utils.fused_apply_mrope_thd + + def wrapped_fused_apply_mrope_thd(*args, **kwargs): + nonlocal fused_calls + fused_calls += 1 + return orig_fused_apply_mrope_thd(*args, **kwargs) + + def unexpected_pack(*args, **kwargs): + raise AssertionError("raw THD mRoPE fusion should not materialize packed freqs") + + monkeypatch.setattr(rope_utils, "fused_apply_mrope_thd", wrapped_fused_apply_mrope_thd) + monkeypatch.setattr(rope_utils, "_pack_thd_raw_mrope_freqs", unexpected_pack) + out = apply_rotary_pos_emb(t_fused, freqs, config, cu_seqlens, cp_group=cp_group) + assert fused_calls == 1 + + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +def test_apply_rotary_pos_emb_thd_eval_uses_triton_without_te(interleaved_mrope, monkeypatch): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs(interleaved_mrope=interleaved_mrope) + config = _make_mrope_config(t.shape[1], mrope_section, interleaved_mrope) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_thd(t, cu_seqlens, emb, cp_group=FakeCPGroup()) + + fused_calls = 0 + orig_fused_apply_mrope_thd = rope_utils.fused_apply_mrope_thd + + def wrapped_fused_apply_mrope_thd(*args, **kwargs): + nonlocal fused_calls + fused_calls += 1 + return orig_fused_apply_mrope_thd(*args, **kwargs) + + def unexpected_pack(*args, **kwargs): + raise AssertionError("raw THD mRoPE fusion should not materialize packed freqs") + + monkeypatch.setattr(rope_utils, "fused_apply_rotary_pos_emb_thd", None) + monkeypatch.setattr(rope_utils, "fused_apply_mrope_thd", wrapped_fused_apply_mrope_thd) + monkeypatch.setattr(rope_utils, "_pack_thd_raw_mrope_freqs", unexpected_pack) + with torch.no_grad(), warnings.catch_warnings(record=True) as recorded_warnings: + warnings.simplefilter("always") + out = apply_rotary_pos_emb(t, freqs, config, cu_seqlens, cp_group=FakeCPGroup()) + + assert fused_calls == 1 + assert not _fallback_warnings(recorded_warnings) + assert not out.requires_grad + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_apply_rotary_pos_emb_thd_fused_dispatch_does_not_read_cuda_scalars(monkeypatch): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + config = _make_mrope_config(t.shape[1], mrope_section) + + def unexpected_item(_tensor): + raise AssertionError("fused raw THD mRoPE dispatch should not call Tensor.item()") + + monkeypatch.setattr(torch.Tensor, "item", unexpected_item) + out = apply_rotary_pos_emb(t, freqs, config, cu_seqlens, cp_group=FakeCPGroup()) + + assert out.shape == t.shape + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize( + "fallback_kwargs, config_kwargs, expected_warning_key, warning_text", + [ + ( + {"mscale": 1.25}, + {}, + "triton-mrope-thd-mscale", + "mscale=1.25 is not supported by Triton fused mRoPE for THD layout", + ), + ( + {"inverse": True}, + {}, + "triton-mrope-thd-inverse", + "inverse RoPE is not supported by Triton fused mRoPE for THD layout", + ), + ( + {"mla_rotary_interleaved": True}, + {}, + "triton-mrope-thd-mla-rotary-interleaved", + "does not support MLA-style interleaving", + ), + ( + {}, + {"rotary_interleaved": True}, + "triton-mrope-thd-rotary-interleaved", + "currently supports rotary_interleaved=False", + ), + ], +) +def test_apply_rotary_pos_emb_thd_raw_mrope_fallback_emits_option_warning( + fallback_kwargs, config_kwargs, expected_warning_key, warning_text +): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + config = _make_mrope_config(t.shape[1], mrope_section, **config_kwargs) + + with warnings.catch_warnings(record=True) as recorded_warnings: + warnings.simplefilter("always") + out = apply_rotary_pos_emb( + t, freqs, config, cu_seqlens, cp_group=FakeCPGroup(), **fallback_kwargs + ) + + fallback_warnings = _fallback_warnings(recorded_warnings) + assert len(fallback_warnings) == 1 + assert warning_text in str(fallback_warnings[0].message) + assert _ROPE_FUSION_FALLBACK_WARNINGS == {expected_warning_key} + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, rotary_interleaved=config.rotary_interleaved + ) + ref = _apply_rotary_pos_emb_thd( + t, + cu_seqlens, + emb, + rotary_interleaved=config.rotary_interleaved, + cp_group=FakeCPGroup(), + **fallback_kwargs, + ) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_thd_raw_mrope_rejects_sequence_length_mismatch(): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + config = _make_mrope_config(t.shape[1], mrope_section) + bad_freqs = freqs[:, :, :-1, :].contiguous() + + with pytest.raises(ValueError, match="sequence length must match local tokens"): + apply_rotary_pos_emb(t, bad_freqs, config, cu_seqlens, cp_group=FakeCPGroup()) + + +def test_thd_raw_mrope_rejects_global_sequence_length_not_divisible_by_cp(): + t = torch.randn(2, 3, 20, dtype=torch.float32) + freqs = torch.randn(3, 1, 5, 8, dtype=torch.float32) + cu_seqlens = torch.tensor([0, 5], dtype=torch.int32) + config = SimpleNamespace( + apply_rope_fusion=True, + mrope_section=[2, 3, 3], + mrope_interleaved=False, + rotary_interleaved=False, + ) + + with pytest.raises(ValueError, match="divisible by context parallel size"): + apply_rotary_pos_emb(t, freqs, config, cu_seqlens, cp_group=FakeCPGroup(size=2)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_thd_raw_mrope_unavailable_reason_rejects_global_sequence_length_not_divisible_by_cp(): + t = torch.randn(3, 3, 20, dtype=torch.bfloat16, device="cuda") + freqs = torch.randn(3, 1, 5, 8, dtype=torch.float32, device="cuda") + cu_seqlens = torch.tensor([0, 5], dtype=torch.int32, device="cuda") + + assert "divisible by context parallel size" in get_fused_mrope_thd_unavailable_reason( + t, cu_seqlens, freqs, cp_size=2, cp_rank=0 + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("cp_rank", [0, 1]) +def test_thd_raw_mrope_cp_odd_local_sequence_lengths_match_manual_reference(cp_rank): + cp_size = 2 + t_ref, freqs, cu_seqlens, mrope_section = _make_thd_inputs( + requires_grad=True, cp_size=cp_size, padded_seq_lens=(10, 14) + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + config = _make_mrope_config(t_ref.shape[1], mrope_section) + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + + cu_seqlens_cpu = cu_seqlens.cpu().tolist() + _assert_thd_cp_freq_index_coverage(cu_seqlens_cpu, cp_size) + packed_freqs = emb[_thd_cp_freq_indices(cu_seqlens_cpu, cp_size, cp_rank)] + + ref = _apply_rotary_pos_emb_bshd(t_ref.unsqueeze(1), packed_freqs).squeeze(1) + out = apply_rotary_pos_emb( + t_fused, freqs, config, cu_seqlens, cp_group=FakeCPGroup(size=cp_size, rank=cp_rank) + ) + + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("cp_rank", [0, 1]) +def test_thd_raw_mrope_fallback_supports_odd_local_sequence_lengths(cp_rank): + cp_size = 2 + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs( + cp_size=cp_size, padded_seq_lens=(10, 14) + ) + config = _make_mrope_config(t.shape[1], mrope_section) + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + cu_seqlens_cpu = cu_seqlens.cpu().tolist() + packed_freqs = emb[_thd_cp_freq_indices(cu_seqlens_cpu, cp_size, cp_rank)] + + with pytest.warns(UserWarning, match="mscale=1.25.*Using unfused implementation"): + out = apply_rotary_pos_emb( + t, + freqs, + config, + cu_seqlens, + mscale=1.25, + cp_group=FakeCPGroup(size=cp_size, rank=cp_rank), + ) + + ref = _apply_rotary_pos_emb_bshd(t.unsqueeze(1), packed_freqs, mscale=1.25).squeeze(1) + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_thd_raw_mrope_rejects_batch_dimension_greater_than_one(): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + config = _make_mrope_config(t.shape[1], mrope_section) + bad_freqs = freqs.expand(-1, 2, -1, -1).contiguous() + + with pytest.raises(ValueError, match="singleton batch dimension"): + apply_rotary_pos_emb(t, bad_freqs, config, cu_seqlens, cp_group=FakeCPGroup()) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_thd_raw_mrope_rejects_non_thd_tensor_shape(): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + config = _make_mrope_config(t.shape[1], mrope_section) + + with pytest.raises(ValueError, match="raw mRoPE THD expects t"): + apply_rotary_pos_emb( + t[..., :8].unsqueeze(1), freqs, config, cu_seqlens, cp_group=FakeCPGroup() + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_fused_mrope_thd_public_api_matches_unfused(): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs() + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + + ref = _apply_rotary_pos_emb_thd(t, cu_seqlens, emb, cp_group=FakeCPGroup()) + assert ( + get_fused_mrope_thd_unavailable_reason(t, cu_seqlens, freqs, cp_size=1, cp_rank=0) is None + ) + out = fused_apply_mrope_thd(t, cu_seqlens, freqs, mrope_section) + + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("padded_seq_lens", [(28,), (8, 10, 10)]) +def test_fused_mrope_thd_matches_unfused_for_different_sequence_counts(padded_seq_lens): + t, freqs, cu_seqlens, mrope_section = _make_thd_inputs(padded_seq_lens=padded_seq_lens) + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + + ref = _apply_rotary_pos_emb_thd(t, cu_seqlens, emb, cp_group=FakeCPGroup()) + out = fused_apply_mrope_thd(t, cu_seqlens, freqs, mrope_section) + + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +def test_fused_mrope_thd_fp32_compute_matches_explicit_cast_forward_backward(): + t_ref, freqs, cu_seqlens, mrope_section = _make_thd_inputs(requires_grad=True) + t_fused = t_ref.detach().clone().requires_grad_(True) + emb = mrope_freqs_to_rotary_emb(freqs, mrope_section, rotary_interleaved=False) + + ref = _apply_rotary_pos_emb_thd(t_ref.float(), cu_seqlens, emb, cp_group=FakeCPGroup()).to( + t_ref.dtype + ) + out = fused_apply_mrope_thd(t_fused, cu_seqlens, freqs, mrope_section, fp32_compute=True) + + torch.testing.assert_close(ref.float(), out.float(), **_dtype_tols(t_ref.dtype)) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **_dtype_tols(t_ref.dtype)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.parametrize("return_raw_freqs", [False, True]) +def test_mrope_packed_seq_keeps_global_freqs_with_context_parallel(return_raw_freqs): + class FakeCPGroup2: + def size(self): + return 2 + + def rank(self): + return 0 + + seq = 16 + batch = 1 + head_dim = 20 + rotary_dim = 16 + mrope_section = [2, 3, 3] + cp_group = FakeCPGroup2() + position_ids = _make_position_ids(seq, batch) + rope = MultimodalRotaryEmbedding( + head_dim, rotary_percent=rotary_dim / head_dim, cp_group=cp_group + ) + + unpacked_freqs = rope( + position_ids, + mrope_section, + cp_group=cp_group, + return_raw_freqs=return_raw_freqs, + packed_seq=False, + ) + packed_freqs = rope( + position_ids, + mrope_section, + cp_group=cp_group, + return_raw_freqs=return_raw_freqs, + packed_seq=True, + ) + + seq_dim = 2 if return_raw_freqs else 0 + assert unpacked_freqs.shape[seq_dim] == seq // cp_group.size() + assert packed_freqs.shape[seq_dim] == seq + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.skipif(Utils.world_size < 2, reason="CP test requires at least 2 distributed ranks") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +def test_raw_mrope_fusion_matches_unfused_with_context_parallel(interleaved_mrope): + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=2) + try: + cp_group = parallel_state.get_context_parallel_group() + seq = 32 + batch = 2 + heads = 3 + head_dim = 20 + rotary_dim = 16 + mrope_section = [3, 3, 2] if interleaved_mrope else [2, 3, 3] + position_ids = _make_position_ids(seq, batch) + + rope = MultimodalRotaryEmbedding( + head_dim, + rotary_percent=rotary_dim / head_dim, + cp_group=cp_group, + interleaved_mrope=interleaved_mrope, + ) + raw_freqs = rope(position_ids, mrope_section, cp_group=cp_group, return_raw_freqs=True) + materialized_emb = rope(position_ids, mrope_section, cp_group=cp_group) + raw_freqs_emb = mrope_freqs_to_rotary_emb( + raw_freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + torch.testing.assert_close(raw_freqs_emb, materialized_emb) + + local_seq = seq // cp_group.size() + assert raw_freqs.shape == (3, batch, local_seq, rotary_dim // 2) + assert materialized_emb.shape == (local_seq, batch, 1, rotary_dim) + + generator = torch.Generator(device="cuda").manual_seed(4321) + t_ref = torch.randn( + local_seq, + batch, + heads, + head_dim, + dtype=torch.bfloat16, + device="cuda", + generator=generator, + requires_grad=True, + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + + config = TransformerConfig( + num_attention_heads=heads, + num_layers=1, + context_parallel_size=cp_group.size(), + apply_rope_fusion=True, + mrope_section=mrope_section, + mrope_interleaved=interleaved_mrope, + ) + + ref = _apply_rotary_pos_emb_bshd(t_ref, materialized_emb, rotary_interleaved=False) + out = apply_rotary_pos_emb(t_fused, raw_freqs, config, cp_group=cp_group) + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + finally: + Utils.destroy_model_parallel() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.skipif( + Utils.world_size < 2, reason="THD CP test requires at least 2 distributed ranks" +) +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +def test_raw_mrope_thd_fusion_matches_unfused_with_context_parallel(interleaved_mrope): + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=2) + try: + cp_group = parallel_state.get_context_parallel_group() + t_ref, _, cu_seqlens, mrope_section = _make_thd_inputs( + requires_grad=True, interleaved_mrope=interleaved_mrope, cp_size=cp_group.size() + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + config = TransformerConfig( + num_attention_heads=t_ref.shape[1], + num_layers=1, + context_parallel_size=cp_group.size(), + apply_rope_fusion=True, + mrope_section=mrope_section, + mrope_interleaved=interleaved_mrope, + ) + total_seq = int(cu_seqlens[-1].item()) + position_ids = _make_position_ids(total_seq, 1) + rope = MultimodalRotaryEmbedding( + t_ref.shape[-1], + rotary_percent=16 / t_ref.shape[-1], + cp_group=cp_group, + interleaved_mrope=interleaved_mrope, + ) + freqs = rope( + position_ids, mrope_section, cp_group=cp_group, return_raw_freqs=True, packed_seq=True + ) + emb = rope(position_ids, mrope_section, cp_group=cp_group, packed_seq=True) + + raw_freqs_emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + assert freqs.shape == (3, 1, total_seq, 8) + assert emb.shape == (total_seq, 1, 1, 16) + torch.testing.assert_close(raw_freqs_emb, emb) + + ref = _apply_rotary_pos_emb_thd(t_ref, cu_seqlens, emb, cp_group=cp_group) + out = apply_rotary_pos_emb(t_fused, freqs, config, cu_seqlens, cp_group=cp_group) + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + finally: + Utils.destroy_model_parallel() + + +# --------------------------------------------------------------------------- +# Real Qwen3.5-VL deployment shapes. +# +# The parametrized tests above use head_dim=16/20 with rotary_dim=16 (~80% of +# channels rotated). The real Qwen3.5-VL config is head_dim=256 with +# rotary_percent=0.25 -> rotary_dim=64 (only 25% rotated, 75% pass-through) and +# mrope_section=[11,11,10] (interleaved). Exercise those exact shapes so a kernel +# regression in the large-pass-through / large-section regime is caught. +# --------------------------------------------------------------------------- + +# (head_dim, rotary_dim, mrope_section, interleaved_mrope) +_REAL_BSHD_SHAPES = [ + (256, 64, [11, 11, 10], True), # Qwen3.5-VL LLM decoder (75% pass-through) + (256, 64, [10, 11, 11], False), # same, section (non-interleaved) layout + (256, 256, [43, 43, 42], True), # full rotary (no pass-through) +] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("head_dim,rotary_dim,mrope_section,interleaved_mrope", _REAL_BSHD_SHAPES) +def test_fused_mrope_matches_unfused_real_shapes( + head_dim, rotary_dim, mrope_section, interleaved_mrope +): + t_ref, freqs, mrope_section = _make_inputs( + requires_grad=True, + head_dim=head_dim, + rotary_dim=rotary_dim, + mrope_section=mrope_section, + interleaved_mrope=interleaved_mrope, + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_bshd(t_ref, emb, rotary_interleaved=False) + out = fused_apply_mrope( + t_fused, freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.skipif(not is_fused_mrope_available(), reason="Triton fused mRoPE not available") +@pytest.mark.parametrize("interleaved_mrope", [False, True]) +def test_fused_mrope_thd_matches_unfused_real_shapes(interleaved_mrope): + # Real Qwen3.5-VL head_dim=256, rotary_dim=64 in THD packed layout. + section = [11, 11, 10] if interleaved_mrope else [10, 11, 11] + t_ref, freqs, cu_seqlens, mrope_section = _make_thd_inputs( + requires_grad=True, + interleaved_mrope=interleaved_mrope, + head_dim=256, + rotary_dim=64, + mrope_section=section, + ) + t_fused = t_ref.detach().clone().requires_grad_(True) + cp_group = FakeCPGroup(size=1, rank=0) + + emb = mrope_freqs_to_rotary_emb( + freqs, mrope_section, interleaved_mrope=interleaved_mrope, rotary_interleaved=False + ) + ref = _apply_rotary_pos_emb_thd(t_ref, cu_seqlens, emb, cp_group=cp_group) + out = fused_apply_mrope_thd( + t_fused, cu_seqlens, freqs, mrope_section, + interleaved_mrope=interleaved_mrope, rotary_interleaved=False, cp_size=1, cp_rank=0, + ) + + tols = _dtype_tols(t_ref.dtype) + torch.testing.assert_close(ref.float(), out.float(), **tols) + + grad = torch.randn_like(ref) + ref.backward(grad) + out.backward(grad) + torch.testing.assert_close(t_ref.grad.float(), t_fused.grad.float(), **tols) + + +def test_thd_unavailable_reason_rejects_non_cp_divisible_subsequence(): + # Per-sequence CP divisibility: total length is divisible by cp_size but an + # individual packed sub-sequence is not. The fused THD launch path + # (apply_rotary_pos_emb -> fused_apply_mrope_thd) must reject this so it falls + # back to the unfused path (which splits per-sequence correctly), instead of + # silently computing wrong CP token indices. + cp_size = 2 + # sub-sequence lengths 10 and 14 -> both even (OK); 9 and 15 -> total 24 even + # but each odd (must be rejected). + cu_seqlens = torch.tensor([0, 9, 24], dtype=torch.int32, device="cuda") + local_tokens = 24 // cp_size + t = torch.randn(local_tokens, 3, 20, dtype=torch.bfloat16, device="cuda") + freqs = torch.randn(3, 1, 24, 8, dtype=torch.float32, device="cuda") + reason = get_fused_mrope_thd_unavailable_reason( + t, cu_seqlens, freqs, rotary_interleaved=False, cp_size=cp_size, cp_rank=0 + ) + assert reason is not None and "sub-sequence" in reason, reason + + # Control: all sub-sequences divisible by cp_size -> launchable (reason None). + cu_ok = torch.tensor([0, 10, 24], dtype=torch.int32, device="cuda") + reason_ok = get_fused_mrope_thd_unavailable_reason( + t, cu_ok, freqs, rotary_interleaved=False, cp_size=cp_size, cp_rank=0 + ) + assert reason_ok is None, reason_ok