Diarisation des locuteurs et pause/reprise (intégration des travaux de main)

- diarizer.py (ECAPA-TDNN + clustering incrémental) repris tel quel de main
- Hooks dans le pipeline : colonne speaker, identification par segment dans
  un thread, locuteur dans l'interface (étiquettes colorées) et les exports
- Pause/reprise : bouton ⏸, le client coupe l'envoi, le serveur clôt le
  segment en cours
- Case « Locuteurs » pour désactiver la diarisation par réunion
- DIARIZATION=on par défaut dans les compose, PyTorch CPU dans l'image

https://claude.ai/code/session_01YHMp3EKzr4s6o8w1ygxuUe
This commit is contained in:
Claude committed 2026-06-12 17:50:51 +00:00
1 parent 21866ec334
commit db73e8db02
10 files changed
+352 -20

No files matched your search

+3 -1
View File
@@ -7,7 +7,9 @@ RUN apt-get update \
WORKDIR /srv
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
# PyTorch CPU-only pour la diarisation (~200 Mo au lieu de ~2 Go avec CUDA)
RUN pip install --no-cache-dir torch torchaudio --index-url https://download.pytorch.org/whl/cpu \
&& pip install --no-cache-dir -r requirements.txt
COPY . .
RUN chmod +x entrypoint.sh
+164
View File
@@ -0,0 +1,164 @@
"""Identification des locuteurs par embeddings ECAPA-TDNN + clustering incrémental.
Chaque segment de parole (PCM 16 kHz mono) est projeté dans un espace de
représentation de dimension 192 via ECAPA-TDNN (SpeechBrain). On maintient
un profil par locuteur (moyenne mobile exponentielle des embeddings) et on
attribue chaque nouveau segment au locuteur le plus proche par similarité
cosinus, ou on crée un nouveau locuteur si le score est sous le seuil.
"""
from __future__ import annotations
import struct
from dataclasses import dataclass, field
import numpy as np
# Imports lourds (torch, speechbrain) chargés paresseusement dans load_model()
# pour ne pas pénaliser le démarrage quand la diarisation est désactivée.
@dataclass
class _SpeakerProfile:
"""Profil incrémental d'un locuteur au sein d'une session."""
label: str
embedding: np.ndarray # centroïde courant (moyenne mobile)
count: int = 0 # nombre d'observations
_ema_alpha: float = field(default=0.3, repr=False)
def update(self, new_emb: np.ndarray) -> None:
"""Met à jour le centroïde via une moyenne mobile exponentielle."""
self.count += 1
if self.count == 1:
self.embedding = new_emb.copy()
else:
self.embedding = (
self._ema_alpha * new_emb
+ (1 - self._ema_alpha) * self.embedding
)
# Renormaliser pour que la similarité cosinus reste cohérente
norm = np.linalg.norm(self.embedding)
if norm > 0:
self.embedding /= norm
def _cosine_similarity(a: np.ndarray, b: np.ndarray) -> float:
"""Similarité cosinus entre deux vecteurs unitaires."""
return float(np.dot(a, b))
class EmbeddingModel:
"""Encapsule le modèle SpeechBrain ECAPA-TDNN (singleton partagé).
Chargé une seule fois au démarrage de l'application, puis réutilisé
par chaque session WebSocket via des instances de SpeakerDiarizer.
"""
def __init__(self) -> None:
import torch # noqa: F811 — import local volontaire
from speechbrain.inference.speaker import EncoderClassifier
self._device = "cpu"
self._model = EncoderClassifier.from_hparams(
source="speechbrain/spkrec-ecapa-voxceleb",
savedir="/data/models/ecapa-tdnn",
run_opts={"device": self._device},
)
self._torch = torch
def extract(self, pcm: bytes, sample_rate: int = 16000) -> np.ndarray:
"""Extrait un embedding 192-d à partir d'un segment PCM 16-bit mono.
Returns:
np.ndarray de forme (192,), normalisé L2.
Raises:
ValueError: si le segment audio est trop court (< 200 ms).
"""
n_samples = len(pcm) // 2
min_samples = sample_rate // 5 # 200 ms minimum
if n_samples < min_samples:
raise ValueError(
f"Segment trop court ({n_samples} samples, min {min_samples})"
)
# PCM 16-bit little-endian → float32 [-1, 1]
samples = struct.unpack(f"<{n_samples}h", pcm)
waveform = self._torch.tensor(samples, dtype=self._torch.float32) / 32768.0
waveform = waveform.unsqueeze(0) # (1, T)
with self._torch.no_grad():
embedding = self._model.encode_batch(waveform)
emb = embedding.squeeze().cpu().numpy() # (192,)
# Normaliser L2
norm = np.linalg.norm(emb)
if norm > 0:
emb /= norm
return emb
class SpeakerDiarizer:
"""Identifie le locuteur de chaque segment audio au sein d'une session.
Chaque instance correspond à une réunion/session et maintient ses propres
profils de locuteurs. Le modèle d'embedding est partagé (EmbeddingModel).
"""
def __init__(
self,
model: EmbeddingModel,
threshold: float = 0.70,
max_speakers: int = 8,
) -> None:
self._model = model
self._threshold = threshold
self._max_speakers = max_speakers
self._profiles: list[_SpeakerProfile] = []
def identify(self, pcm: bytes, sample_rate: int = 16000) -> str:
"""Identifie le locuteur d'un segment PCM.
Returns:
Label du locuteur ("Locuteur 1", "Locuteur 2", etc.)
ou chaîne vide si le segment est trop court pour être analysé.
"""
try:
embedding = self._model.extract(pcm, sample_rate)
except ValueError:
# Segment trop court — on ne peut pas identifier le locuteur
return ""
# Comparer avec les profils existants
best_score = -1.0
best_profile: _SpeakerProfile | None = None
for profile in self._profiles:
score = _cosine_similarity(embedding, profile.embedding)
if score > best_score:
best_score = score
best_profile = profile
if best_profile is not None and best_score >= self._threshold:
best_profile.update(embedding)
return best_profile.label
# Nouveau locuteur (sauf si on a atteint le max)
if len(self._profiles) >= self._max_speakers:
# Forcer l'attribution au profil le plus proche
if best_profile is not None:
best_profile.update(embedding)
return best_profile.label
# Cas dégénéré : aucun profil et max atteint (ne devrait pas arriver)
return "Locuteur 1"
new_label = f"Locuteur {len(self._profiles) + 1}"
new_profile = _SpeakerProfile(label=new_label, embedding=embedding)
new_profile.update(embedding)
self._profiles.append(new_profile)
return new_label
def reset(self) -> None:
"""Réinitialise les profils pour une nouvelle session."""
self._profiles.clear()
+69 -14
View File
@@ -42,6 +42,11 @@ def clean_transcript(text: str) -> str:
text = text.rstrip(" 。,、!?") # ponctuation chinoise en queue
return text.strip()
# Diarisation (identification des locuteurs) : embeddings ECAPA-TDNN sur CPU
DIARIZATION = os.environ.get("DIARIZATION", "off").strip().lower() == "on"
DIARIZATION_THRESHOLD = float(os.environ.get("DIARIZATION_THRESHOLD", "0.70"))
DIARIZATION_MAX_SPEAKERS = int(os.environ.get("DIARIZATION_MAX_SPEAKERS", "8"))
LIVEFLOW_USER = os.environ.get("LIVEFLOW_USER", "admin")
LIVEFLOW_PASSWORD = os.environ.get("LIVEFLOW_PASSWORD", "admin")
SESSION_TTL = 7 * 24 * 3600 # 7 jours
@@ -50,6 +55,7 @@ SESSION_COOKIE = "liveflow_session"
db: aiosqlite.Connection | None = None
http: httpx.AsyncClient | None = None
session_secret: bytes = b""
embedding_model = None # EmbeddingModel, chargé au démarrage si DIARIZATION=on
def load_session_secret() -> bytes:
@@ -99,12 +105,25 @@ async def lifespan(app: FastAPI):
meeting_id INTEGER NOT NULL REFERENCES meetings(id) ON DELETE CASCADE,
t0 REAL NOT NULL,
t1 REAL NOT NULL,
text TEXT NOT NULL
text TEXT NOT NULL,
speaker TEXT NOT NULL DEFAULT ''
);
"""
)
# Migration : ajouter la colonne speaker aux bases existantes
try:
await db.execute("ALTER TABLE segments ADD COLUMN speaker TEXT NOT NULL DEFAULT ''")
except Exception:
pass # la colonne existe déjà
await db.execute("PRAGMA foreign_keys = ON")
await db.commit()
global embedding_model
if DIARIZATION and embedding_model is None:
print("Chargement du modèle de diarisation ECAPA-TDNN...", flush=True)
from diarizer import EmbeddingModel
embedding_model = await asyncio.to_thread(EmbeddingModel)
print("Modèle de diarisation prêt.", flush=True)
http = httpx.AsyncClient(timeout=120)
yield
await http.aclose()
@@ -230,6 +249,8 @@ async def ws_transcribe(ws: WebSocket):
return
title = (start.get("title") or "").strip() or datetime.now().strftime("Réunion du %d/%m/%Y %H:%M")
# Diarisation active si le serveur la supporte ET que le client la demande
session_diarization = DIARIZATION and start.get("diarization", True)
cur = await db.execute(
"INSERT INTO meetings (title, created_at) VALUES (?, ?)",
(title, datetime.now(timezone.utc).isoformat()),
@@ -242,7 +263,17 @@ async def ws_transcribe(ws: WebSocket):
queue: asyncio.Queue[Segment | None] = asyncio.Queue()
async def worker():
"""Transcrit les segments dans l'ordre et pousse le texte au client."""
"""Identifie le locuteur puis transcrit les segments, dans l'ordre."""
# Chaque session a son propre diarizer (profils de locuteurs isolés)
diarizer = None
if session_diarization and embedding_model is not None:
from diarizer import SpeakerDiarizer
diarizer = SpeakerDiarizer(
model=embedding_model,
threshold=DIARIZATION_THRESHOLD,
max_speakers=DIARIZATION_MAX_SPEAKERS,
)
while True:
seg = await queue.get()
if seg is None:
@@ -251,6 +282,15 @@ async def ws_transcribe(ws: WebSocket):
print(f"[réunion {meeting_id}] segment {seg.t0:.1f}s → {seg.t1:.1f}s "
f"({len(seg.pcm) / 2 / SAMPLE_RATE:.1f}s d'audio, niveau crête {peak}"
f"{', amplifié' if pcm is not seg.pcm else ''}), transcription...", flush=True)
# Diarisation (~30-80 ms CPU, dans un thread pour ne pas bloquer)
speaker = ""
if diarizer is not None:
try:
speaker = await asyncio.to_thread(diarizer.identify, pcm)
except Exception as exc:
print(f"[réunion {meeting_id}] diarisation échouée : {exc}", flush=True)
try:
text = await transcribe(pcm)
except Exception as exc:
@@ -261,15 +301,17 @@ async def ws_transcribe(ws: WebSocket):
if cleaned != text:
print(f"[réunion {meeting_id}] filtré (CJK) : {text[:60]!r} -> {cleaned[:60]!r}", flush=True)
text = cleaned
print(f"[réunion {meeting_id}] texte : {text[:80]!r}", flush=True)
print(f"[réunion {meeting_id}] {speaker or 'texte'} : {text[:80]!r}", flush=True)
if not text:
continue
await db.execute(
"INSERT INTO segments (meeting_id, t0, t1, text) VALUES (?, ?, ?, ?)",
(meeting_id, seg.t0, seg.t1, text),
"INSERT INTO segments (meeting_id, t0, t1, text, speaker) VALUES (?, ?, ?, ?, ?)",
(meeting_id, seg.t0, seg.t1, text, speaker),
)
await db.commit()
await ws.send_json({"type": "segment", "t0": seg.t0, "t1": seg.t1, "text": text})
await ws.send_json(
{"type": "segment", "t0": seg.t0, "t1": seg.t1, "text": text, "speaker": speaker}
)
worker_task = asyncio.create_task(worker())
try:
@@ -282,7 +324,8 @@ async def ws_transcribe(ws: WebSocket):
queue.put_nowait(seg)
elif msg.get("text"):
control = json.loads(msg["text"])
if control.get("type") == "stop":
cmd = control.get("type")
if cmd == "stop":
if (last := segmenter.flush()) is not None:
queue.put_nowait(last)
queue.put_nowait(None)
@@ -290,6 +333,10 @@ async def ws_transcribe(ws: WebSocket):
worker_task = None
await ws.send_json({"type": "done", "meeting_id": meeting_id})
break
elif cmd == "pause":
# clôt le segment en cours ; le client cesse d'envoyer l'audio
if (last := segmenter.flush()) is not None:
queue.put_nowait(last)
except WebSocketDisconnect:
pass
finally:
@@ -331,7 +378,8 @@ async def get_meeting_or_404(meeting_id: int) -> dict:
async def get_meeting(meeting_id: int):
meeting = await get_meeting_or_404(meeting_id)
rows = await db.execute_fetchall(
"SELECT t0, t1, text FROM segments WHERE meeting_id = ? ORDER BY id", (meeting_id,)
"SELECT t0, t1, text, speaker FROM segments WHERE meeting_id = ? ORDER BY id",
(meeting_id,),
)
meeting["segments"] = [dict(r) for r in rows]
return meeting
@@ -368,16 +416,23 @@ async def export_meeting(meeting_id: int, format: str = "txt"):
body, mime, ext = json.dumps(meeting, ensure_ascii=False, indent=2), "application/json", "json"
elif format == "md":
lines = [f"# {title}", ""]
lines += [f"**[{fmt_ts(s['t0'])}]** {s['text']}" for s in segs]
for s in segs:
prefix = f"**{s['speaker']} —** " if s.get("speaker") else ""
lines.append(f"**[{fmt_ts(s['t0'])}]** {prefix}{s['text']}")
body, mime, ext = "\n\n".join(lines) + "\n", "text/markdown", "md"
elif format == "srt":
blocks = [
f"{i}\n{fmt_ts(s['t0'], srt=True)} --> {fmt_ts(s['t1'], srt=True)}\n{s['text']}"
for i, s in enumerate(segs, 1)
]
blocks = []
for i, s in enumerate(segs, 1):
speaker_line = f"<i>{s['speaker']}</i>\n" if s.get("speaker") else ""
blocks.append(
f"{i}\n{fmt_ts(s['t0'], srt=True)} --> {fmt_ts(s['t1'], srt=True)}\n"
f"{speaker_line}{s['text']}"
)
body, mime, ext = "\n\n".join(blocks) + "\n", "application/x-subrip", "srt"
elif format == "txt":
body, mime, ext = "\n".join(s["text"] for s in segs) + "\n", "text/plain", "txt"
def txt_line(s: dict) -> str:
return f"[{s['speaker']}] {s['text']}" if s.get("speaker") else s["text"]
body, mime, ext = "\n".join(txt_line(s) for s in segs) + "\n", "text/plain", "txt"
else:
raise HTTPException(400, "Format inconnu (txt, md, srt, json)")
+3
View File
@@ -3,3 +3,6 @@ uvicorn[standard]~=0.34
httpx~=0.28
aiosqlite~=0.21
webrtcvad-wheels~=2.0
speechbrain>=1.0
torch>=2.0,<3.0
torchaudio>=2.0,<3.0
+58 -5
View File
@@ -16,8 +16,18 @@ const state = {
sourceNode: null,
autoMicQueue: null,
switching: false,
paused: false,
};
// couleurs des étiquettes de locuteurs (classes .speaker-0 à .speaker-7)
const speakerColors = {};
function speakerColorIndex(speaker) {
if (!(speaker in speakerColors)) {
speakerColors[speaker] = Object.keys(speakerColors).length % 8;
}
return speakerColors[speaker];
}
const BATCH_SAMPLES = 4096; // ~256 ms de PCM 16 kHz par message WebSocket
// fetch avec redirection vers la page de connexion si la session a expiré
@@ -124,7 +134,11 @@ async function startRecording() {
const proto = location.protocol === 'https:' ? 'wss' : 'ws';
state.ws = new WebSocket(`${proto}://${location.host}/ws`);
state.ws.onopen = () => state.ws.send(JSON.stringify({ type: 'start', title: $('title').value }));
state.ws.onopen = () => state.ws.send(JSON.stringify({
type: 'start',
title: $('title').value,
diarization: $('diarization-cb').checked,
}));
state.ws.onmessage = onServerMessage;
state.ws.onclose = (e) => {
if (e.code === 4401) { location.href = '/login'; return; }
@@ -142,6 +156,7 @@ async function startRecording() {
listMics();
state.recording = true;
state.paused = false;
state.startedAt = Date.now();
state.lastLoudAt = Date.now();
state.silentWarned = false;
@@ -149,12 +164,18 @@ async function startRecording() {
state.timerInterval = setInterval(updateTimer, 500);
$('record-btn').textContent = '■ Arrêter';
$('record-btn').classList.add('recording');
$('pause-btn').classList.remove('hidden', 'paused');
$('pause-btn').textContent = '⏸ Pause';
$('title').disabled = true;
setStatus('rec', 'Enregistrement…');
clearTranscript();
}
function queuePcm(samples) {
if (state.paused) {
state.lastLoudAt = Date.now(); // pas de chasse au micro pendant la pause
return;
}
updateVu(samples);
if (!state.recording || !state.ws || state.ws.readyState !== WebSocket.OPEN) return;
state.sendBuffer.push(samples);
@@ -200,10 +221,32 @@ function stopRecording(abrupt = false) {
$('record-btn').textContent = '● Démarrer';
$('record-btn').classList.remove('recording');
$('pause-btn').classList.add('hidden');
state.paused = false;
$('title').disabled = false;
$('vu-bar').style.width = '0%';
}
function togglePause() {
if (!state.recording || !state.ws || state.ws.readyState !== WebSocket.OPEN) return;
state.paused = !state.paused;
const btn = $('pause-btn');
if (state.paused) {
flushPcm();
state.ws.send(JSON.stringify({ type: 'pause' }));
btn.textContent = '▶ Reprendre';
btn.classList.add('paused');
setStatus('busy', 'En pause');
$('vu-bar').style.width = '0%';
} else {
state.ws.send(JSON.stringify({ type: 'resume' }));
state.lastLoudAt = Date.now();
btn.textContent = '⏸ Pause';
btn.classList.remove('paused');
setStatus('rec', 'Enregistrement…');
}
}
function onServerMessage(event) {
const msg = JSON.parse(event.data);
if (msg.type === 'ready') {
@@ -238,7 +281,7 @@ function updateTimer() {
String(Math.floor(s / 60)).padStart(2, '0') + ':' + String(s % 60).padStart(2, '0');
// micro muet depuis 5 s : on essaie automatiquement les autres micros
if (!state.recording) return;
if (!state.recording || state.paused) return;
const silent = Date.now() - state.lastLoudAt > 5000;
if (silent) {
state.silentWarned = true;
@@ -270,7 +313,12 @@ function clearTranscript() {
function appendSegment(seg) {
const div = document.createElement('div');
div.className = 'segment';
div.innerHTML = `<span class="ts">${fmtTs(seg.t0)}</span><span class="text"></span>`;
let speakerHtml = '';
if (seg.speaker) {
speakerHtml = `<span class="speaker speaker-${speakerColorIndex(seg.speaker)}"></span>`;
}
div.innerHTML = `<span class="ts">${fmtTs(seg.t0)}</span>${speakerHtml}<span class="text"></span>`;
if (seg.speaker) div.querySelector('.speaker').textContent = seg.speaker;
div.querySelector('.text').textContent = seg.text;
$('transcript').appendChild(div);
$('transcript').scrollTop = $('transcript').scrollHeight;
@@ -328,8 +376,12 @@ async function deleteCurrentMeeting() {
}
async function copyTranscript() {
const text = [...document.querySelectorAll('#transcript .segment .text')]
.map((el) => el.textContent).join('\n');
const text = [...document.querySelectorAll('#transcript .segment')]
.map((el) => {
const speaker = el.querySelector('.speaker');
const t = el.querySelector('.text').textContent;
return speaker ? `[${speaker.textContent}] ${t}` : t;
}).join('\n');
await navigator.clipboard.writeText(text);
$('copy-btn').textContent = '✓ Copié';
setTimeout(() => ($('copy-btn').textContent = '📋 Copier'), 1500);
@@ -338,6 +390,7 @@ async function copyTranscript() {
// --------------------------------------------------------------------- init
$('record-btn').onclick = () => (state.recording ? stopRecording() : startRecording());
$('pause-btn').onclick = togglePause;
$('logout-btn').onclick = async () => { await fetch('/api/logout', { method: 'POST' }); location.href = '/login'; };
$('mic-select').onchange = async () => {
const id = $('mic-select').value;
+4
View File
@@ -27,8 +27,12 @@
<option value="">Micro par défaut</option>
</select>
<button id="record-btn" class="record">● Démarrer</button>
<button id="pause-btn" class="pause hidden">⏸ Pause</button>
<span id="timer">00:00</span>
<div id="vu" title="Niveau du micro"><div id="vu-bar"></div></div>
<label id="diarization-toggle" class="toggle" title="Identifier qui parle (Locuteur 1, 2...)">
<input type="checkbox" id="diarization-cb" checked> Locuteurs
</label>
</section>
<section id="transcript-panel">
+32
View File
@@ -109,6 +109,38 @@ button, .controls a {
}
button.record { background: var(--accent); font-weight: 600; min-width: 130px; }
button.record.recording { background: var(--rec); }
button.pause { background: var(--panel-2); font-weight: 600; }
button.pause.paused { background: #2a2410; color: #ffd166; }
button.pause.hidden { display: none; }
.toggle {
display: flex;
align-items: center;
gap: 6px;
font-size: 0.82rem;
color: var(--muted);
cursor: pointer;
white-space: nowrap;
}
.toggle input { accent-color: var(--accent); }
.speaker {
font-size: 0.72rem;
font-weight: 600;
padding: 2px 8px;
border-radius: 999px;
align-self: flex-start;
margin-top: 2px;
flex-shrink: 0;
}
.speaker-0 { background: #1d3a5f; color: #7cb8ff; }
.speaker-1 { background: #1f4a2c; color: #6fdc8c; }
.speaker-2 { background: #4a2a1f; color: #ffab70; }
.speaker-3 { background: #3d1f4a; color: #d8a6ff; }
.speaker-4 { background: #4a1f2e; color: #ff9bb3; }
.speaker-5 { background: #1f4a47; color: #6fd9d2; }
.speaker-6 { background: #4a431f; color: #e3d56f; }
.speaker-7 { background: #2e2e4a; color: #a3a8ff; }
#timer { font-variant-numeric: tabular-nums; color: var(--muted); min-width: 48px; }
#vu {