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
2 changes: 0 additions & 2 deletions caul-core/caul_core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
8 changes: 2 additions & 6 deletions caul/caul/tasks/inference/parakeet.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -47,7 +45,6 @@ def _from_config(
) -> FromConfig:
return cls(
model_name=config.model_name,
return_timestamps=config.return_timestamps,
**extras,
)

Expand All @@ -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
Expand Down Expand Up @@ -92,7 +89,7 @@ def _transcribe(
"""
return self._model.transcribe(
audio_inputs,
timestamps=self._return_timestamps,
timestamps=True,
override_config=self._transcribe_config,
)

Expand All @@ -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

Expand Down
30 changes: 24 additions & 6 deletions caul/caul/tasks/inference/parakeet_trt.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -26,6 +25,7 @@

if TYPE_CHECKING:
import torch
from nemo.collections.asr.parts.utils.rnnt_utils import Hypothesis


@cache
Expand Down Expand Up @@ -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)
Expand All @@ -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,
)

Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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,
)
2 changes: 1 addition & 1 deletion caul/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
]

Expand Down
3 changes: 1 addition & 2 deletions test/unit/mock.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
74 changes: 68 additions & 6 deletions test/unit/test_parakeet_trt.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,17 @@

import pytest
import torch
from omegaconf import OmegaConf

from caul.exception import TrtEngineLoadError
from caul.tasks.inference.parakeet_trt import ParakeetTrtInferenceRunner
from caul_core import PARAKEET_MODEL_REF

_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)
Expand Down Expand Up @@ -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")
)
Expand All @@ -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")
)
Expand All @@ -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")
)
Expand All @@ -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")
)
Expand All @@ -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,
Expand Down Expand Up @@ -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")
)
Expand Down
Loading