no online conflict when multiple users

This commit is contained in:
Quentin Fuxa 2025-01-03 14:48:45 +01:00
parent b7a2d23a18
commit aa0ba598f0
2 changed files with 16 additions and 11 deletions

View file

@ -9,7 +9,7 @@ from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import HTMLResponse from fastapi.responses import HTMLResponse
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from whisper_online import asr_factory, add_shared_args from whisper_online import backend_factory, online_factory, add_shared_args
app = FastAPI() app = FastAPI()
app.add_middleware( app.add_middleware(
@ -40,7 +40,7 @@ parser.add_argument(
add_shared_args(parser) add_shared_args(parser)
args = parser.parse_args() args = parser.parse_args()
asr, online = asr_factory(args) asr, tokenizer = backend_factory(args)
# Load demo HTML for the root endpoint # Load demo HTML for the root endpoint
with open("src/live_transcription.html", "r") as f: with open("src/live_transcription.html", "r") as f:
@ -85,6 +85,9 @@ async def websocket_endpoint(websocket: WebSocket):
ffmpeg_process = await start_ffmpeg_decoder() ffmpeg_process = await start_ffmpeg_decoder()
pcm_buffer = bytearray() pcm_buffer = bytearray()
print("Loading online.")
online = online_factory(args, asr, tokenizer)
print("Online loaded.")
# Continuously read decoded PCM from ffmpeg stdout in a background task # Continuously read decoded PCM from ffmpeg stdout in a background task
async def ffmpeg_stdout_reader(): async def ffmpeg_stdout_reader():

View file

@ -920,11 +920,7 @@ def add_shared_args(parser):
default="DEBUG", default="DEBUG",
) )
def backend_factory(args):
def asr_factory(args, logfile=sys.stderr):
"""
Creates and configures an ASR and ASR Online instance based on the specified backend and arguments.
"""
backend = args.backend backend = args.backend
if backend == "openai-api": if backend == "openai-api":
logger.debug("Using OpenAI API.") logger.debug("Using OpenAI API.")
@ -967,10 +963,10 @@ def asr_factory(args, logfile=sys.stderr):
tokenizer = create_tokenizer(tgt_language) tokenizer = create_tokenizer(tgt_language)
else: else:
tokenizer = None tokenizer = None
return asr, tokenizer
# Create the OnlineASRProcessor def online_factory(args, asr, tokenizer, logfile=sys.stderr):
if args.vac: if args.vac:
online = VACOnlineASRProcessor( online = VACOnlineASRProcessor(
args.min_chunk_size, args.min_chunk_size,
asr, asr,
@ -985,10 +981,16 @@ def asr_factory(args, logfile=sys.stderr):
logfile=logfile, logfile=logfile,
buffer_trimming=(args.buffer_trimming, args.buffer_trimming_sec), buffer_trimming=(args.buffer_trimming, args.buffer_trimming_sec),
) )
return online
def asr_factory(args, logfile=sys.stderr):
"""
Creates and configures an ASR and ASR Online instance based on the specified backend and arguments.
"""
asr, tokenizer = backend_factory(args)
online = online_factory(args, asr, tokenizer, logfile=logfile)
return asr, online return asr, online
def set_logging(args, logger, other="_server"): def set_logging(args, logger, other="_server"):
logging.basicConfig(format="%(levelname)s\t%(message)s") # format='%(name)s logging.basicConfig(format="%(levelname)s\t%(message)s") # format='%(name)s
logger.setLevel(args.log_level) logger.setLevel(args.log_level)