Merge pull request #51 from QuentinFuxa/diart_integration_improvements
Diart integration improvements
This commit is contained in:
commit
450c93fef8
3 changed files with 121 additions and 63 deletions
|
|
@ -5,6 +5,11 @@ from rx.subject import Subject
|
||||||
import threading
|
import threading
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import re
|
||||||
|
|
||||||
|
def extract_number(s):
|
||||||
|
match = re.search(r'\d+', s)
|
||||||
|
return int(match.group()) if match else None
|
||||||
|
|
||||||
class WebSocketAudioSource(AudioSource):
|
class WebSocketAudioSource(AudioSource):
|
||||||
"""
|
"""
|
||||||
|
|
@ -44,37 +49,48 @@ def create_pipeline(SAMPLE_RATE):
|
||||||
return inference, ws_source
|
return inference, ws_source
|
||||||
|
|
||||||
|
|
||||||
def init_diart(SAMPLE_RATE):
|
def init_diart(SAMPLE_RATE, diar_instance):
|
||||||
inference, ws_source = create_pipeline(SAMPLE_RATE)
|
diar_pipeline = SpeakerDiarization()
|
||||||
|
ws_source = WebSocketAudioSource(uri="websocket_source", sample_rate=SAMPLE_RATE)
|
||||||
|
inference = StreamingInference(
|
||||||
|
pipeline=diar_pipeline,
|
||||||
|
source=ws_source,
|
||||||
|
do_plot=False,
|
||||||
|
show_progress=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
l_speakers_queue = asyncio.Queue()
|
||||||
|
|
||||||
def diar_hook(result):
|
def diar_hook(result):
|
||||||
"""
|
"""
|
||||||
Hook called each time Diart processes a chunk.
|
Hook called each time Diart processes a chunk.
|
||||||
result is (annotation, audio).
|
result is (annotation, audio).
|
||||||
We store the label of the last segment in 'current_speaker'.
|
For each detected speaker segment, push its info to the queue and update processed_time.
|
||||||
"""
|
"""
|
||||||
global l_speakers
|
|
||||||
l_speakers = []
|
|
||||||
annotation, audio = result
|
annotation, audio = result
|
||||||
|
if annotation._labels:
|
||||||
for speaker in annotation._labels:
|
for speaker in annotation._labels:
|
||||||
segments_beg = annotation._labels[speaker].segments_boundaries_[0]
|
segments_beg = annotation._labels[speaker].segments_boundaries_[0]
|
||||||
segments_end = annotation._labels[speaker].segments_boundaries_[-1]
|
segments_end = annotation._labels[speaker].segments_boundaries_[-1]
|
||||||
|
if segments_end > diar_instance.processed_time:
|
||||||
|
diar_instance.processed_time = segments_end
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
l_speakers_queue.put({"speaker": speaker, "beg": segments_beg, "end": segments_end})
|
l_speakers_queue.put({"speaker": speaker, "beg": segments_beg, "end": segments_end})
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
audio_duration = audio.extent.end
|
||||||
|
if audio_duration > diar_instance.processed_time:
|
||||||
|
diar_instance.processed_time = audio_duration
|
||||||
|
|
||||||
l_speakers_queue = asyncio.Queue()
|
|
||||||
inference.attach_hooks(diar_hook)
|
inference.attach_hooks(diar_hook)
|
||||||
|
|
||||||
# Launch Diart in a background thread
|
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
diar_future = loop.run_in_executor(None, inference)
|
diar_future = loop.run_in_executor(None, inference)
|
||||||
return inference, l_speakers_queue, ws_source
|
return inference, l_speakers_queue, ws_source
|
||||||
|
|
||||||
|
class DiartDiarization:
|
||||||
class DiartDiarization():
|
|
||||||
def __init__(self, SAMPLE_RATE):
|
def __init__(self, SAMPLE_RATE):
|
||||||
self.inference, self.l_speakers_queue, self.ws_source = init_diart(SAMPLE_RATE)
|
self.processed_time = 0
|
||||||
|
self.inference, self.l_speakers_queue, self.ws_source = init_diart(SAMPLE_RATE, self)
|
||||||
self.segment_speakers = []
|
self.segment_speakers = []
|
||||||
|
|
||||||
async def diarize(self, pcm_array):
|
async def diarize(self, pcm_array):
|
||||||
|
|
@ -86,16 +102,17 @@ class DiartDiarization():
|
||||||
def close(self):
|
def close(self):
|
||||||
self.ws_source.close()
|
self.ws_source.close()
|
||||||
|
|
||||||
|
|
||||||
def assign_speakers_to_chunks(self, chunks):
|
def assign_speakers_to_chunks(self, chunks):
|
||||||
"""
|
"""
|
||||||
Go through each chunk and see which speaker(s) overlap
|
For each chunk (a dict with keys "beg" and "end"), assign a speaker label.
|
||||||
that chunk's time range in the Diart annotation.
|
|
||||||
Then store the speaker label(s) (or choose the most overlapping).
|
- If a chunk overlaps with a detected speaker segment, assign that label.
|
||||||
This modifies `chunks` in-place or returns a new list with assigned speakers.
|
- If the chunk's end time is within the processed time and no speaker was assigned,
|
||||||
|
mark it as "No speaker".
|
||||||
|
- If the chunk's time hasn't been fully processed yet, leave it (or mark as "Processing").
|
||||||
"""
|
"""
|
||||||
if not self.segment_speakers:
|
for ch in chunks:
|
||||||
return chunks
|
ch["speaker"] = ch.get("speaker", -1)
|
||||||
|
|
||||||
for segment in self.segment_speakers:
|
for segment in self.segment_speakers:
|
||||||
seg_beg = segment["beg"]
|
seg_beg = segment["beg"]
|
||||||
|
|
@ -104,7 +121,10 @@ class DiartDiarization():
|
||||||
for ch in chunks:
|
for ch in chunks:
|
||||||
if seg_end <= ch["beg"] or seg_beg >= ch["end"]:
|
if seg_end <= ch["beg"] or seg_beg >= ch["end"]:
|
||||||
continue
|
continue
|
||||||
# We have overlap. Let's just pick the speaker (could be more precise in a more complex implementation)
|
ch["speaker"] = extract_number(speaker) + 1
|
||||||
ch["speaker"] = speaker
|
if self.processed_time > 0:
|
||||||
|
for ch in chunks:
|
||||||
|
if ch["end"] <= self.processed_time and ch["speaker"] == -1:
|
||||||
|
ch["speaker"] = -2
|
||||||
|
|
||||||
return chunks
|
return chunks
|
||||||
|
|
@ -179,8 +179,9 @@
|
||||||
The server might send:
|
The server might send:
|
||||||
{
|
{
|
||||||
"lines": [
|
"lines": [
|
||||||
{"speaker": 0, "text": "Hello."},
|
{"speaker": 0, "text": "Hello.", "beg": "00:00", "end": "00:01"},
|
||||||
{"speaker": 1, "text": "Bonjour."},
|
{"speaker": -2, "text": "Hi, no speaker here.", "beg": "00:01", "end": "00:02"},
|
||||||
|
{"speaker": -1, "text": "...", "beg": "00:02", "end": "00:03" },
|
||||||
...
|
...
|
||||||
],
|
],
|
||||||
"buffer": "..."
|
"buffer": "..."
|
||||||
|
|
@ -198,14 +199,27 @@
|
||||||
linesTranscriptDiv.innerHTML = "";
|
linesTranscriptDiv.innerHTML = "";
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
// Build the HTML
|
|
||||||
// The buffer is appended to the last line if it's non-empty
|
|
||||||
const linesHtml = lines.map((item, idx) => {
|
const linesHtml = lines.map((item, idx) => {
|
||||||
|
let speakerLabel = "";
|
||||||
|
if (item.speaker === -2) {
|
||||||
|
speakerLabel = "No speaker";
|
||||||
|
} else if (item.speaker !== -1) {
|
||||||
|
speakerLabel = `Speaker ${item.speaker}`;
|
||||||
|
}
|
||||||
|
|
||||||
|
let timeInfo = "";
|
||||||
|
if (item.beg !== undefined && item.end !== undefined) {
|
||||||
|
timeInfo = ` [${item.beg}, ${item.end}]`;
|
||||||
|
}
|
||||||
|
|
||||||
let textContent = item.text;
|
let textContent = item.text;
|
||||||
if (idx === lines.length - 1 && buffer) {
|
if (idx === lines.length - 1 && buffer) {
|
||||||
textContent += `<span class="buffer">${buffer}</span>`;
|
textContent += `<span class="buffer">${buffer}</span>`;
|
||||||
}
|
}
|
||||||
return `<p><strong>Speaker ${item.speaker}:</strong> ${textContent}</p>`;
|
|
||||||
|
return speakerLabel
|
||||||
|
? `<p><strong>${speakerLabel}${timeInfo}</strong> ${textContent}</p>`
|
||||||
|
: `<p>${textContent}</p>`;
|
||||||
}).join("");
|
}).join("");
|
||||||
|
|
||||||
linesTranscriptDiv.innerHTML = linesHtml;
|
linesTranscriptDiv.innerHTML = linesHtml;
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@ import argparse
|
||||||
import asyncio
|
import asyncio
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import ffmpeg
|
import ffmpeg
|
||||||
from time import time
|
from time import time, sleep
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||||
|
|
@ -12,9 +12,12 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
|
||||||
from src.whisper_streaming.whisper_online import backend_factory, online_factory, add_shared_args
|
from src.whisper_streaming.whisper_online import backend_factory, online_factory, add_shared_args
|
||||||
|
|
||||||
import subprocess
|
|
||||||
import math
|
import math
|
||||||
import logging
|
import logging
|
||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
def format_time(seconds):
|
||||||
|
return str(timedelta(seconds=int(seconds)))
|
||||||
|
|
||||||
|
|
||||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||||
|
|
@ -48,6 +51,12 @@ parser.add_argument(
|
||||||
help="Whether to enable speaker diarization.",
|
help="Whether to enable speaker diarization.",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--transcription",
|
||||||
|
type=bool,
|
||||||
|
default=True,
|
||||||
|
help="To disable to only see live diarization results.",
|
||||||
|
)
|
||||||
|
|
||||||
add_shared_args(parser)
|
add_shared_args(parser)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
@ -68,7 +77,10 @@ if args.diarization:
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI):
|
async def lifespan(app: FastAPI):
|
||||||
global asr, tokenizer
|
global asr, tokenizer
|
||||||
|
if args.transcription:
|
||||||
asr, tokenizer = backend_factory(args)
|
asr, tokenizer = backend_factory(args)
|
||||||
|
else:
|
||||||
|
asr, tokenizer = None, None
|
||||||
yield
|
yield
|
||||||
|
|
||||||
app = FastAPI(lifespan=lifespan)
|
app = FastAPI(lifespan=lifespan)
|
||||||
|
|
@ -117,7 +129,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||||
|
|
||||||
ffmpeg_process = None
|
ffmpeg_process = None
|
||||||
pcm_buffer = bytearray()
|
pcm_buffer = bytearray()
|
||||||
online = online_factory(args, asr, tokenizer)
|
online = online_factory(args, asr, tokenizer) if args.transcription else None
|
||||||
diarization = DiartDiarization(SAMPLE_RATE) if args.diarization else None
|
diarization = DiartDiarization(SAMPLE_RATE) if args.diarization else None
|
||||||
|
|
||||||
async def restart_ffmpeg():
|
async def restart_ffmpeg():
|
||||||
|
|
@ -130,7 +142,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||||
logger.warning(f"Error killing FFmpeg process: {e}")
|
logger.warning(f"Error killing FFmpeg process: {e}")
|
||||||
ffmpeg_process = await start_ffmpeg_decoder()
|
ffmpeg_process = await start_ffmpeg_decoder()
|
||||||
pcm_buffer = bytearray()
|
pcm_buffer = bytearray()
|
||||||
online = online_factory(args, asr, tokenizer)
|
online = online_factory(args, asr, tokenizer) if args.transcription else None
|
||||||
if args.diarization:
|
if args.diarization:
|
||||||
diarization = DiartDiarization(SAMPLE_RATE)
|
diarization = DiartDiarization(SAMPLE_RATE)
|
||||||
logger.info("FFmpeg process started.")
|
logger.info("FFmpeg process started.")
|
||||||
|
|
@ -142,7 +154,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
full_transcription = ""
|
full_transcription = ""
|
||||||
beg = time()
|
beg = time()
|
||||||
|
beg_loop = time()
|
||||||
chunk_history = [] # Will store dicts: {beg, end, text, speaker}
|
chunk_history = [] # Will store dicts: {beg, end, text, speaker}
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
|
|
@ -184,45 +196,57 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||||
/ 32768.0
|
/ 32768.0
|
||||||
)
|
)
|
||||||
pcm_buffer = pcm_buffer[MAX_BYTES_PER_SEC:]
|
pcm_buffer = pcm_buffer[MAX_BYTES_PER_SEC:]
|
||||||
|
|
||||||
|
if args.transcription:
|
||||||
logger.info(f"{len(online.audio_buffer) / online.SAMPLING_RATE} seconds of audio will be processed by the model.")
|
logger.info(f"{len(online.audio_buffer) / online.SAMPLING_RATE} seconds of audio will be processed by the model.")
|
||||||
online.insert_audio_chunk(pcm_array)
|
online.insert_audio_chunk(pcm_array)
|
||||||
transcription = online.process_iter()
|
transcription = online.process_iter()
|
||||||
|
if transcription.start:
|
||||||
if transcription:
|
|
||||||
chunk_history.append({
|
chunk_history.append({
|
||||||
"beg": transcription.start,
|
"beg": transcription.start,
|
||||||
"end": transcription.end,
|
"end": transcription.end,
|
||||||
"text": transcription.text,
|
"text": transcription.text,
|
||||||
"speaker": "0"
|
|
||||||
})
|
})
|
||||||
|
|
||||||
full_transcription += transcription.text if transcription else ""
|
full_transcription += transcription.text if transcription else ""
|
||||||
buffer = online.get_buffer()
|
buffer = online.get_buffer()
|
||||||
|
|
||||||
if buffer in full_transcription: # With VAC, the buffer is not updated until the next chunk is processed
|
if buffer in full_transcription: # With VAC, the buffer is not updated until the next chunk is processed
|
||||||
buffer = ""
|
buffer = ""
|
||||||
|
else:
|
||||||
lines = [
|
chunk_history.append({
|
||||||
{
|
"beg": time() - beg_loop,
|
||||||
"speaker": "0",
|
"end": time() - beg_loop + 0.1,
|
||||||
"text": "",
|
"text": '',
|
||||||
}
|
})
|
||||||
]
|
sleep(0.1)
|
||||||
|
buffer = ''
|
||||||
|
|
||||||
if args.diarization:
|
if args.diarization:
|
||||||
await diarization.diarize(pcm_array)
|
await diarization.diarize(pcm_array)
|
||||||
diarization.assign_speakers_to_chunks(chunk_history)
|
diarization.assign_speakers_to_chunks(chunk_history)
|
||||||
|
|
||||||
|
|
||||||
|
current_speaker = -1
|
||||||
|
lines = [{
|
||||||
|
"beg": 0,
|
||||||
|
"end": 0,
|
||||||
|
"speaker": current_speaker,
|
||||||
|
"text": ""
|
||||||
|
}]
|
||||||
for ch in chunk_history:
|
for ch in chunk_history:
|
||||||
if args.diarization and ch["speaker"] and ch["speaker"][-1] != lines[-1]["speaker"]:
|
if args.diarization and ch["speaker"] and ch["speaker"] != current_speaker:
|
||||||
|
new_speaker = ch["speaker"]
|
||||||
lines.append(
|
lines.append(
|
||||||
{
|
{
|
||||||
"speaker": ch["speaker"][-1],
|
"speaker": new_speaker,
|
||||||
"text": ch['text']
|
"text": ch['text'],
|
||||||
|
"beg": format_time(ch['beg']),
|
||||||
|
"end": format_time(ch['end']),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
current_speaker = new_speaker
|
||||||
else:
|
else:
|
||||||
lines[-1]["text"] += ch['text']
|
lines[-1]["text"] += ch['text']
|
||||||
|
lines[-1]["end"] = format_time(ch['end'])
|
||||||
|
|
||||||
response = {"lines": lines, "buffer": buffer}
|
response = {"lines": lines, "buffer": buffer}
|
||||||
await websocket.send_json(response)
|
await websocket.send_json(response)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue