"""Agrupamento de falas que representam retakes da mesma fala ou ideia. Trabalha com uma janela temporal: compara cada fala com as ``n`` seguintes (e só as da mesma faixa de áudio), evitando agrupar frases iguais que aparecem em partes muito distantes do vídeo. Não decide a classificação — apenas forma grupos candidatos e reúne as métricas. """ from __future__ import annotations from typing import Iterable from .analisador_de_intervalos import AnalisadorDeIntervalos from .comparador_de_falas import ComparadorDeFalas from .comparador_semantico import ComparadorSemantico from .detector_de_reinicios import DetectorDeReinicios from .modelos_de_retakes import ( AgrupamentoDeFalas, Fala, ParAgrupado, ResultadoDoIntervalo, SinalDeReinicio, ) # Limiares de candidatura de um par a pertencer a um grupo de retake. _SIMILARIDADE_MINIMA_CANDIDATO = 0.55 _SIMILARIDADE_FORTE = 0.75 _PROXIMIDADE_AUXILIAR = 0.60 # Janela padrão de falas a serem comparadas com cada uma. _JANELA_PADRAO = 5 # Região de dúvida que dispara a comparação semântica. _ZONA_DE_DUVIDA_INFERIOR = 0.60 _ZONA_DE_DUVIDA_SUPERIOR = 0.85 class AgrupadorDeTakes: """Gera grupos candidatos conectando falas semelhantes e próximas.""" def __init__( self, comparador: ComparadorDeFalas, comparador_semantico: ComparadorSemantico, analisador_de_intervalos: AnalisadorDeIntervalos, detector_de_reinicios: DetectorDeReinicios, janela: int = _JANELA_PADRAO, ) -> None: self.comparador = comparador self.comparador_semantico = comparador_semantico self.analisador_de_intervalos = analisador_de_intervalos self.detector_de_reinicios = detector_de_reinicios self.janela = janela def agrupar(self, falas: Iterable[Fala]) -> list[AgrupamentoDeFalas]: falas = sorted(falas, key=lambda item: (item.ordem, item.inicio)) sinais = {sinal.fala_id: sinal for sinal in self.detector_de_reinicios.detectar(falas)} arestas: list[tuple[int, int, ParAgrupado]] = [] for indice, fala_a in enumerate(falas): limite = min(len(falas), indice + 1 + self.janela) for jindice in range(indice + 1, limite): fala_b = falas[jindice] if fala_a.faixa_id != fala_b.faixa_id: continue par = self._avaliar_par(fala_a, fala_b, sinais) if par is not None: arestas.append((indice, jindice, par)) return self._agrupar_por_arestas(falas, arestas, sinais) def _avaliar_par( self, fala_a: Fala, fala_b: Fala, sinais: dict[str, SinalDeReinicio], ) -> ParAgrupado | None: comparacao = self.comparador.comparar(fala_a, fala_b) intervalo = self.analisador_de_intervalos.analisar(fala_a, fala_b) tem_reinicio = sinais.get(fala_b.id) is not None semantica = 0.0 # Comparação semântica só na zona de dúvida, para evitar custo. texto = comparacao.similaridade_final if _ZONA_DE_DUVIDA_INFERIOR <= texto <= _ZONA_DE_DUVIDA_SUPERIOR: semantica = self.comparador_semantico.comparar(fala_a, fala_b) melhor = max(comparacao.similaridade_final, semantica) candidato = ( melhor >= _SIMILARIDADE_FORTE or (melhor >= _SIMILARIDADE_MINIMA_CANDIDATO and intervalo.proximidade_temporal >= _PROXIMIDADE_AUXILIAR and (tem_reinicio or intervalo.indicio_de_nova_tentativa)) ) if candidato: return ParAgrupado(fala_a, fala_b, comparacao, intervalo, round(semantica, 4)) return None @staticmethod def _agrupar_por_arestas( falas: list[Fala], arestas: list[tuple[int, int, ParAgrupado]], sinais: dict[str, SinalDeReinicio], ) -> list[AgrupamentoDeFalas]: pai = list(range(len(falas))) def encontrar(x: int) -> int: while pai[x] != x: pai[x] = pai[pai[x]] x = pai[x] return x arestas.sort(key=lambda item: (item[0], item[1])) for a, b, _ in arestas: raiz_a, raiz_b = encontrar(a), encontrar(b) if raiz_a != raiz_b: pai[raiz_b] = raiz_a componentes: dict[int, list[int]] = {} for indice in range(len(falas)): componentes.setdefault(encontrar(indice), []).append(indice) grupos: list[AgrupamentoDeFalas] = [] for inds in componentes.values(): if len(inds) < 2: continue inds.sort() fala_ids = tuple(falas[indice].id for indice in inds) pares_por_chave: dict[tuple[str, str], ParAgrupado] = { (par.fala_a.id, par.fala_b.id): par for _, _, par in arestas if par.fala_a.id in fala_ids and par.fala_b.id in fala_ids } pares: list[ParAgrupado] = [] for a, b in zip(inds, inds[1:]): chave = (falas[a].id, falas[b].id) par = pares_por_chave.get(chave) if par is None: par = pares_por_chave.get((falas[b].id, falas[a].id)) if par is not None: pares.append(par) sinais_do_grupo = tuple( sinais[fala_id] for fala_id in fala_ids if fala_id in sinais) if pares: grupos.append(AgrupamentoDeFalas(fala_ids, tuple(pares), sinais_do_grupo)) return grupos