Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
0387f14
perf(npu): optimize Qwen3 TTS on 310P
zyz111222 Jul 2, 2026
bcb1131
perf(npu): move common Qwen3 TTS ops to mainline
zyz111222 Jul 7, 2026
a9b04a6
docs(npu): clarify Qwen3 TTS weight packing
zyz111222 Jul 7, 2026
e284abc
refactor(npu): simplify Qwen3 TTS patch
zyz111222 Jul 7, 2026
7d86030
refactor(npu): remove redundant setup work
zyz111222 Jul 7, 2026
d4f8fcb
refactor(npu): hoist torch_npu import
zyz111222 Jul 10, 2026
ee096df
refactor(npu): dispatch code predictor ops
zyz111222 Jul 10, 2026
9867232
fix(npu): handle nested code predictor graphs
zyz111222 Jul 10, 2026
4855949
refactor(npu): simplify code predictor attention
zyz111222 Jul 11, 2026
12504db
refactor(npu): remove nested graph fallback
zyz111222 Jul 11, 2026
b0bbd94
refactor(npu): patch qwen3 tokenizer ops
zyz111222 Jul 11, 2026
693eaf2
refactor(npu): inline code predictor fusion
zyz111222 Jul 11, 2026
c379dab
refactor(npu): patch code2wav weights
zyz111222 Jul 11, 2026
171232a
refactor(310p): simplify code predictor patch
zyz111222 Jul 11, 2026
bc4f318
fix(310p): preserve stored sampling
zyz111222 Jul 11, 2026
e1ce4fa
perf(npu): remove redundant tensor transforms
zyz111222 Jul 13, 2026
646f2ec
revert: remove redundant tensor transforms
zyz111222 Jul 13, 2026
41d1ce6
perf(310p): fuse code predictor qkv
zyz111222 Jul 13, 2026
ecebe94
Merge branch 'main' into main
amy-why-3459 Jul 14, 2026
0ad18c3
test(npu): cover Qwen3 TTS 310P patches
zyz111222 Jul 10, 2026
a22297a
test(npu): tighten Qwen3 TTS coverage
zyz111222 Jul 13, 2026
850cc2a
Merge branch 'main' into main
zyz111222 Jul 14, 2026
6a1e6e6
Merge branch 'main' into main
zyz111222 Jul 14, 2026
0888cd2
Merge remote-tracking branch 'origin/main' into ut-main-refresh
zyz111222 Jul 14, 2026
ba73142
Merge remote-tracking branch 'origin/pr-4841' into ut-main-refresh
zyz111222 Jul 14, 2026
9a6a71c
Merge remote-tracking branch 'origin/main' into ut-main-refresh
zyz111222 Jul 14, 2026
38c4d50
fix(310p): align Qwen3 TTS runtime paths
zyz111222 Jul 14, 2026
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 @@ -49,6 +49,11 @@ def _load_module(name: str, filename: str):

def _build_mock_modules(mocker: MockerFixture) -> dict[str, object]:
"""Build the dict of modules to inject into sys.modules."""

class NativeCustomOp(torch.nn.Module):
def forward(self, *args, **kwargs):
return self.forward_native(*args, **kwargs)

platforms_mock = mocker.MagicMock()
platforms_mock.current_omni_platform.supports_torch_inductor.return_value = False
platforms_mock.current_omni_platform.is_npu.return_value = False
Expand Down Expand Up @@ -76,6 +81,8 @@ def _build_mock_modules(mocker: MockerFixture) -> dict[str, object]:

vllm_parallel_mock = mocker.MagicMock()
vllm_parallel_mock.VocabParallelEmbedding = torch.nn.Embedding
custom_op_mock = types.ModuleType("vllm_omni.diffusion.layers.custom_op")
custom_op_mock.CustomOp = NativeCustomOp

return {
"vllm_omni": mocker.MagicMock(),
Expand All @@ -85,6 +92,7 @@ def _build_mock_modules(mocker: MockerFixture) -> dict[str, object]:
"vllm.config.vllm": vllm_config_mod,
"vllm.model_executor.model_loader.weight_utils": weight_utils_mock,
"vllm.model_executor.layers.vocab_parallel_embedding": vllm_parallel_mock,
"vllm_omni.diffusion.layers.custom_op": custom_op_mock,
"vllm_omni.model_executor": types.ModuleType("vllm_omni.model_executor"),
"vllm_omni.model_executor.models": models_pkg,
"vllm_omni.model_executor.models.common": common_pkg,
Expand Down Expand Up @@ -168,6 +176,38 @@ def _make_vllm_config(mocker: MockerFixture, max_num_seqs: int = 4):
return vllm_config


def test_npu_custom_ops_use_fused_norm_and_cached_rope(mocker: MockerFixture, loaded_target_classes) -> None:
common_mod = sys.modules["vllm_omni.model_executor.models.common.qwen3_code_predictor"]
cp_config, _ = _make_tiny_config(loaded_target_classes)
rms_calls = []

def npu_rms_norm(hidden_states, weight, epsilon):
rms_calls.append((hidden_states, weight, epsilon))
return hidden_states + 1, None

mocker.patch.object(common_mod.current_omni_platform, "is_npu", return_value=True)
mocker.patch.object(
common_mod,
"torch_npu",
types.SimpleNamespace(npu_rms_norm=npu_rms_norm),
create=True,
)

hidden_states = torch.zeros(1, 2, cp_config.hidden_size, dtype=torch.float16)
norm = common_mod._RMSNorm(cp_config.hidden_size, eps=cp_config.rms_norm_eps)
torch.testing.assert_close(norm.forward_npu(hidden_states), hidden_states + 1)
assert len(rms_calls) == 1
assert rms_calls[0][2] == cp_config.rms_norm_eps

rotary = common_mod._RotaryEmbedding(cp_config)
assert rotary.cos_cached.shape == (cp_config.num_code_groups + 1, cp_config.head_dim)
assert rotary.sin_cached.shape == rotary.cos_cached.shape
position_ids = torch.tensor([[0, 2, 4]])
cos, sin = rotary.forward_npu(hidden_states, position_ids)
torch.testing.assert_close(cos, rotary.cos_cached[position_ids].to(torch.float16))
torch.testing.assert_close(sin, rotary.sin_cached[position_ids].to(torch.float16))


class TestCodePredictorDtypeAlignment:
"""Test that code predictor buffers match model parameter dtype."""

Expand Down
Loading
Loading