Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions slime_plugins/models/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,7 @@ def forward(
initial_state=None,
output_final_state=False,
use_qk_l2norm_in_kernel=True,
cu_seqlens=cu_seqlens,
)

z_shape_og = z.shape
Expand Down
1 change: 1 addition & 0 deletions slime_plugins/models/qwen3_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,7 @@ def forward(
initial_state=None,
output_final_state=False,
use_qk_l2norm_in_kernel=True,
cu_seqlens=cu_seqlens,
)

z_shape_og = z.shape
Expand Down
165 changes: 165 additions & 0 deletions tests/test_qwen3_linear_attention_cu_seqlens.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
from __future__ import annotations

import importlib
import sys
import types
from types import SimpleNamespace

import pytest
import torch
import torch.nn as nn


def install_megatron_stubs() -> None:
if "megatron" in sys.modules:
return

megatron_mod = types.ModuleType("megatron")
core_mod = types.ModuleType("megatron.core")
models_mod = types.ModuleType("megatron.core.models")
gpt_mod = types.ModuleType("megatron.core.models.gpt")
gpt_layer_specs_mod = types.ModuleType("megatron.core.models.gpt.gpt_layer_specs")
inference_mod = types.ModuleType("megatron.core.inference")
inference_contexts_mod = types.ModuleType("megatron.core.inference.contexts")
packed_seq_mod = types.ModuleType("megatron.core.packed_seq_params")
transformer_mod = types.ModuleType("megatron.core.transformer")
transformer_module_mod = types.ModuleType("megatron.core.transformer.module")
spec_utils_mod = types.ModuleType("megatron.core.transformer.spec_utils")
transformer_block_mod = types.ModuleType("megatron.core.transformer.transformer_block")
transformer_layer_mod = types.ModuleType("megatron.core.transformer.transformer_layer")

class PackedSeqParams:
def __init__(self, **kwargs):
for key, value in kwargs.items():
setattr(self, key, value)

class MegatronModule(nn.Module):
def __init__(self, config=None):
super().__init__()
self.config = config

class ModuleSpec:
def __init__(self, module=None, params=None):
self.module = module
self.params = params or {}

mpu_stub = types.SimpleNamespace(
get_context_parallel_world_size=lambda: 1,
get_context_parallel_group=lambda: None,
get_context_parallel_rank=lambda: 0,
get_tensor_model_parallel_group=lambda: None,
)
tensor_parallel_stub = types.SimpleNamespace(
gather_from_sequence_parallel_region=lambda x, group=None: x,
scatter_to_sequence_parallel_region=lambda x, group=None: x,
)

gpt_layer_specs_mod.get_gpt_decoder_block_spec = lambda *args, **kwargs: None
inference_contexts_mod.BaseInferenceContext = type("BaseInferenceContext", (), {})
packed_seq_mod.PackedSeqParams = PackedSeqParams
transformer_module_mod.MegatronModule = MegatronModule
spec_utils_mod.ModuleSpec = ModuleSpec
transformer_block_mod.get_num_layers_to_build = lambda *args, **kwargs: 0
transformer_layer_mod.get_transformer_layer_offset = lambda *args, **kwargs: 0

core_mod.mpu = mpu_stub
core_mod.tensor_parallel = tensor_parallel_stub

sys.modules["megatron"] = megatron_mod
sys.modules["megatron.core"] = core_mod
sys.modules["megatron.core.models"] = models_mod
sys.modules["megatron.core.models.gpt"] = gpt_mod
sys.modules["megatron.core.models.gpt.gpt_layer_specs"] = gpt_layer_specs_mod
sys.modules["megatron.core.inference"] = inference_mod
sys.modules["megatron.core.inference.contexts"] = inference_contexts_mod
sys.modules["megatron.core.packed_seq_params"] = packed_seq_mod
sys.modules["megatron.core.transformer"] = transformer_mod
sys.modules["megatron.core.transformer.module"] = transformer_module_mod
sys.modules["megatron.core.transformer.spec_utils"] = spec_utils_mod
sys.modules["megatron.core.transformer.transformer_block"] = transformer_block_mod
sys.modules["megatron.core.transformer.transformer_layer"] = transformer_layer_mod


class FakeShortConvolution(nn.Module):
def __init__(self, *args, **kwargs):
super().__init__()

def forward(self, x, cu_seqlens=None, **kwargs):
return x, None


class FakeFusedRMSNormGated(nn.Module):
def __init__(self, *args, **kwargs):
super().__init__()

def forward(self, x, z):
return x


def make_config() -> SimpleNamespace:
return SimpleNamespace(
hidden_size=32,
linear_num_value_heads=4,
linear_num_key_heads=2,
linear_key_head_dim=8,
linear_value_head_dim=8,
linear_conv_kernel_dim=4,
hidden_act="silu",
rms_norm_eps=1e-6,
dtype=torch.float32,
)


def load_module(module_name: str):
install_megatron_stubs()
sys.modules.pop("slime_plugins.models.hf_attention", None)
sys.modules.pop(module_name, None)
return importlib.import_module(module_name)


@pytest.mark.unit
@pytest.mark.parametrize(
("module_name", "class_name"),
[
("slime_plugins.models.qwen3_5", "Qwen3_5GatedDeltaNet"),
("slime_plugins.models.qwen3_next", "Qwen3NextGatedDeltaNet"),
],
)
def test_linear_attention_forwards_cu_seqlens_to_chunk_kernel(monkeypatch, module_name: str, class_name: str):
module = load_module(module_name)

monkeypatch.setattr(module.torch.cuda, "current_device", lambda: "cpu")
monkeypatch.setattr(module, "ShortConvolution", FakeShortConvolution)
monkeypatch.setattr(module, "FusedRMSNormGated", FakeFusedRMSNormGated)

chunk_calls = []

def fake_chunk_gated_delta_rule(
q,
k,
v,
*,
g,
beta,
initial_state,
output_final_state,
use_qk_l2norm_in_kernel,
cu_seqlens=None,
**kwargs,
):
chunk_calls.append(cu_seqlens.clone() if cu_seqlens is not None else None)
assert q.shape[0] == 1
assert cu_seqlens is not None
return torch.zeros_like(v), None

monkeypatch.setattr(module, "chunk_gated_delta_rule", fake_chunk_gated_delta_rule)

layer = getattr(module, class_name)(make_config(), layer_idx=0)
hidden_states = torch.randn(1, 7, 32)
cu_seqlens = torch.tensor([0, 3, 7], dtype=torch.int32)

output = layer(hidden_states, cu_seqlens=cu_seqlens)

assert output.shape == hidden_states.shape
assert len(chunk_calls) == 1
assert torch.equal(chunk_calls[0], cu_seqlens)