no online conflict when multiple users
This commit is contained in:
parent
b7a2d23a18
commit
aa0ba598f0
2 changed files with 16 additions and 11 deletions
|
|
@ -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():
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue