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
Original file line number Diff line number Diff line change
Expand Up @@ -535,7 +535,7 @@ def test_upsample_forward_only_fuses_nearest_2x() -> None:


def test_rms_norm_vae_substitute_is_not_matched() -> None:
from vllm_omni.diffusion.layers.norm import RMSNormVAE
from vllm_omni.diffusion.models.wan2_2.norm import RMSNormVAE

assert not fastpath_forwards.is_diffusers_rms_norm(RMSNormVAE(8, images=False))
_, vae = _build_pair(TINY_RESIDUAL, torch.float32)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -195,6 +195,17 @@ def test_no_propagation_when_tf_quant_config_is_none(self, mocker: MockerFixture
class TestPatchWanRmsNorm:
"""Test that patch_wan_rms_norm doesn't raise on concurrent module registration."""

@pytest.fixture(autouse=True)
def restore_wan_rms_norm(self):
originals = [
(module, module.__dict__["WanRMS_norm"])
for module in list(sys.modules.values())
if getattr(module, "__dict__", None) is not None and "WanRMS_norm" in module.__dict__
]
yield
for module, original in originals:
module.__dict__["WanRMS_norm"] = original

def test_patches_modules_with_wan_rms_norm(self):
from vllm_omni.diffusion.models.wan2_2.norm import RMSNormVAE
from vllm_omni.diffusion.models.wan2_2.patch_diffusers import patch_wan_rms_norm
Expand Down
5 changes: 4 additions & 1 deletion vllm_omni/diffusion/models/wan2_2/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from vllm_omni.platforms import current_omni_platform

from .patch_diffusers import patch_wan_rms_norm
from .pipeline_wan2_2 import (
Wan22Pipeline,
Expand Down Expand Up @@ -53,4 +55,5 @@
"WanVACETransformer3DModel",
]

patch_wan_rms_norm()
if current_omni_platform.is_npu():
patch_wan_rms_norm()
Loading