Team Ai
Apppublic

flrtemis/https-huggingface-co-spaces-ftrtemis-moshi

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
stt_from_file_rust_server.py136 linesDownload Raw Back to scripts
1# /// script2# requires-python = ">=3.12"3# dependencies = [4#     "msgpack",5#     "numpy",6#     "sphn",7#     "websockets",8# ]9# ///10import argparse11import asyncio12import time13 14import msgpack15import numpy as np16import sphn17import websockets18 19SAMPLE_RATE = 2400020FRAME_SIZE = 1920  # Send data in chunks21 22 23def load_and_process_audio(file_path):24    """Load an MP3 file, resample to 24kHz, convert to mono, and extract PCM float32 data."""25    pcm_data, _ = sphn.read(file_path, sample_rate=SAMPLE_RATE)26    return pcm_data[0]27 28 29async def receive_messages(websocket):30    transcript = []31 32    async for message in websocket:33        data = msgpack.unpackb(message, raw=False)34        if data["type"] == "Step":35            # This message contains the signal from the semantic VAD, and tells us how36            # much audio the server has already processed. We don't use either here.37            continue38        if data["type"] == "Word":39            print(data["text"], end=" ", flush=True)40            transcript.append(41                {42                    "text": data["text"],43                    "timestamp": [data["start_time"], data["start_time"]],44                }45            )46        if data["type"] == "EndWord":47            if len(transcript) > 0:48                transcript[-1]["timestamp"][1] = data["stop_time"]49        if data["type"] == "Marker":50            # Received marker, stopping stream51            break52 53    return transcript54 55 56async def send_messages(websocket, rtf: float):57    audio_data = load_and_process_audio(args.in_file)58 59    async def send_audio(audio: np.ndarray):60        await websocket.send(61            msgpack.packb(62                {"type": "Audio", "pcm": [float(x) for x in audio]},63                use_single_float=True,64            )65        )66 67    # Start with a second of silence.68    # This is needed for the 2.6B model for technical reasons.69    await send_audio([0.0] * SAMPLE_RATE)70 71    start_time = time.time()72    for i in range(0, len(audio_data), FRAME_SIZE):73        await send_audio(audio_data[i : i + FRAME_SIZE])74 75        expected_send_time = start_time + (i + 1) / SAMPLE_RATE / rtf76        current_time = time.time()77        if current_time < expected_send_time:78            await asyncio.sleep(expected_send_time - current_time)79        else:80            await asyncio.sleep(0.001)81 82    for _ in range(5):83        await send_audio([0.0] * SAMPLE_RATE)84 85    # Send a marker to indicate the end of the stream.86    await websocket.send(87        msgpack.packb({"type": "Marker", "id": 0}, use_single_float=True)88    )89 90    # We'll get back the marker once the corresponding audio has been transcribed,91    # accounting for the delay of the model. That's why we need to send some silence92    # after the marker, because the model will not return the marker immediately.93    for _ in range(35):94        await send_audio([0.0] * SAMPLE_RATE)95 96 97async def stream_audio(url: str, api_key: str, rtf: float):98    """Stream audio data to a WebSocket server."""99    headers = {"kyutai-api-key": api_key}100 101    # Instead of using the header, you can authenticate by adding `?auth_id={api_key}` to the URL102    async with websockets.connect(url, additional_headers=headers) as websocket:103        send_task = asyncio.create_task(send_messages(websocket, rtf))104        receive_task = asyncio.create_task(receive_messages(websocket))105        _, transcript = await asyncio.gather(send_task, receive_task)106 107    return transcript108 109 110if __name__ == "__main__":111    parser = argparse.ArgumentParser()112    parser.add_argument("in_file")113    parser.add_argument(114        "--url",115        help="The url of the server to which to send the audio",116        default="ws://127.0.0.1:8080",117    )118    parser.add_argument("--api-key", default="public_token")119    parser.add_argument(120        "--rtf",121        type=float,122        default=1.01,123        help="The real-time factor of how fast to feed in the audio.",124    )125    args = parser.parse_args()126 127    url = f"{args.url}/api/asr-streaming"128    transcript = asyncio.run(stream_audio(url, args.api_key, args.rtf))129 130    print()131    print()132    for word in transcript:133        print(134            f"{word['timestamp'][0]:7.2f} -{word['timestamp'][1]:7.2f}  {word['text']}"135        )136