Team Ai
Apppublic

Teapack1/Assistant-Audio-Intent-Classification

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
6likes
server.py133 linesDownload Raw Back to root
1from fastapi import FastAPI, WebSocket, Request, WebSocketDisconnect2from fastapi.staticfiles import StaticFiles3from fastapi.responses import HTMLResponse4from fastapi.templating import Jinja2Templates5 6import numpy as np7from transformers import pipeline8import torch9from transformers.pipelines.audio_utils import ffmpeg_microphone_live10 11device = "cuda:0" if torch.cuda.is_available() else "cpu"12 13classifier = pipeline(14    "audio-classification", model="MIT/ast-finetuned-speech-commands-v2", device=device15)16intent_class_pipe = pipeline(17    "audio-classification", model="anton-l/xtreme_s_xlsr_minds14", device=device18)19 20 21async def launch_fn(22    wake_word="marvin",23    prob_threshold=0.5,24    chunk_length_s=2.0,25    stream_chunk_s=0.25,26    debug=False,27):28    if wake_word not in classifier.model.config.label2id.keys():29        raise ValueError(30            f"Wake word {wake_word} not in set of valid class labels, pick a wake word in the set {classifier.model.config.label2id.keys()}."31        )32 33    sampling_rate = classifier.feature_extractor.sampling_rate34 35    mic = ffmpeg_microphone_live(36        sampling_rate=sampling_rate,37        chunk_length_s=chunk_length_s,38        stream_chunk_s=stream_chunk_s,39    )40 41    print("Listening for wake word...")42    for prediction in classifier(mic):43        prediction = prediction[0]44        if debug:45            print(prediction)46        if prediction["label"] == wake_word:47            if prediction["score"] > prob_threshold:48                return True49 50 51async def listen(websocket, chunk_length_s=2.0, stream_chunk_s=2.0):52    sampling_rate = intent_class_pipe.feature_extractor.sampling_rate53 54    mic = ffmpeg_microphone_live(55        sampling_rate=sampling_rate,56        chunk_length_s=chunk_length_s,57        stream_chunk_s=stream_chunk_s,58    )59    audio_buffer = []60    61    print("Listening")62    for i in range(4):63        audio_chunk = next(mic)64        audio_buffer.append(audio_chunk["raw"])65        66        prediction = intent_class_pipe(audio_chunk["raw"])67        print(prediction)68        await websocket.send_text(f"chunk: {prediction[0]['label']} | {i+1} / 4")69    70        if await is_silence(audio_chunk["raw"], threshold=0.7):71            print("Silence detected, processing audio.")72            break73 74    combined_audio = np.concatenate(audio_buffer)75    prediction = intent_class_pipe(combined_audio)76    top_3_predictions = prediction[:3]77    formatted_predictions = "\n".join([f"{pred['label']}: {pred['score'] * 100:.2f}%" for pred in top_3_predictions])78    await websocket.send_text(f"classes: \n{formatted_predictions}")79    return80 81 82async def is_silence(audio_chunk, threshold):83    silence = intent_class_pipe(audio_chunk)84    if silence[0]["label"] == "silence" and silence[0]["score"] > threshold:85        return True86    else:87        return False88 89 90# Initialize FastAPI app91app = FastAPI()92 93# Set up static file directory94app.mount("/static", StaticFiles(directory="static"), name="static")95 96# Jinja2 Template for HTML rendering97templates = Jinja2Templates(directory="templates")98 99 100@app.get("/", response_class=HTMLResponse)101async def get_home(request: Request):102    return templates.TemplateResponse("index.html", {"request": request})103 104 105@app.websocket("/ws")106async def websocket_endpoint(websocket: WebSocket):107    await websocket.accept()108    try:109        process_active = False  # Flag to track the state of the process110 111        while True:112            message = await websocket.receive_text()113 114            if message == "start" and not process_active:115                process_active = True116                await websocket.send_text("Listening for wake word...")117                wake_word_detected = await launch_fn(debug=True)118                if wake_word_detected:119                    await websocket.send_text("Wake word detected. Listening for your query...")120                    await listen(websocket) 121                    process_active = False  # Reset the process flag122 123            elif message == "stop":124                if process_active:125                    # Implement logic to stop the ongoing process126                    # This might involve setting a flag that your launch_fn and listen functions check127                    process_active = False128                    await websocket.send_text("Process stopped. Ready to restart.")129                    break  # Or keep the loop running if you want to allow restarting without reconnecting130 131    except WebSocketDisconnect:132        print("Client disconnected.")133