File size: 5,860 Bytes
41d27d9
 
 
 
c1f7cc8
41d27d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cbb0ac8
 
41d27d9
cbb0ac8
 
 
 
 
 
 
41d27d9
 
 
 
 
 
cbb0ac8
41d27d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cbb0ac8
41d27d9
 
 
cbb0ac8
41d27d9
cbb0ac8
41d27d9
 
cbb0ac8
41d27d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cbb0ac8
41d27d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c1f7cc8
 
41d27d9
 
 
 
cbb0ac8
41d27d9
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
"""
Optimised NeMo Parakeet-TDT streaming demo for CPU-only Hugging Face Spaces
"""

import os, time, threading, queue, logging, re
import numpy as np
import gradio as gr
from scipy import signal
import torch
from nemo.collections.asr.models import ASRModel

# ────────────────────────────────────────────────
# General CPU settings (2 vCPU space)
# ────────────────────────────────────────────────
os.environ["OMP_NUM_THREADS"] = "2"          # One MKL/OpenMP thread per vCPU
torch.set_num_threads(2)
torch.backends.quantized.engine = "fbgemm"   # Fastest INT8 kernels on x86

# ────────────────────────────────────────────────
# Logging
# ────────────────────────────────────────────────
logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s - %(message)s",
    datefmt="%H:%M:%S",
)
logger = logging.getLogger("asr_app")

# ────────────────────────────────────────────────
# Constants
# ────────────────────────────────────────────────
SR              = 16_000
CHUNK_SECONDS   = 4
CHUNK_SAMPLES   = SR * CHUNK_SECONDS

# ────────────────────────────────────────────────
# Prepare UI Description data from README.md file
# ────────────────────────────────────────────────
with open('README.md', 'r', encoding='utf-8') as file:
    content = file.read()
README_CONTENT_without_YAML = re.sub(r'^---.*?---\s*', '', content, flags=re.DOTALL)

# ────────────────────────────────────────────────
# ASR Application
# ────────────────────────────────────────────────
class ASRApp:
    def __init__(self):
        self.audio_queue      = queue.Queue(maxsize=8)
        self.transcript_queue = queue.Queue()
        self.transcript_list  = []
        self._load_model()
        self._start_worker()

    # ---------- helpers ----------
    def _log(self, func: str, msg: str):
        logger.info(
            f"{func} | audio_q={self.audio_queue.qsize():02}, "
            f"txt_q={self.transcript_queue.qsize():02} | {msg}"
        )

    # ---------- model ----------
    def _load_model(self):
        self._log("load_model", "loading Parakeet-TDT-0.6B-V2 (CPU)…")
        t0 = time.time()
        model = ASRModel.from_pretrained(
            model_name="nvidia/parakeet-tdt-0.6b-v2",
            map_location="cpu",
        )
        model.eval()
        self.asr_model = model
        self._log("load_model", f"model ready in {time.time()-t0:.1f}s")
        with torch.inference_mode():
            _ = self.asr_model.transcribe([np.zeros(SR, dtype=np.float32)])
        self._log("load_model", "warm-up done")
    
    # ---------- threading ----------
    def _start_worker(self):
        threading.Thread(target=self._worker, daemon=True).start()

    def _worker(self):
        buf = np.array([], dtype=np.float32)
        while True:
            try:
                while len(buf) < CHUNK_SAMPLES:
                    buf = np.concatenate([buf, self.audio_queue.get()])
                chunk, buf = buf[:CHUNK_SAMPLES], buf[CHUNK_SAMPLES:]
                self._log("_worker", f"β†’ transcribe {len(chunk)} samples")
                t0 = time.time()
                with torch.inference_mode():
                    out = self.asr_model.transcribe([chunk])
                dur = time.time() - t0
                text = out[0].text
                self._log("_worker", f"inference {dur:.2f}s β†’ β€œ{text}”")
                self.transcript_queue.put(text)
            except Exception as e:
                self._log("_worker", f"ASR error: {e}")

    # ---------- audio preprocessing ----------
    def _preprocess(self, audio):
        sr, y = audio
        y = signal.resample_poly(y, SR, sr)
        y = y.astype(np.float32)
        y /= (np.abs(y).max() + 1e-9)
        return y

    # ---------- Gradio stream callback ----------
    def stream_fn(self, audio):
        self._log("stream_fn", "audio arrived")
        self.audio_queue.put(self._preprocess(audio))
        while not self.transcript_queue.empty():
            self.transcript_list.append(self.transcript_queue.get())
        return (
            " ".join(self.transcript_list)
            if self.transcript_list
            else "…listening…"
        )

# ────────────────────────────────────────────────
# Gradio UI
# ────────────────────────────────────────────────
asr_app = ASRApp()
with gr.Blocks() as demo:
    mic = gr.Audio(
        sources=["microphone"],
        type="numpy",
        streaming=True,
        label="Microphone",
    )
    out = gr.Textbox(label="Transcription")
    gr.Markdown(README_CONTENT_without_YAML)
    
    mic.stream(
        fn=asr_app.stream_fn,
        inputs=mic,
        outputs=out,
        stream_every=0.5,
    )

asr_app._log("main", "launching UI")
demo.launch()