Team Ai
Modelpublic

mazesmazes/tiny-audio

sourceHugging Facemitupdated 16h agoView on Hugging Face
1likes177downloads
handler.py113 linesDownload Raw Back to root
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