diff --git a/README.md b/README.md index 4c12eda..1f812bc 100644 --- a/README.md +++ b/README.md @@ -66,6 +66,18 @@ docker compose -f docker-compose.unraid.yml up -d app > Première utilisation : le paquet ghcr.io doit être **public** (GitHub → > page du dépôt → Packages → liveflow → Package settings → Change visibility). +## Diarisation (qui parle ?) + +Avec `DIARIZATION=on` (par défaut), chaque segment est attribué à un locuteur +(« Locuteur 1 », « Locuteur 2 »...) par empreinte vocale (ECAPA-TDNN sur CPU, +modèle ~80 Mo téléchargé au premier démarrage dans `data/models`). Les +locuteurs apparaissent dans l'interface et les exports. Réglages : +`DIARIZATION_MAX_SPEAKERS` (défaut 8) et `DIARIZATION_THRESHOLD` (défaut 0.70, +baisser si deux personnes sont confondues, monter si une personne est coupée +en deux). La case « Locuteurs » de l'interface permet de désactiver la +diarisation pour une réunion donnée. Pendant l'enregistrement, le bouton +**⏸ Pause** suspend la capture sans clore la réunion. + ## Authentification L'interface est protégée par un identifiant/mot de passe (**admin / admin** diff --git a/app/Dockerfile b/app/Dockerfile index 93393a4..8a1615c 100644 --- a/app/Dockerfile +++ b/app/Dockerfile @@ -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 diff --git a/app/diarizer.py b/app/diarizer.py new file mode 100644 index 0000000..6f6c266 --- /dev/null +++ b/app/diarizer.py @@ -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() diff --git a/app/main.py b/app/main.py index 8391d5a..5f756b5 100644 --- a/app/main.py +++ b/app/main.py @@ -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"{s['speaker']}\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)") diff --git a/app/requirements.txt b/app/requirements.txt index f018be3..5d58c24 100644 --- a/app/requirements.txt +++ b/app/requirements.txt @@ -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 diff --git a/app/static/app.js b/app/static/app.js index a71b407..95ea437 100644 --- a/app/static/app.js +++ b/app/static/app.js @@ -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 = `${fmtTs(seg.t0)}`; + let speakerHtml = ''; + if (seg.speaker) { + speakerHtml = ``; + } + div.innerHTML = `${fmtTs(seg.t0)}${speakerHtml}`; + 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; diff --git a/app/static/index.html b/app/static/index.html index c7462ee..2791675 100644 --- a/app/static/index.html +++ b/app/static/index.html @@ -27,8 +27,12 @@ + 00:00