diff --git a/app/diarizer.py b/app/diarizer.py index 6f6c266..45da4e4 100644 --- a/app/diarizer.py +++ b/app/diarizer.py @@ -25,10 +25,14 @@ class _SpeakerProfile: 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) + _ema_alpha: float = field(default=0.15, repr=False) def update(self, new_emb: np.ndarray) -> None: - """Met à jour le centroïde via une moyenne mobile exponentielle.""" + """Met à jour le centroïde via une moyenne mobile exponentielle. + + Alpha faible : le centroïde reste stable et ne dérive pas vers les + variations d'un segment à l'autre (qui scinderaient un même locuteur). + """ self.count += 1 if self.count == 1: self.embedding = new_emb.copy() @@ -106,59 +110,77 @@ class SpeakerDiarizer: profils de locuteurs. Le modèle d'embedding est partagé (EmbeddingModel). """ + # Un nouveau locuteur n'est créé que si le meilleur score est NETTEMENT + # sous le seuil (zone morte) ET que le segment est assez long pour donner + # une empreinte fiable. Sinon le segment est rattaché au meilleur profil. + # Cela évite qu'un même locuteur soit éclaté en Locuteur 3, 4, 5... + NEW_SPEAKER_MARGIN = 0.20 # écart sous le seuil pour autoriser un nouveau + NEW_SPEAKER_MIN_S = 1.2 # durée mini d'un segment pour créer un locuteur + STICKY_MARGIN = 0.05 # adhérence au dernier locuteur si scores proches + def __init__( self, model: EmbeddingModel, - threshold: float = 0.70, + threshold: float = 0.35, max_speakers: int = 8, ) -> None: self._model = model self._threshold = threshold self._max_speakers = max_speakers self._profiles: list[_SpeakerProfile] = [] + self._last: _SpeakerProfile | None = None 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é. + Label du locuteur ("Locuteur 1", "Locuteur 2"...) 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 + duration_s = (len(pcm) // 2) / sample_rate - for profile in self._profiles: - score = _cosine_similarity(embedding, profile.embedding) - if score > best_score: - best_score = score - best_profile = profile + # Score de chaque profil existant + scores = [(_cosine_similarity(embedding, p.embedding), p) for p in self._profiles] + best_score, best_profile = max(scores, default=(-1.0, None), key=lambda x: x[0]) + # Adhérence : si le dernier locuteur est presque aussi proche que le + # meilleur, on le garde (réduit le clignotement entre deux profils). + if self._last is not None and best_profile is not self._last: + last_score = next((s for s, p in scores if p is self._last), -1.0) + if last_score >= best_score - self.STICKY_MARGIN: + best_score, best_profile = last_score, self._last + + # Rattachement à un locuteur connu if best_profile is not None and best_score >= self._threshold: best_profile.update(embedding) + self._last = best_profile 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" + # Création d'un nouveau locuteur : conditions strictes + can_create = ( + len(self._profiles) < self._max_speakers + and duration_s >= self.NEW_SPEAKER_MIN_S + and best_score < self._threshold - self.NEW_SPEAKER_MARGIN + ) + if not can_create and best_profile is not None: + # Zone ambiguë ou segment court : on rattache au plus proche + best_profile.update(embedding) + self._last = best_profile + return best_profile.label 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) + self._last = new_profile return new_label def reset(self) -> None: """Réinitialise les profils pour une nouvelle session.""" self._profiles.clear() + self._last = None diff --git a/app/main.py b/app/main.py index 5f756b5..a74855e 100644 --- a/app/main.py +++ b/app/main.py @@ -44,7 +44,7 @@ def clean_transcript(text: str) -> str: # 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_THRESHOLD = float(os.environ.get("DIARIZATION_THRESHOLD", "0.35")) DIARIZATION_MAX_SPEAKERS = int(os.environ.get("DIARIZATION_MAX_SPEAKERS", "8")) LIVEFLOW_USER = os.environ.get("LIVEFLOW_USER", "admin") diff --git a/app/segmenter.py b/app/segmenter.py index 4c8bb89..f52be69 100644 --- a/app/segmenter.py +++ b/app/segmenter.py @@ -9,6 +9,8 @@ silence prolongé ou une durée maximale. from collections import deque from dataclasses import dataclass +import audioop + import webrtcvad SAMPLE_RATE = 16000 @@ -22,6 +24,15 @@ SILENCE_END_MS = 700 # silence qui clôt un segment MIN_SPEECH_MS = 300 # en dessous, le segment est ignoré (bruit) MAX_SEGMENT_S = 25 # coupe forcée pour garder une latence raisonnable +# Contrôle de gain automatique appliqué AVANT la détection de parole : sans lui, +# un micro faible (casque...) reste sous le seuil du VAD et aucune phrase n'est +# détectée. On amène le niveau crête vers une cible exploitable. +AGC_TARGET_PEAK = 11000 # niveau crête visé (échelle int16) +AGC_NOISE_FLOOR = 120 # en dessous : du silence, on n'amplifie pas +AGC_MAX_GAIN = 60.0 # amplification maximale +AGC_ATTACK = 0.5 # vitesse de montée du gain (0-1) +AGC_RELEASE = 0.2 # vitesse de descente du gain (0-1) + @dataclass class Segment: @@ -30,9 +41,29 @@ class Segment: t1: float +class _AutoGain: + """Amplification adaptative du flux pour fiabiliser la détection de parole.""" + + def __init__(self): + self.gain = 1.0 + + def process(self, frame: bytes) -> bytes: + peak = audioop.max(frame, 2) + if peak >= AGC_NOISE_FLOOR: + desired = min(AGC_MAX_GAIN, AGC_TARGET_PEAK / peak) + rate = AGC_ATTACK if desired > self.gain else AGC_RELEASE + self.gain += (desired - self.gain) * rate + self.gain = max(1.0, min(AGC_MAX_GAIN, self.gain)) + if self.gain > 1.01: + return audioop.mul(frame, 2, self.gain) + return frame + + + class SpeechSegmenter: def __init__(self): self._vad = webrtcvad.Vad(VAD_AGGRESSIVENESS) + self._agc = _AutoGain() self._pending = bytearray() self._ring: deque[tuple[bytes, bool]] = deque(maxlen=PREROLL_FRAMES) self._frame_index = 0 @@ -49,6 +80,7 @@ class SpeechSegmenter: while len(self._pending) >= FRAME_BYTES: frame = bytes(self._pending[:FRAME_BYTES]) del self._pending[:FRAME_BYTES] + frame = self._agc.process(frame) # amplifie avant le VAD seg = self._process_frame(frame) if seg is not None: segments.append(seg) diff --git a/app/static/app.js b/app/static/app.js index cae6878..62d5d83 100644 --- a/app/static/app.js +++ b/app/static/app.js @@ -15,8 +15,6 @@ const state = { heardSound: false, silentWarned: false, sourceNode: null, - gainNode: null, - lastGainAt: 0, paused: false, }; @@ -78,8 +76,7 @@ async function switchMic(deviceId) { if (state.mediaStream) state.mediaStream.getTracks().forEach((t) => t.stop()); state.mediaStream = stream; state.sourceNode = state.audioContext.createMediaStreamSource(stream); - state.sourceNode.connect(state.gainNode); - state.gainNode.gain.value = savedGain(); + state.sourceNode.connect(state.workletNode); state.lastLoudAt = Date.now(); return true; } @@ -125,12 +122,10 @@ async function startRecording() { await state.audioContext.audioWorklet.addModule('worklet.js'); state.workletNode = new AudioWorkletNode(state.audioContext, 'pcm-downsampler'); state.workletNode.port.onmessage = (e) => queuePcm(new Int16Array(e.data)); - // gain automatique : compense les micros trop faibles (casques...) - state.gainNode = state.audioContext.createGain(); - state.gainNode.gain.value = savedGain(); // repart du gain mémorisé - state.gainNode.connect(state.workletNode); + // L'amplification des micros faibles se fait côté serveur (AGC avant le + // VAD) : le client envoie l'audio brut, sans gain qui se battrait avec. state.sourceNode = state.audioContext.createMediaStreamSource(state.mediaStream); - state.sourceNode.connect(state.gainNode); + state.sourceNode.connect(state.workletNode); // les libellés des micros ne sont disponibles qu'une fois la permission accordée listMics(); @@ -169,38 +164,13 @@ function updateVu(samples) { const v = Math.abs(samples[i]); if (v > peak) peak = v; } - $('vu-bar').style.width = Math.min(100, (peak / 32768) * 140) + '%'; - if (peak > 1500) { + // Affichage amplifié (x8) : le serveur amplifie le signal réel, le vumètre + // reflète l'activité même avec un micro faible. + $('vu-bar').style.width = Math.min(100, (peak / 32768) * 140 * 8) + '%'; + if (peak > 150) { state.lastLoudAt = Date.now(); state.heardSound = true; } - autoGain(peak); -} - -// Gain automatique : monte progressivement le volume des micros faibles -// jusqu'à un niveau exploitable par la détection de parole, redescend en cas -// de saturation, et mémorise le gain trouvé pour les prochaines réunions. -function savedGain() { - const g = parseFloat(localStorage.getItem('liveflow-gain')); - return Number.isFinite(g) ? Math.min(80, Math.max(1, g)) : 1; -} - -function autoGain(peak) { - if (!state.gainNode || state.paused) return; - const now = Date.now(); - if (now - state.lastGainAt < 250) return; - state.lastGainAt = now; - const g = state.gainNode.gain.value; - let next = g; - if (peak > 27000) { - next = Math.max(1, g * 0.7); // proche saturation - } else if (peak > 60 && peak < 12000) { - next = Math.min(80, g * 1.3); // signal trop faible - } - if (next !== g) { - state.gainNode.gain.value = next; - localStorage.setItem('liveflow-gain', String(next)); - } } function flushPcm() {