Files
meeting-assistant/MeetingAssistant.Tests/ResemblyzerSpeakerIdentificationServiceTests.cs
T

597 lines
24 KiB
C#

using MeetingAssistant;
using MeetingAssistant.MeetingNotes;
using MeetingAssistant.Speakers;
using MeetingAssistant.Transcription;
using Microsoft.EntityFrameworkCore;
using Microsoft.Extensions.Logging.Abstractions;
using Microsoft.Extensions.Options;
namespace MeetingAssistant.Tests;
public sealed class ResemblyzerSpeakerIdentificationServiceTests
{
[Fact]
public async Task LiveMatchRelabelsSpeakerAndStoresFiveVectorsWithoutWavSnippets()
{
await using var fixture = await Fixture.CreateAsync();
var identity = await fixture.AddIdentityAsync("Chris", [UnitVector(0), UnitVector(0, 20, 0.02f)]);
fixture.Encoder.Vectors = Enumerable.Range(0, 5)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
var result = await service.IdentifyKnownSpeakersAsync(
fixture.CreateRequest("Guest03", sampleCount: 5),
CancellationToken.None);
Assert.Equal("Chris", result.Segments.Single().Speaker);
Assert.Equal("Chris", result.SpeakerMappings["Guest03"]);
var saved = await fixture.LoadIdentityAsync(identity.Id);
Assert.Equal(7, saved.VoiceVectors.Count);
Assert.Empty(saved.Snippets);
Assert.Single(saved.References);
Assert.Single(fixture.Encoder.Requests);
Assert.Equal(5, fixture.Encoder.Requests[0].Count);
}
[Fact]
public async Task SummaryOverrideCreatesNamedIdentityFromThreeAvailableVectors()
{
await using var fixture = await Fixture.CreateAsync();
fixture.Encoder.Vectors = Enumerable.Range(0, 3)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
var request = fixture.CreateRequest("Guest-01", sampleCount: 3);
await service.ApplySpeakerOverrideAsync(
request,
"Guest-01",
"Sabrina",
CancellationToken.None);
var saved = await fixture.LoadOnlyIdentityAsync();
Assert.Equal("Sabrina", saved.CanonicalName);
Assert.Equal(3, saved.VoiceVectors.Count);
Assert.Empty(saved.Snippets);
Assert.Single(saved.References);
Assert.Equal(3, fixture.Encoder.Requests.Single().Count);
}
[Fact]
public async Task FinalProcessingLearnsUnmatchedSpeakerFromFiveVectorsAndAttendees()
{
await using var fixture = await Fixture.CreateAsync();
fixture.Encoder.Vectors = Enumerable.Range(0, 5)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
await service.ProcessFinishedTranscriptAsync(
fixture.CreateRequest("Guest-01", sampleCount: 5, attendees: ["John", "Mike"]),
CancellationToken.None);
var saved = await fixture.LoadOnlyIdentityAsync();
Assert.Null(saved.CanonicalName);
Assert.Equal(["John", "Mike"], saved.CandidateNames.Select(candidate => candidate.Name).Order());
Assert.Equal(5, saved.VoiceVectors.Count);
Assert.Empty(saved.Snippets);
Assert.Single(saved.References);
}
[Fact]
public async Task FiveVectorThresholdDoesNotCapQualifyingCurrentRunEvidence()
{
await using var fixture = await Fixture.CreateAsync();
fixture.Encoder.Vectors = Enumerable.Range(0, 8)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
await service.ProcessFinishedTranscriptAsync(
fixture.CreateRequest("Guest-01", sampleCount: 8, attendees: ["John", "Mike"]),
CancellationToken.None);
var saved = await fixture.LoadOnlyIdentityAsync();
Assert.Equal(8, saved.VoiceVectors.Count);
Assert.Equal(8, fixture.Encoder.Requests.Single().Count);
}
[Fact]
public async Task FinalProcessingPrunesForeignVectorsFromNewIdentityAtMinimum()
{
await using var fixture = await Fixture.CreateAsync(options =>
{
options.OutlierPruningMinimumVectors = 20;
options.OutlierPruningNeighborSimilarity = 0.90;
options.OutlierPruningMinimumNeighbors = 3;
options.OutlierPruningMinimumClusterRatio = 0.60;
});
fixture.Encoder.Vectors = Enumerable.Range(0, 16)
.Select(index => UnitVector(0, index + 2, 0.04f))
.Concat(Enumerable.Range(0, 4)
.Select(index => UnitVector(1, index + 30, 0.04f)))
.ToList();
var service = fixture.CreateService();
await service.ProcessFinishedTranscriptAsync(
fixture.CreateRequest("Guest-01", sampleCount: 20, attendees: ["John", "Mike"]),
CancellationToken.None);
var saved = await fixture.LoadOnlyIdentityAsync();
Assert.Equal(16, saved.VoiceVectors.Count);
Assert.All(
saved.VoiceVectors,
stored => Assert.True(SpeakerVoiceVectors.Decode(stored)[0] > 0.9f));
}
[Theory]
[InlineData(false)]
[InlineData(true)]
public async Task FinalProcessingStoresVectorsCollectedAfterLiveMatch(bool transcriptAlreadyRelabeled)
{
await using var fixture = await Fixture.CreateAsync();
var identity = await fixture.AddIdentityAsync(
"Chris",
[UnitVector(0), UnitVector(0, 20, 0.02f)]);
var currentRunVectors = Enumerable.Range(0, 8)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
fixture.Encoder.Vectors = currentRunVectors.Take(5).ToList();
var service = fixture.CreateService();
var liveResult = await service.IdentifyKnownSpeakersAsync(
fixture.CreateRequest("Guest03", sampleCount: 5),
CancellationToken.None);
fixture.Encoder.Vectors = currentRunVectors;
var finalRequest = fixture.CreateRequest("Guest03", sampleCount: 8) with
{
KnownSpeakerMappings = liveResult.SpeakerMappings
};
if (transcriptAlreadyRelabeled)
{
finalRequest = finalRequest with { Segments = liveResult.Segments };
}
await service.ProcessFinishedTranscriptAsync(finalRequest, CancellationToken.None);
var saved = await fixture.LoadIdentityAsync(identity.Id);
Assert.Equal(10, saved.VoiceVectors.Count);
Assert.Equal([5, 8], fixture.Encoder.Requests.Select(request => request.Count));
}
[Fact]
public async Task RepeatedMatchDeduplicatesVectorsAndHonorsConfiguredLimit()
{
await using var fixture = await Fixture.CreateAsync(options => options.MaxVectorsPerIdentity = 6);
var identity = await fixture.AddIdentityAsync("Chris", [UnitVector(0), UnitVector(0, 20, 0.02f)]);
fixture.Encoder.Vectors =
[
UnitVector(0),
UnitVector(0, 2, 0.01f),
UnitVector(0, 3, 0.02f),
UnitVector(0, 4, 0.03f),
UnitVector(0, 5, 0.04f)
];
var service = fixture.CreateService();
var request = fixture.CreateRequest("Guest03", sampleCount: 5);
await service.IdentifyKnownSpeakersAsync(request, CancellationToken.None);
await service.IdentifyKnownSpeakersAsync(request, CancellationToken.None);
var saved = await fixture.LoadIdentityAsync(identity.Id);
Assert.Equal(6, saved.VoiceVectors.Count);
Assert.Single(saved.References);
}
[Fact]
public async Task FinishedMatchingExtractsFiveNonOverlappingSamplesWhenLiveSamplesAreMissing()
{
await using var fixture = await Fixture.CreateAsync();
await fixture.AddIdentityAsync("Chris", [UnitVector(0), UnitVector(0, 20, 0.02f)]);
fixture.Encoder.Vectors = Enumerable.Range(0, 5)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
var result = await service.IdentifyFinishedSpeakersAsync(
fixture.CreateRequest("Guest03", sampleCount: 0, segmentCount: 5),
CancellationToken.None);
Assert.Equal("Chris", result.Segments[0].Speaker);
Assert.Equal(5, fixture.SnippetExtractor.Requests.Count);
Assert.All(
fixture.SnippetExtractor.Requests.Zip(fixture.SnippetExtractor.Requests.Skip(1)),
pair => Assert.True(pair.First[^1].End <= pair.Second[0].Start));
Assert.Equal(5, fixture.Encoder.Requests.Single().Count);
}
[Fact]
public async Task FinalMatchPromotesTheRemainingCandidateAndAuditsTheTranscript()
{
await using var fixture = await Fixture.CreateAsync();
var identity = await fixture.AddIdentityAsync(
null,
[UnitVector(0), UnitVector(0, 20, 0.02f)],
candidates: ["John", "Mike"]);
fixture.Encoder.Vectors = Enumerable.Range(0, 5)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
var result = await service.ProcessFinishedTranscriptAsync(
fixture.CreateRequest("Guest03", sampleCount: 5, attendees: ["Jane", "John", "Chris"]),
CancellationToken.None);
var saved = await fixture.LoadIdentityAsync(identity.Id);
Assert.Equal("John", saved.CanonicalName);
Assert.Equal(["John"], saved.CandidateNames.Select(candidate => candidate.Name));
Assert.Equal("John", result.Segments.Single().Speaker);
Assert.All(saved.References, reference =>
Assert.Contains("Guest03 was identified as John", File.ReadAllText(reference.TranscriptPath)));
}
[Fact]
public async Task FinalMatchResetsCandidatesWhenAttendeesDoNotIntersect()
{
await using var fixture = await Fixture.CreateAsync();
var identity = await fixture.AddIdentityAsync(
null,
[UnitVector(0), UnitVector(0, 20, 0.02f)],
candidates: ["John", "Mike"]);
fixture.Encoder.Vectors = Enumerable.Range(0, 5)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
await service.ProcessFinishedTranscriptAsync(
fixture.CreateRequest("Guest03", sampleCount: 5, attendees: ["Jane", "Chris"]),
CancellationToken.None);
var saved = await fixture.LoadIdentityAsync(identity.Id);
Assert.Null(saved.CanonicalName);
Assert.Equal(["Chris", "Jane"], saved.CandidateNames.Select(candidate => candidate.Name).Order());
}
[Fact]
public async Task FinalUnmatchedSpeakerWithOneCandidateIsAuditedWhenLearned()
{
await using var fixture = await Fixture.CreateAsync();
fixture.Encoder.Vectors = Enumerable.Range(0, 5)
.Select(index => UnitVector(0, index + 2, 0.03f))
.ToList();
var service = fixture.CreateService();
await service.ProcessFinishedTranscriptAsync(
fixture.CreateRequest("Guest-01", sampleCount: 5, attendees: ["Manuel"]),
CancellationToken.None);
var saved = await fixture.LoadOnlyIdentityAsync();
Assert.Equal("Manuel", saved.CanonicalName);
Assert.Contains("Guest-01 was identified as Manuel", File.ReadAllText(saved.References.Single().TranscriptPath));
}
[Fact]
public async Task CandidateLimitCountsOnlyIdentitiesWithCompatibleVectors()
{
await using var fixture = await Fixture.CreateAsync(
configureSpeaker: options => options.MaxMatchCandidates = 1);
await fixture.AddIdentityAsync("Legacy WAV identity", []);
await fixture.AddIdentityAsync("Chris", [UnitVector(0)]);
fixture.Encoder.Vectors = Enumerable.Repeat(UnitVector(0), 5).ToList();
var service = fixture.CreateService();
var result = await service.IdentifyKnownSpeakersAsync(
fixture.CreateRequest("Guest03", sampleCount: 5, attendees: []),
CancellationToken.None);
Assert.Equal("Chris", result.Segments.Single().Speaker);
}
[Fact]
public async Task SummaryOverrideMergePreservesAllEvidenceFromCurrentRunCandidate()
{
await using var fixture = await Fixture.CreateAsync();
var target = await fixture.AddIdentityAsync("Sabrina", [UnitVector(0)]);
var request = fixture.CreateRequest("Guest-01", sampleCount: 3, attendees: ["Sabrina"]);
await fixture.AddCurrentRunCandidateAsync(request, "Sabrina");
fixture.Encoder.Vectors = Enumerable.Range(0, 3)
.Select(index => UnitVector(0, index + 2, 0.02f))
.ToList();
var service = fixture.CreateService();
await service.ApplySpeakerOverrideAsync(
request,
"Guest-01",
"Sabrina",
CancellationToken.None);
var saved = await fixture.LoadIdentityAsync(target.Id);
Assert.Single(saved.Snippets);
Assert.Equal(5, saved.VoiceVectors.Count);
Assert.Contains(saved.VoiceVectors, vector => vector.ModelId == "older-model");
await using var context = new TestDbContextFactory(fixture.DatabasePath).CreateDbContext();
Assert.Single(await context.SpeakerIdentities.ToListAsync());
}
private static float[] UnitVector(
int primaryDimension,
int? secondaryDimension = null,
float secondaryValue = 0)
{
var vector = new float[256];
vector[primaryDimension] = 1;
if (secondaryDimension is { } dimension)
{
vector[dimension] = 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 sealed class Fixture : IAsyncDisposable
{
private readonly string directory;
private readonly string databasePath;
private readonly SpeakerIdentificationOptions speakerOptions;
private Fixture(string directory, string databasePath, SpeakerIdentificationOptions speakerOptions)
{
this.directory = directory;
this.databasePath = databasePath;
this.speakerOptions = speakerOptions;
}
public FakeEncoder Encoder { get; } = new();
public FakeSnippetExtractor SnippetExtractor { get; } = new();
public string DatabasePath => databasePath;
public static async Task<Fixture> CreateAsync(
Action<ResemblyzerSpeakerRecognitionOptions>? configure = null,
Action<SpeakerIdentificationOptions>? configureSpeaker = null)
{
var directory = Path.Combine(
Path.GetTempPath(),
"meeting-assistant-tests",
Guid.NewGuid().ToString("N"));
Directory.CreateDirectory(directory);
var databasePath = Path.Combine(directory, "speaker-identities.db");
var speakerOptions = new SpeakerIdentificationOptions
{
DatabasePath = databasePath,
MatchBatchSize = 6,
MaxMatchCandidates = 100,
MatchIdentityActiveAge = TimeSpan.FromDays(365),
MinimumSampleSpeechDuration = TimeSpan.Zero,
Resemblyzer = new ResemblyzerSpeakerRecognitionOptions
{
Enabled = true,
RequiredVectorsPerSpeaker = 5,
MaxVectorsPerIdentity = 1000,
MinimumClusterCohesion = 0.75,
MinimumIdentitySimilarity = 0.75,
MinimumSimilarityMargin = 0.05,
ModelId = "resemblyzer-0.1.4-pretrained"
}
};
configure?.Invoke(speakerOptions.Resemblyzer);
configureSpeaker?.Invoke(speakerOptions);
await using var context = new SpeakerIdentityDbContext(
new DbContextOptionsBuilder<SpeakerIdentityDbContext>()
.UseSqlite($"Data Source={databasePath};Pooling=False")
.Options);
await SpeakerIdentitySchema.EnsureCreatedOrUpdatedAsync(context, CancellationToken.None);
return new Fixture(directory, databasePath, speakerOptions);
}
public ResemblyzerSpeakerIdentificationService CreateService()
{
var appOptions = new MeetingAssistantOptions { SpeakerIdentification = speakerOptions };
return new ResemblyzerSpeakerIdentificationService(
new TestDbContextFactory(databasePath),
SnippetExtractor,
Encoder,
new ResemblyzerVoiceClusterMatcher(
speakerOptions.Resemblyzer,
NullLogger<ResemblyzerVoiceClusterMatcher>.Instance),
new ResemblyzerVoiceVectorOutlierPruner(
speakerOptions.Resemblyzer,
NullLogger<ResemblyzerVoiceVectorOutlierPruner>.Instance),
Options.Create(appOptions),
NullLogger<ResemblyzerSpeakerIdentificationService>.Instance);
}
public SpeakerIdentificationRequest CreateRequest(
string speaker,
int sampleCount,
IReadOnlyList<string>? attendees = null,
int segmentCount = 1)
{
var transcriptPath = Path.Combine(directory, "transcript.md");
File.WriteAllText(transcriptPath, "Transcript");
var segments = Enumerable.Range(0, segmentCount)
.Select(index => new TranscriptionSegment(
TimeSpan.FromSeconds(index * 30),
TimeSpan.FromSeconds((index + 1) * 30),
speaker,
"enough useful words for a speaker sample"))
.ToList();
var segment = segments[0];
return new SpeakerIdentificationRequest(
Path.Combine(directory, "meeting.wav"),
new MeetingNote(
Path.Combine(directory, "meeting.md"),
new MeetingNoteFrontmatter
{
Title = "Test",
Attendees = attendees?.ToList() ?? ["Chris"],
Transcript = transcriptPath,
AssistantContext = Path.Combine(directory, "context.md"),
Summary = Path.Combine(directory, "summary.md")
},
""),
segments,
Enumerable.Range(0, sampleCount)
.Select(index => new SpeakerAudioSample(speaker, segment, [(byte)(index + 1)], 100 - index))
.ToList());
}
public async Task<SpeakerIdentity> AddIdentityAsync(
string? name,
IReadOnlyList<float[]> vectors,
IReadOnlyList<string>? candidates = null,
IReadOnlyList<string>? aliases = null)
{
await using var context = new TestDbContextFactory(databasePath).CreateDbContext();
var now = DateTimeOffset.UtcNow;
var identity = new SpeakerIdentity
{
CanonicalName = name,
CreatedAt = now,
UpdatedAt = now,
CandidateNames = candidates?.Select(candidate => new SpeakerCandidateName { Name = candidate }).ToList() ?? [],
Aliases = aliases?.Select(alias => new SpeakerAlias { Name = alias }).ToList() ?? [],
VoiceVectors = vectors.Select((vector, index) => new SpeakerVoiceVector
{
ModelId = speakerOptions.Resemblyzer.ModelId,
Dimensions = 256,
VectorBytes = ToBytes(vector),
Fingerprint = $"known-{index}",
CreatedAt = now.AddMinutes(index)
}).ToList()
};
context.SpeakerIdentities.Add(identity);
await context.SaveChangesAsync();
return identity;
}
public async Task AddCurrentRunCandidateAsync(
SpeakerIdentificationRequest request,
string candidateName)
{
await using var context = new TestDbContextFactory(databasePath).CreateDbContext();
var now = DateTimeOffset.UtcNow;
context.SpeakerIdentities.Add(new SpeakerIdentity
{
CreatedAt = now,
UpdatedAt = now,
CandidateNames = [new SpeakerCandidateName { Name = candidateName }],
Snippets = [new SpeakerSnippet { WavBytes = [9, 8, 7], CreatedAt = now }],
VoiceVectors =
[
new SpeakerVoiceVector
{
ModelId = "older-model",
Dimensions = 256,
VectorBytes = ToBytes(UnitVector(1)),
Fingerprint = "older-model-vector",
CreatedAt = now
}
],
References =
[
SpeakerIdentityReferences.Create(
request.MeetingNote.Path,
request.MeetingNote.Frontmatter.Transcript,
now)
]
});
await context.SaveChangesAsync();
}
public async Task<SpeakerIdentity> LoadIdentityAsync(int id)
{
await using var context = new TestDbContextFactory(databasePath).CreateDbContext();
return await context.SpeakerIdentities
.Include(identity => identity.Snippets)
.Include(identity => identity.VoiceVectors)
.Include(identity => identity.CandidateNames)
.Include(identity => identity.Aliases)
.Include(identity => identity.References)
.SingleAsync(identity => identity.Id == id);
}
public async Task<SpeakerIdentity> LoadOnlyIdentityAsync()
{
await using var context = new TestDbContextFactory(databasePath).CreateDbContext();
return await context.SpeakerIdentities
.Include(identity => identity.Snippets)
.Include(identity => identity.VoiceVectors)
.Include(identity => identity.CandidateNames)
.Include(identity => identity.References)
.SingleAsync();
}
public ValueTask DisposeAsync()
{
if (Directory.Exists(directory))
{
Directory.Delete(directory, recursive: true);
}
return ValueTask.CompletedTask;
}
}
private sealed class FakeEncoder : IResemblyzerVoiceEncoder
{
public IReadOnlyList<float[]> Vectors { get; set; } = [];
public List<IReadOnlyList<byte[]>> Requests { get; } = [];
public Task<IReadOnlyList<float[]>> EncodeAsync(
IReadOnlyList<byte[]> wavSamples,
CancellationToken cancellationToken)
{
Requests.Add(wavSamples.Select(sample => sample.ToArray()).ToList());
return Task.FromResult<IReadOnlyList<float[]>>(Vectors.Take(wavSamples.Count).ToList());
}
public Task WarmUpAsync(CancellationToken cancellationToken)
{
return Task.CompletedTask;
}
}
private sealed class FakeSnippetExtractor : ISpeakerSnippetExtractor
{
public List<IReadOnlyList<TranscriptionSegment>> Requests { get; } = [];
public Task<byte[]> ExtractSnippetAsync(
string audioPath,
IReadOnlyList<TranscriptionSegment> speakerSegments,
CancellationToken cancellationToken)
{
Requests.Add(speakerSegments.ToList());
return Task.FromResult<byte[]>([checked((byte)Requests.Count)]);
}
}
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);
}
}
}