Files
FunASR/funasr/models/fun_asr_nano/inference_vllm_pipeline.py
freedakgmail 6116b1f3c6
Update API Documentation / build-api-docs (push) Has been cancelled
Initial commit: FunASR Speech Recognition Toolkit
Add complete FunASR codebase including models, runtime, and documentation.
2026-07-09 22:38:58 +08:00

373 lines
14 KiB
Python

#!/usr/bin/env python3
# -*- encoding: utf-8 -*-
# Copyright FunASR (https://github.com/alibaba-damo-academy/FunASR). All Rights Reserved.
# MIT License (https://opensource.org/licenses/MIT)
"""
Fun-ASR-Nano vLLM Pipeline: VAD + ASR(vLLM) + Speaker Diarization.
Replicates AutoModel's inference_with_vad pipeline but uses vLLM for
the LLM decoding step, enabling batch processing of all VAD segments
in a single generate() call.
Usage:
from funasr.models.fun_asr_nano.inference_vllm_pipeline import FunASRNanoVLLMPipeline
model = FunASRNanoVLLMPipeline(
model="FunAudioLLM/Fun-ASR-Nano-2512",
vad_model="fsmn-vad",
spk_model="cam++",
tensor_parallel_size=2,
)
results = model.generate("long_meeting.wav", language="中文")
"""
import logging
import os
import re
import time
from typing import List, Optional, Union
import numpy as np
import torch
import torch.nn as nn
logger = logging.getLogger(__name__)
dtype_map = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}
def _clean_text(text: str) -> str:
"""Remove tags, fillers, and garbage from output."""
text = re.sub(r"<[^>]*>|</[^>]*>", "", text)
text = re.sub(r"(>.{2,8}?)\1{3,}", "", text)
text = re.sub(r"\[breath\]|\[noise\]|/sil|endofbreak|FFFF", "", text)
text = re.sub(r"\s+", " ", text)
text = text.replace("", "").lstrip(">")
return text.strip()
class FunASRNanoVLLMPipeline:
"""VAD + ASR(vLLM) + Speaker pipeline.
Pipeline:
1. VAD: segment long audio into speech regions (torch)
2. ASR: batch ALL segments through vLLM in single generate() call
3. Speaker: extract embeddings per segment, cluster (torch)
4. Combine: merge text + timestamps + speaker labels
Args:
model: Fun-ASR-Nano model name or path.
vad_model: VAD model name (e.g. "fsmn-vad"). None to disable.
vad_kwargs: VAD config (e.g. {"max_single_segment_time": 30000}).
spk_model: Speaker model name (e.g. "cam++"). None to disable.
hub: "ms" or "hf".
device: Device for audio encoder + VAD + speaker.
dtype: Compute dtype for ASR.
tensor_parallel_size: GPUs for vLLM.
gpu_memory_utilization: GPU memory fraction for vLLM.
max_model_len: Maximum sequence length for vLLM.
"""
def __init__(
self,
model: str = "FunAudioLLM/Fun-ASR-Nano-2512",
vad_model: str = None,
vad_kwargs: dict = None,
spk_model: str = None,
spk_kwargs: dict = None,
hub: str = "ms",
device: str = "cuda:0",
dtype: str = "bf16",
tensor_parallel_size: int = 1,
gpu_memory_utilization: float = 0.8,
max_model_len: int = 4096,
enforce_eager: bool = False,
**kwargs,
):
from funasr.models.fun_asr_nano.inference_vllm import FunASRNanoVLLM
# ASR engine (vLLM)
self.asr_engine = FunASRNanoVLLM.from_pretrained(
model=model, hub=hub, device=device, dtype=dtype,
tensor_parallel_size=tensor_parallel_size,
gpu_memory_utilization=gpu_memory_utilization,
max_model_len=max_model_len, enforce_eager=enforce_eager,
**kwargs,
)
# VAD model (torch)
self.vad_model = None
if vad_model is not None:
from funasr import AutoModel
vad_kw = vad_kwargs or {}
self.vad_model = AutoModel(
model=vad_model, device=device, disable_update=True, **vad_kw
)
# Speaker model (torch)
self.spk_model = None
self.cb_model = None
if spk_model is not None:
from funasr import AutoModel
from funasr.models.campplus.cluster_backend import ClusterBackend
spk_kw = spk_kwargs or {}
self.spk_model = AutoModel(
model=spk_model, device=device, disable_update=True, **spk_kw
)
cb_kwargs = spk_kw.get("cb_kwargs", {})
self.cb_model = ClusterBackend(**cb_kwargs).to(device)
self.device = device
self.sample_rate = 16000
def generate(
self,
input: Union[str, List[str]],
hotwords: List[str] = None,
language: str = None,
itn: bool = True,
max_new_tokens: int = 512,
batch_size_s: int = 300,
return_spk_res: bool = True,
**kwargs,
) -> List[dict]:
"""Run the full pipeline: VAD → ASR(vLLM) → Speaker.
Args:
input: Audio file path(s).
hotwords: Hotwords for ASR.
language: Language hint.
itn: Inverse text normalization.
max_new_tokens: Max tokens per segment.
batch_size_s: Max batch duration in seconds (for memory control).
return_spk_res: Whether to return speaker info.
Returns:
List of dicts: [{"key", "text", "timestamp", "sentence_info"}]
"""
if isinstance(input, str):
input = [input]
results_all = []
for audio_path in input:
result = self._process_one(
audio_path, hotwords=hotwords, language=language,
itn=itn, max_new_tokens=max_new_tokens,
batch_size_s=batch_size_s, return_spk_res=return_spk_res,
**kwargs,
)
results_all.append(result)
return results_all
def _process_one(self, audio_path, **kwargs):
"""Process a single audio file through the full pipeline."""
from funasr.utils.load_utils import load_audio_text_image_video
from funasr.utils.vad_utils import slice_padding_audio_samples
key = os.path.splitext(os.path.basename(audio_path))[0]
# Load audio
audio_data = load_audio_text_image_video(audio_path, fs=self.sample_rate)
if isinstance(audio_data, torch.Tensor):
audio_np = audio_data.numpy()
else:
audio_np = np.array(audio_data)
speech_length = len(audio_np)
# Step 1: VAD
if self.vad_model is not None:
vad_res = self.vad_model.generate(input=audio_path, cache={}, is_final=True)
vad_segments = vad_res[0]["value"] # [[start_ms, end_ms], ...]
else:
vad_segments = [[0, int(speech_length / self.sample_rate * 1000)]]
if not vad_segments:
return {"key": key, "text": "", "timestamp": []}
n_segments = len(vad_segments)
logger.info(f"VAD: {n_segments} segments for {key}")
# Step 2: Slice audio by VAD segments and encode
segment_audios = []
for seg in vad_segments:
start_sample = int(seg[0] * self.sample_rate / 1000)
end_sample = int(seg[1] * self.sample_rate / 1000)
end_sample = min(end_sample, speech_length)
segment_audios.append(audio_np[start_sample:end_sample])
# Step 3: Batch ASR via vLLM
# Encode all segments and build prompts
from vllm import SamplingParams
try:
from vllm.inputs import EmbedsPrompt
except ImportError:
from vllm.inputs.data import EmbedsPrompt
from funasr.models.fun_asr_nano.vllm_utils import resolve_repetition_penalty
prompts = []
for seg_audio in segment_audios:
seg_tensor = torch.from_numpy(seg_audio).float()
adaptor_out, adaptor_out_lens, _, _ = self.asr_engine._encode_audio(seg_tensor)
input_embeds = self.asr_engine._build_input_embeds(
adaptor_out, adaptor_out_lens,
hotwords=kwargs.get("hotwords"),
language=kwargs.get("language"),
itn=kwargs.get("itn", True),
)
prompts.append(EmbedsPrompt(prompt_embeds=input_embeds.float()))
params = SamplingParams(
max_tokens=kwargs.get("max_new_tokens", 512),
temperature=0.0,
# Prompt-embeds mode has no token IDs to penalize; see #2948.
repetition_penalty=resolve_repetition_penalty(
kwargs.get("repetition_penalty", 1.0)
),
skip_special_tokens=True,
)
# Single batch generate for ALL segments
t0 = time.perf_counter()
outputs = self.asr_engine.vllm_engine.generate(prompts, params, use_tqdm=False)
t1 = time.perf_counter()
logger.info(f"vLLM batch ASR: {n_segments} segments in {t1-t0:.3f}s")
# Decode results
asr_results = []
for output in outputs:
text = output.outputs[0].text
if not text and output.outputs[0].token_ids:
text = self.asr_engine.tokenizer.decode(
list(output.outputs[0].token_ids), skip_special_tokens=True
)
text = _clean_text(text)
asr_results.append(text)
# Step 4: Speaker embeddings (if spk_model configured)
spk_embeddings = None
if self.spk_model is not None and kwargs.get("return_spk_res", True):
from funasr.models.campplus.utils import sv_chunk, postprocess, distribute_spk
all_segments = []
all_spk_embs = []
for i, seg_audio in enumerate(segment_audios):
vad_seg = [
[vad_segments[i][0] / 1000.0, vad_segments[i][1] / 1000.0, seg_audio]
]
chunks = sv_chunk(vad_seg)
all_segments.extend(chunks)
speech_chunks = [c[2] for c in chunks]
spk_res = self.spk_model.generate(input=speech_chunks, cache={}, is_final=True)
embs = torch.cat([r["spk_embedding"] for r in spk_res], dim=0)
all_spk_embs.append(embs)
if all_spk_embs:
spk_embeddings = torch.cat(all_spk_embs, dim=0)
# Step 5: Combine results
# Merge text with timestamps
full_text = ""
all_timestamps = []
for i, (seg, text) in enumerate(zip(vad_segments, asr_results)):
if text:
if full_text:
full_text += " "
full_text += text
# Simple word-level timestamp from VAD boundaries
all_timestamps.append([int(seg[0]), int(seg[1])])
result = {"key": key, "text": full_text}
# Add timestamps if available from CTC
if self.asr_engine.ctc_decoder is not None:
try:
detailed_timestamps = self._compute_all_timestamps(
segment_audios, vad_segments, asr_results
)
if detailed_timestamps:
result["timestamp"] = detailed_timestamps
except Exception as e:
logger.debug(f"Timestamp computation failed: {e}")
result["timestamp"] = all_timestamps
else:
result["timestamp"] = all_timestamps
# Add speaker info
if spk_embeddings is not None and self.cb_model is not None:
from funasr.models.campplus.utils import postprocess, distribute_spk
all_segments_sorted = sorted(all_segments, key=lambda x: x[0])
labels = self.cb_model(
spk_embeddings.cpu(),
oracle_num=kwargs.get("preset_spk_num", None),
)
sv_output = postprocess(all_segments_sorted, None, labels, spk_embeddings.cpu())
# Build sentence_info
sentence_list = []
for i, (seg, text) in enumerate(zip(vad_segments, asr_results)):
if text:
sentence_list.append({
"start": seg[0],
"end": seg[1],
"text": text,
"timestamp": [[int(seg[0]), int(seg[1])]],
})
distribute_spk(sentence_list, sv_output)
result["sentence_info"] = sentence_list
return result
def _compute_all_timestamps(self, segment_audios, vad_segments, asr_results):
"""Compute CTC timestamps for all segments with VAD offsets."""
from funasr.models.fun_asr_nano.tools.utils import forced_align
all_timestamps = []
for seg_audio, vad_seg, text in zip(segment_audios, vad_segments, asr_results):
if not text:
continue
try:
seg_tensor = torch.from_numpy(seg_audio).float()
from funasr.utils.load_utils import extract_fbank
speech, speech_lengths = extract_fbank(
seg_tensor, data_type="sound",
frontend=self.asr_engine.frontend, is_final=True
)
speech = speech.to(self.device, dtype=torch.float32)
speech_lengths = speech_lengths.to(self.device)
with torch.no_grad():
enc_out, enc_lens = self.asr_engine.audio_encoder(speech, speech_lengths)
dec_out, dec_lens = self.asr_engine.ctc_decoder(enc_out, enc_lens)
ctc_logits = self.asr_engine.ctc.log_softmax(dec_out)
x = ctc_logits[0, :enc_lens[0].item(), :]
target_ids = torch.tensor(
self.asr_engine.ctc_tokenizer.encode(text), dtype=torch.int64
)
if len(target_ids) == 0:
continue
timestamps = forced_align(x, target_ids, self.asr_engine.blank_id)
vad_offset_ms = int(vad_seg[0])
for ts in timestamps:
ts["token"] = self.asr_engine.ctc_tokenizer.decode([ts["token"]])
ts["start_time"] = ts["start_time"] * 6 * 10 / 1000 + vad_offset_ms / 1000
ts["end_time"] = ts["end_time"] * 6 * 10 / 1000 + vad_offset_ms / 1000
all_timestamps.extend(timestamps)
except Exception as e:
logger.debug(f"Timestamp failed for segment: {e}")
all_timestamps.append({
"start_time": vad_seg[0] / 1000,
"end_time": vad_seg[1] / 1000,
"token": text,
})
return all_timestamps
@classmethod
def from_pretrained(cls, model="FunAudioLLM/Fun-ASR-Nano-2512", **kwargs):
"""Convenience constructor."""
return cls(model=model, **kwargs)