Skip to content

Commit 23138d5

Browse files
author
CodexIntegrationCheck
committed
Address latest Qwen ASR review feedback
Signed-off-by: CodexIntegrationCheck <codex-integration-check@invalid.example>
1 parent a0ef7d8 commit 23138d5

3 files changed

Lines changed: 3 additions & 70 deletions

File tree

nemo_curator/models/asr/qwen_asr.py

Lines changed: 0 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -23,14 +23,12 @@
2323
from __future__ import annotations
2424

2525
import gc
26-
import inspect
2726
from copy import deepcopy
2827
from dataclasses import dataclass, field
2928
from typing import Any
3029

3130
import numpy as np
3231
import torch
33-
import transformers
3432
from huggingface_hub import snapshot_download
3533
from loguru import logger
3634

@@ -54,30 +52,6 @@ def _qwen_asr_model_cls() -> Any: # noqa: ANN401
5452
return Qwen3ASRModel
5553

5654

57-
def _patch_transformers_compat() -> None:
58-
"""Accept qwen-asr's decorator-factory syntax on newer Transformers.
59-
60-
This matches the compatibility patch used by the nkoluguri Qwen-ASR
61-
reference. Newer Transformers exposes ``check_model_inputs`` as a plain
62-
decorator, while qwen-asr 0.0.6 still invokes it with parentheses.
63-
"""
64-
try:
65-
original = getattr(transformers, "check_model_inputs", None)
66-
if original is None:
67-
return
68-
parameters = list(inspect.signature(original).parameters.values())
69-
if parameters and parameters[0].name == "func":
70-
71-
def compat_check_model_inputs(*args: Any) -> Any: # noqa: ANN401
72-
if args and callable(args[0]):
73-
return original(args[0])
74-
return original
75-
76-
transformers.check_model_inputs = compat_check_model_inputs
77-
except Exception: # noqa: BLE001, S110
78-
pass
79-
80-
8155
@dataclass
8256
class QwenASRAdapter:
8357
"""Run vLLM-backed Qwen3-ASR over Curator waveform items.
@@ -171,7 +145,6 @@ def load_model(self, *, num_gpus: int) -> None:
171145
)
172146
if model_kwargs["revision"] is None:
173147
del model_kwargs["revision"]
174-
_patch_transformers_compat()
175148
try:
176149
self._model = _qwen_asr_model_cls().LLM(**model_kwargs)
177150
except Exception:

tests/models/asr/test_qwen_asr.py

Lines changed: 2 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@
2727
import torch
2828

2929
from nemo_curator.models.asr.base import ASRAdapter
30-
from nemo_curator.models.asr.qwen_asr import _MIN_SAMPLES, QwenASRAdapter, _patch_transformers_compat
30+
from nemo_curator.models.asr.qwen_asr import _MIN_SAMPLES, QwenASRAdapter
3131
from nemo_curator.stages.audio.inference.asr.stage import ASRStage
3232
from nemo_curator.tasks import AudioTask
3333

@@ -107,17 +107,7 @@ def test_qwen_adapter_copies_nested_vllm_kwargs() -> None:
107107

108108
@pytest.mark.parametrize(
109109
"reserved_key",
110-
[
111-
"model",
112-
"revision",
113-
"gpu_memory_utilization",
114-
"max_inference_batch_size",
115-
"max_new_tokens",
116-
"trust_remote_code",
117-
"enforce_eager",
118-
"enable_prefix_caching",
119-
"prefix_caching_hash_algo",
120-
],
110+
QwenASRAdapter()._model_owned_vllm_kwargs(),
121111
)
122112
def test_qwen_adapter_rejects_adapter_owned_vllm_kwargs(reserved_key: str) -> None:
123113
adapter = QwenASRAdapter(vllm_kwargs={reserved_key: object()})
@@ -190,34 +180,6 @@ def test_load_model_forwards_additional_vllm_kwargs() -> None:
190180
assert model_cls.LLM.call_args.kwargs["max_model_len"] == 8192
191181

192182

193-
def test_load_model_applies_nkoluguri_transformers_compat_patch() -> None:
194-
model_cls = MagicMock()
195-
adapter = QwenASRAdapter()
196-
197-
with (
198-
patch("nemo_curator.models.asr.qwen_asr._patch_transformers_compat") as compat,
199-
patch("nemo_curator.models.asr.qwen_asr._qwen_asr_model_cls", return_value=model_cls),
200-
):
201-
adapter.load_model(num_gpus=1)
202-
203-
compat.assert_called_once_with()
204-
205-
206-
def test_transformers_compat_wraps_plain_decorator_for_factory_syntax() -> None:
207-
def plain_decorator(func): # noqa: ANN001, ANN202
208-
return func
209-
210-
transformers_stub = SimpleNamespace(check_model_inputs=plain_decorator)
211-
with patch("nemo_curator.models.asr.qwen_asr.transformers", transformers_stub):
212-
_patch_transformers_compat()
213-
214-
def decorated() -> str:
215-
return "ok"
216-
217-
assert transformers_stub.check_model_inputs()(decorated)() == "ok"
218-
assert transformers_stub.check_model_inputs(decorated)() == "ok"
219-
220-
221183
def test_load_model_without_qwen_asr_names_required_extras() -> None:
222184
adapter = QwenASRAdapter()
223185
with (

tests/utils/test_vllm_utils.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@
3232
)
3333

3434

35-
class TestValidateVllmKwargs:
35+
class TestVllmKwargs:
3636
def test_accepts_non_conflicting_kwargs(self) -> None:
3737
validate_vllm_kwargs(
3838
{"max_model_len": 8192},
@@ -51,8 +51,6 @@ def test_rejects_conflicts_in_sorted_order(self) -> None:
5151
owner_description="adapter-owned arguments",
5252
)
5353

54-
55-
class TestMergeVllmKwargs:
5654
def test_merges_owned_kwargs_without_mutating_user_kwargs(self) -> None:
5755
user_kwargs = {"max_model_len": 8192, "compilation_config": {"cudagraph_mode": "NONE"}}
5856

0 commit comments

Comments
 (0)