diff --git a/tests/model_executor/models/qwen3_tts/test_code_predictor_dtype.py b/tests/model_executor/models/qwen3_tts/test_code_predictor_dtype.py index 519940ec310..5e31c2cbd60 100644 --- a/tests/model_executor/models/qwen3_tts/test_code_predictor_dtype.py +++ b/tests/model_executor/models/qwen3_tts/test_code_predictor_dtype.py @@ -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 @@ -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(), @@ -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, @@ -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.""" diff --git a/tests/platforms/npu/test_310p_patches.py b/tests/platforms/npu/test_310p_patches.py index 255b092c860..e138a6d63da 100644 --- a/tests/platforms/npu/test_310p_patches.py +++ b/tests/platforms/npu/test_310p_patches.py @@ -67,10 +67,37 @@ class FakeCodePredictorDecoderLayer(torch.nn.Module): class FakeCodePredictorBaseModel(torch.nn.Module): pass - class FakeMimiEuclideanCodebook(torch.nn.Module): - @property - def embed(self): - return self._embed + class FakeProjection(torch.nn.Linear): + def __init__(self): + super().__init__(4, 4, bias=False) + self.call_shapes: list[tuple[int, ...]] = [] + + def forward(self, hidden_states): + self.call_shapes.append(tuple(hidden_states.shape)) + return super().forward(hidden_states) + + class FakeCodePredictorWrapper(torch.nn.Module): + def __init__(self, *args, **kwargs): + del args, kwargs + super().__init__() + self.model = torch.nn.Module() + self.model.codec_embedding = torch.nn.ModuleList([torch.nn.Embedding(8, 4), torch.nn.Embedding(8, 4)]) + self.model.linear = torch.nn.Linear(4, 4, bias=False) + self.lm_head = torch.nn.ModuleList([torch.nn.Linear(4, 8, bias=False), torch.nn.Linear(4, 8, bias=False)]) + self.small_to_mtp_projection = FakeProjection() + with torch.no_grad(): + self.small_to_mtp_projection.weight.copy_(torch.diag(torch.tensor([1.0, 2.0, 3.0, 4.0]))) + self._wrapper_config = SimpleNamespace(use_parallel_embedding=False, sampling_mode="per_call") + self._projected_codec_embed_weight = None + + def load_weights(self, weights): + del weights + return {"loaded"} + + class FakeCode2WavBase(torch.nn.Module): + def __init__(self, *args, **kwargs): + del args, kwargs + super().__init__() class FakeEncoder: def __init__(self): @@ -136,8 +163,21 @@ class FakePromptEmbedsBuilder: CodePredictorAttention=FakeCodePredictorAttention, CodePredictorDecoderLayer=FakeCodePredictorDecoderLayer, CodePredictorBaseModel=FakeCodePredictorBaseModel, + CodePredictorWrapper=FakeCodePredictorWrapper, _rotate_half=lambda x: x, ) + fake_qwen3_tts_code_predictor_vllm = _install_fake_module( + monkeypatch, + "vllm_omni.model_executor.models.qwen3_tts.qwen3_tts_code_predictor_vllm", + Qwen3TTSTalkerCodePredictorForConditionalGenerationVLLM=FakeCodePredictorWrapper, + Qwen3TTSTalkerCodePredictorModelVLLM=FakeCodePredictorBaseModel, + CodePredictorWrapper=FakeCodePredictorWrapper, + ) + fake_qwen3_tts_code2wav = _install_fake_module( + monkeypatch, + "vllm_omni.model_executor.models.qwen3_tts.qwen3_tts_code2wav", + Qwen3TTSCode2Wav=FakeCode2WavBase, + ) fake_prompt_builder = _install_fake_module( monkeypatch, "vllm_omni.model_executor.models.qwen3_tts.prompt_embeds_builder", @@ -150,22 +190,40 @@ class FakePromptEmbedsBuilder: Qwen3TTSTalkerForConditionalGeneration=FakeTalkerBase, Qwen3TTSPromptEmbedsBuilder=FakePromptEmbedsBuilder, ) - fake_modeling_mimi = _install_fake_module( - monkeypatch, - "transformers.models.mimi.modeling_mimi", - MimiEuclideanCodebook=FakeMimiEuclideanCodebook, - ) - fake_mimi = _install_fake_module( - monkeypatch, - "transformers.models.mimi", - modeling_mimi=fake_modeling_mimi, - ) _install_fake_module(monkeypatch, "vllm") _install_fake_module(monkeypatch, "vllm.multimodal") _install_fake_module(monkeypatch, "vllm.multimodal.audio", AudioResampler=FakeAudioResampler) - _install_fake_module(monkeypatch, "transformers") - _install_fake_module(monkeypatch, "transformers.models", mimi=fake_mimi) + _install_fake_module(monkeypatch, "torch_npu", npu_format_cast=lambda weight, _fmt: weight) + _install_fake_module(monkeypatch, "vllm_ascend") + _install_fake_module(monkeypatch, "vllm_ascend._310p") + _install_fake_module(monkeypatch, "vllm_ascend._310p.attention") + _install_fake_module( + monkeypatch, + "vllm_ascend._310p.attention.attention_mask", + AttentionMaskBuilder310=SimpleNamespace( + gen_causal_additive_mask=lambda max_seq, device: torch.zeros( + max_seq, + max_seq, + device=device, + ) + ), + ) + _install_fake_module(monkeypatch, "vllm_ascend.sample") + _install_fake_module( + monkeypatch, + "vllm_ascend.sample.sampler", + apply_top_k_top_p=lambda logits, **_kwargs: logits, + random_sample=lambda probs, _generators: probs.argmax(dim=-1, keepdim=True), + ) + _install_fake_module( + monkeypatch, + "vllm_ascend.utils", + ACL_FORMAT_FRACTAL_NZ=29, + aligned_16=lambda tensor: tensor, + maybe_trans_nz=lambda weight: weight, + nd_to_nz_2d=lambda tensor: tensor, + ) _install_fake_module(monkeypatch, "vllm_omni") _install_fake_module(monkeypatch, "vllm_omni.model_executor") _install_fake_module(monkeypatch, "vllm_omni.model_executor.models") @@ -178,9 +236,17 @@ class FakePromptEmbedsBuilder: monkeypatch, "vllm_omni.model_executor.models.qwen3_tts", prompt_embeds_builder=fake_prompt_builder, + qwen3_tts_code2wav=fake_qwen3_tts_code2wav, + qwen3_tts_code_predictor_vllm=fake_qwen3_tts_code_predictor_vllm, qwen3_tts_talker=fake_talker, ) - return fake_qwen3_code_predictor, fake_prompt_builder, fake_talker + return ( + fake_qwen3_code_predictor, + fake_qwen3_tts_code_predictor_vllm, + fake_qwen3_tts_code2wav, + fake_prompt_builder, + fake_talker, + ) def _load_qwen3_tts_patch(monkeypatch: pytest.MonkeyPatch): @@ -193,7 +259,7 @@ def _load_qwen3_tts_patch(monkeypatch: pytest.MonkeyPatch): def test_registry_applies_worker_once_and_model_patch_lazily(monkeypatch: pytest.MonkeyPatch) -> None: registry_path = _repo_root() / "vllm_omni" / "platforms" / "npu" / "_310p" / "patch" / "__init__.py" registry = _load_source_module("vllm_omni_test_310p_patch_registry", registry_path) - calls = {"worker": 0, "talker": 0} + calls = {"worker": 0, "talker": 0, "code2wav": 0} _install_fake_module( monkeypatch, @@ -204,52 +270,35 @@ def test_registry_applies_worker_once_and_model_patch_lazily(monkeypatch: pytest monkeypatch, "vllm_omni.platforms.npu._310p.patch.qwen3_tts", apply_talker_patches=lambda: calls.__setitem__("talker", calls["talker"] + 1), + apply_code2wav_patches=lambda: calls.__setitem__("code2wav", calls["code2wav"] + 1), ) registry.apply_patches() registry.apply_patches() registry.apply_model_patches(SimpleNamespace(model_arch="OtherModel")) registry.apply_model_patches(SimpleNamespace(model_arch="Qwen3TTSTalkerForConditionalGeneration")) + registry.apply_model_patches(SimpleNamespace(model_arch="Qwen3TTSCode2Wav")) - assert calls == {"worker": 1, "talker": 1} - - -def test_worker_patch_replaces_base_and_runs_disable_jit(monkeypatch: pytest.MonkeyPatch) -> None: - calls: list[str] = [] - - class FakeOmniNPUWorkerBase: - def _init_device(self): - calls.append("parent") - return "npu:0" - - fake_worker_base = _install_fake_module( - monkeypatch, - "vllm_omni.platforms.npu.worker.base", - OmniNPUWorkerBase=FakeOmniNPUWorkerBase, - ) - _install_fake_module(monkeypatch, "vllm_omni") - _install_fake_module(monkeypatch, "vllm_omni.platforms") - _install_fake_module(monkeypatch, "vllm_omni.platforms.npu") - _install_fake_module( - monkeypatch, - "vllm_omni.platforms.npu._310p", - disable_jit_compile=lambda: calls.append("disable_jit"), - ) - _install_fake_module(monkeypatch, "vllm_omni.platforms.npu.worker", base=fake_worker_base) - - path = _repo_root() / "vllm_omni" / "platforms" / "npu" / "_310p" / "patch" / "worker.py" - module = _load_source_module("vllm_omni_test_310p_worker_patch", path) - module.apply_patch() + assert calls == {"worker": 1, "talker": 1, "code2wav": 1} - assert fake_worker_base.OmniNPUWorkerBase is module._OmniNPUWorkerBase310P - assert fake_worker_base.OmniNPUWorkerBase()._init_device() == "npu:0" - assert calls == ["parent", "disable_jit"] +def test_qwen3_tts_patch_registers_target_classes(monkeypatch: pytest.MonkeyPatch) -> None: + ( + module, + ( + fake_code_predictor, + fake_code_predictor_vllm, + fake_code2wav, + fake_prompt_builder, + fake_talker, + ), + ) = _load_qwen3_tts_patch(monkeypatch) -def test_qwen3_tts_patch_replaces_target_classes(monkeypatch: pytest.MonkeyPatch) -> None: - module, (fake_code_predictor, fake_prompt_builder, fake_talker) = _load_qwen3_tts_patch(monkeypatch) + original_common_wrapper = fake_code_predictor.CodePredictorWrapper + original_vllm_wrapper = fake_code_predictor_vllm.CodePredictorWrapper module.apply_talker_patches() + module.apply_code2wav_patches() assert fake_talker.Qwen3TTSTalkerForConditionalGeneration is module._Qwen3TTSTalker310P assert fake_talker.Qwen3TTSPromptEmbedsBuilder is module._Qwen3TTSPromptEmbedsBuilder310P @@ -257,6 +306,112 @@ def test_qwen3_tts_patch_replaces_target_classes(monkeypatch: pytest.MonkeyPatch assert fake_code_predictor.CodePredictorAttention is module._Qwen3CodePredictorAttention310P assert fake_code_predictor.CodePredictorDecoderLayer is module._Qwen3CodePredictorDecoderLayer310P assert fake_code_predictor.CodePredictorBaseModel is module._Qwen3CodePredictorBaseModel310P + assert ( + fake_code_predictor_vllm.Qwen3TTSTalkerCodePredictorForConditionalGenerationVLLM + is module._Qwen3TTSTalkerCodePredictor310P + ) + assert fake_code_predictor.CodePredictorWrapper is original_common_wrapper + assert fake_code_predictor_vllm.CodePredictorWrapper is original_vllm_wrapper + assert fake_code2wav.Qwen3TTSCode2Wav is module._Qwen3TTSCode2Wav310P + + code2wav = module._Qwen3TTSCode2Wav310P( + vllm_config=SimpleNamespace(device_config=SimpleNamespace(device=torch.device("cpu"))) + ) + + assert code2wav._npu_decoder_runtime_dtype(torch.device("cpu")) is torch.float16 + + +def test_qwen3_tts_code_predictor_forward_uses_projected_embedding_and_sampling( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module, _ = _load_qwen3_tts_patch(monkeypatch) + predictor = module._Qwen3TTSTalkerCodePredictor310P( + vllm_config=object(), + config=object(), + talker_config=object(), + ) + codec_weight = predictor.model.codec_embedding[0].weight.detach().clone() + projection_weight = predictor.small_to_mtp_projection.weight.detach().clone() + + assert predictor.load_weights(iter(())) == {"loaded"} + torch.testing.assert_close( + predictor._projected_codec_embed_weight[0], + torch.nn.functional.linear(codec_weight, projection_weight), + ) + assert predictor.small_to_mtp_projection.call_shapes == [(8, 4), (8, 4)] + predictor.small_to_mtp_projection.call_shapes.clear() + predictor._num_groups = 3 + predictor._model_dtype = torch.float32 + predictor._prefix_graphs_enabled = False + predictor._bucket_pos_ids = {} + predictor._device_graphs = {} + predictor._lm_heads_list = [ + lambda hidden: torch.nn.functional.one_hot( + torch.full((hidden.shape[0],), 2, dtype=torch.long), + num_classes=8, + ).to(torch.float32), + lambda hidden: torch.nn.functional.one_hot( + torch.full((hidden.shape[0],), 3, dtype=torch.long), + num_classes=8, + ).to(torch.float32), + ] + predictor._setup_compile = lambda: None + predictor._padded_bsz = lambda bsz: bsz + + def ensure_buffers(device, dtype, padded_bsz): + predictor._proj_buf = torch.zeros(padded_bsz, 4, 4, device=device, dtype=dtype) + + predictor._ensure_buffers = ensure_buffers + predictor._compiled_model_fwd = lambda embeds, _positions: torch.zeros_like(embeds) + + codes = predictor.forward( + torch.tensor([1]), + torch.zeros(1, 4), + torch.zeros(1, 4), + do_sample=False, + ) + + assert codes.tolist() == [[1, 2, 3]] + assert predictor.small_to_mtp_projection.call_shapes == [(1, 2, 4)] + torch.testing.assert_close( + predictor._proj_buf[0, 2], + predictor._projected_codec_embed_weight[0, 2], + ) + + filter_calls = [] + sample_calls = [] + + def apply_top_k_top_p(logits, *, p, k, top_k): + filter_calls.append((p, k, top_k)) + return logits + + def random_sample(_probs, generators): + sample_calls.append(generators) + return torch.tensor([[4]]) + + monkeypatch.setattr(module, "apply_top_k_top_p", apply_top_k_top_p) + monkeypatch.setattr(module, "random_sample", random_sample) + predictor._num_groups = 2 + predictor._wrapper_config.sampling_mode = "stored" + predictor._top_k = 2 + predictor._top_p = 0.8 + predictor._lm_heads_list = predictor._lm_heads_list[:1] + generator = torch.Generator().manual_seed(1234) + + sampled_codes = predictor.forward( + torch.tensor([1]), + torch.zeros(1, 4), + torch.zeros(1, 4), + generator=generator, + ) + + assert sampled_codes.tolist() == [[1, 4]] + assert len(filter_calls) == 1 + top_p_tensor, top_k_tensor, top_k_hint = filter_calls[0] + torch.testing.assert_close(top_p_tensor, torch.tensor([0.8])) + assert top_k_tensor.tolist() == [2] + assert top_k_hint == 2 + assert sample_calls == [{0: generator}] def test_qwen3_tts_talker_patch_uses_fp16_runtime_dtype(monkeypatch: pytest.MonkeyPatch) -> None: @@ -265,12 +420,14 @@ def test_qwen3_tts_talker_patch_uses_fp16_runtime_dtype(monkeypatch: pytest.Monk assert talker._embedding_dtype is torch.float16 assert talker._prompt_builder._embedding_dtype is torch.float16 + assert talker.talker_mtp_graph_safe is False + assert talker.talker_mtp_accepts_per_row_generators is True assert talker.load_weights([]) == {"loaded"} - assert talker.encoder.to_calls[-1] == {"dtype": torch.float16} + assert talker.encoder.to_calls[-1] == {"device": torch.device("cpu"), "dtype": torch.float32} codes = talker._encode_ref_audio_batch([np.zeros(8, dtype=np.float32)], 24000, device=torch.device("cpu")) - assert talker.encoder.last_input_dtype is torch.float16 + assert talker.encoder.last_input_dtype is torch.float32 assert len(codes) == 1 assert codes[0].dtype is torch.long assert codes[0].shape == (4, 2) @@ -311,28 +468,154 @@ def forward(self, mels): assert speaker.dtype is torch.float16 -def test_qwen3_tts_mimi_codebook_quantize_uses_cpu_fp32_cdist(monkeypatch: pytest.MonkeyPatch) -> None: - module, _ = _load_qwen3_tts_patch(monkeypatch) - real_cdist = torch.cdist - captured = {} +def test_qwen3_tts_tokenizer_npu_patch_dispatches_fused_ops(monkeypatch: pytest.MonkeyPatch) -> None: + rotary_calls = [] + rms_calls = [] + + def rotary_mul(hidden_states, cos, sin): + rotary_calls.append((hidden_states, cos, sin)) + return hidden_states + 1 + + def rms_norm(hidden_states, weight, *, epsilon): + rms_calls.append((hidden_states, weight, epsilon)) + return hidden_states * weight, None + + _install_fake_module( + monkeypatch, + "torch_npu", + npu_rotary_mul=rotary_mul, + npu_rms_norm=rms_norm, + ) + _install_fake_module(monkeypatch, "vllm") + _install_fake_module(monkeypatch, "vllm.logger", init_logger=lambda _name: SimpleNamespace(debug=lambda *_: None)) + + class FakeRMSNorm(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([2.0, 3.0])) + self.variance_epsilon = 1e-5 + + def forward(self, hidden_states): + return hidden_states + + def original_rope(q, k, cos, sin): + del cos, sin + return q, k + + tokenizer = _install_fake_module( + monkeypatch, + "vllm_omni.model_executor.models.qwen3_tts.tokenizer_12hz.modeling_qwen3_tts_tokenizer_v2", + Qwen3TTSTokenizerV2DecoderRMSNorm=FakeRMSNorm, + apply_rotary_pos_emb=original_rope, + ) + tokenizer_package = _install_fake_module( + monkeypatch, + "vllm_omni.model_executor.models.qwen3_tts.tokenizer_12hz", + modeling_qwen3_tts_tokenizer_v2=tokenizer, + ) + _install_fake_module(monkeypatch, "vllm_omni") + _install_fake_module(monkeypatch, "vllm_omni.model_executor") + _install_fake_module(monkeypatch, "vllm_omni.model_executor.models") + _install_fake_module(monkeypatch, "vllm_omni.model_executor.models.qwen3_tts") + monkeypatch.setitem( + sys.modules, + "vllm_omni.model_executor.models.qwen3_tts.tokenizer_12hz", + tokenizer_package, + ) + + path = _repo_root() / "vllm_omni" / "platforms" / "npu" / "models" / "qwen3_tts_tokenizer_v2.py" + module = _load_source_module("vllm_omni_test_qwen3_tts_tokenizer_npu_patch", path) + module.apply_qwen3_tts_tokenizer_v2_patch() + + q = torch.zeros(1, 2, 3, 2) + k = torch.ones_like(q) + cos = torch.zeros(1, 3, 2) + sin = torch.ones_like(cos) + q_out, k_out = tokenizer.apply_rotary_pos_emb(q, k, cos, sin) + norm = FakeRMSNorm() + norm_out = norm(torch.ones(1, 2)) + + assert len(rotary_calls) == 2 + assert rotary_calls[0][1].shape == (1, 1, 3, 2) + torch.testing.assert_close(q_out, q + 1) + torch.testing.assert_close(k_out, k + 1) + assert len(rms_calls) == 1 + assert rms_calls[0][2] == pytest.approx(1e-5) + torch.testing.assert_close(norm_out, torch.tensor([[2.0, 3.0]])) + + +def test_qwen3_tts_code2wav_npu_patch_prepares_loaded_decoder(monkeypatch: pytest.MonkeyPatch) -> None: + linear_weights = [] + conv_weights = [] + + def maybe_trans_nz(weight): + linear_weights.append(weight) + return weight + + def format_cast(weight, fmt): + conv_weights.append((weight, fmt)) + return weight + + class FakeDecoder(torch.nn.Module): + def __init__(self): + super().__init__() + self.linear = torch.nn.Linear(4, 4) + self.conv = torch.nn.Conv1d(4, 4, 3) + self.deconv = torch.nn.ConvTranspose1d(4, 4, 4) + self.grouped_conv = torch.nn.Conv1d(4, 4, 3, groups=2) + self.cache_precompute_calls = 0 + + def precompute_snake_caches(self): + self.cache_precompute_calls += 1 - def fake_cdist(x1, x2, p=2): - captured["x1_device"] = x1.device - captured["x2_device"] = x2.device - captured["x1_dtype"] = x1.dtype - captured["x2_dtype"] = x2.dtype - return real_cdist(x1, x2, p=p) + class FakeCode2Wav: + def __init__(self, *, vllm_config, prefix=""): + self.vllm_config = vllm_config + self.prefix = prefix + self.decoder = FakeDecoder() - monkeypatch.setattr(torch, "cdist", fake_cdist) - codebook = object.__new__(module._MimiEuclideanCodebook310P) - codebook._embed = torch.tensor([[0.0, 0.0], [2.0, 0.0]], dtype=torch.float16) + def _npu_decoder_runtime_dtype(self, _device): + return torch.float16 - indices = codebook.quantize(torch.tensor([[0.1, 0.0], [1.8, 0.0]], dtype=torch.float16)) + def load_weights(self, weights): + assert list(weights) == [] + return {"loaded"} - assert captured == { - "x1_device": torch.device("cpu"), - "x2_device": torch.device("cpu"), - "x1_dtype": torch.float32, - "x2_dtype": torch.float32, + logger = SimpleNamespace(info=lambda *_: None, debug=lambda *_: None) + current_platform = SimpleNamespace(is_npu=lambda: False) + target = _install_fake_module( + monkeypatch, + "vllm_omni.model_executor.models.qwen3_tts.qwen3_tts_code2wav", + Qwen3TTSCode2Wav=FakeCode2Wav, + ) + _install_fake_module(monkeypatch, "torch_npu", npu_format_cast=format_cast) + _install_fake_module(monkeypatch, "vllm") + _install_fake_module(monkeypatch, "vllm.config", VllmConfig=object) + _install_fake_module(monkeypatch, "vllm.logger", init_logger=lambda _name: logger) + _install_fake_module(monkeypatch, "vllm_ascend") + _install_fake_module(monkeypatch, "vllm_ascend.utils", maybe_trans_nz=maybe_trans_nz) + _install_fake_module(monkeypatch, "vllm_omni") + _install_fake_module(monkeypatch, "vllm_omni.platforms", current_omni_platform=current_platform) + _install_fake_module(monkeypatch, "vllm_omni.model_executor") + _install_fake_module(monkeypatch, "vllm_omni.model_executor.models") + _install_fake_module(monkeypatch, "vllm_omni.model_executor.models.qwen3_tts") + + path = _repo_root() / "vllm_omni" / "platforms" / "npu" / "models" / "qwen3_tts_code2wav.py" + module = _load_source_module("vllm_omni_test_qwen3_tts_code2wav_npu_patch", path) + module.apply_qwen3_tts_code2wav_patch() + + model = target.Qwen3TTSCode2Wav( + vllm_config=SimpleNamespace(device_config=SimpleNamespace(device=torch.device("cpu"))), + prefix="stage1", + ) + assert model.load_weights(iter(())) == {"loaded"} + + assert model.prefix == "stage1" + assert model.decoder.linear.weight.dtype is torch.float16 + assert [weight.data_ptr() for weight in linear_weights] == [model.decoder.linear.weight.data_ptr()] + assert {weight.data_ptr() for weight, _ in conv_weights} == { + model.decoder.conv.weight.data_ptr(), + model.decoder.deconv.weight.data_ptr(), } - assert indices.tolist() == [0, 1] + assert all(fmt == module._ACL_FORMAT_FRACTAL_Z for _, fmt in conv_weights) + assert model.decoder.cache_precompute_calls == 1 diff --git a/vllm_omni/model_executor/models/common/qwen3_code_predictor.py b/vllm_omni/model_executor/models/common/qwen3_code_predictor.py index a8a259504a3..7ac2ba123d8 100644 --- a/vllm_omni/model_executor/models/common/qwen3_code_predictor.py +++ b/vllm_omni/model_executor/models/common/qwen3_code_predictor.py @@ -3,7 +3,7 @@ Shared by Qwen3-Omni and Qwen3-TTS talker models. * SDPA attention (F.scaled_dot_product_attention) with native GQA support -* HF-compatible numerics (float32 RMSNorm, float32 RoPE, separate linear layers) +* HF-compatible CPU/CUDA numerics with NPU-only fused norm/RoPE fast paths * Per-call embedding buffer to avoid cross-request aliasing * Pre-allocated position_ids (read-only, safe to persist) * torch.compile (epilogue_fusion=False) on inner transformer by default @@ -24,6 +24,7 @@ from vllm.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm_omni.diffusion.layers.custom_op import CustomOp from vllm_omni.platforms import current_omni_platform logger = init_logger(__name__) @@ -31,33 +32,48 @@ _GeneratorLike = torch.Generator | Sequence[torch.Generator | None] | None _UNIFORM_EPS = 1e-20 +if current_omni_platform.is_npu(): + import torch_npu + # =================================================================== -# HF-numerics-compatible layers for code predictor +# Portable layers for code predictor # =================================================================== # # These use plain PyTorch ops (nn.Linear, manual RMSNorm in float32, # rotate_half RoPE) to produce outputs numerically identical to the -# HuggingFace reference. vLLM's fused kernels (RMSNorm, QKVParallel, -# get_rope) introduce small precision differences that compound across -# the autoregressive steps of the code predictor, causing severe -# audio quality degradation. +# HuggingFace reference on CPU/CUDA. vLLM's fused kernels (RMSNorm, +# QKVParallel, get_rope) introduce small precision differences that compound +# across the autoregressive steps of the code predictor, causing severe audio +# quality degradation. The Ascend fused norm/RoPE kernels below are dispatched +# by the current device platform. # # See: https://github.com/vllm-project/vllm-omni/issues/2274 -class _RMSNorm(nn.Module): - """RMSNorm matching HuggingFace's implementation exactly. - - Computes variance in float32 to avoid bfloat16 precision loss. - """ +class _RMSNorm(CustomOp): + """RMSNorm with HuggingFace-compatible CPU/CUDA math and an NPU fast path.""" def __init__(self, hidden_size: int, eps: float = 1e-6) -> None: super().__init__() self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + def forward_npu(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states, _ = torch_npu.npu_rms_norm( + hidden_states, + self.weight, + self.variance_epsilon, + ) + return hidden_states + + def forward_cuda(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.forward_native(hidden_states) + + def forward_xpu(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.forward_native(hidden_states) + + def forward_native(self, hidden_states: torch.Tensor) -> torch.Tensor: input_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) variance = hidden_states.pow(2).mean(-1, keepdim=True) @@ -72,11 +88,8 @@ def _rotate_half(x: torch.Tensor) -> torch.Tensor: return torch.cat((-x2, x1), dim=-1) -class _RotaryEmbedding(nn.Module): - """RoPE matching HuggingFace's implementation exactly. - - Forces float32 computation for cos/sin, matching HF's torch.autocast(enabled=False). - """ +class _RotaryEmbedding(CustomOp): + """RoPE with HuggingFace-compatible CPU/CUDA math and cached NPU tables.""" def __init__(self, config) -> None: super().__init__() @@ -88,8 +101,24 @@ def __init__(self, config) -> None: rope_theta = getattr(config, "rope_theta", 10000.0) inv_freq = 1.0 / (rope_theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) + if current_omni_platform.is_npu(): + max_seq = int(getattr(config, "num_code_groups", 0) or 0) + 1 + positions = torch.arange(max_seq, dtype=torch.float32) + freqs = torch.outer(positions, inv_freq) + emb = torch.cat((freqs, freqs), dim=-1) + self.register_buffer("cos_cached", emb.cos(), persistent=False) + self.register_buffer("sin_cached", emb.sin(), persistent=False) - def forward(self, x: torch.Tensor, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + def forward_npu(self, x: torch.Tensor, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + return self.cos_cached[position_ids].to(dtype=x.dtype), self.sin_cached[position_ids].to(dtype=x.dtype) + + def forward_cuda(self, x: torch.Tensor, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + return self.forward_native(x, position_ids) + + def forward_xpu(self, x: torch.Tensor, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + return self.forward_native(x, position_ids) + + def forward_native(self, x: torch.Tensor, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: # position_ids: [batch, seq_len] inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) position_ids_expanded = position_ids[:, None, :].float() @@ -145,7 +174,6 @@ def __init__(self, config, *, prefix: str = "") -> None: self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) self.q_norm = _RMSNorm(self.head_dim, eps=config.rms_norm_eps) self.k_norm = _RMSNorm(self.head_dim, eps=config.rms_norm_eps) - if current_omni_platform.is_npu(): if self.max_seq > 2048: raise ValueError( @@ -168,8 +196,6 @@ def _forward_npu_attention( bsz: int, seq_len: int, ) -> torch.Tensor: - import torch_npu - q_f, k_f, v_f = q, k, v if self.is_gqa: k_f = ( @@ -230,10 +256,13 @@ def forward( # cos/sin are [batch, seq_len, head_dim], need unsqueeze at dim=1 for heads cos = cos.unsqueeze(1) # [batch, 1, seq_len, head_dim] sin = sin.unsqueeze(1) - q = (q * cos) + (_rotate_half(q) * sin) - k = (k * cos) + (_rotate_half(k) * sin) - - if not current_omni_platform.is_npu(): + if current_omni_platform.is_npu(): + q = torch_npu.npu_rotary_mul(q, cos, sin) + k = torch_npu.npu_rotary_mul(k, cos, sin) + attn_out = self._forward_npu_attention(q, k, v, bsz, seq_len) + else: + q = (q * cos) + (_rotate_half(q) * sin) + k = (k * cos) + (_rotate_half(k) * sin) attn_out = F.scaled_dot_product_attention( q, k, @@ -242,8 +271,6 @@ def forward( is_causal=True, enable_gqa=self.is_gqa, ) - else: - attn_out = self._forward_npu_attention(q, k, v, bsz, seq_len) attn_out = attn_out.transpose(1, 2).reshape(bsz, seq_len, -1) return self.o_proj(attn_out) @@ -290,10 +317,17 @@ def forward( residual = hidden_states hidden_states = self.input_layernorm(hidden_states) hidden_states = self.self_attn(hidden_states, position_embeddings) - hidden_states = residual + hidden_states - - residual = hidden_states - hidden_states = self.post_attention_layernorm(hidden_states) + if current_omni_platform.is_npu(): + hidden_states, _, residual = torch_npu.npu_add_rms_norm( + hidden_states, + residual, + self.post_attention_layernorm.weight, + self.post_attention_layernorm.variance_epsilon, + ) + else: + hidden_states = residual + hidden_states + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) hidden_states = residual + hidden_states return hidden_states @@ -875,7 +909,6 @@ def forward( # Use captured device graph if available, otherwise call compiled fn. device_graph_entry = self._device_graphs.get(graph_key) - # Run transformer (device graph replay or compiled forward) if device_graph_entry is not None: device_graph_entry[0].replay() hidden_out = device_graph_entry[1] @@ -939,6 +972,18 @@ def forward( # Weight loading # ------------------------------------------------------------------ + def _prepare_npu_weights(self) -> None: + from vllm_ascend.utils import maybe_trans_nz + + linear_count = 0 + with torch.no_grad(): + # Pack linear weights once for NPU matmul. + for module in self.modules(): + if isinstance(module, nn.Linear): + module.weight.data = maybe_trans_nz(module.weight.data) + linear_count += 1 + logger.info("Prepared NPU code predictor weights: linear=%d", linear_count) + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: """Load weights directly (no fused projection remapping needed).""" loaded: set[str] = set() @@ -965,4 +1010,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: weight_loader(param, w) loaded.add(name) + if current_omni_platform.is_npu(): + self._prepare_npu_weights() + return loaded diff --git a/vllm_omni/platforms/npu/_310p/patch/__init__.py b/vllm_omni/platforms/npu/_310p/patch/__init__.py index b3ab637ff2e..5bfd424294f 100644 --- a/vllm_omni/platforms/npu/_310p/patch/__init__.py +++ b/vllm_omni/platforms/npu/_310p/patch/__init__.py @@ -7,6 +7,7 @@ _WORKER_PATCHED = False _QWEN3_TTS_TALKER_ARCH = "Qwen3TTSTalkerForConditionalGeneration" +_QWEN3_TTS_CODE2WAV_ARCH = "Qwen3TTSCode2Wav" def apply_patches() -> None: @@ -22,9 +23,13 @@ def apply_patches() -> None: def apply_model_patches(model_config) -> None: - if getattr(model_config, "model_arch", None) != _QWEN3_TTS_TALKER_ARCH: - return + model_arch = getattr(model_config, "model_arch", None) + if model_arch == _QWEN3_TTS_TALKER_ARCH: + from vllm_omni.platforms.npu._310p.patch.qwen3_tts import apply_talker_patches - from vllm_omni.platforms.npu._310p.patch.qwen3_tts import apply_talker_patches + apply_talker_patches() + return + elif model_arch == _QWEN3_TTS_CODE2WAV_ARCH: + from vllm_omni.platforms.npu._310p.patch.qwen3_tts import apply_code2wav_patches - apply_talker_patches() + apply_code2wav_patches() diff --git a/vllm_omni/platforms/npu/_310p/patch/qwen3_tts.py b/vllm_omni/platforms/npu/_310p/patch/qwen3_tts.py index c43cdd3cf42..713b67f8136 100644 --- a/vllm_omni/platforms/npu/_310p/patch/qwen3_tts.py +++ b/vllm_omni/platforms/npu/_310p/patch/qwen3_tts.py @@ -5,26 +5,30 @@ from __future__ import annotations +from collections.abc import Sequence + import numpy as np import torch -from transformers.models.mimi import modeling_mimi +import torch.nn as nn +import torch.nn.functional as F +import torch_npu from vllm.multimodal.audio import AudioResampler +from vllm_ascend._310p.attention.attention_mask import AttentionMaskBuilder310 +from vllm_ascend.sample.sampler import apply_top_k_top_p, random_sample +from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, aligned_16, maybe_trans_nz, nd_to_nz_2d from vllm_omni.model_executor.models.common import qwen3_code_predictor -from vllm_omni.model_executor.models.qwen3_tts import prompt_embeds_builder, qwen3_tts_talker +from vllm_omni.model_executor.models.qwen3_tts import ( + prompt_embeds_builder, + qwen3_tts_code2wav, + qwen3_tts_code_predictor_vllm, + qwen3_tts_talker, +) _RUNTIME_DTYPE = torch.float16 _CPU_DEVICE = torch.device("cpu") - - -class _MimiEuclideanCodebook310P(modeling_mimi.MimiEuclideanCodebook): - def quantize(self, hidden_states): - # 310P does not support torch.cdist on NPU. - device = hidden_states.device - dists = torch.cdist( - hidden_states[None].to(_CPU_DEVICE, torch.float32), self.embed[None].to(_CPU_DEVICE, torch.float32), p=2 - )[0] - return dists.argmin(dim=-1).to(device=device) +_PATCHED = False +_CODE2WAV_PATCHED = False class _Qwen3TTSTalker310P(qwen3_tts_talker.Qwen3TTSTalkerForConditionalGeneration): @@ -32,10 +36,18 @@ def __init__(self, *, vllm_config, prefix: str = "") -> None: super().__init__(vllm_config=vllm_config, prefix=prefix) self._embedding_dtype = _RUNTIME_DTYPE self._prompt_builder._embedding_dtype = _RUNTIME_DTYPE + # 310P random sampling operators cannot be captured by ACL graphs. + # Keep Talker-MTP eager so CodePredictor can replay its own NPU graphs + # without changing sampling semantics. + self.talker_mtp_graph_safe = False + self.talker_mtp_accepts_per_row_generators = True def load_weights(self, weights): loaded = super().load_weights(weights) - self.encoder.to(dtype=_RUNTIME_DTYPE) + # The Mimi tokenizer encoder is used only while building ref_audio + # prompts. Keep this preprocessing-only module on CPU because the 310P + # NPU path is sensitive to changing reference-audio shapes. + self.encoder.to(device=_CPU_DEVICE, dtype=torch.float32) return loaded def _encode_ref_audio_batch( @@ -51,7 +63,7 @@ def _encode_ref_audio_batch( resampler = AudioResampler(target_sr=target_sr) wavs = [resampler.resample(w.astype(np.float32), orig_sr=int(sr)) for w in wavs] - inputs = fe(raw_audio=wavs, sampling_rate=target_sr, return_tensors="pt").to(device).to(_RUNTIME_DTYPE) + inputs = fe(raw_audio=wavs, sampling_rate=target_sr, return_tensors="pt").to(torch.float32) with torch.inference_mode(): encoded = self.encoder.encode( @@ -100,10 +112,56 @@ def extract_speaker_embedding(self, wav: np.ndarray, sr: int) -> torch.Tensor: return spk.to(dtype=dtype) +# =================================================================== +# Code2Wav layer patches +# =================================================================== +# +# Code2Wav runs under the 310P graph path after the stage-0 Talker has +# produced codec tokens. Common NPU weight packing and fused tokenizer ops +# live in the shared model code; the 310P patch only selects the runtime dtype. + + +class _Qwen3TTSCode2Wav310P(qwen3_tts_code2wav.Qwen3TTSCode2Wav): + """Qwen3-TTS Code2Wav specialized for the 310P NPU path.""" + + def _npu_decoder_runtime_dtype(self, device: torch.device) -> torch.dtype: + return _RUNTIME_DTYPE + + +# =================================================================== +# CodePredictor layer patches +# =================================================================== +# +# Keep the portable implementation in common/qwen3_code_predictor.py. +# The overrides below are installed only by the 310P platform patch because +# the short CodePredictor loop is graph-captured on 310P and needs the 310P +# flash-attention mask layout plus loop-local projection and sampling changes. + + class _Qwen3CodePredictorAttention310P(qwen3_code_predictor.CodePredictorAttention): + """Attention override using 310P RoPE and flash-attention kernels. + + The shared attention path is written in portable PyTorch. This override + keeps the 310P-specific RoPE op, token alignment, and FRACTAL_NZ mask path + scoped to the platform patch. + """ + def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) self._buffers.pop("_fusion_causal_mask", None) + self._q_size = self.num_heads * self.head_dim + self._kv_size = self.num_kv_heads * self.head_dim + self._fused_qkv_weight = None + self._fused_qkv_bias = None + + def prepare_qkv_weights(self) -> None: + # Pack QKV once so each graph replay uses one matmul and consumes the + # weight directly in the 310P matmul layout. + self._fused_qkv_weight = maybe_trans_nz( + torch.cat((self.q_proj.weight, self.k_proj.weight, self.v_proj.weight), dim=0).contiguous() + ) + if self.q_proj.bias is not None: + self._fused_qkv_bias = torch.cat((self.q_proj.bias, self.k_proj.bias, self.v_proj.bias), dim=0) def forward( self, @@ -111,32 +169,29 @@ def forward( position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: torch.Tensor | None = None, ) -> torch.Tensor: - if hidden_states.device.type != "npu" or attention_mask is None: - return super().forward(hidden_states, position_embeddings) - bsz, seq_len, _ = hidden_states.shape - hidden_shape_q = (bsz, seq_len, self.num_heads, self.head_dim) - hidden_shape_kv = (bsz, seq_len, self.num_kv_heads, self.head_dim) - - q = self.q_norm(self.q_proj(hidden_states).view(hidden_shape_q)).transpose(1, 2) - k = self.k_norm(self.k_proj(hidden_states).view(hidden_shape_kv)).transpose(1, 2) - v = self.v_proj(hidden_states).view(hidden_shape_kv).transpose(1, 2) + qkv = F.linear(hidden_states, self._fused_qkv_weight, self._fused_qkv_bias) + q, k, v = qkv.split((self._q_size, self._kv_size, self._kv_size), dim=-1) + q = self.q_norm(q.view(bsz, seq_len, self.num_heads, self.head_dim)).transpose(1, 2) + k = self.k_norm(k.view(bsz, seq_len, self.num_kv_heads, self.head_dim)).transpose(1, 2) + v = v.view(bsz, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) cos, sin = position_embeddings cos = cos.unsqueeze(1) sin = sin.unsqueeze(1) - q = (q * cos) + (qwen3_code_predictor._rotate_half(q) * sin) - k = (k * cos) + (qwen3_code_predictor._rotate_half(k) * sin) + # Use the fused Ascend RoPE op instead of expanding RoPE into + # elementwise mul/add/rotate-half kernels. + q = torch_npu.npu_rotary_mul(q, cos, sin) + k = torch_npu.npu_rotary_mul(k, cos, sin) real_tokens = int(bsz) * int(seq_len) output_dtype = q.dtype - import torch_npu - from vllm_ascend.utils import aligned_16 - - q_f = aligned_16(q.to(torch.float16).transpose(1, 2).reshape(real_tokens, self.num_heads, self.head_dim)) - k_f = aligned_16(k.to(torch.float16).transpose(1, 2).reshape(real_tokens, self.num_kv_heads, self.head_dim)) - v_f = aligned_16(v.to(torch.float16).transpose(1, 2).reshape(real_tokens, self.num_kv_heads, self.head_dim)) + # 310P flash attention consumes token-major fp16 inputs with 16-token + # alignment; seq_lens carries the padding information. + q_f = aligned_16(q.transpose(1, 2).reshape(real_tokens, self.num_heads, self.head_dim)) + k_f = aligned_16(k.transpose(1, 2).reshape(real_tokens, self.num_kv_heads, self.head_dim)) + v_f = aligned_16(v.transpose(1, 2).reshape(real_tokens, self.num_kv_heads, self.head_dim)) aligned_tokens = int(q_f.shape[0]) seq_lens = torch.full((int(bsz),), int(seq_len), dtype=torch.int32, device="cpu") @@ -160,23 +215,34 @@ def forward( class _Qwen3CodePredictorDecoderLayer310P(qwen3_code_predictor.CodePredictorDecoderLayer): + """Decoder layer override that passes the 310P attention mask.""" + def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None = None, + attention_mask: torch.Tensor, ) -> torch.Tensor: residual = hidden_states hidden_states = self.input_layernorm(hidden_states) hidden_states = self.self_attn(hidden_states, position_embeddings, attention_mask=attention_mask) - hidden_states = residual + hidden_states - - residual = hidden_states - hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states, _, residual = torch_npu.npu_add_rms_norm( + hidden_states, + residual, + self.post_attention_layernorm.weight, + self.post_attention_layernorm.variance_epsilon, + ) return residual + self.mlp(hidden_states) class _Qwen3CodePredictorBaseModel310P(qwen3_code_predictor.CodePredictorBaseModel): + """Base model override with a cached 310P causal mask. + + The 310P flash-attention path consumes the additive causal mask in + FRACTAL_NZ format. The CodePredictor sequence length is fixed and short, + so the mask is built once and reused across graph replays. + """ + def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) self._attention_mask_310p = None @@ -188,14 +254,9 @@ def forward( inputs_embeds: torch.Tensor, position_ids: torch.Tensor, ) -> torch.Tensor: - if inputs_embeds.device.type != "npu": - return super().forward(inputs_embeds, position_ids) - if self._attention_mask_310p is None or self._attention_mask_310p_device != inputs_embeds.device: - import torch_npu - from vllm_ascend._310p.attention.attention_mask import AttentionMaskBuilder310 - from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, nd_to_nz_2d - + # Store the additive causal mask in the format consumed by the + # 310P flash-attention kernel. mask = AttentionMaskBuilder310.gen_causal_additive_mask( self._attention_mask_310p_max_seq, inputs_embeds.device, @@ -205,23 +266,239 @@ def forward( input_dtype = inputs_embeds.dtype hidden_states = inputs_embeds - with torch.amp.autocast(inputs_embeds.device.type, enabled=False, dtype=torch.float32): - position_embeddings = self.rotary_emb(hidden_states, position_ids) - for layer in self.layers: - hidden_states = layer( - hidden_states, - position_embeddings, - attention_mask=self._attention_mask_310p, - ) - hidden_states = self.norm(hidden_states) + position_embeddings = self.rotary_emb(hidden_states, position_ids) + for layer in self.layers: + hidden_states = layer( + hidden_states, + position_embeddings, + attention_mask=self._attention_mask_310p, + ) + hidden_states = self.norm(hidden_states) return hidden_states.to(input_dtype) +class _Qwen3TTSTalkerCodePredictor310P( + qwen3_tts_code_predictor_vllm.Qwen3TTSTalkerCodePredictorForConditionalGenerationVLLM +): + """Qwen3-TTS code predictor specialized for the 310P NPU path.""" + + def __init__( + self, + *, + vllm_config, + config, + talker_config, + prefix: str = "code_predictor", + ) -> None: + super().__init__( + vllm_config=vllm_config, + config=config, + talker_config=talker_config, + prefix=prefix, + ) + self._projected_codec_embed_weight = None + + def _prepare_npu_weights(self) -> None: + qkv_projections = set() + with torch.no_grad(): + for layer in self.model.layers: + attention = layer.self_attn + attention.prepare_qkv_weights() + qkv_projections.update((attention.q_proj, attention.k_proj, attention.v_proj)) + + for module in self.modules(): + if isinstance(module, nn.Linear) and module not in qkv_projections: + module.weight.data = maybe_trans_nz(module.weight.data) + + def load_weights(self, weights): + loaded = super().load_weights(weights) + with torch.no_grad(): + self._projected_codec_embed_weight = torch.stack( + [self.small_to_mtp_projection(embed.weight).detach() for embed in self.model.codec_embedding], + dim=0, + ).contiguous() + return loaded + + @torch.inference_mode() + def forward( + self, + layer0_code: torch.Tensor, + layer0_embed: torch.Tensor, + last_talker_hidden: torch.Tensor, + do_sample: bool = True, + temperature: float = 0.9, + top_k: int = 50, + top_p: float = 1.0, + generator: torch.Generator | None = None, + generators: Sequence[torch.Generator | None] | None = None, + ) -> torch.Tensor: + bsz = int(layer0_code.shape[0]) + if generators is not None and len(generators) != bsz: + raise ValueError(f"generators must have one entry per row: got {len(generators)} for batch {bsz}") + num_groups = self._num_groups + device = layer0_code.device + + self._setup_compile() + dtype = self._model_dtype + + padded_bsz = self._padded_bsz(bsz) + self._ensure_buffers(device, dtype, padded_bsz) + + proj_buf = self._proj_buf + max_seq = num_groups + 1 + projection = self.small_to_mtp_projection + model_fwd = self._compiled_model_fwd + lm_heads = self._lm_heads_list + if generators is not None: + npu_generators = {i: row_generator for i, row_generator in enumerate(generators) if row_generator} + elif generator is not None: + npu_generators = {i: generator for i in range(bsz)} + else: + npu_generators = {} + + proj_buf[:padded_bsz].zero_() + initial_embeds = torch.cat( + ( + last_talker_hidden.reshape(bsz, 1, -1), + layer0_embed.reshape(bsz, 1, -1), + ), + dim=1, + ) + if initial_embeds.dtype != dtype: + initial_embeds = initial_embeds.to(dtype) + proj_buf[:bsz, :2, :].copy_(projection(initial_embeds)) + + stored_mode = self._wrapper_config.sampling_mode == "stored" + if stored_mode: + s_top_k = self._top_k + s_top_p = self._top_p + else: + use_sampling = do_sample and temperature > 0 + inv_temperature = 1.0 / max(temperature, 1e-6) if use_sampling else 0.0 + if use_sampling and top_p != 1.0: + raise NotImplementedError( + "top_p sampling is not implemented for the vLLM-native code predictor; please set top_p=1.0." + ) + + top_k_tensor = None + top_p_tensor = None + if stored_mode: + top_k_hint = s_top_k if s_top_k > 0 else None + if s_top_k > 0: + top_k_tensor = torch.full((bsz,), s_top_k, dtype=torch.int32, device=device) + if s_top_p < 1.0: + top_p_tensor = torch.full((bsz,), s_top_p, dtype=dtype, device=device) + elif use_sampling: + top_k_hint = top_k if top_k > 0 else None + if top_k > 0: + top_k_tensor = torch.full((bsz,), top_k, dtype=torch.int32, device=device) + else: + top_k_hint = None + + all_codes = torch.empty(bsz, num_groups, dtype=torch.long, device=device) + all_codes[:, 0] = layer0_code.reshape(bsz) + + for step in range(1, num_groups): + graph_key: int | tuple[int, int] = padded_bsz + seq_len = max_seq + if self._prefix_graphs_enabled: + prefix_key = (padded_bsz, step + 1) + if prefix_key in self._device_graphs: + graph_key = prefix_key + seq_len = step + 1 + pos_ids = self._bucket_pos_ids.get(graph_key) + if pos_ids is None: + pos_ids = ( + torch.arange(seq_len, device=device, dtype=torch.long) + .unsqueeze(0) + .expand(padded_bsz, -1) + .contiguous() + ) + + device_graph_entry = self._device_graphs.get(graph_key) + if device_graph_entry is not None: + device_graph_entry[0].replay() + hidden_out = device_graph_entry[1] + else: + hidden_out = model_fwd(proj_buf[:padded_bsz, :seq_len, :], pos_ids) + + logits = lm_heads[step - 1](hidden_out[:bsz, step, :]) + + if stored_mode: + if top_k_tensor is not None or top_p_tensor is not None: + logits = apply_top_k_top_p(logits, p=top_p_tensor, k=top_k_tensor, top_k=top_k_hint) + candidate_indices = None + if isinstance(logits, tuple): + logits, candidate_indices = logits + probs = F.softmax(logits, dim=-1, dtype=torch.float32) + code = random_sample(probs, npu_generators) + if candidate_indices is not None: + code = candidate_indices.gather(1, code.unsqueeze(1)).squeeze(1) + else: + if use_sampling: + scaled = logits * inv_temperature + if top_k_tensor is not None: + scaled = apply_top_k_top_p(scaled, p=None, k=top_k_tensor, top_k=top_k_hint) + candidate_indices = None + if isinstance(scaled, tuple): + scaled, candidate_indices = scaled + probs = F.softmax(scaled, dim=-1, dtype=torch.float32) + code = random_sample(probs, npu_generators) + if candidate_indices is not None: + code = candidate_indices.gather(1, code.unsqueeze(1)).squeeze(1) + else: + code = logits.argmax(dim=-1, keepdim=True) + + all_codes[:, step] = code.reshape(bsz) + if step < num_groups - 1: + proj_buf[:bsz, step + 1, :].copy_( + F.embedding(code.reshape(-1), self._projected_codec_embed_weight[step - 1]) + ) + + return all_codes + + +# =================================================================== +# Patch registration +# =================================================================== + + def apply_talker_patches() -> None: - modeling_mimi.MimiEuclideanCodebook = _MimiEuclideanCodebook310P + """Install Qwen3-TTS Talker and CodePredictor 310P patches. + + The generic model modules stay unchanged. Patch registration swaps in the + 310P-specialized CodePredictor classes and wrapper methods only when the + 310P platform applies the Talker patch. + """ + + global _PATCHED + + if _PATCHED: + return + qwen3_tts_talker.Qwen3TTSTalkerForConditionalGeneration = _Qwen3TTSTalker310P qwen3_tts_talker.Qwen3TTSPromptEmbedsBuilder = _Qwen3TTSPromptEmbedsBuilder310P + qwen3_tts_talker.Qwen3TTSTalkerCodePredictorForConditionalGenerationVLLM = _Qwen3TTSTalkerCodePredictor310P + qwen3_tts_code_predictor_vllm.Qwen3TTSTalkerCodePredictorForConditionalGenerationVLLM = ( + _Qwen3TTSTalkerCodePredictor310P + ) + qwen3_tts_code_predictor_vllm.CodePredictorBaseModel = _Qwen3CodePredictorBaseModel310P + qwen3_tts_code_predictor_vllm.Qwen3TTSTalkerCodePredictorModelVLLM = _Qwen3CodePredictorBaseModel310P prompt_embeds_builder.Qwen3TTSPromptEmbedsBuilder = _Qwen3TTSPromptEmbedsBuilder310P qwen3_code_predictor.CodePredictorAttention = _Qwen3CodePredictorAttention310P qwen3_code_predictor.CodePredictorDecoderLayer = _Qwen3CodePredictorDecoderLayer310P qwen3_code_predictor.CodePredictorBaseModel = _Qwen3CodePredictorBaseModel310P + + _PATCHED = True + + +def apply_code2wav_patches() -> None: + """Install the 310P Code2Wav runtime patch.""" + global _CODE2WAV_PATCHED + + if _CODE2WAV_PATCHED: + return + + qwen3_tts_code2wav.Qwen3TTSCode2Wav = _Qwen3TTSCode2Wav310P + + _CODE2WAV_PATCHED = True diff --git a/vllm_omni/platforms/npu/_310p/patch/worker.py b/vllm_omni/platforms/npu/_310p/patch/worker.py index 3116b54da43..49ee6052cb6 100644 --- a/vllm_omni/platforms/npu/_310p/patch/worker.py +++ b/vllm_omni/platforms/npu/_310p/patch/worker.py @@ -12,6 +12,9 @@ from __future__ import annotations +from vllm.v1.sample.sampler import Sampler +from vllm_ascend._310p.sample.sampler import AscendSampler310 + from vllm_omni.platforms.npu._310p import disable_jit_compile from vllm_omni.platforms.npu.worker import base as worker_base @@ -24,4 +27,6 @@ def _init_device(self): def apply_patch() -> None: + # Triton-Ascend does not target 310P; use vLLM's native penalty path. + AscendSampler310.apply_penalties = staticmethod(Sampler.apply_penalties) worker_base.OmniNPUWorkerBase = _OmniNPUWorkerBase310P diff --git a/vllm_omni/platforms/npu/models/qwen3_tts_code2wav.py b/vllm_omni/platforms/npu/models/qwen3_tts_code2wav.py index 282689bca47..bd74e9dc6ad 100644 --- a/vllm_omni/platforms/npu/models/qwen3_tts_code2wav.py +++ b/vllm_omni/platforms/npu/models/qwen3_tts_code2wav.py @@ -1,15 +1,18 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Monkey-patch ``Qwen3TTSCode2Wav.__init__`` for NPU Code2Wav runtime knobs.""" +"""Patch Qwen3-TTS Code2Wav NPU runtime setup and weight preparation.""" from __future__ import annotations from typing import TYPE_CHECKING import torch +import torch.nn as nn +import torch_npu from vllm.config import VllmConfig from vllm.logger import init_logger +from vllm_ascend.utils import maybe_trans_nz if TYPE_CHECKING: pass @@ -18,6 +21,8 @@ _PATCHED = False _original_init = None +_original_load_weights = None +_ACL_FORMAT_FRACTAL_Z = 4 def _prepare_npu_code2wav_runtime() -> None: @@ -35,14 +40,43 @@ def _patched_init(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: _original_init(self, vllm_config=vllm_config, prefix=prefix) +def _prepare_npu_decoder_weights(decoder: nn.Module) -> None: + linear_count = 0 + conv_count = 0 + with torch.no_grad(): + for module in decoder.modules(): + if isinstance(module, nn.Linear): + module.weight.data = maybe_trans_nz(module.weight.data) + linear_count += 1 + elif isinstance(module, (nn.Conv1d, nn.ConvTranspose1d)) and module.groups == 1: + module.weight.data = torch_npu.npu_format_cast(module.weight.data.contiguous(), _ACL_FORMAT_FRACTAL_Z) + conv_count += 1 + + logger.info("Prepared NPU Code2Wav weights: linear=%d conv=%d", linear_count, conv_count) + + +def _patched_load_weights(self, weights): + assert _original_load_weights is not None + loaded = _original_load_weights(self, weights) + device = self.vllm_config.device_config.device + runtime_dtype = getattr(self, "_npu_decoder_runtime_dtype", lambda _: torch.float32)(device) + self.decoder.to(device=device, dtype=runtime_dtype) + _prepare_npu_decoder_weights(self.decoder) + if runtime_dtype != torch.float32 and hasattr(self.decoder, "precompute_snake_caches"): + self.decoder.precompute_snake_caches() + return loaded + + def apply_qwen3_tts_code2wav_patch() -> None: - global _PATCHED, _original_init + global _PATCHED, _original_init, _original_load_weights if _PATCHED: return from vllm_omni.model_executor.models.qwen3_tts.qwen3_tts_code2wav import Qwen3TTSCode2Wav _original_init = Qwen3TTSCode2Wav.__init__ + _original_load_weights = Qwen3TTSCode2Wav.load_weights Qwen3TTSCode2Wav.__init__ = _patched_init # type: ignore[method-assign] + Qwen3TTSCode2Wav.load_weights = _patched_load_weights # type: ignore[method-assign] _PATCHED = True - logger.debug("Applied NPU patch for Qwen3TTSCode2Wav.__init__") + logger.debug("Applied NPU patch for Qwen3TTSCode2Wav") diff --git a/vllm_omni/platforms/npu/models/qwen3_tts_tokenizer_v2.py b/vllm_omni/platforms/npu/models/qwen3_tts_tokenizer_v2.py new file mode 100644 index 00000000000..f89ecded9e7 --- /dev/null +++ b/vllm_omni/platforms/npu/models/qwen3_tts_tokenizer_v2.py @@ -0,0 +1,47 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +"""NPU patches for the Qwen3-TTS 12Hz tokenizer decoder.""" + +from __future__ import annotations + +import torch +import torch_npu +from vllm.logger import init_logger + +logger = init_logger(__name__) + +_PATCHED = False + + +def _apply_rotary_pos_emb_npu( + q: torch.Tensor, + k: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, + position_ids=None, + unsqueeze_dim: int = 1, +) -> tuple[torch.Tensor, torch.Tensor]: + del position_ids + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + return torch_npu.npu_rotary_mul(q, cos, sin), torch_npu.npu_rotary_mul(k, cos, sin) + + +def _rms_norm_forward_npu(self, hidden_states: torch.Tensor) -> torch.Tensor: + return torch_npu.npu_rms_norm(hidden_states, self.weight, epsilon=self.variance_epsilon)[0] + + +def apply_qwen3_tts_tokenizer_v2_patch() -> None: + global _PATCHED + if _PATCHED: + return + + from vllm_omni.model_executor.models.qwen3_tts.tokenizer_12hz import ( + modeling_qwen3_tts_tokenizer_v2, + ) + + modeling_qwen3_tts_tokenizer_v2.apply_rotary_pos_emb = _apply_rotary_pos_emb_npu + modeling_qwen3_tts_tokenizer_v2.Qwen3TTSTokenizerV2DecoderRMSNorm.forward = _rms_norm_forward_npu + _PATCHED = True + logger.debug("Applied NPU patch for Qwen3-TTS 12Hz tokenizer decoder") diff --git a/vllm_omni/platforms/npu/platform.py b/vllm_omni/platforms/npu/platform.py index 30aafdb324b..fe49455d506 100644 --- a/vllm_omni/platforms/npu/platform.py +++ b/vllm_omni/platforms/npu/platform.py @@ -37,8 +37,12 @@ def __init__(self) -> None: from vllm_omni.platforms.npu.models.qwen3_tts_code2wav import ( apply_qwen3_tts_code2wav_patch, ) + from vllm_omni.platforms.npu.models.qwen3_tts_tokenizer_v2 import ( + apply_qwen3_tts_tokenizer_v2_patch, + ) apply_qwen3_tts_code2wav_patch() + apply_qwen3_tts_tokenizer_v2_patch() apply_310p_patches() @classmethod