From 363c69b1cb6ea850da1c072c8d89c2c6e0521d80 Mon Sep 17 00:00:00 2001 From: Lion Summerbell Date: Wed, 30 Sep 2026 20:08:10 +0200 Subject: [PATCH 1/2] fix(parakeet-trt): enable timestamps in decoder --- caul-core/caul_core/config.py | 2 - caul/caul/tasks/inference/parakeet.py | 8 +-- caul/caul/tasks/inference/parakeet_trt.py | 30 +++++++-- test/unit/mock.py | 3 +- test/unit/test_parakeet_trt.py | 74 +++++++++++++++++++++-- 5 files changed, 95 insertions(+), 22 deletions(-) diff --git a/caul-core/caul_core/config.py b/caul-core/caul_core/config.py index 6224a43..3b88531 100644 --- a/caul-core/caul_core/config.py +++ b/caul-core/caul_core/config.py @@ -193,14 +193,12 @@ class ParakeetPreprocessorConfig(BasePreprocessorConfig): class ParakeetInferenceRunnerConfig(BaseInferenceRunnerConfig): model: ClassVar[str] = Field(frozen=True, default=ASRModel.PARAKEET) model_name: str = PARAKEET_MODEL_REF - return_timestamps: bool = True class ParakeetTrtInferenceRunnerConfig(BaseInferenceRunnerConfig): model_path: Path | str = None engine_path: Path | str = None model: ClassVar[str] = Field(frozen=True, default=ASRModel.PARAKEET_TRT) - return_timestamps: bool = True class ParakeetPostprocessorConfig(BasePostprocessorConfig): diff --git a/caul/caul/tasks/inference/parakeet.py b/caul/caul/tasks/inference/parakeet.py index 7b02fa7..f832700 100644 --- a/caul/caul/tasks/inference/parakeet.py +++ b/caul/caul/tasks/inference/parakeet.py @@ -31,12 +31,10 @@ def __init__( self, model_name: str = PARAKEET_MODEL_REF, device: TorchDevice = TorchDevice.CPU, - return_timestamps: bool = True, batch_size: int = 4, ): super().__init__(device) self.model_name = model_name - self._return_timestamps = return_timestamps self._model = None self.__transcribe_config = None self._batch_size = batch_size @@ -47,7 +45,6 @@ def _from_config( ) -> FromConfig: return cls( model_name=config.model_name, - return_timestamps=config.return_timestamps, **extras, ) @@ -63,7 +60,7 @@ def _transcribe_config(self): self.__transcribe_config = TranscribeConfig( use_lhotse=False, batch_size=self._batch_size, - timestamps=self._return_timestamps, + timestamps=True, return_hypotheses=True, # Bug in Nemo's AudioToBPEDataset—by default TranscribeConfig spawns 2 # DataLoader workers, but AudioToBPEDataset defines a class TokenizerWrapper @@ -92,7 +89,7 @@ def _transcribe( """ return self._model.transcribe( audio_inputs, - timestamps=self._return_timestamps, + timestamps=True, override_config=self._transcribe_config, ) @@ -113,7 +110,6 @@ def process( # pylint: disable=too-many-locals audios = [str(i.path) for i in input_batch] hypotheses = self._transcribe(audios) - # Get timestamped segments if available, otherwise default to whole text for idx, hyps in enumerate(hypotheses): best_hyp = hyps diff --git a/caul/caul/tasks/inference/parakeet_trt.py b/caul/caul/tasks/inference/parakeet_trt.py index 5f2947a..0bea88f 100644 --- a/caul/caul/tasks/inference/parakeet_trt.py +++ b/caul/caul/tasks/inference/parakeet_trt.py @@ -11,7 +11,6 @@ TorchDevice, ) from icij_common.registrable import FromConfig -from torchaudio.models import Hypothesis from ...exception import TrtEngineLoadError from ...trt import import_trt @@ -26,6 +25,7 @@ if TYPE_CHECKING: import torch + from nemo.collections.asr.parts.utils.rnnt_utils import Hypothesis @cache @@ -76,13 +76,11 @@ def __init__( model_path: Path | str, engine_path: Path | str, device: TorchDevice = TorchDevice.CPU, - return_timestamps: bool = True, batch_size: int = 4, ): ParakeetInferenceRunner.__init__( self, device=device, - return_timestamps=return_timestamps, batch_size=batch_size, ) TrtInferenceMixin.__init__(self) @@ -97,7 +95,6 @@ def _from_config( return cls( model_path=config.model_path, engine_path=config.engine_path, - return_timestamps=config.return_timestamps, **extras, ) @@ -113,8 +110,17 @@ def __enter__(self): import nemo.collections.asr as nemo_asr # pylint: disable=import-outside-toplevel + from omegaconf import open_dict # pylint: disable=import-outside-toplevel + + config = nemo_asr.models.ASRModel.restore_from( + self._model_path, return_config=True + ) + with open_dict(config.decoding): + config.decoding.compute_timestamps = True + self._decoder = nemo_asr.models.ASRModel.restore_from( self._model_path, + override_config_path=config, map_location=self._torch_device, save_restore_connector=_decoder_joint_connector(), strict=False, @@ -129,7 +135,7 @@ def _transcribe( self, audio_inputs: "torch.Tensor | str | Path | Iterable[torch.Tensor | str | Path]", trt_device: TorchDevice = None, - ) -> list[Hypothesis] | list[list[Hypothesis]]: + ) -> "list[Hypothesis] | list[list[Hypothesis]]": """Transcribe audio tensors :param audio_inputs: audio tensors or audio file paths @@ -178,6 +184,18 @@ def _transcribe( enc_len = enc_len.to(self._torch_device) with torch.no_grad(): - return self._decoder.decoding.rnnt_decoder_predictions_tensor( + hypotheses = self._decoder.decoding.rnnt_decoder_predictions_tensor( enc_out, enc_len, return_hypotheses=True ) + + # pylint: disable=import-outside-toplevel + from nemo.collections.asr.parts.utils.timestamp_utils import ( + process_timestamp_outputs, + ) + + # convert frame offsets to seconds, as the model's transcribe() would + return process_timestamp_outputs( + hypotheses, + self._decoder.encoder.subsampling_factor, + self._decoder.cfg.preprocessor.window_stride, + ) diff --git a/test/unit/mock.py b/test/unit/mock.py index 418c0cf..3fb1732 100644 --- a/test/unit/mock.py +++ b/test/unit/mock.py @@ -69,10 +69,9 @@ def __init__( self, model_name: str, device: TorchDevice | torch.device = TorchDevice.CPU, - return_timestamps: bool = True, batch_size: int = 4, ): - super().__init__(model_name, device, return_timestamps, batch_size) + super().__init__(model_name, device, batch_size) def __enter__(self): self._model = MockParakeetModel() diff --git a/test/unit/test_parakeet_trt.py b/test/unit/test_parakeet_trt.py index 6d8c889..daac4fc 100644 --- a/test/unit/test_parakeet_trt.py +++ b/test/unit/test_parakeet_trt.py @@ -2,6 +2,7 @@ import pytest import torch +from omegaconf import OmegaConf from caul.exception import TrtEngineLoadError from caul.tasks.inference.parakeet_trt import ParakeetTrtInferenceRunner @@ -9,6 +10,9 @@ _ENGINE_PATH = "/fake/encoder.trt" _INFERENCE_HANDLER_PATH = "caul.tasks.inference.parakeet_trt.TrtInferenceHandler" +_PROCESS_TIMESTAMPS_PATH = ( + "nemo.collections.asr.parts.utils.timestamp_utils.process_timestamp_outputs" +) _BATCH_SIZE = 2 _SIGNAL_LEN = 16000 _AUDIO_INPUT = torch.zeros(_BATCH_SIZE, _SIGNAL_LEN) @@ -36,7 +40,10 @@ def test__builds_length_tensor_from_audio_shape(self): mock_inference_runner = _mock_inference_runner() mock_trt_handler = _mock_trt_handler(_ENC_OUT, _ENC_OUT_LEN) - with patch(_INFERENCE_HANDLER_PATH, mock_trt_handler): + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH), + ): mock_inference_runner._transcribe( _AUDIO_INPUT, trt_device=torch.device("cpu") ) @@ -50,7 +57,10 @@ def test__encoder_outputs_forwarded_to_decoder(self): mock_inference_runner = _mock_inference_runner() mock_trt_handler = _mock_trt_handler(_ENC_OUT, _ENC_OUT_LEN) - with patch(_INFERENCE_HANDLER_PATH, mock_trt_handler): + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH), + ): mock_inference_runner._transcribe( _AUDIO_INPUT, trt_device=torch.device("cpu") ) @@ -71,18 +81,45 @@ def test__returns_decoder_predictions(self): expected ) - with patch(_INFERENCE_HANDLER_PATH, mock_trt_handler): + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH, side_effect=lambda hyps, *_: hyps), + ): result = mock_inference_runner._transcribe( _AUDIO_INPUT, trt_device=torch.device("cpu") ) assert result is expected + def test__converts_decoder_timestamps_to_seconds(self): + mock_inference_runner = _mock_inference_runner() + mock_trt_handler = _mock_trt_handler(_ENC_OUT, _ENC_OUT_LEN) + decoder = mock_inference_runner._decoder + decoder.encoder.subsampling_factor = 8 + decoder.cfg.preprocessor.window_stride = 0.01 + hypotheses = [MagicMock(), MagicMock()] + decoder.decoding.rnnt_decoder_predictions_tensor.return_value = hypotheses + processed = [MagicMock(), MagicMock()] + + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH, return_value=processed) as mock_process, + ): + result = mock_inference_runner._transcribe( + _AUDIO_INPUT, trt_device=torch.device("cpu") + ) + + mock_process.assert_called_once_with(hypotheses, 8, 0.01) + assert result is processed + def test__requests_hypotheses_with_timestamps_from_decoder(self): mock_inference_runner = _mock_inference_runner() mock_trt_handler = _mock_trt_handler(_ENC_OUT, _ENC_OUT_LEN) - with patch(_INFERENCE_HANDLER_PATH, mock_trt_handler): + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH), + ): mock_inference_runner._transcribe( _AUDIO_INPUT, trt_device=torch.device("cpu") ) @@ -98,7 +135,10 @@ def test__builds_length_tensor_from_original_audio_not_padded_audio(self): short_audio = torch.zeros(_SIGNAL_LEN // 2) long_audio = torch.zeros(_SIGNAL_LEN) - with patch(_INFERENCE_HANDLER_PATH, mock_trt_handler): + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH), + ): mock_inference_runner._transcribe( [short_audio, long_audio], trt_device=torch.device("cpu") ) @@ -118,6 +158,7 @@ def test__loads_path_inputs_as_tensors(self): with ( patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH), patch( "caul.tasks.inference.parakeet_trt.load_audio", side_effect=loaded ) as mock_load, @@ -153,12 +194,33 @@ def test__raises_when_engine_fails_to_deserialize(self): assert runner._decoder is None + def test__restores_decoder_with_compute_timestamps(self): + runner = ParakeetTrtInferenceRunner(PARAKEET_MODEL_REF, _ENGINE_PATH) + config = OmegaConf.create({"decoding": {"strategy": "greedy_batch"}}) + OmegaConf.set_struct(config, True) + mock_restore = MagicMock(side_effect=[config, MagicMock()]) + + with ( + patch("caul.tasks.inference.parakeet_trt.import_trt"), + patch("builtins.open", mock_open(read_data=b"engine")), + patch("nemo.collections.asr.models.ASRModel.restore_from", mock_restore), + ): + runner.__enter__() + + config_call, model_call = mock_restore.call_args_list + assert config_call.kwargs["return_config"] is True + override = model_call.kwargs["override_config_path"] + assert override.decoding.compute_timestamps is True + def test__pads_short_batches_to_engine_minimum(self): mock_inference_runner = _mock_inference_runner() mock_trt_handler = _mock_trt_handler(_ENC_OUT, _ENC_OUT_LEN) short_audio = [torch.ones(1600), torch.ones(8000)] - with patch(_INFERENCE_HANDLER_PATH, mock_trt_handler): + with ( + patch(_INFERENCE_HANDLER_PATH, mock_trt_handler), + patch(_PROCESS_TIMESTAMPS_PATH), + ): mock_inference_runner._transcribe( short_audio, trt_device=torch.device("cpu") ) From 0d88364e335c2f54e00b27ec7c7014966b84d083 Mon Sep 17 00:00:00 2001 From: Lion Summerbell Date: Wed, 30 Sep 2026 20:08:42 +0200 Subject: [PATCH 2/2] bump caul-core --- caul/pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/caul/pyproject.toml b/caul/pyproject.toml index 5a48a3b..075cd86 100644 --- a/caul/pyproject.toml +++ b/caul/pyproject.toml @@ -14,7 +14,7 @@ dependencies = [ "pydantic-extra-types[pycountry]~=2.11", "typer~=0.24", "huggingface-hub~=1.11", - "caul-core~=0.6.5", + "caul-core~=0.6.6", "tiktoken~=0.13", ]