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
23 changes: 23 additions & 0 deletions tests/compile/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ def test_something(mock_cuda_platform):
def _mock_platform(is_cuda: bool = True, capability: tuple[int, int] | None = None):
mock_platform = MagicMock()
mock_platform.is_cuda.return_value = is_cuda
mock_platform.is_xpu.return_value = False
device_capability = (
DeviceCapability(*capability) if capability is not None else None
)
Expand All @@ -46,3 +47,25 @@ def is_device_capability_family(
yield mock_platform

return _mock_platform


@pytest.fixture
def mock_xpu_platform():
"""
Fixture that returns a factory for creating mocked XPU platforms.

Usage:
def test_something(mock_xpu_platform):
with mock_xpu_platform():
# test code
"""

@contextmanager
def _mock_platform():
mock_platform = MagicMock()
mock_platform.is_cuda.return_value = False
mock_platform.is_xpu.return_value = True
with patch("vllm.platforms.current_platform", mock_platform):
yield mock_platform

return _mock_platform
82 changes: 82 additions & 0 deletions tests/compile/test_sequence_parallelism_threshold.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,3 +108,85 @@ def test_hidden_size_boundary(self, mock_cuda_platform):
element_size=2,
)
assert result is not None


# XPU-specific constants (must match sequence_parallelism.py values)
_XPU_MIN_HIDDEN_SIZE = 4096
_XPU_MIN_PER_GPU_SIZE_MB = 8.0


class TestGetSequenceParallelismThresholdXPU:
"""Tests for get_sequence_parallelism_threshold on XPU platform."""

def test_xpu_small_hidden_size_returns_none(self, mock_xpu_platform):
"""XPU with hidden_size below threshold should return None."""
with mock_xpu_platform():
result = get_sequence_parallelism_threshold(
hidden_size=_XPU_MIN_HIDDEN_SIZE - 1,
tp_size=2,
element_size=2,
)
assert result is None

def test_xpu_large_model_returns_threshold(self, mock_xpu_platform):
"""XPU with hidden_size >= threshold should return calculated value."""
with mock_xpu_platform():
hidden_size = _XPU_MIN_HIDDEN_SIZE
tp_size = 2
element_size = 2
result = get_sequence_parallelism_threshold(
hidden_size=hidden_size,
tp_size=tp_size,
element_size=element_size,
)
# (8 * 2 * 1024 * 1024) // (4096 * 2) = 2048
MiB = 1024 * 1024
expected = int(
(_XPU_MIN_PER_GPU_SIZE_MB * tp_size * MiB) // (hidden_size * element_size)
)
assert result == expected
assert result == 2048

@pytest.mark.parametrize(
"hidden_size,tp_size,element_size,expected",
[
# (8 * 1 * 1024 * 1024) // (4096 * 2) = 1024
(4096, 1, 2, 1024),
# (8 * 4 * 1024 * 1024) // (4096 * 2) = 4096
(4096, 4, 2, 4096),
# (8 * 2 * 1024 * 1024) // (8192 * 2) = 1024
(8192, 2, 2, 1024),
# (8 * 2 * 1024 * 1024) // (4096 * 4) = 1024
(4096, 2, 4, 1024),
],
)
def test_xpu_threshold_calculation_variations(
self, mock_xpu_platform, hidden_size, tp_size, element_size, expected
):
"""Test XPU threshold calculation with various parameter combinations."""
with mock_xpu_platform():
result = get_sequence_parallelism_threshold(
hidden_size=hidden_size,
tp_size=tp_size,
element_size=element_size,
)
assert result == expected

def test_xpu_hidden_size_boundary(self, mock_xpu_platform):
"""Test behavior at the exact XPU hidden_size boundary."""
with mock_xpu_platform():
# Just below threshold
result = get_sequence_parallelism_threshold(
hidden_size=_XPU_MIN_HIDDEN_SIZE - 1,
tp_size=2,
element_size=2,
)
assert result is None

# Exactly at threshold
result = get_sequence_parallelism_threshold(
hidden_size=_XPU_MIN_HIDDEN_SIZE,
tp_size=2,
element_size=2,
)
assert result is not None
37 changes: 20 additions & 17 deletions vllm/compilation/passes/fusion/sequence_parallelism.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,24 +72,27 @@ def get_sequence_parallelism_threshold(
"""
from vllm.platforms import current_platform

if not current_platform.is_cuda():
return None

capability = current_platform.get_device_capability()
if capability is None:
return None

# Collapse Blackwell variants (sm100/sm103/...) into one policy bucket.
if current_platform.is_device_capability_family(100):
device_capability = 100
if current_platform.is_xpu():
min_hidden_size = 4096
min_per_gpu_size_mb = 8.0
elif current_platform.is_cuda():
capability = current_platform.get_device_capability()
if capability is None:
return None

# Collapse Blackwell variants (sm100/sm103/...) into one policy bucket.
if current_platform.is_device_capability_family(100):
device_capability = 100
else:
device_capability = capability.to_int()

# Check if device has configured thresholds
_hidden = SP_MIN_HIDDEN_SIZE.get(device_capability)
_gpu_mb = SP_MIN_PER_GPU_SIZE_MB.get(device_capability)
if _hidden is None or _gpu_mb is None:
return None
min_hidden_size, min_per_gpu_size_mb = _hidden, _gpu_mb
else:
device_capability = capability.to_int()

# Check if device has configured thresholds
min_hidden_size = SP_MIN_HIDDEN_SIZE.get(device_capability)
min_per_gpu_size_mb = SP_MIN_PER_GPU_SIZE_MB.get(device_capability)

if min_hidden_size is None or min_per_gpu_size_mb is None:
return None

# Only apply sequence parallelism for models meeting the size threshold
Expand Down
4 changes: 3 additions & 1 deletion vllm/compilation/passes/pass_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,9 @@
RocmAiterTritonAddRMSNormPadFusionPass,
)

if current_platform.is_cuda_alike() or current_platform.is_xpu():
from .fusion.sequence_parallelism import SequenceParallelismPass

Comment thread
chaojun-zhang marked this conversation as resolved.
if current_platform.is_cuda_alike():
from .fusion.act_quant_fusion import ActivationQuantFusionPass
from .fusion.attn_quant_fusion import AttnQuantFusionPass
Expand All @@ -37,7 +40,6 @@
from .fusion.qk_norm_rope_fusion import QKNormRoPEFusionPass
from .fusion.rms_quant_fusion import RMSNormQuantFusionPass
from .fusion.rope_kvcache_fusion import RopeKVCacheFusionPass
from .fusion.sequence_parallelism import SequenceParallelismPass
from .utility.scatter_split_replace import ScatterSplitReplacementPass
from .utility.split_coalescing import SplitCoalescingPass

Expand Down
1 change: 0 additions & 1 deletion vllm/platforms/xpu.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,7 +208,6 @@ def check_and_update_config(cls, vllm_config: VllmConfig) -> None:

pass_config = compilation_config.pass_config
fusion_passes_to_disable = {
"enable_sp": "Sequence parallelism",
"fuse_gemm_comms": "Async TP",
"fuse_allreduce_rms": "AllReduce + RMSNorm fusion",
"fuse_attn_quant": "Attention + quant fusion",
Expand Down
Loading