Files
wehub-resource-sync eec33d25b2
pre-commit / pre-commit (push) Failing after 1s
Build Wheel / build (3.11) (push) Failing after 1s
Build Wheel / build (3.12) (push) Failing after 0s
chore: import upstream snapshot with attribution
2026-07-13 12:29:08 +08:00

387 lines
15 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for HSDP (Hybrid Sharded Data Parallel) configuration and utilities.
These tests verify HSDP configuration logic without requiring a distributed environment.
"""
import pytest
import torch.nn as nn
from vllm_omni.diffusion.data import DiffusionParallelConfig
from vllm_omni.diffusion.distributed.hsdp import HSDPInferenceConfig
pytestmark = [pytest.mark.diffusion, pytest.mark.parallel, pytest.mark.cpu, pytest.mark.core_model]
class TestHSDPInferenceConfig:
"""Tests for HSDPInferenceConfig dataclass."""
def test_default_values(self):
"""Test default configuration values."""
config = HSDPInferenceConfig()
assert config.enabled is False
assert config.hsdp_replicate_size == 1
assert config.hsdp_shard_size == -1
def test_custom_values(self):
"""Test custom configuration values."""
config = HSDPInferenceConfig(
enabled=True,
hsdp_replicate_size=2,
hsdp_shard_size=4,
)
assert config.enabled is True
assert config.hsdp_replicate_size == 2
assert config.hsdp_shard_size == 4
class TestDiffusionParallelConfigHSDP:
"""Tests for HSDP settings in DiffusionParallelConfig."""
def test_hsdp_disabled_by_default(self):
"""HSDP should be disabled by default."""
config = DiffusionParallelConfig()
assert config.use_hsdp is False
assert config.hsdp_shard_size == -1
assert config.hsdp_replicate_size == 1
def test_hsdp_auto_shard_size(self):
"""Test auto-calculation of hsdp_shard_size when use_hsdp=True."""
# ulysses_degree=4 -> world_size=4
# hsdp_shard_size should be auto-calculated as 4 // 1 = 4
config = DiffusionParallelConfig(
ulysses_degree=4,
use_hsdp=True,
)
assert config.world_size == 4
assert config.hsdp_shard_size == 4
assert config.hsdp_replicate_size == 1
def test_hsdp_auto_shard_size_fails_standalone(self):
"""Test that auto-calculate fails when other parallelism is all 1."""
# When all other parallelism is 1, cannot auto-calculate
# User must specify hsdp_shard_size explicitly
with pytest.raises(ValueError, match="Cannot auto-calculate hsdp_shard_size"):
DiffusionParallelConfig(
use_hsdp=True,
# All other parallelism defaults to 1
)
def test_hsdp_standalone_mode(self):
"""Test standalone HSDP (HSDP without other parallelism)."""
# Standalone HSDP: all other parallelism=1, explicit shard_size
config = DiffusionParallelConfig(
use_hsdp=True,
hsdp_shard_size=4, # Explicit shard size
hsdp_replicate_size=1,
)
# world_size should be determined by HSDP
assert config.world_size == 4
assert config.hsdp_shard_size == 4
assert config.hsdp_replicate_size == 1
def test_hsdp_standalone_with_replicate(self):
"""Test standalone HSDP with replication."""
config = DiffusionParallelConfig(
use_hsdp=True,
hsdp_shard_size=4,
hsdp_replicate_size=2,
)
# world_size = shard_size * replicate_size
assert config.world_size == 8
assert config.hsdp_shard_size == 4
assert config.hsdp_replicate_size == 2
def test_hsdp_with_replicate(self):
"""Test HSDP with replication (hybrid mode) combined with other parallelism."""
# world_size=8, replicate=2 -> shard_size should be 4
config = DiffusionParallelConfig(
ulysses_degree=8,
use_hsdp=True,
hsdp_replicate_size=2,
)
assert config.world_size == 8
assert config.hsdp_shard_size == 4
assert config.hsdp_replicate_size == 2
def test_hsdp_explicit_shard_size_valid(self):
"""Test explicit hsdp_shard_size that matches world_size."""
config = DiffusionParallelConfig(
ulysses_degree=4,
use_hsdp=True,
hsdp_shard_size=4,
hsdp_replicate_size=1,
)
assert config.hsdp_shard_size == 4
def test_hsdp_explicit_shard_size_invalid(self):
"""Test that invalid HSDP dimensions raise an error when combined with other parallelism."""
with pytest.raises(ValueError, match="HSDP dimensions"):
DiffusionParallelConfig(
ulysses_degree=4, # world_size=4
use_hsdp=True,
hsdp_shard_size=3, # 1 * 3 != 4
hsdp_replicate_size=1,
)
def test_hsdp_replicate_size_exceeds_world_size(self):
"""Test that replicate_size > world_size raises an error."""
with pytest.raises(ValueError, match="replicate_size.*must evenly divide world_size"):
DiffusionParallelConfig(
ulysses_degree=4, # world_size=4
use_hsdp=True,
hsdp_replicate_size=8, # 8 > 4, invalid
)
def test_hsdp_combined_world_size(self):
"""Test that combined HSDP matches other parallelism world_size."""
config_no_hsdp = DiffusionParallelConfig(ulysses_degree=4)
config_with_hsdp = DiffusionParallelConfig(
ulysses_degree=4,
use_hsdp=True,
hsdp_shard_size=4,
)
# When combined with other parallelism, world_size should match
assert config_no_hsdp.world_size == config_with_hsdp.world_size == 4
def test_hsdp_standalone_world_size(self):
"""Test that standalone HSDP determines world_size."""
config_hsdp = DiffusionParallelConfig(
use_hsdp=True,
hsdp_shard_size=8,
)
# Standalone HSDP: world_size is determined by HSDP
assert config_hsdp.world_size == 8
def test_hsdp_cannot_use_with_tp(self):
"""Test that HSDP and Tensor Parallelism cannot be used together."""
with pytest.raises(ValueError, match="cannot be used with TP or DP"):
DiffusionParallelConfig(
tensor_parallel_size=2,
use_hsdp=True,
hsdp_shard_size=4,
)
def test_hsdp_cannot_use_with_dp(self):
"""Test that HSDP and Data Parallelism cannot be used together."""
with pytest.raises(ValueError, match="cannot be used with TP or DP"):
DiffusionParallelConfig(
data_parallel_size=2,
use_hsdp=True,
hsdp_shard_size=4,
)
def test_from_dict_with_hsdp(self):
"""Test creating config from dict with HSDP settings."""
config = DiffusionParallelConfig.from_dict(
{
"ulysses_degree": 4,
"use_hsdp": True,
"hsdp_replicate_size": 2,
}
)
assert config.use_hsdp is True
assert config.hsdp_replicate_size == 2
assert config.hsdp_shard_size == 2 # auto: 4 // 2
class TestStandaloneHSDPDetection:
"""Tests for standalone HSDP detection and dit_parallel_size calculation.
These tests verify the logic used in initialize_model_parallel() to detect
standalone HSDP mode and calculate effective parallel sizes.
Standalone HSDP is when all non-HSDP parallelism dimensions are 1.
"""
@staticmethod
def compute_standalone_hsdp_params(
data_parallel_size: int = 1,
cfg_parallel_size: int = 1,
sequence_parallel_size: int = 1,
pipeline_parallel_size: int = 1,
tensor_parallel_size: int = 1,
fully_shard_degree: int = 1,
hsdp_replicate_size: int = 1,
) -> dict:
"""Compute standalone HSDP detection parameters.
This mirrors the logic in initialize_model_parallel().
"""
dit_parallel_size = (
data_parallel_size
* cfg_parallel_size
* sequence_parallel_size
* pipeline_parallel_size
* tensor_parallel_size
)
# Check for standalone HSDP: all non-HSDP parallelism dimensions are 1
is_standalone_hsdp = dit_parallel_size == 1 and fully_shard_degree > 1
# For standalone HSDP: use (fully_shard_degree * hsdp_replicate_size)
if is_standalone_hsdp:
effective_dit_parallel_size = fully_shard_degree * hsdp_replicate_size
else:
effective_dit_parallel_size = dit_parallel_size
effective_dp_size = (fully_shard_degree * hsdp_replicate_size) if is_standalone_hsdp else data_parallel_size
return {
"original_dit_parallel_size": dit_parallel_size,
"is_standalone_hsdp": is_standalone_hsdp,
"effective_dit_parallel_size": effective_dit_parallel_size,
"effective_dp_size": effective_dp_size,
}
def test_standalone_hsdp_basic(self):
"""Test basic standalone HSDP detection (shard_size=4, replicate=1)."""
result = self.compute_standalone_hsdp_params(
fully_shard_degree=4,
hsdp_replicate_size=1,
)
assert result["original_dit_parallel_size"] == 1
assert result["is_standalone_hsdp"] is True
assert result["effective_dit_parallel_size"] == 4
assert result["effective_dp_size"] == 4
def test_standalone_hsdp_with_replicate(self):
"""Test standalone HSDP with replication (shard_size=4, replicate=2)."""
result = self.compute_standalone_hsdp_params(
fully_shard_degree=4,
hsdp_replicate_size=2,
)
assert result["original_dit_parallel_size"] == 1
assert result["is_standalone_hsdp"] is True
assert result["effective_dit_parallel_size"] == 8 # 4 * 2
assert result["effective_dp_size"] == 8
def test_combined_hsdp_sp_not_standalone(self):
"""Test HSDP combined with SP is NOT detected as standalone.
This is a regression test for the bug where the condition
`dit_parallel_size == fully_shard_degree` incorrectly matched
combined modes like SP=4 + HSDP=4.
"""
result = self.compute_standalone_hsdp_params(
sequence_parallel_size=4,
fully_shard_degree=4,
hsdp_replicate_size=1,
)
assert result["original_dit_parallel_size"] == 4
assert result["is_standalone_hsdp"] is False
# Should NOT override dp_size for combined mode
assert result["effective_dp_size"] == 1 # original data_parallel_size
def test_combined_hsdp_cfg_not_standalone(self):
"""Test HSDP combined with CFG is NOT detected as standalone."""
result = self.compute_standalone_hsdp_params(
cfg_parallel_size=2,
fully_shard_degree=4,
hsdp_replicate_size=1,
)
assert result["original_dit_parallel_size"] == 2
assert result["is_standalone_hsdp"] is False
assert result["effective_dp_size"] == 1
def test_combined_hsdp_pp_not_standalone(self):
"""Test HSDP combined with PP is NOT detected as standalone."""
result = self.compute_standalone_hsdp_params(
pipeline_parallel_size=2,
fully_shard_degree=4,
hsdp_replicate_size=1,
)
assert result["original_dit_parallel_size"] == 2
assert result["is_standalone_hsdp"] is False
assert result["effective_dp_size"] == 1
def test_no_hsdp_not_standalone(self):
"""Test that no HSDP (fully_shard_degree=1) is NOT standalone."""
result = self.compute_standalone_hsdp_params(
fully_shard_degree=1,
)
assert result["original_dit_parallel_size"] == 1
assert result["is_standalone_hsdp"] is False
assert result["effective_dp_size"] == 1
def test_combined_multiple_parallelism_not_standalone(self):
"""Test HSDP combined with multiple parallelism is NOT standalone."""
result = self.compute_standalone_hsdp_params(
sequence_parallel_size=2,
cfg_parallel_size=2,
fully_shard_degree=4,
hsdp_replicate_size=1,
)
assert result["original_dit_parallel_size"] == 4 # 2 * 2
assert result["is_standalone_hsdp"] is False
assert result["effective_dp_size"] == 1
def test_standalone_hsdp_large_shard(self):
"""Test standalone HSDP with large shard size."""
result = self.compute_standalone_hsdp_params(
fully_shard_degree=8,
hsdp_replicate_size=1,
)
assert result["is_standalone_hsdp"] is True
assert result["effective_dit_parallel_size"] == 8
assert result["effective_dp_size"] == 8
def test_standalone_hsdp_large_replicate(self):
"""Test standalone HSDP with large replicate size."""
result = self.compute_standalone_hsdp_params(
fully_shard_degree=4,
hsdp_replicate_size=4,
)
assert result["is_standalone_hsdp"] is True
assert result["effective_dit_parallel_size"] == 16 # 4 * 4
assert result["effective_dp_size"] == 16
class TestHSDPShardConditions:
"""Tests for _hsdp_shard_conditions matching logic."""
@staticmethod
def _is_transformer_block(name: str, module: nn.Module) -> bool:
"""Example shard condition matching transformer blocks."""
return "blocks" in name and name.split(".")[-1].isdigit()
def test_condition_matches_blocks(self):
"""Test that condition matches transformer block patterns."""
cond = self._is_transformer_block
# Should match
assert cond("blocks.0", nn.Linear(10, 10)) is True
assert cond("blocks.15", nn.Linear(10, 10)) is True
assert cond("transformer.blocks.0", nn.Linear(10, 10)) is True
# Should not match
assert cond("blocks", nn.Linear(10, 10)) is False
assert cond("blocks.norm", nn.Linear(10, 10)) is False
assert cond("embeddings", nn.Linear(10, 10)) is False
def test_model_with_shard_conditions(self):
"""Test model with _hsdp_shard_conditions attribute."""
class MockModel(nn.Module):
@staticmethod
def _is_block(name: str, module: nn.Module) -> bool:
return name.startswith("blocks.") and name.split(".")[-1].isdigit()
_hsdp_shard_conditions = [_is_block]
def __init__(self):
super().__init__()
self.blocks = nn.ModuleList([nn.Linear(10, 10) for _ in range(2)])
model = MockModel()
conditions = getattr(model, "_hsdp_shard_conditions", None)
assert conditions is not None
assert len(conditions) == 1
# Verify conditions work on actual model modules
matched = []
for name, module in model.named_modules():
if any(cond(name, module) for cond in conditions):
matched.append(name)
assert "blocks.0" in matched
assert "blocks.1" in matched