Files
meeting-assistant/MeetingAssistant/Speakers/ResemblyzerVoiceVectorOutlierPruner.cs

209 lines
6.3 KiB
C#

namespace MeetingAssistant.Speakers;
public sealed record ResemblyzerVoiceVectorUpdateResult(
int AddedCount,
int RemovedCount)
{
public bool Changed => AddedCount > 0 || RemovedCount > 0;
}
public sealed class ResemblyzerVoiceVectorOutlierPruner
{
private readonly ResemblyzerSpeakerRecognitionOptions options;
private readonly ILogger<ResemblyzerVoiceVectorOutlierPruner> logger;
public ResemblyzerVoiceVectorOutlierPruner(
ResemblyzerSpeakerRecognitionOptions options,
ILogger<ResemblyzerVoiceVectorOutlierPruner> logger)
{
this.options = options;
this.logger = logger;
}
public int Prune(SpeakerIdentity identity)
{
var compatible = new List<DecodedSpeakerVoiceVector>();
foreach (var entry in SpeakerVoiceVectors.DecodeCompatibleEntries(
identity,
options.ModelId,
logger))
{
try
{
compatible.Add(new DecodedSpeakerVoiceVector(
entry.Stored,
SpeakerVoiceVectors.Normalize(
entry.Vector,
$"Stored voice vector {entry.Stored.Id}")));
}
catch (InvalidDataException exception)
{
logger.LogWarning(
exception,
"Preserving invalid stored voice vector {VectorId} while pruning identity {IdentityId}",
entry.Stored.Id,
identity.Id);
}
}
var minimumVectors = options.OutlierPruningMinimumVectors;
if (compatible.Count < minimumVectors)
{
return 0;
}
var clusters = FindDensityClusters(compatible.Select(item => item.Vector).ToList());
var ranked = clusters
.GroupBy(cluster => cluster)
.Where(group => group.Key > 0)
.Select(group => new { Id = group.Key, Count = group.Count() })
.OrderByDescending(group => group.Count)
.ThenBy(group => group.Id)
.ToList();
if (ranked.Count == 0)
{
return Skip(identity, "no dense cluster was found");
}
if (ranked.Count > 1 && ranked[0].Count == ranked[1].Count)
{
return Skip(identity, "no uniquely largest dense cluster was found");
}
var dominant = ranked[0];
var ratio = (double)dominant.Count / compatible.Count;
if (ratio < options.OutlierPruningMinimumClusterRatio)
{
return Skip(
identity,
$"largest dense cluster ratio {ratio:F4} is below {options.OutlierPruningMinimumClusterRatio:F4}");
}
var removed = compatible
.Where((_, index) => clusters[index] != dominant.Id)
.Select(item => item.Stored)
.ToList();
foreach (var vector in removed)
{
identity.VoiceVectors.Remove(vector);
}
if (removed.Count > 0)
{
identity.UpdatedAt = DateTimeOffset.UtcNow;
}
logger.LogInformation(
"Pruned {RemovedCount} Resemblyzer voice-vector outliers from identity {IdentityId}; retained dominant cluster of {RetainedCount}/{CompatibleCount} vectors",
removed.Count,
identity.Id,
dominant.Count,
compatible.Count);
return removed.Count;
}
public ResemblyzerVoiceVectorUpdateResult AddAndPrune(
SpeakerIdentity identity,
IReadOnlyList<float[]> vectors,
DateTimeOffset createdAt)
{
var removedBeforeAdding = Prune(identity);
var added = SpeakerVoiceVectors.AddDistinct(
identity,
vectors,
options.ModelId,
options.MaxVectorsPerIdentity,
createdAt);
var removedAfterAdding = Prune(identity);
var result = new ResemblyzerVoiceVectorUpdateResult(
added,
removedBeforeAdding + removedAfterAdding);
if (result.Changed)
{
identity.UpdatedAt = createdAt;
}
return result;
}
private int[] FindDensityClusters(IReadOnlyList<float[]> vectors)
{
var neighbors = Enumerable.Range(0, vectors.Count)
.Select(index => Enumerable.Range(0, vectors.Count)
.Where(candidate => SpeakerVoiceVectors.Cosine(vectors[index], vectors[candidate]) >= options.OutlierPruningNeighborSimilarity)
.ToList())
.ToList();
var minimumNeighbors = options.OutlierPruningMinimumNeighbors;
var labels = new int[vectors.Count];
var clusterId = 0;
for (var point = 0; point < vectors.Count; point++)
{
if (labels[point] != 0)
{
continue;
}
if (neighbors[point].Count < minimumNeighbors)
{
labels[point] = -1;
continue;
}
clusterId++;
ExpandCluster(point, clusterId, neighbors, labels, minimumNeighbors);
}
return labels;
}
private static void ExpandCluster(
int seed,
int clusterId,
IReadOnlyList<List<int>> neighbors,
int[] labels,
int minimumNeighbors)
{
labels[seed] = clusterId;
var pending = new Queue<int>(neighbors[seed]);
var queued = neighbors[seed].ToHashSet();
while (pending.TryDequeue(out var point))
{
if (labels[point] == -1)
{
labels[point] = clusterId;
}
if (labels[point] != 0)
{
continue;
}
labels[point] = clusterId;
if (neighbors[point].Count < minimumNeighbors)
{
continue;
}
foreach (var neighbor in neighbors[point])
{
if (queued.Add(neighbor))
{
pending.Enqueue(neighbor);
}
}
}
}
private int Skip(
SpeakerIdentity identity,
string reason)
{
logger.LogInformation(
"Skipped Resemblyzer voice-vector pruning for identity {IdentityId}: {Reason}",
identity.Id,
reason);
return 0;
}
}