Files
meeting-assistant/MeetingAssistant.Tests/ResemblyzerSpeakerIdentityMergeServiceTests.cs

175 lines
6.7 KiB
C#

using MeetingAssistant;
using MeetingAssistant.Speakers;
using Microsoft.EntityFrameworkCore;
using Microsoft.Extensions.Logging.Abstractions;
using Microsoft.Extensions.Options;
namespace MeetingAssistant.Tests;
public sealed class ResemblyzerSpeakerIdentityMergeServiceTests
{
[Fact]
public async Task MergeRecentIdentitiesRequiresTwoDisjointMatchingVectorClusters()
{
var directory = Path.Combine(Path.GetTempPath(), "meeting-assistant-tests", Guid.NewGuid().ToString("N"));
Directory.CreateDirectory(directory);
try
{
var databasePath = Path.Combine(directory, "identities.db");
var factory = new TestDbContextFactory(databasePath);
await using (var context = factory.CreateDbContext())
{
await SpeakerIdentitySchema.EnsureCreatedOrUpdatedAsync(context, CancellationToken.None);
context.SpeakerIdentities.Add(CreateIdentity(
"Chris",
DateTimeOffset.UtcNow.AddMonths(-2),
Enumerable.Range(0, 10)
.Select(index => UnitVector(0, index + 20, 0.02f))
.ToList(),
directory));
context.SpeakerIdentities.Add(CreateIdentity(
"Chris duplicate",
DateTimeOffset.UtcNow.AddDays(-1),
Enumerable.Range(0, 10)
.Select(index => UnitVector(0, index + 2, 0.02f))
.Concat(Enumerable.Range(0, 4)
.Select(index => UnitVector(1, index + 40, 0.02f)))
.ToList(),
directory));
context.SpeakerIdentities.Add(CreateIdentity(
"Decoy",
DateTimeOffset.UtcNow.AddMonths(-1),
[UnitVector(0)],
directory));
await context.SaveChangesAsync();
}
var resemblyzerOptions = new ResemblyzerSpeakerRecognitionOptions
{
Enabled = true,
RequiredVectorsPerSpeaker = 5,
MaxVectorsPerIdentity = 1000,
MinimumClusterCohesion = 0.75,
MinimumIdentitySimilarity = 0.75,
MinimumSimilarityMargin = 0.05,
ModelId = "resemblyzer-0.1.4-pretrained"
};
var service = new ResemblyzerSpeakerIdentityMergeService(
factory,
new ResemblyzerVoiceClusterMatcher(
resemblyzerOptions,
NullLogger<ResemblyzerVoiceClusterMatcher>.Instance),
new ResemblyzerVoiceVectorOutlierPruner(
resemblyzerOptions,
NullLogger<ResemblyzerVoiceVectorOutlierPruner>.Instance),
Options.Create(new MeetingAssistantOptions
{
SpeakerIdentification = new SpeakerIdentificationOptions
{
MergeRecentIdentityAge = TimeSpan.FromDays(14),
MaxMatchCandidates = 1,
MaxSnippetsPerSpeaker = 3,
Resemblyzer = resemblyzerOptions
}
}),
NullLogger<ResemblyzerSpeakerIdentityMergeService>.Instance);
var result = await service.MergeRecentIdentitiesAsync(TimeSpan.FromDays(14), CancellationToken.None);
Assert.Equal(2, result.MatchAttempts);
Assert.Equal(1, result.MergedPairs);
await using var verification = factory.CreateDbContext();
var saved = await verification.SpeakerIdentities
.Include(identity => identity.Aliases)
.Include(identity => identity.VoiceVectors)
.SingleAsync(identity => identity.CanonicalName == "Chris");
Assert.Equal(2, await verification.SpeakerIdentities.CountAsync());
Assert.Equal("Chris", saved.CanonicalName);
Assert.Contains(saved.Aliases, alias => alias.Name == "Chris duplicate");
Assert.Equal(20, saved.VoiceVectors.Count);
Assert.All(saved.VoiceVectors, vector => Assert.True(ToVector(vector.VectorBytes)[0] > 0.9f));
}
finally
{
Directory.Delete(directory, recursive: true);
}
}
private static SpeakerIdentity CreateIdentity(
string name,
DateTimeOffset createdAt,
IReadOnlyList<float[]> vectors,
string directory)
{
var transcriptPath = Path.Combine(directory, $"{Guid.NewGuid():N}.md");
File.WriteAllText(transcriptPath, "Transcript");
return new SpeakerIdentity
{
CanonicalName = name,
CreatedAt = createdAt,
UpdatedAt = createdAt,
VoiceVectors = vectors.Select((vector, index) => new SpeakerVoiceVector
{
ModelId = "resemblyzer-0.1.4-pretrained",
Dimensions = 256,
VectorBytes = ToBytes(vector),
Fingerprint = $"{name}-{index}",
CreatedAt = createdAt.AddMinutes(index)
}).ToList(),
References =
[
new SpeakerIdentityReference
{
MeetingNotePath = Path.Combine(directory, $"{Guid.NewGuid():N}.md"),
TranscriptPath = transcriptPath,
CreatedAt = createdAt
}
]
};
}
private static float[] UnitVector(int primary, int? secondary = null, float secondaryValue = 0)
{
var vector = new float[256];
vector[primary] = 1;
if (secondary is { } index)
{
vector[index] = secondaryValue;
}
return vector;
}
private static byte[] ToBytes(float[] vector)
{
var bytes = new byte[vector.Length * sizeof(float)];
Buffer.BlockCopy(vector, 0, bytes, 0, bytes.Length);
return bytes;
}
private static float[] ToVector(byte[] bytes)
{
var vector = new float[bytes.Length / sizeof(float)];
Buffer.BlockCopy(bytes, 0, vector, 0, bytes.Length);
return vector;
}
private sealed class TestDbContextFactory : IDbContextFactory<SpeakerIdentityDbContext>
{
private readonly string databasePath;
public TestDbContextFactory(string databasePath)
{
this.databasePath = databasePath;
}
public SpeakerIdentityDbContext CreateDbContext()
{
return new SpeakerIdentityDbContext(
new DbContextOptionsBuilder<SpeakerIdentityDbContext>()
.UseSqlite($"Data Source={databasePath};Pooling=False")
.Options);
}
}
}