Teapack1/Assistant-Audio-Intent-Classification
6
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 