mazesmazes/tiny-audio
1177
1"""Custom inference handler for HuggingFace Inference Endpoints."""2 3import base644import binascii5import os6from pathlib import Path7from typing import TYPE_CHECKING, Any8 9if TYPE_CHECKING:10 from .alignment import QwenForcedAligner11 from .asr_modeling import ASRModel12 from .asr_pipeline import ASRPipeline13 from .diarization import NemotronDiarizer14 from .diarization import get_device as _best_device15else:16 try:17 # For remote execution, imports are relative18 from .alignment import QwenForcedAligner19 from .asr_modeling import ASRModel20 from .asr_pipeline import ASRPipeline21 from .diarization import NemotronDiarizer22 from .diarization import get_device as _best_device23 except ImportError:24 # For local execution, imports are not relative25 from alignment import QwenForcedAligner26 from asr_modeling import ASRModel27 from asr_pipeline import ASRPipeline28 from diarization import NemotronDiarizer29 from diarization import get_device as _best_device30 31 32def decode_inputs(inputs: Any) -> Any:33 """A JSON request's base64 audio as bytes; paths, raw bytes and arrays unchanged.34 35 The toolkit hands `inputs` over as raw bytes for an `audio/*` body, but a36 JSON body can only carry audio as a base64 string, which the pipeline37 would otherwise read as a file path.38 """39 if not isinstance(inputs, str) or Path(inputs).exists():40 return inputs41 try:42 return base64.b64decode(inputs, validate=True)43 except (binascii.Error, ValueError):44 return inputs45 46 47class EndpointHandler:48 """HuggingFace Inference Endpoints handler for ASR model.49 50 Handles model loading, warmup, and inference requests for deployment51 on HuggingFace Inference Endpoints or similar services.52 """53 54 def __init__(self, path: str = ""):55 """Initialize the endpoint handler.56 57 Args:58 path: Path to model directory or HuggingFace model ID59 """60 os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")61 62 # `ASRModel.from_pretrained` constructs its own submodules and forwards63 # **kwargs to `ASRModel.__init__`, which discards them -- no loader64 # reads `device_map`, `torch_dtype` or `low_cpu_mem_usage`, and it sets65 # `_is_loading_from_pretrained` precisely to keep `device_map="auto"`66 # out of the sub-model loaders. Passing them here did nothing, so the67 # model stayed on CPU and a GPU Inference Endpoint silently decoded on68 # CPU. Place it explicitly instead. dtype comes from69 # `config.model_dtype` and the attention backend from70 # `config.attn_implementation`, which already downgrades FA2 when71 # flash_attn is missing.72 self.model = ASRModel.from_pretrained(path)73 self.device = _best_device()74 # PreTrainedModel.to is functools.wraps'd, which pyright cannot bind as a method.75 self.model.to(self.device) # pyright: ignore[reportArgumentType]76 self.model.eval()77 78 self.pipe = ASRPipeline(79 model=self.model,80 feature_extractor=self.model.feature_extractor,81 tokenizer=self.model.tokenizer,82 device=self.device,83 )84 # Load the forced aligner (return_timestamps) and Nemotron diarizer85 # (return_speakers) at boot, as the Spaces demo does, so the first86 # request that asks for them doesn't pay the download and load.87 QwenForcedAligner.get_instance()88 NemotronDiarizer.get_instance()89 90 def __call__(self, data: dict[str, Any]) -> dict[str, Any] | list[dict[str, Any]]:91 """Process an inference request.92 93 Args:94 data: Request data containing 'inputs' (audio path, bytes or base6495 string) and optional 'parameters' for the pipeline, e.g.96 `return_timestamps`, `return_speakers`, `num_speakers`,97 `max_speakers` -- the same options as the Spaces demo98 99 Returns:100 Transcription result with 'text' key; 'words' with timestamps, plus101 speaker labels and 'speaker_segments' when diarizing102 """103 inputs = data.get("inputs")104 if inputs is None:105 msg = "Missing 'inputs' in request data"106 raise ValueError(msg)107 108 # Pass through any parameters from request, let model config provide defaults109 params = data.get("parameters", {})110 111 result: dict[str, Any] | list[dict[str, Any]] = self.pipe(decode_inputs(inputs), **params)112 return result113 