From 019b5780d638094467da534875a81aa7dc2d6871 Mon Sep 17 00:00:00 2001 From: KakaruHayate Date: Tue, 4 Aug 2026 15:25:05 +0800 Subject: [PATCH 1/3] Implement variance retake with hard compose --- .../DiffSinger/DiffSingerVariance.cs | 137 +++++++++-- .../DiffSinger/DiffSingerVariancePatch.cs | 215 ++++++++++-------- .../DiffSinger/DiffSingerVariancePatchTest.cs | 113 ++++++--- 3 files changed, 321 insertions(+), 144 deletions(-) diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs index 6d4193458..403aec863 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs @@ -259,10 +259,13 @@ void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { new DenseTensor(new long[] { speedup }, new int[] { 1 },false))); } //Speaker + float[]? speakerEmbed = null; if(dsConfig.speakers != null) { var speakerEmbedManager = getSpeakerEmbedManager(); var spkEmbedTensor = speakerEmbedManager.PhraseSpeakerEmbedByFrame(phrase, ph_dur, frameMs, totalFrames, headFrames, tailFrames); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("spk_embed", spkEmbedTensor)); + speakerEmbed = spkEmbedTensor.ToArray(); + // Speaker embedding is a retake-able frame-level condition. + AddVarianceInput(NamedOnnxValue.CreateFromTensor("spk_embed", spkEmbedTensor), includeInPatchKey: false); } ulong? variancePatchKey = null; if (Preferences.Default.DiffSingerTensorCache && @@ -270,16 +273,58 @@ void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { var baseHash = new DiffSingerCache(varianceHash, variancePatchInputs).Hash; variancePatchKey = DiffSingerVariancePatch.BuildStateKey(baseHash, phrase.position, phrase.end); } - Onnx.VerifyInputNames(varianceModel, varianceInputs); - var varianceCache = Preferences.Default.DiffSingerTensorCache + var fullVarianceCache = Preferences.Default.DiffSingerTensorCache ? new DiffSingerCache(varianceHash, varianceInputs) : null; - var varianceOutputs = varianceCache?.Load(); - if (varianceOutputs is null) { - varianceOutputs = varianceModel.Run(varianceInputs).Cast().ToList(); - varianceCache?.Save(varianceOutputs); - phrase.AddCacheFile(varianceCache?.Filename); + var fullVarianceOutputs = fullVarianceCache?.Load(); + if (fullVarianceOutputs != null) { + var cachedResult = ParseVarianceResult(fullVarianceOutputs, frameMs, headFrames, tailFrames, totalFrames); + if (variancePatchKey.HasValue) { + variancePatchStates[variancePatchKey.Value] = + new VariancePatchState(pitch, speakerEmbed, cachedResult); + } + return cachedResult; + } + VariancePatchState? previous = null; + bool[]? retakeMask = null; + if (variancePatchKey.HasValue && variancePatchStates.TryGetValue(variancePatchKey.Value, out var cachedState) && + DiffSingerVariancePatch.IsMetadataCompatible(cachedState.result, new VarianceResult { + frameMs = frameMs, + headFrames = headFrames, + tailFrames = tailFrames, + totalFrames = totalFrames, + })) { + previous = cachedState; + var pitchMask = DiffSingerVariancePatch.BuildChangedFrameMask(cachedState.pitch, pitch, 1e-4f); + var speakerMask = DiffSingerVariancePatch.BuildChangedFrameMask( + cachedState.speakerEmbed ?? Array.Empty(), + speakerEmbed ?? Array.Empty(), + totalFrames, + 1e-4f); + retakeMask = new bool[totalFrames]; + for (int i = 0; i < retakeMask.Length; i++) { + retakeMask[i] = (i < pitchMask.Length && pitchMask[i]) || + (i < speakerMask.Length && speakerMask[i]); + } + if (!retakeMask.Any(x => x)) { + return DiffSingerVariancePatch.CloneResult(cachedState.result); + } + if (retakeMask.All(x => x)) { + previous = null; + } else { + ReplaceVarianceInputsWithPrevious(varianceInputs, cachedState.result); + } } + if (retakeMask != null) { + var retakeTensorValues = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); + var retakeInput = varianceInputs.First(x => x.Name == "retake"); + varianceInputs[varianceInputs.IndexOf(retakeInput)] = NamedOnnxValue.CreateFromTensor( + "retake", + new DenseTensor(retakeTensorValues, new[] { retakeTensorValues.Length }, false) + .Reshape(new[] { 1, totalFrames, numVariances })); + } + Onnx.VerifyInputNames(varianceModel, varianceInputs); + var varianceOutputs = varianceModel.Run(varianceInputs).Cast().ToList(); Tensor? energy_pred = dsConfig.predict_energy ? varianceOutputs .Where(o => o.Name == "energy_pred") @@ -314,22 +359,76 @@ void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { tailFrames = tailFrames, totalFrames = totalFrames, }; + if (previous != null && retakeMask != null) { + var channelMask = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); + result = DiffSingerVariancePatch.HardCompose(previous.result, result, channelMask, numVariances); + } + if (fullVarianceCache != null) { + fullVarianceCache.Save(BuildVarianceOutputs(result)); + phrase.AddCacheFile(fullVarianceCache.Filename); + } if (variancePatchKey.HasValue) { - result = ApplyVariancePatch(variancePatchKey.Value, pitch, result); + variancePatchStates[variancePatchKey.Value] = + new VariancePatchState(pitch, speakerEmbed, result); } return result; } - VarianceResult ApplyVariancePatch(ulong patchKey, float[] pitch, VarianceResult result) { - try { - variancePatchStates.TryGetValue(patchKey, out var previous); - var merged = DiffSingerVariancePatch.Merge(previous, pitch, result); - variancePatchStates[patchKey] = new VariancePatchState(pitch, merged); - return merged; - } catch (Exception e) { - Log.Warning(e, "Failed to apply DiffSinger variance local pitch patch."); - variancePatchStates[patchKey] = new VariancePatchState(pitch, result); - return result; + VarianceResult ParseVarianceResult( + ICollection outputs, + float frameMs, + int headFrames, + int tailFrames, + int totalFrames) { + return new VarianceResult { + energy = dsConfig.predict_energy ? outputs.First(o => o.Name == "energy_pred").AsTensor().ToArray() : null, + breathiness = dsConfig.predict_breathiness ? outputs.First(o => o.Name == "breathiness_pred").AsTensor().ToArray() : null, + voicing = dsConfig.predict_voicing ? outputs.First(o => o.Name == "voicing_pred").AsTensor().ToArray() : null, + tension = dsConfig.predict_tension ? outputs.First(o => o.Name == "tension_pred").AsTensor().ToArray() : null, + frameMs = frameMs, + headFrames = headFrames, + tailFrames = tailFrames, + totalFrames = totalFrames, + }; + } + + List BuildVarianceOutputs(VarianceResult result) { + var outputs = new List(); + void Add(string name, float[]? values) { + if (values != null) { + outputs.Add(NamedOnnxValue.CreateFromTensor( + name, + new DenseTensor(values, new[] { values.Length }, false) + .Reshape(new[] { 1, values.Length }))); + } + } + Add("energy_pred", result.energy); + Add("breathiness_pred", result.breathiness); + Add("voicing_pred", result.voicing); + Add("tension_pred", result.tension); + return outputs; + } + + static void ReplaceVarianceInputsWithPrevious( + List inputs, + VarianceResult previous) { + var channels = new[] { + ("energy", previous.energy), + ("breathiness", previous.breathiness), + ("voicing", previous.voicing), + ("tension", previous.tension), + }; + foreach (var (name, values) in channels) { + if (values == null) continue; + var input = inputs.FirstOrDefault(x => x.Name == name); + if (input == null) continue; + var current = input.AsTensor().ToArray(); + if (current.Length != values.Length) continue; + Array.Copy(values, current, values.Length); + inputs[inputs.IndexOf(input)] = NamedOnnxValue.CreateFromTensor( + name, + new DenseTensor(current, new[] { current.Length }, false) + .Reshape(new[] { 1, current.Length })); } } diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs index bd74de2de..eeb8fc544 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs @@ -3,30 +3,19 @@ using System.Linq; namespace OpenUtau.Core.DiffSinger { - internal readonly struct VariancePatchRange { - public readonly int start; - public readonly int end; - - public VariancePatchRange(int start, int end) { - this.start = start; - this.end = end; - } - } - - internal class VariancePatchState { + internal sealed class VariancePatchState { public readonly float[] pitch; + public readonly float[]? speakerEmbed; public readonly VarianceResult result; - public VariancePatchState(float[] pitch, VarianceResult result) { + public VariancePatchState(float[] pitch, float[]? speakerEmbed, VarianceResult result) { this.pitch = pitch.ToArray(); + this.speakerEmbed = speakerEmbed?.ToArray(); this.result = DiffSingerVariancePatch.CloneResult(result); } } internal static class DiffSingerVariancePatch { - const float PitchEpsilon = 1e-4f; - const float CrossfadeMs = 50f; - public static ulong BuildStateKey(ulong baseHash, int phrasePosition, int phraseEnd) { unchecked { ulong hash = baseHash; @@ -36,90 +25,145 @@ public static ulong BuildStateKey(ulong baseHash, int phrasePosition, int phrase } } - public static VarianceResult Merge( - VariancePatchState? previous, - float[] currentPitch, - VarianceResult current) { - if (previous == null || - previous.pitch.Length != currentPitch.Length || - !IsMetadataCompatible(previous.result, current)) { - return CloneResult(current); - } - var ranges = FindChangedRanges(previous.pitch, currentPitch, PitchEpsilon); - if (ranges.Count == 0) { - return CloneResult(previous.result); + internal static bool[] BuildChangedFrameMask( + IReadOnlyList previous, + IReadOnlyList current, + float epsilon) { + int length = Math.Max(previous.Count, current.Count); + var mask = new bool[length]; + for (int i = 0; i < length; i++) { + mask[i] = i >= previous.Count || i >= current.Count || + Math.Abs(previous[i] - current[i]) > epsilon; } - int crossfadeFrames = Math.Clamp((int)Math.Round(CrossfadeMs / current.frameMs), 1, 20); - var weights = BuildWeights(currentPitch.Length, ranges, crossfadeFrames); - return new VarianceResult { - energy = Blend(previous.result.energy, current.energy, weights), - breathiness = Blend(previous.result.breathiness, current.breathiness, weights), - voicing = Blend(previous.result.voicing, current.voicing, weights), - tension = Blend(previous.result.tension, current.tension, weights), - frameMs = current.frameMs, - headFrames = current.headFrames, - tailFrames = current.tailFrames, - totalFrames = current.totalFrames, - }; + return mask; } - internal static List FindChangedRanges( - IReadOnlyList previousPitch, - IReadOnlyList currentPitch, + internal static bool[] BuildChangedFrameMask( + IReadOnlyList previous, + IReadOnlyList current, + int frameCount, float epsilon) { - var ranges = new List(); - int length = Math.Min(previousPitch.Count, currentPitch.Count); - int start = -1; - for (int i = 0; i < length; ++i) { - bool changed = Math.Abs(previousPitch[i] - currentPitch[i]) > epsilon; - if (changed && start < 0) { - start = i; - } else if (!changed && start >= 0) { - ranges.Add(new VariancePatchRange(start, i)); - start = -1; - } + if (frameCount <= 0) { + return Array.Empty(); + } + if (previous.Count != current.Count || previous.Count % frameCount != 0) { + return Enumerable.Repeat(true, frameCount).ToArray(); } - if (start >= 0) { - ranges.Add(new VariancePatchRange(start, length)); + int valuesPerFrame = previous.Count / frameCount; + var mask = new bool[frameCount]; + for (int frame = 0; frame < frameCount; frame++) { + int offset = frame * valuesPerFrame; + for (int i = 0; i < valuesPerFrame; i++) { + if (Math.Abs(previous[offset + i] - current[offset + i]) > epsilon) { + mask[frame] = true; + break; + } + } } - return ranges; + return mask; } - internal static float[] BuildWeights(int length, IReadOnlyList ranges, int crossfadeFrames) { - var weights = new float[length]; - foreach (var range in ranges) { - int start = Math.Clamp(range.start, 0, length); - int end = Math.Clamp(range.end, start, length); - for (int i = start; i < end; ++i) { - weights[i] = 1f; - } - int leftStart = Math.Max(0, start - crossfadeFrames); - int leftLength = start - leftStart; - for (int i = leftStart; i < start; ++i) { - float weight = (float)(i - leftStart + 1) / (leftLength + 1); - weights[i] = Math.Max(weights[i], weight); - } - int rightEnd = Math.Min(length, end + crossfadeFrames); - int rightLength = rightEnd - end; - for (int i = end; i < rightEnd; ++i) { - float weight = 1f - (float)(i - end + 1) / (rightLength + 1); - weights[i] = Math.Max(weights[i], weight); + internal static bool[] ExpandToChannels( + IReadOnlyList frameMask, + int channelCount) { + if (channelCount < 0) { + throw new ArgumentOutOfRangeException(nameof(channelCount)); + } + var mask = new bool[frameMask.Count * channelCount]; + for (int frame = 0; frame < frameMask.Count; frame++) { + if (!frameMask[frame]) continue; + for (int channel = 0; channel < channelCount; channel++) { + mask[frame * channelCount + channel] = true; } } - return weights; + return mask; } - internal static float[]? Blend(float[]? previous, float[]? current, IReadOnlyList weights) { - if (previous == null || current == null || previous.Length != current.Length || previous.Length != weights.Count) { - return current?.ToArray(); + internal static VarianceResult HardCompose( + VarianceResult previous, + VarianceResult predicted, + IReadOnlyList retakeMask, + int channelCount) { + if (!IsCompatible(previous, predicted) || + retakeMask.Count != previous.totalFrames * channelCount) { + return CloneResult(predicted); } - var result = new float[current.Length]; - for (int i = 0; i < result.Length; ++i) { - result[i] = previous[i] * (1f - weights[i]) + current[i] * weights[i]; + int channel = 0; + var energy = ComposeEnabledChannel(previous.energy, predicted.energy, retakeMask, previous.totalFrames, ref channel, channelCount); + var breathiness = ComposeEnabledChannel(previous.breathiness, predicted.breathiness, retakeMask, previous.totalFrames, ref channel, channelCount); + var voicing = ComposeEnabledChannel(previous.voicing, predicted.voicing, retakeMask, previous.totalFrames, ref channel, channelCount); + var tension = ComposeEnabledChannel(previous.tension, predicted.tension, retakeMask, previous.totalFrames, ref channel, channelCount); + return new VarianceResult { + energy = energy, + breathiness = breathiness, + voicing = voicing, + tension = tension, + frameMs = predicted.frameMs, + headFrames = predicted.headFrames, + tailFrames = predicted.tailFrames, + totalFrames = predicted.totalFrames, + }; + } + + static float[]? ComposeEnabledChannel( + float[]? previous, + float[]? predicted, + IReadOnlyList mask, + int frameCount, + ref int channel, + int channelCount) { + if (previous == null && predicted == null) { + return null; + } + int currentChannel = channel++; + return ComposeChannel(previous, predicted, mask, frameCount, currentChannel, channelCount); + } + + static float[]? ComposeChannel( + float[]? previous, + float[]? predicted, + IReadOnlyList mask, + int frameCount, + int channel, + int channelCount) { + if (previous == null || predicted == null) { + return predicted?.ToArray(); + } + if (previous.Length != frameCount || predicted.Length != frameCount) { + return predicted.ToArray(); + } + var result = previous.ToArray(); + for (int frame = 0; frame < frameCount; frame++) { + if (mask[frame * channelCount + channel]) { + result[frame] = predicted[frame]; + } } return result; } + internal static bool IsMetadataCompatible(VarianceResult previous, VarianceResult current) { + return previous.totalFrames == current.totalFrames && + previous.headFrames == current.headFrames && + previous.tailFrames == current.tailFrames && + Math.Abs(previous.frameMs - current.frameMs) < 1e-4f; + } + + internal static bool IsCompatible(VarianceResult previous, VarianceResult current) { + return IsMetadataCompatible(previous, current) && + previous.totalFrames == current.totalFrames && + previous.headFrames == current.headFrames && + previous.tailFrames == current.tailFrames && + Math.Abs(previous.frameMs - current.frameMs) < 1e-4f && + SameLength(previous.energy, current.energy) && + SameLength(previous.breathiness, current.breathiness) && + SameLength(previous.voicing, current.voicing) && + SameLength(previous.tension, current.tension); + } + + static bool SameLength(float[]? a, float[]? b) { + return (a == null) == (b == null) && (a == null || a.Length == b!.Length); + } + internal static VarianceResult CloneResult(VarianceResult result) { return new VarianceResult { energy = result.energy?.ToArray(), @@ -132,12 +176,5 @@ internal static VarianceResult CloneResult(VarianceResult result) { totalFrames = result.totalFrames, }; } - - static bool IsMetadataCompatible(VarianceResult previous, VarianceResult current) { - return previous.totalFrames == current.totalFrames && - previous.headFrames == current.headFrames && - previous.tailFrames == current.tailFrames && - Math.Abs(previous.frameMs - current.frameMs) < 1e-4f; - } } } diff --git a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs index 7b05b0b21..44a700a5f 100644 --- a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs +++ b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs @@ -1,64 +1,105 @@ -using System.Linq; using OpenUtau.Core.DiffSinger; using Xunit; namespace OpenUtau.Core { public class DiffSingerVariancePatchTest { [Fact] - public void FindChangedRangesGroupsContiguousPitchChanges() { - var previous = new[] { 1f, 1f, 1f, 1f, 1f, 1f }; - var current = new[] { 1f, 2f, 2f, 1f, 2f, 1f }; + public void BuildChangedFrameMaskMarksOnlyChangedFrames() { + var mask = DiffSingerVariancePatch.BuildChangedFrameMask( + new[] { 1f, 1f, 1f, 2f }, + new[] { 1f, 2f, 1f, 2f }, + 1e-4f); - var ranges = DiffSingerVariancePatch.FindChangedRanges(previous, current, 1e-4f); + Assert.Equal(new[] { false, true, false, false }, mask); + } + + [Fact] + public void BuildChangedFrameMaskGroupsSpeakerEmbeddingByFrame() { + var mask = DiffSingerVariancePatch.BuildChangedFrameMask( + new[] { 1f, 2f, 3f, 4f, 5f, 6f }, + new[] { 1f, 2f, 3f, 40f, 5f, 6f }, + 3, + 1e-4f); - Assert.Equal(2, ranges.Count); - Assert.Equal(1, ranges[0].start); - Assert.Equal(3, ranges[0].end); - Assert.Equal(4, ranges[1].start); - Assert.Equal(5, ranges[1].end); + Assert.Equal(new[] { false, true, false }, mask); } [Fact] - public void MergeKeepsPreviousResultWhenPitchDoesNotChange() { - var previousResult = Result(new[] { 1f, 2f, 3f }); - var currentResult = Result(new[] { 10f, 20f, 30f }); - var previous = new VariancePatchState(new[] { 60f, 61f, 62f }, previousResult); + public void BuildChangedFrameMaskMarksAllFramesForIncompatibleEmbeddingShape() { + var mask = DiffSingerVariancePatch.BuildChangedFrameMask( + new[] { 1f, 2f, 3f, 4f }, + new[] { 1f, 2f, 3f }, + 2, + 1e-4f); - var merged = DiffSingerVariancePatch.Merge(previous, new[] { 60f, 61f, 62f }, currentResult); + Assert.Equal(new[] { true, true }, mask); + } - Assert.Equal(previousResult.energy!, merged.energy!); + [Fact] + public void ExpandToChannelsUsesSharedFrameMask() { + var mask = DiffSingerVariancePatch.ExpandToChannels( + new[] { false, true, false }, 3); + + Assert.Equal( + new[] { false, false, false, true, true, true, false, false, false }, + mask); } [Fact] - public void MergeBlendsOnlyChangedPitchRange() { - var previousResult = Result(Enumerable.Repeat(0f, 6).ToArray(), frameMs: 50); - var currentResult = Result(Enumerable.Repeat(10f, 6).ToArray(), frameMs: 50); - var previous = new VariancePatchState( - new[] { 60f, 60f, 60f, 60f, 60f, 60f }, - previousResult); - - var merged = DiffSingerVariancePatch.Merge( - previous, - new[] { 60f, 60f, 61f, 61f, 60f, 60f }, - currentResult); - - Assert.Equal(new[] { 0f, 5f, 10f, 10f, 5f, 0f }, merged.energy!); + public void HardComposePreservesUnmaskedFramesExactly() { + var previous = Result( + new[] { 1f, 2f, 3f, 4f }, + new[] { 5f, 6f, 7f, 8f }); + var predicted = Result( + new[] { 10f, 20f, 30f, 40f }, + new[] { 50f, 60f, 70f, 80f }); + var mask = DiffSingerVariancePatch.ExpandToChannels( + new[] { false, true, false, true }, 2); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 2); + + Assert.Equal(new[] { 1f, 20f, 3f, 40f }, result.energy); + Assert.Equal(new[] { 5f, 60f, 7f, 80f }, result.breathiness); } [Fact] - public void MergeFallsBackToCurrentResultWhenMetadataChanges() { - var previousResult = Result(new[] { 1f, 2f, 3f }, frameMs: 50); - var currentResult = Result(new[] { 10f, 20f, 30f }, frameMs: 60); - var previous = new VariancePatchState(new[] { 60f, 61f, 62f }, previousResult); + public void HardComposeDoesNotLeakModelChangesOutsideMask() { + var previous = Result(new[] { 1f, 2f, 3f }); + var predicted = Result(new[] { 100f, 200f, 300f }); + var mask = DiffSingerVariancePatch.ExpandToChannels( + new[] { false, true, false }, 1); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 1); - var merged = DiffSingerVariancePatch.Merge(previous, new[] { 60f, 62f, 62f }, currentResult); + Assert.Equal(new[] { 1f, 200f, 3f }, result.energy); + } + + [Fact] + public void HardComposeFallsBackToPredictedForIncompatibleMetadata() { + var previous = Result(new[] { 1f, 2f, 3f }, frameMs: 50); + var predicted = Result(new[] { 10f, 20f, 30f }, frameMs: 60); + var mask = new[] { true, false, true }; + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 1); + + Assert.Equal(predicted.energy, result.energy); + } + + [Fact] + public void IsMetadataCompatibleRejectsFrameLayoutChanges() { + var previous = Result(new[] { 1f, 2f, 3f }); + var changed = Result(new[] { 1f, 2f, 3f, 4f }); - Assert.Equal(currentResult.energy!, merged.energy!); + Assert.False(DiffSingerVariancePatch.IsMetadataCompatible(previous, changed)); } - static VarianceResult Result(float[] energy, float frameMs = 50) { + static VarianceResult Result( + float[] energy, + float[]? breathiness = null, + float frameMs = 50) { return new VarianceResult { energy = energy, + breathiness = breathiness, frameMs = frameMs, headFrames = 1, tailFrames = 1, From c9238a8d866b8dd44c10361da55ba786d920df1f Mon Sep 17 00:00:00 2001 From: KakaruHayate Date: Tue, 4 Aug 2026 17:26:20 +0800 Subject: [PATCH 2/3] Fix variance retake review findings --- .../DiffSinger/DiffSingerVariance.cs | 36 +++++--- .../DiffSinger/DiffSingerVariancePatch.cs | 43 +++++++++ .../DiffSinger/DiffSingerVariancePatchTest.cs | 87 +++++++++++++++++++ 3 files changed, 153 insertions(+), 13 deletions(-) diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs index 403aec863..d8ca15374 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs @@ -35,7 +35,9 @@ public class DsVariance : IDisposable{ IG2p g2p; float frameMs; DiffSingerSpeakerEmbedManager speakerEmbedManager; - readonly Dictionary variancePatchStates = new Dictionary(); + const int VariancePatchStateCapacity = 16; + readonly VariancePatchStateCache variancePatchStates = + new VariancePatchStateCache(VariancePatchStateCapacity); public float FrameMs => frameMs; @@ -273,15 +275,22 @@ void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { var baseHash = new DiffSingerCache(varianceHash, variancePatchInputs).Hash; variancePatchKey = DiffSingerVariancePatch.BuildStateKey(baseHash, phrase.position, phrase.end); } - var fullVarianceCache = Preferences.Default.DiffSingerTensorCache - ? new DiffSingerCache(varianceHash, varianceInputs) + // Cache the final pipeline result in a separate namespace from raw predictor outputs. + var resultCacheInputs = new List(varianceInputs) { + NamedOnnxValue.CreateFromTensor( + "result_cache_version", + new DenseTensor(new long[] { 1 }, new int[] { 1 }, false)), + }; + var resultCache = Preferences.Default.DiffSingerTensorCache + ? new DiffSingerCache(varianceHash, resultCacheInputs) : null; - var fullVarianceOutputs = fullVarianceCache?.Load(); - if (fullVarianceOutputs != null) { - var cachedResult = ParseVarianceResult(fullVarianceOutputs, frameMs, headFrames, tailFrames, totalFrames); + var cachedOutputs = resultCache?.Load(); + if (cachedOutputs != null) { + var cachedResult = ParseVarianceResult(cachedOutputs, frameMs, headFrames, tailFrames, totalFrames); if (variancePatchKey.HasValue) { - variancePatchStates[variancePatchKey.Value] = - new VariancePatchState(pitch, speakerEmbed, cachedResult); + variancePatchStates.Set( + variancePatchKey.Value, + new VariancePatchState(pitch, speakerEmbed, cachedResult)); } return cachedResult; } @@ -363,13 +372,14 @@ void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { var channelMask = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); result = DiffSingerVariancePatch.HardCompose(previous.result, result, channelMask, numVariances); } - if (fullVarianceCache != null) { - fullVarianceCache.Save(BuildVarianceOutputs(result)); - phrase.AddCacheFile(fullVarianceCache.Filename); + if (resultCache != null) { + resultCache.Save(BuildVarianceOutputs(result)); + phrase.AddCacheFile(resultCache.Filename); } if (variancePatchKey.HasValue) { - variancePatchStates[variancePatchKey.Value] = - new VariancePatchState(pitch, speakerEmbed, result); + variancePatchStates.Set( + variancePatchKey.Value, + new VariancePatchState(pitch, speakerEmbed, result)); } return result; } diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs index eeb8fc544..c92d1d923 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs @@ -15,6 +15,49 @@ public VariancePatchState(float[] pitch, float[]? speakerEmbed, VarianceResult r } } + internal sealed class VariancePatchStateCache { + readonly int capacity; + readonly Dictionary> entries = new(); + readonly LinkedList<(ulong key, VariancePatchState state)> recency = new(); + + internal VariancePatchStateCache(int capacity) { + if (capacity <= 0) { + throw new ArgumentOutOfRangeException(nameof(capacity)); + } + this.capacity = capacity; + } + + internal int Count => entries.Count; + + internal bool TryGetValue(ulong key, out VariancePatchState state) { + if (!entries.TryGetValue(key, out var node)) { + state = null!; + return false; + } + recency.Remove(node); + recency.AddFirst(node); + state = node.Value.state; + return true; + } + + internal void Set(ulong key, VariancePatchState state) { + if (entries.TryGetValue(key, out var existing)) { + existing.Value = (key, state); + recency.Remove(existing); + recency.AddFirst(existing); + return; + } + var node = recency.AddFirst((key, state)); + entries.Add(key, node); + if (entries.Count <= capacity) { + return; + } + var oldest = recency.Last!; + recency.RemoveLast(); + entries.Remove(oldest.Value.key); + } + } + internal static class DiffSingerVariancePatch { public static ulong BuildStateKey(ulong baseHash, int phrasePosition, int phraseEnd) { unchecked { diff --git a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs index 44a700a5f..e4b171368 100644 --- a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs +++ b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs @@ -74,6 +74,67 @@ public void HardComposeDoesNotLeakModelChangesOutsideMask() { Assert.Equal(new[] { 1f, 200f, 3f }, result.energy); } + [Fact] + public void HardComposeHandlesNullMiddleChannel() { + var previous = Result( + new[] { 1f, 2f }, + voicing: new[] { 3f, 4f }); + var predicted = Result( + new[] { 10f, 20f }, + voicing: new[] { 30f, 40f }); + var mask = new[] { false, false, true, true }; + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 2); + + Assert.Equal(new[] { 1f, 20f }, result.energy); + Assert.Null(result.breathiness); + Assert.Equal(new[] { 3f, 40f }, result.voicing); + } + + [Fact] + public void HardComposeHandlesAllChannels() { + var previous = Result( + new[] { 1f, 2f }, + new[] { 3f, 4f }, + new[] { 5f, 6f }, + new[] { 7f, 8f }); + var predicted = Result( + new[] { 10f, 20f }, + new[] { 30f, 40f }, + new[] { 50f, 60f }, + new[] { 70f, 80f }); + var mask = DiffSingerVariancePatch.ExpandToChannels(new[] { false, true }, 4); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 4); + + Assert.Equal(new[] { 1f, 20f }, result.energy); + Assert.Equal(new[] { 3f, 40f }, result.breathiness); + Assert.Equal(new[] { 5f, 60f }, result.voicing); + Assert.Equal(new[] { 7f, 80f }, result.tension); + } + + [Fact] + public void HardComposePreservesAllPreviousChannelsForFalseMask() { + var previous = Result( + new[] { 1f, 2f }, + new[] { 3f, 4f }, + new[] { 5f, 6f }, + new[] { 7f, 8f }); + var predicted = Result( + new[] { 10f, 20f }, + new[] { 30f, 40f }, + new[] { 50f, 60f }, + new[] { 70f, 80f }); + var mask = DiffSingerVariancePatch.ExpandToChannels(new[] { false, false }, 4); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 4); + + Assert.Equal(previous.energy, result.energy); + Assert.Equal(previous.breathiness, result.breathiness); + Assert.Equal(previous.voicing, result.voicing); + Assert.Equal(previous.tension, result.tension); + } + [Fact] public void HardComposeFallsBackToPredictedForIncompatibleMetadata() { var previous = Result(new[] { 1f, 2f, 3f }, frameMs: 50); @@ -85,6 +146,21 @@ public void HardComposeFallsBackToPredictedForIncompatibleMetadata() { Assert.Equal(predicted.energy, result.energy); } + [Fact] + public void VariancePatchStateCacheEvictsLeastRecentlyUsedState() { + var cache = new VariancePatchStateCache(2); + cache.Set(1, State(1)); + cache.Set(2, State(2)); + Assert.True(cache.TryGetValue(1, out _)); + + cache.Set(3, State(3)); + + Assert.Equal(2, cache.Count); + Assert.True(cache.TryGetValue(1, out _)); + Assert.False(cache.TryGetValue(2, out _)); + Assert.True(cache.TryGetValue(3, out _)); + } + [Fact] public void IsMetadataCompatibleRejectsFrameLayoutChanges() { var previous = Result(new[] { 1f, 2f, 3f }); @@ -93,13 +169,24 @@ public void IsMetadataCompatibleRejectsFrameLayoutChanges() { Assert.False(DiffSingerVariancePatch.IsMetadataCompatible(previous, changed)); } + static VariancePatchState State(float value) { + return new VariancePatchState( + new[] { value }, + null, + Result(new[] { value })); + } + static VarianceResult Result( float[] energy, float[]? breathiness = null, + float[]? voicing = null, + float[]? tension = null, float frameMs = 50) { return new VarianceResult { energy = energy, breathiness = breathiness, + voicing = voicing, + tension = tension, frameMs = frameMs, headFrames = 1, tailFrames = 1, From 87e5b4b09dcbabb461e2ac3bf09673572eac0177 Mon Sep 17 00:00:00 2001 From: Kakaru <97896816+KakaruHayate@users.noreply.github.com> Date: Tue, 4 Aug 2026 18:25:43 +0800 Subject: [PATCH 3/3] Validate variance retake channel layout --- .../DiffSinger/DiffSingerVariance.cs | 931 +++++++++--------- .../DiffSinger/DiffSingerVariancePatch.cs | 459 ++++----- .../DiffSinger/DiffSingerVariancePatchTest.cs | 432 ++++---- 3 files changed, 940 insertions(+), 882 deletions(-) diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs index d8ca15374..5cd8e7f21 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs @@ -1,462 +1,469 @@ -using System; -using System.Collections.Generic; -using System.IO; -using System.Linq; -using System.Text; -using K4os.Hash.xxHash; -using Serilog; -using Microsoft.ML.OnnxRuntime; -using Microsoft.ML.OnnxRuntime.Tensors; - -using OpenUtau.Api; -using OpenUtau.Core.Render; -using OpenUtau.Core.Util; - -namespace OpenUtau.Core.DiffSinger{ - public struct VarianceResult{ - public float[]? energy; - public float[]? breathiness; - public float[]? voicing; - public float[]? tension; - public float frameMs; - public int headFrames; - public int tailFrames; - public int totalFrames; - } - public class DsVariance : IDisposable{ - string rootPath; - DsConfig dsConfig; - Dictionary languageIds = new Dictionary(); - Dictionary phonemeTokens; - ulong linguisticHash; - ulong varianceHash; - InferenceSession linguisticModel; - InferenceSession varianceModel; - IG2p g2p; - float frameMs; - DiffSingerSpeakerEmbedManager speakerEmbedManager; - const int VariancePatchStateCapacity = 16; - readonly VariancePatchStateCache variancePatchStates = - new VariancePatchStateCache(VariancePatchStateCapacity); - - public float FrameMs => frameMs; - - public DsVariance(string rootPath) - { - this.rootPath = rootPath; - var dsconfigPath = Path.Combine(rootPath, "dsconfig.yaml"); - try { - dsConfig = Yaml.DefaultDeserializer.Deserialize( - File.ReadAllText(dsconfigPath, Encoding.UTF8)); - } catch (Exception e) { - throw new Exception($"Failed to load {dsconfigPath}", e); - } - if(dsConfig.variance == null){ - throw new Exception("This voicebank doesn't contain a variance model"); - } - //Load language id if needed - if(dsConfig.use_lang_id){ - if(dsConfig.languages == null){ - throw new Exception("\"languages\" field is not specified in dsconfig.yaml"); - } - var langIdPath = Path.Join(rootPath, dsConfig.languages); - try { - languageIds = DiffSingerUtils.LoadLanguageIds(langIdPath); - } catch (Exception e) { - Log.Error(e, $"failed to load language id from {langIdPath}"); - throw new Exception($"Failed to load {langIdPath}", e); - } - } - //Load phonemes list - if (dsConfig.phonemes == null) { - throw new Exception("Configuration key \"phonemes\" is null."); - } - string phonemesPath = Path.Combine(rootPath, dsConfig.phonemes); - phonemeTokens = DiffSingerUtils.LoadPhonemes(phonemesPath); - //Load models - if (dsConfig.linguistic == null) { - throw new Exception("Configuration key \"linguistic\" is null."); - } - var linguisticModelPath = Path.Join(rootPath, dsConfig.linguistic); - var linguisticModelBytes = File.ReadAllBytes(linguisticModelPath); - linguisticHash = XXH64.DigestOf(linguisticModelBytes); - linguisticModel = Onnx.getInferenceSession(linguisticModelBytes); - var varianceModelPath = Path.Join(rootPath, dsConfig.variance); - var varianceModelBytes = File.ReadAllBytes(varianceModelPath); - varianceHash = XXH64.DigestOf(varianceModelBytes); - varianceModel = Onnx.getInferenceSession(varianceModelBytes); - frameMs = 1000f * dsConfig.hop_size / dsConfig.sample_rate; - //Load g2p - g2p = LoadG2p(rootPath); - } - - protected IG2p LoadG2p(string rootPath) { - // Load dictionary from singer folder. - string file = Path.Combine(rootPath, "dsdict.yaml"); - if(!File.Exists(file)){ - throw new Exception($"File not found: {file}"); - } - try { - var g2pBuilder = G2pDictionary.NewBuilder().Load(File.ReadAllText(file)); - //SP and AP should always be vowel - g2pBuilder.AddSymbol("SP", true); - g2pBuilder.AddSymbol("AP", true); - return g2pBuilder.Build(); - } catch (Exception e) { - throw new Exception($"Failed to load {file}", e); - } - } - - public DiffSingerSpeakerEmbedManager getSpeakerEmbedManager(){ - if(speakerEmbedManager is null) { - speakerEmbedManager = new DiffSingerSpeakerEmbedManager(dsConfig, rootPath); - } - return speakerEmbedManager; - } - - int PhonemeTokenize(string phoneme){ - bool success = phonemeTokens.TryGetValue(phoneme, out int token); - if(!success){ - throw new Exception($"Phoneme \"{phoneme}\" isn't supported by variance model. Please check {Path.Combine(rootPath, dsConfig.phonemes)}"); - } - return token; - } - - public VarianceResult Process(RenderPhrase phrase){ - int headFrames = DiffSingerUtils.headFrames; - int tailFrames = DiffSingerUtils.tailFrames; - if (dsConfig.predict_dur) { - //Check if all phonemes are defined in dsdict.yaml (for their types) - foreach (var phone in phrase.phones) { - if (!g2p.IsValidSymbol(phone.phoneme)) { - throw new InvalidDataException( - $"Type definition of symbol \"{phone.phoneme}\" not found. Consider adding it to dsdict.yaml of the variance predictor."); - } - } - } - //Linguistic Encoder - var linguisticInputs = new List(); - var tokens = phrase.phones.Select(p => p.phoneme) - .Prepend("SP") - .Append("SP") - .Select(x => (Int64)PhonemeTokenize(x)) - .ToArray(); - var ph_dur = DiffSingerUtils.PaddedPhoneDurations(phrase, frameMs, headFrames, tailFrames); - int totalFrames = ph_dur.Sum(); - linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("tokens", - new DenseTensor(tokens, new int[] { tokens.Length }, false) - .Reshape(new int[] { 1, tokens.Length }))); - if(dsConfig.predict_dur){ - //if predict_dur is true, use word encode mode - var (word_div, word_dur) = DiffSingerUtils.PaddedWordDivAndDur(phrase, ph_dur, g2p.IsVowel); - linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("word_div", - new DenseTensor(word_div, new int[] { word_div.Length }, false) - .Reshape(new int[] { 1, word_div.Length }))); - linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("word_dur", - new DenseTensor(word_dur, new int[] { word_dur.Length }, false) - .Reshape(new int[] { 1, word_dur.Length }))); - }else{ - //if predict_dur is false, use phoneme encode mode - linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("ph_dur", - new DenseTensor(ph_dur.Select(x=>(Int64)x).ToArray(), new int[] { ph_dur.Length }, false) - .Reshape(new int[] { 1, ph_dur.Length }))); - } - //Language id - if(dsConfig.use_lang_id){ - var langIdByPhone = phrase.phones - .Select(p => (long)languageIds.GetValueOrDefault( - DiffSingerUtils.PhonemeLanguage(p.phoneme),0 - )) - .Prepend(0) - .Append(0) - .ToArray(); - var langIdTensor = new DenseTensor(langIdByPhone, new int[] { langIdByPhone.Length }, false) - .Reshape(new int[] { 1, langIdByPhone.Length }); - linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("languages", langIdTensor)); - } - - Onnx.VerifyInputNames(linguisticModel, linguisticInputs); - var linguisticCache = Preferences.Default.DiffSingerTensorCache - ? new DiffSingerCache(linguisticHash, linguisticInputs) - : null; - var linguisticOutputs = linguisticCache?.Load(); - if (linguisticOutputs is null) { - linguisticOutputs = linguisticModel.Run(linguisticInputs).Cast().ToList(); - linguisticCache?.Save(linguisticOutputs); - phrase.AddCacheFile(linguisticCache?.Filename); - } - Tensor encoder_out = linguisticOutputs - .Where(o => o.Name == "encoder_out") - .First() - .AsTensor(); - - //Variance Predictor - var pitch = DiffSingerUtils.SampleCurve(phrase, phrase.pitches, 0, frameMs, totalFrames, headFrames, tailFrames, - x => x * 0.01).Select(f => (float)f).ToArray(); - var toneShift = DiffSingerUtils.SampleCurve(phrase, phrase.toneShift, 0, frameMs, totalFrames, headFrames, tailFrames, - x => x * 0.01).Select(f => (float)f).ToArray(); - pitch = pitch.Zip(toneShift, (x, d) => x + d).ToArray(); - - var varianceInputs = new List(); - var variancePatchInputs = new List(); - void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { - varianceInputs.Add(input); - if (includeInPatchKey) { - variancePatchInputs.Add(input); - } - } - AddVarianceInput(NamedOnnxValue.CreateFromTensor("encoder_out", encoder_out)); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("ph_dur", - new DenseTensor(ph_dur.Select(x=>(Int64)x).ToArray(), new int[] { ph_dur.Length }, false) - .Reshape(new int[] { 1, ph_dur.Length }))); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("pitch", - new DenseTensor(pitch, new int[] { pitch.Length }, false) - .Reshape(new int[] { 1, totalFrames })), includeInPatchKey: false); - if (dsConfig.predict_energy) { - var energy = Enumerable.Repeat(0f, totalFrames).ToArray(); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("energy", - new DenseTensor(energy, new int[] { energy.Length }, false) - .Reshape(new int[] { 1, totalFrames }))); - } - if (dsConfig.predict_breathiness) { - var breathiness = Enumerable.Repeat(0f, totalFrames).ToArray(); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("breathiness", - new DenseTensor(breathiness, new int[] { breathiness.Length }, false) - .Reshape(new int[] { 1, totalFrames }))); - } - if (dsConfig.predict_voicing) { - var voicing = Enumerable.Repeat(0f, totalFrames).ToArray(); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("voicing", - new DenseTensor(voicing, new int[] { voicing.Length }, false) - .Reshape(new int[] { 1, totalFrames }))); - } - if (dsConfig.predict_tension) { - var tension = Enumerable.Repeat(0f, totalFrames).ToArray(); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("tension", - new DenseTensor(tension, new int[] { tension.Length }, false) - .Reshape(new int[] { 1, totalFrames }))); - } - - var numVariances = new[] { - dsConfig.predict_energy, - dsConfig.predict_breathiness, - dsConfig.predict_voicing, - dsConfig.predict_tension, - }.Sum(Convert.ToInt32); - var retake = Enumerable.Repeat(true, totalFrames * numVariances).ToArray(); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("retake", - new DenseTensor(retake, new int[] { retake.Length }, false) - .Reshape(new int[] { 1, totalFrames, numVariances }))); - var steps = Preferences.Default.DiffSingerStepsVariance; - if (dsConfig.useContinuousAcceleration) { - AddVarianceInput(NamedOnnxValue.CreateFromTensor("steps", - new DenseTensor(new long[] { steps }, new int[] { 1 }, false))); - } else { - // find a largest integer speedup that are less than 1000 / steps and is a factor of 1000 - long speedup = Math.Max(1, 1000 / steps); - while (1000 % speedup != 0 && speedup > 1) { - speedup--; - } - AddVarianceInput(NamedOnnxValue.CreateFromTensor("speedup", - new DenseTensor(new long[] { speedup }, new int[] { 1 },false))); - } - //Speaker - float[]? speakerEmbed = null; - if(dsConfig.speakers != null) { - var speakerEmbedManager = getSpeakerEmbedManager(); - var spkEmbedTensor = speakerEmbedManager.PhraseSpeakerEmbedByFrame(phrase, ph_dur, frameMs, totalFrames, headFrames, tailFrames); - speakerEmbed = spkEmbedTensor.ToArray(); - // Speaker embedding is a retake-able frame-level condition. - AddVarianceInput(NamedOnnxValue.CreateFromTensor("spk_embed", spkEmbedTensor), includeInPatchKey: false); - } - ulong? variancePatchKey = null; - if (Preferences.Default.DiffSingerTensorCache && - Preferences.Default.DiffSingerVarianceLocalPitchPatch) { - var baseHash = new DiffSingerCache(varianceHash, variancePatchInputs).Hash; - variancePatchKey = DiffSingerVariancePatch.BuildStateKey(baseHash, phrase.position, phrase.end); - } - // Cache the final pipeline result in a separate namespace from raw predictor outputs. - var resultCacheInputs = new List(varianceInputs) { - NamedOnnxValue.CreateFromTensor( - "result_cache_version", - new DenseTensor(new long[] { 1 }, new int[] { 1 }, false)), - }; - var resultCache = Preferences.Default.DiffSingerTensorCache - ? new DiffSingerCache(varianceHash, resultCacheInputs) - : null; - var cachedOutputs = resultCache?.Load(); - if (cachedOutputs != null) { - var cachedResult = ParseVarianceResult(cachedOutputs, frameMs, headFrames, tailFrames, totalFrames); - if (variancePatchKey.HasValue) { - variancePatchStates.Set( - variancePatchKey.Value, - new VariancePatchState(pitch, speakerEmbed, cachedResult)); - } - return cachedResult; - } - VariancePatchState? previous = null; - bool[]? retakeMask = null; - if (variancePatchKey.HasValue && variancePatchStates.TryGetValue(variancePatchKey.Value, out var cachedState) && - DiffSingerVariancePatch.IsMetadataCompatible(cachedState.result, new VarianceResult { - frameMs = frameMs, - headFrames = headFrames, - tailFrames = tailFrames, - totalFrames = totalFrames, - })) { - previous = cachedState; - var pitchMask = DiffSingerVariancePatch.BuildChangedFrameMask(cachedState.pitch, pitch, 1e-4f); - var speakerMask = DiffSingerVariancePatch.BuildChangedFrameMask( - cachedState.speakerEmbed ?? Array.Empty(), - speakerEmbed ?? Array.Empty(), - totalFrames, - 1e-4f); - retakeMask = new bool[totalFrames]; - for (int i = 0; i < retakeMask.Length; i++) { - retakeMask[i] = (i < pitchMask.Length && pitchMask[i]) || - (i < speakerMask.Length && speakerMask[i]); - } - if (!retakeMask.Any(x => x)) { - return DiffSingerVariancePatch.CloneResult(cachedState.result); - } - if (retakeMask.All(x => x)) { - previous = null; - } else { - ReplaceVarianceInputsWithPrevious(varianceInputs, cachedState.result); - } - } - if (retakeMask != null) { - var retakeTensorValues = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); - var retakeInput = varianceInputs.First(x => x.Name == "retake"); - varianceInputs[varianceInputs.IndexOf(retakeInput)] = NamedOnnxValue.CreateFromTensor( - "retake", - new DenseTensor(retakeTensorValues, new[] { retakeTensorValues.Length }, false) - .Reshape(new[] { 1, totalFrames, numVariances })); - } - Onnx.VerifyInputNames(varianceModel, varianceInputs); - var varianceOutputs = varianceModel.Run(varianceInputs).Cast().ToList(); - Tensor? energy_pred = dsConfig.predict_energy - ? varianceOutputs - .Where(o => o.Name == "energy_pred") - .First() - .AsTensor() - : null; - Tensor? breathiness_pred = dsConfig.predict_breathiness - ? varianceOutputs - .Where(o => o.Name == "breathiness_pred") - .First() - .AsTensor() - : null; - Tensor? voicing_pred = dsConfig.predict_voicing - ? varianceOutputs - .Where(o => o.Name == "voicing_pred") - .First() - .AsTensor() - : null; - Tensor? tension_pred = dsConfig.predict_tension - ? varianceOutputs - .Where(o => o.Name == "tension_pred") - .First() - .AsTensor() - : null; - var result = new VarianceResult{ - energy = energy_pred?.ToArray(), - breathiness = breathiness_pred?.ToArray(), - voicing = voicing_pred?.ToArray(), - tension = tension_pred?.ToArray(), - frameMs = frameMs, - headFrames = headFrames, - tailFrames = tailFrames, - totalFrames = totalFrames, - }; - if (previous != null && retakeMask != null) { - var channelMask = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); - result = DiffSingerVariancePatch.HardCompose(previous.result, result, channelMask, numVariances); - } - if (resultCache != null) { - resultCache.Save(BuildVarianceOutputs(result)); - phrase.AddCacheFile(resultCache.Filename); - } - if (variancePatchKey.HasValue) { - variancePatchStates.Set( - variancePatchKey.Value, - new VariancePatchState(pitch, speakerEmbed, result)); - } - return result; - } - - VarianceResult ParseVarianceResult( - ICollection outputs, - float frameMs, - int headFrames, - int tailFrames, - int totalFrames) { - return new VarianceResult { - energy = dsConfig.predict_energy ? outputs.First(o => o.Name == "energy_pred").AsTensor().ToArray() : null, - breathiness = dsConfig.predict_breathiness ? outputs.First(o => o.Name == "breathiness_pred").AsTensor().ToArray() : null, - voicing = dsConfig.predict_voicing ? outputs.First(o => o.Name == "voicing_pred").AsTensor().ToArray() : null, - tension = dsConfig.predict_tension ? outputs.First(o => o.Name == "tension_pred").AsTensor().ToArray() : null, - frameMs = frameMs, - headFrames = headFrames, - tailFrames = tailFrames, - totalFrames = totalFrames, - }; - } - - List BuildVarianceOutputs(VarianceResult result) { - var outputs = new List(); - void Add(string name, float[]? values) { - if (values != null) { - outputs.Add(NamedOnnxValue.CreateFromTensor( - name, - new DenseTensor(values, new[] { values.Length }, false) - .Reshape(new[] { 1, values.Length }))); - } - } - Add("energy_pred", result.energy); - Add("breathiness_pred", result.breathiness); - Add("voicing_pred", result.voicing); - Add("tension_pred", result.tension); - return outputs; - } - - static void ReplaceVarianceInputsWithPrevious( - List inputs, - VarianceResult previous) { - var channels = new[] { - ("energy", previous.energy), - ("breathiness", previous.breathiness), - ("voicing", previous.voicing), - ("tension", previous.tension), - }; - foreach (var (name, values) in channels) { - if (values == null) continue; - var input = inputs.FirstOrDefault(x => x.Name == name); - if (input == null) continue; - var current = input.AsTensor().ToArray(); - if (current.Length != values.Length) continue; - Array.Copy(values, current, values.Length); - inputs[inputs.IndexOf(input)] = NamedOnnxValue.CreateFromTensor( - name, - new DenseTensor(current, new[] { current.Length }, false) - .Reshape(new[] { 1, current.Length })); - } - } - - private bool disposedValue; - - protected virtual void Dispose(bool disposing) { - if (!disposedValue) { - if (disposing) { - linguisticModel?.Dispose(); - varianceModel?.Dispose(); - } - disposedValue = true; - } - } - - public void Dispose() { - Dispose(disposing: true); - GC.SuppressFinalize(this); - } - } -} +using System; +using System.Collections.Generic; +using System.IO; +using System.Linq; +using System.Text; +using K4os.Hash.xxHash; +using Serilog; +using Microsoft.ML.OnnxRuntime; +using Microsoft.ML.OnnxRuntime.Tensors; + +using OpenUtau.Api; +using OpenUtau.Core.Render; +using OpenUtau.Core.Util; + +namespace OpenUtau.Core.DiffSinger{ + public struct VarianceResult{ + public float[]? energy; + public float[]? breathiness; + public float[]? voicing; + public float[]? tension; + public float frameMs; + public int headFrames; + public int tailFrames; + public int totalFrames; + } + public class DsVariance : IDisposable{ + string rootPath; + DsConfig dsConfig; + Dictionary languageIds = new Dictionary(); + Dictionary phonemeTokens; + ulong linguisticHash; + ulong varianceHash; + InferenceSession linguisticModel; + InferenceSession varianceModel; + IG2p g2p; + float frameMs; + DiffSingerSpeakerEmbedManager speakerEmbedManager; + const int VariancePatchStateCapacity = 16; + readonly VariancePatchStateCache variancePatchStates = + new VariancePatchStateCache(VariancePatchStateCapacity); + + public float FrameMs => frameMs; + + public DsVariance(string rootPath) + { + this.rootPath = rootPath; + var dsconfigPath = Path.Combine(rootPath, "dsconfig.yaml"); + try { + dsConfig = Yaml.DefaultDeserializer.Deserialize( + File.ReadAllText(dsconfigPath, Encoding.UTF8)); + } catch (Exception e) { + throw new Exception($"Failed to load {dsconfigPath}", e); + } + if(dsConfig.variance == null){ + throw new Exception("This voicebank doesn't contain a variance model"); + } + //Load language id if needed + if(dsConfig.use_lang_id){ + if(dsConfig.languages == null){ + throw new Exception("\"languages\" field is not specified in dsconfig.yaml"); + } + var langIdPath = Path.Join(rootPath, dsConfig.languages); + try { + languageIds = DiffSingerUtils.LoadLanguageIds(langIdPath); + } catch (Exception e) { + Log.Error(e, $"failed to load language id from {langIdPath}"); + throw new Exception($"Failed to load {langIdPath}", e); + } + } + //Load phonemes list + if (dsConfig.phonemes == null) { + throw new Exception("Configuration key \"phonemes\" is null."); + } + string phonemesPath = Path.Combine(rootPath, dsConfig.phonemes); + phonemeTokens = DiffSingerUtils.LoadPhonemes(phonemesPath); + //Load models + if (dsConfig.linguistic == null) { + throw new Exception("Configuration key \"linguistic\" is null."); + } + var linguisticModelPath = Path.Join(rootPath, dsConfig.linguistic); + var linguisticModelBytes = File.ReadAllBytes(linguisticModelPath); + linguisticHash = XXH64.DigestOf(linguisticModelBytes); + linguisticModel = Onnx.getInferenceSession(linguisticModelBytes); + var varianceModelPath = Path.Join(rootPath, dsConfig.variance); + var varianceModelBytes = File.ReadAllBytes(varianceModelPath); + varianceHash = XXH64.DigestOf(varianceModelBytes); + varianceModel = Onnx.getInferenceSession(varianceModelBytes); + frameMs = 1000f * dsConfig.hop_size / dsConfig.sample_rate; + //Load g2p + g2p = LoadG2p(rootPath); + } + + protected IG2p LoadG2p(string rootPath) { + // Load dictionary from singer folder. + string file = Path.Combine(rootPath, "dsdict.yaml"); + if(!File.Exists(file)){ + throw new Exception($"File not found: {file}"); + } + try { + var g2pBuilder = G2pDictionary.NewBuilder().Load(File.ReadAllText(file)); + //SP and AP should always be vowel + g2pBuilder.AddSymbol("SP", true); + g2pBuilder.AddSymbol("AP", true); + return g2pBuilder.Build(); + } catch (Exception e) { + throw new Exception($"Failed to load {file}", e); + } + } + + public DiffSingerSpeakerEmbedManager getSpeakerEmbedManager(){ + if(speakerEmbedManager is null) { + speakerEmbedManager = new DiffSingerSpeakerEmbedManager(dsConfig, rootPath); + } + return speakerEmbedManager; + } + + int PhonemeTokenize(string phoneme){ + bool success = phonemeTokens.TryGetValue(phoneme, out int token); + if(!success){ + throw new Exception($"Phoneme \"{phoneme}\" isn't supported by variance model. Please check {Path.Combine(rootPath, dsConfig.phonemes)}"); + } + return token; + } + + public VarianceResult Process(RenderPhrase phrase){ + int headFrames = DiffSingerUtils.headFrames; + int tailFrames = DiffSingerUtils.tailFrames; + if (dsConfig.predict_dur) { + //Check if all phonemes are defined in dsdict.yaml (for their types) + foreach (var phone in phrase.phones) { + if (!g2p.IsValidSymbol(phone.phoneme)) { + throw new InvalidDataException( + $"Type definition of symbol \"{phone.phoneme}\" not found. Consider adding it to dsdict.yaml of the variance predictor."); + } + } + } + //Linguistic Encoder + var linguisticInputs = new List(); + var tokens = phrase.phones.Select(p => p.phoneme) + .Prepend("SP") + .Append("SP") + .Select(x => (Int64)PhonemeTokenize(x)) + .ToArray(); + var ph_dur = DiffSingerUtils.PaddedPhoneDurations(phrase, frameMs, headFrames, tailFrames); + int totalFrames = ph_dur.Sum(); + linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("tokens", + new DenseTensor(tokens, new int[] { tokens.Length }, false) + .Reshape(new int[] { 1, tokens.Length }))); + if(dsConfig.predict_dur){ + //if predict_dur is true, use word encode mode + var (word_div, word_dur) = DiffSingerUtils.PaddedWordDivAndDur(phrase, ph_dur, g2p.IsVowel); + linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("word_div", + new DenseTensor(word_div, new int[] { word_div.Length }, false) + .Reshape(new int[] { 1, word_div.Length }))); + linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("word_dur", + new DenseTensor(word_dur, new int[] { word_dur.Length }, false) + .Reshape(new int[] { 1, word_dur.Length }))); + }else{ + //if predict_dur is false, use phoneme encode mode + linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("ph_dur", + new DenseTensor(ph_dur.Select(x=>(Int64)x).ToArray(), new int[] { ph_dur.Length }, false) + .Reshape(new int[] { 1, ph_dur.Length }))); + } + //Language id + if(dsConfig.use_lang_id){ + var langIdByPhone = phrase.phones + .Select(p => (long)languageIds.GetValueOrDefault( + DiffSingerUtils.PhonemeLanguage(p.phoneme),0 + )) + .Prepend(0) + .Append(0) + .ToArray(); + var langIdTensor = new DenseTensor(langIdByPhone, new int[] { langIdByPhone.Length }, false) + .Reshape(new int[] { 1, langIdByPhone.Length }); + linguisticInputs.Add(NamedOnnxValue.CreateFromTensor("languages", langIdTensor)); + } + + Onnx.VerifyInputNames(linguisticModel, linguisticInputs); + var linguisticCache = Preferences.Default.DiffSingerTensorCache + ? new DiffSingerCache(linguisticHash, linguisticInputs) + : null; + var linguisticOutputs = linguisticCache?.Load(); + if (linguisticOutputs is null) { + linguisticOutputs = linguisticModel.Run(linguisticInputs).Cast().ToList(); + linguisticCache?.Save(linguisticOutputs); + phrase.AddCacheFile(linguisticCache?.Filename); + } + Tensor encoder_out = linguisticOutputs + .Where(o => o.Name == "encoder_out") + .First() + .AsTensor(); + + //Variance Predictor + var pitch = DiffSingerUtils.SampleCurve(phrase, phrase.pitches, 0, frameMs, totalFrames, headFrames, tailFrames, + x => x * 0.01).Select(f => (float)f).ToArray(); + var toneShift = DiffSingerUtils.SampleCurve(phrase, phrase.toneShift, 0, frameMs, totalFrames, headFrames, tailFrames, + x => x * 0.01).Select(f => (float)f).ToArray(); + pitch = pitch.Zip(toneShift, (x, d) => x + d).ToArray(); + + var varianceInputs = new List(); + var variancePatchInputs = new List(); + void AddVarianceInput(NamedOnnxValue input, bool includeInPatchKey = true) { + varianceInputs.Add(input); + if (includeInPatchKey) { + variancePatchInputs.Add(input); + } + } + AddVarianceInput(NamedOnnxValue.CreateFromTensor("encoder_out", encoder_out)); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("ph_dur", + new DenseTensor(ph_dur.Select(x=>(Int64)x).ToArray(), new int[] { ph_dur.Length }, false) + .Reshape(new int[] { 1, ph_dur.Length }))); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("pitch", + new DenseTensor(pitch, new int[] { pitch.Length }, false) + .Reshape(new int[] { 1, totalFrames })), includeInPatchKey: false); + if (dsConfig.predict_energy) { + var energy = Enumerable.Repeat(0f, totalFrames).ToArray(); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("energy", + new DenseTensor(energy, new int[] { energy.Length }, false) + .Reshape(new int[] { 1, totalFrames }))); + } + if (dsConfig.predict_breathiness) { + var breathiness = Enumerable.Repeat(0f, totalFrames).ToArray(); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("breathiness", + new DenseTensor(breathiness, new int[] { breathiness.Length }, false) + .Reshape(new int[] { 1, totalFrames }))); + } + if (dsConfig.predict_voicing) { + var voicing = Enumerable.Repeat(0f, totalFrames).ToArray(); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("voicing", + new DenseTensor(voicing, new int[] { voicing.Length }, false) + .Reshape(new int[] { 1, totalFrames }))); + } + if (dsConfig.predict_tension) { + var tension = Enumerable.Repeat(0f, totalFrames).ToArray(); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("tension", + new DenseTensor(tension, new int[] { tension.Length }, false) + .Reshape(new int[] { 1, totalFrames }))); + } + + var numVariances = new[] { + dsConfig.predict_energy, + dsConfig.predict_breathiness, + dsConfig.predict_voicing, + dsConfig.predict_tension, + }.Sum(Convert.ToInt32); + var retake = Enumerable.Repeat(true, totalFrames * numVariances).ToArray(); + AddVarianceInput(NamedOnnxValue.CreateFromTensor("retake", + new DenseTensor(retake, new int[] { retake.Length }, false) + .Reshape(new int[] { 1, totalFrames, numVariances }))); + var steps = Preferences.Default.DiffSingerStepsVariance; + if (dsConfig.useContinuousAcceleration) { + AddVarianceInput(NamedOnnxValue.CreateFromTensor("steps", + new DenseTensor(new long[] { steps }, new int[] { 1 }, false))); + } else { + // find a largest integer speedup that are less than 1000 / steps and is a factor of 1000 + long speedup = Math.Max(1, 1000 / steps); + while (1000 % speedup != 0 && speedup > 1) { + speedup--; + } + AddVarianceInput(NamedOnnxValue.CreateFromTensor("speedup", + new DenseTensor(new long[] { speedup }, new int[] { 1 },false))); + } + //Speaker + float[]? speakerEmbed = null; + if(dsConfig.speakers != null) { + var speakerEmbedManager = getSpeakerEmbedManager(); + var spkEmbedTensor = speakerEmbedManager.PhraseSpeakerEmbedByFrame(phrase, ph_dur, frameMs, totalFrames, headFrames, tailFrames); + speakerEmbed = spkEmbedTensor.ToArray(); + // Speaker embedding is a retake-able frame-level condition. + AddVarianceInput(NamedOnnxValue.CreateFromTensor("spk_embed", spkEmbedTensor), includeInPatchKey: false); + } + ulong? variancePatchKey = null; + if (Preferences.Default.DiffSingerTensorCache && + Preferences.Default.DiffSingerVarianceLocalPitchPatch) { + var baseHash = new DiffSingerCache(varianceHash, variancePatchInputs).Hash; + variancePatchKey = DiffSingerVariancePatch.BuildStateKey(baseHash, phrase.position, phrase.end); + } + // Cache the final pipeline result in a separate namespace from raw predictor outputs. + var resultCacheInputs = new List(varianceInputs) { + NamedOnnxValue.CreateFromTensor( + "result_cache_version", + new DenseTensor(new long[] { 1 }, new int[] { 1 }, false)), + }; + var resultCache = Preferences.Default.DiffSingerTensorCache + ? new DiffSingerCache(varianceHash, resultCacheInputs) + : null; + var cachedOutputs = resultCache?.Load(); + if (cachedOutputs != null) { + var cachedResult = ParseVarianceResult(cachedOutputs, frameMs, headFrames, tailFrames, totalFrames); + if (variancePatchKey.HasValue) { + variancePatchStates.Set( + variancePatchKey.Value, + new VariancePatchState(pitch, speakerEmbed, cachedResult)); + } + return cachedResult; + } + VariancePatchState? previous = null; + bool[]? retakeMask = null; + if (variancePatchKey.HasValue && variancePatchStates.TryGetValue(variancePatchKey.Value, out var cachedState) && + DiffSingerVariancePatch.IsMetadataCompatible(cachedState.result, new VarianceResult { + frameMs = frameMs, + headFrames = headFrames, + tailFrames = tailFrames, + totalFrames = totalFrames, + }) && + DiffSingerVariancePatch.IsChannelLayoutCompatible( + cachedState.result, + totalFrames, + dsConfig.predict_energy, + dsConfig.predict_breathiness, + dsConfig.predict_voicing, + dsConfig.predict_tension)) { + previous = cachedState; + var pitchMask = DiffSingerVariancePatch.BuildChangedFrameMask(cachedState.pitch, pitch, 1e-4f); + var speakerMask = DiffSingerVariancePatch.BuildChangedFrameMask( + cachedState.speakerEmbed ?? Array.Empty(), + speakerEmbed ?? Array.Empty(), + totalFrames, + 1e-4f); + retakeMask = new bool[totalFrames]; + for (int i = 0; i < retakeMask.Length; i++) { + retakeMask[i] = (i < pitchMask.Length && pitchMask[i]) || + (i < speakerMask.Length && speakerMask[i]); + } + if (!retakeMask.Any(x => x)) { + return DiffSingerVariancePatch.CloneResult(cachedState.result); + } + if (retakeMask.All(x => x)) { + previous = null; + } else { + ReplaceVarianceInputsWithPrevious(varianceInputs, cachedState.result); + } + } + if (retakeMask != null) { + var retakeTensorValues = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); + var retakeInput = varianceInputs.First(x => x.Name == "retake"); + varianceInputs[varianceInputs.IndexOf(retakeInput)] = NamedOnnxValue.CreateFromTensor( + "retake", + new DenseTensor(retakeTensorValues, new[] { retakeTensorValues.Length }, false) + .Reshape(new[] { 1, totalFrames, numVariances })); + } + Onnx.VerifyInputNames(varianceModel, varianceInputs); + var varianceOutputs = varianceModel.Run(varianceInputs).Cast().ToList(); + Tensor? energy_pred = dsConfig.predict_energy + ? varianceOutputs + .Where(o => o.Name == "energy_pred") + .First() + .AsTensor() + : null; + Tensor? breathiness_pred = dsConfig.predict_breathiness + ? varianceOutputs + .Where(o => o.Name == "breathiness_pred") + .First() + .AsTensor() + : null; + Tensor? voicing_pred = dsConfig.predict_voicing + ? varianceOutputs + .Where(o => o.Name == "voicing_pred") + .First() + .AsTensor() + : null; + Tensor? tension_pred = dsConfig.predict_tension + ? varianceOutputs + .Where(o => o.Name == "tension_pred") + .First() + .AsTensor() + : null; + var result = new VarianceResult{ + energy = energy_pred?.ToArray(), + breathiness = breathiness_pred?.ToArray(), + voicing = voicing_pred?.ToArray(), + tension = tension_pred?.ToArray(), + frameMs = frameMs, + headFrames = headFrames, + tailFrames = tailFrames, + totalFrames = totalFrames, + }; + if (previous != null && retakeMask != null) { + var channelMask = DiffSingerVariancePatch.ExpandToChannels(retakeMask, numVariances); + result = DiffSingerVariancePatch.HardCompose(previous.result, result, channelMask, numVariances); + } + if (resultCache != null) { + resultCache.Save(BuildVarianceOutputs(result)); + phrase.AddCacheFile(resultCache.Filename); + } + if (variancePatchKey.HasValue) { + variancePatchStates.Set( + variancePatchKey.Value, + new VariancePatchState(pitch, speakerEmbed, result)); + } + return result; + } + + VarianceResult ParseVarianceResult( + ICollection outputs, + float frameMs, + int headFrames, + int tailFrames, + int totalFrames) { + return new VarianceResult { + energy = dsConfig.predict_energy ? outputs.First(o => o.Name == "energy_pred").AsTensor().ToArray() : null, + breathiness = dsConfig.predict_breathiness ? outputs.First(o => o.Name == "breathiness_pred").AsTensor().ToArray() : null, + voicing = dsConfig.predict_voicing ? outputs.First(o => o.Name == "voicing_pred").AsTensor().ToArray() : null, + tension = dsConfig.predict_tension ? outputs.First(o => o.Name == "tension_pred").AsTensor().ToArray() : null, + frameMs = frameMs, + headFrames = headFrames, + tailFrames = tailFrames, + totalFrames = totalFrames, + }; + } + + List BuildVarianceOutputs(VarianceResult result) { + var outputs = new List(); + void Add(string name, float[]? values) { + if (values != null) { + outputs.Add(NamedOnnxValue.CreateFromTensor( + name, + new DenseTensor(values, new[] { values.Length }, false) + .Reshape(new[] { 1, values.Length }))); + } + } + Add("energy_pred", result.energy); + Add("breathiness_pred", result.breathiness); + Add("voicing_pred", result.voicing); + Add("tension_pred", result.tension); + return outputs; + } + + static void ReplaceVarianceInputsWithPrevious( + List inputs, + VarianceResult previous) { + var channels = new[] { + ("energy", previous.energy), + ("breathiness", previous.breathiness), + ("voicing", previous.voicing), + ("tension", previous.tension), + }; + foreach (var (name, values) in channels) { + if (values == null) continue; + var input = inputs.FirstOrDefault(x => x.Name == name); + if (input == null) continue; + var current = input.AsTensor().ToArray(); + if (current.Length != values.Length) continue; + Array.Copy(values, current, values.Length); + inputs[inputs.IndexOf(input)] = NamedOnnxValue.CreateFromTensor( + name, + new DenseTensor(current, new[] { current.Length }, false) + .Reshape(new[] { 1, current.Length })); + } + } + + private bool disposedValue; + + protected virtual void Dispose(bool disposing) { + if (!disposedValue) { + if (disposing) { + linguisticModel?.Dispose(); + varianceModel?.Dispose(); + } + disposedValue = true; + } + } + + public void Dispose() { + Dispose(disposing: true); + GC.SuppressFinalize(this); + } + } +} diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs index c92d1d923..3ab87f4bf 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs @@ -1,223 +1,236 @@ -using System; -using System.Collections.Generic; -using System.Linq; - -namespace OpenUtau.Core.DiffSinger { - internal sealed class VariancePatchState { - public readonly float[] pitch; - public readonly float[]? speakerEmbed; - public readonly VarianceResult result; - - public VariancePatchState(float[] pitch, float[]? speakerEmbed, VarianceResult result) { - this.pitch = pitch.ToArray(); - this.speakerEmbed = speakerEmbed?.ToArray(); - this.result = DiffSingerVariancePatch.CloneResult(result); - } - } - - internal sealed class VariancePatchStateCache { - readonly int capacity; - readonly Dictionary> entries = new(); - readonly LinkedList<(ulong key, VariancePatchState state)> recency = new(); - - internal VariancePatchStateCache(int capacity) { - if (capacity <= 0) { - throw new ArgumentOutOfRangeException(nameof(capacity)); - } - this.capacity = capacity; - } - - internal int Count => entries.Count; - - internal bool TryGetValue(ulong key, out VariancePatchState state) { - if (!entries.TryGetValue(key, out var node)) { - state = null!; - return false; - } - recency.Remove(node); - recency.AddFirst(node); - state = node.Value.state; - return true; - } - - internal void Set(ulong key, VariancePatchState state) { - if (entries.TryGetValue(key, out var existing)) { - existing.Value = (key, state); - recency.Remove(existing); - recency.AddFirst(existing); - return; - } - var node = recency.AddFirst((key, state)); - entries.Add(key, node); - if (entries.Count <= capacity) { - return; - } - var oldest = recency.Last!; - recency.RemoveLast(); - entries.Remove(oldest.Value.key); - } - } - - internal static class DiffSingerVariancePatch { - public static ulong BuildStateKey(ulong baseHash, int phrasePosition, int phraseEnd) { - unchecked { - ulong hash = baseHash; - hash = (hash ^ (uint)phrasePosition) * 1099511628211UL; - hash = (hash ^ (uint)phraseEnd) * 1099511628211UL; - return hash; - } - } - - internal static bool[] BuildChangedFrameMask( - IReadOnlyList previous, - IReadOnlyList current, - float epsilon) { - int length = Math.Max(previous.Count, current.Count); - var mask = new bool[length]; - for (int i = 0; i < length; i++) { - mask[i] = i >= previous.Count || i >= current.Count || - Math.Abs(previous[i] - current[i]) > epsilon; - } - return mask; - } - - internal static bool[] BuildChangedFrameMask( - IReadOnlyList previous, - IReadOnlyList current, - int frameCount, - float epsilon) { - if (frameCount <= 0) { - return Array.Empty(); - } - if (previous.Count != current.Count || previous.Count % frameCount != 0) { - return Enumerable.Repeat(true, frameCount).ToArray(); - } - int valuesPerFrame = previous.Count / frameCount; - var mask = new bool[frameCount]; - for (int frame = 0; frame < frameCount; frame++) { - int offset = frame * valuesPerFrame; - for (int i = 0; i < valuesPerFrame; i++) { - if (Math.Abs(previous[offset + i] - current[offset + i]) > epsilon) { - mask[frame] = true; - break; - } - } - } - return mask; - } - - internal static bool[] ExpandToChannels( - IReadOnlyList frameMask, - int channelCount) { - if (channelCount < 0) { - throw new ArgumentOutOfRangeException(nameof(channelCount)); - } - var mask = new bool[frameMask.Count * channelCount]; - for (int frame = 0; frame < frameMask.Count; frame++) { - if (!frameMask[frame]) continue; - for (int channel = 0; channel < channelCount; channel++) { - mask[frame * channelCount + channel] = true; - } - } - return mask; - } - - internal static VarianceResult HardCompose( - VarianceResult previous, - VarianceResult predicted, - IReadOnlyList retakeMask, - int channelCount) { - if (!IsCompatible(previous, predicted) || - retakeMask.Count != previous.totalFrames * channelCount) { - return CloneResult(predicted); - } - int channel = 0; - var energy = ComposeEnabledChannel(previous.energy, predicted.energy, retakeMask, previous.totalFrames, ref channel, channelCount); - var breathiness = ComposeEnabledChannel(previous.breathiness, predicted.breathiness, retakeMask, previous.totalFrames, ref channel, channelCount); - var voicing = ComposeEnabledChannel(previous.voicing, predicted.voicing, retakeMask, previous.totalFrames, ref channel, channelCount); - var tension = ComposeEnabledChannel(previous.tension, predicted.tension, retakeMask, previous.totalFrames, ref channel, channelCount); - return new VarianceResult { - energy = energy, - breathiness = breathiness, - voicing = voicing, - tension = tension, - frameMs = predicted.frameMs, - headFrames = predicted.headFrames, - tailFrames = predicted.tailFrames, - totalFrames = predicted.totalFrames, - }; - } - - static float[]? ComposeEnabledChannel( - float[]? previous, - float[]? predicted, - IReadOnlyList mask, - int frameCount, - ref int channel, - int channelCount) { - if (previous == null && predicted == null) { - return null; - } - int currentChannel = channel++; - return ComposeChannel(previous, predicted, mask, frameCount, currentChannel, channelCount); - } - - static float[]? ComposeChannel( - float[]? previous, - float[]? predicted, - IReadOnlyList mask, - int frameCount, - int channel, - int channelCount) { - if (previous == null || predicted == null) { - return predicted?.ToArray(); - } - if (previous.Length != frameCount || predicted.Length != frameCount) { - return predicted.ToArray(); - } - var result = previous.ToArray(); - for (int frame = 0; frame < frameCount; frame++) { - if (mask[frame * channelCount + channel]) { - result[frame] = predicted[frame]; - } - } - return result; - } - - internal static bool IsMetadataCompatible(VarianceResult previous, VarianceResult current) { - return previous.totalFrames == current.totalFrames && - previous.headFrames == current.headFrames && - previous.tailFrames == current.tailFrames && - Math.Abs(previous.frameMs - current.frameMs) < 1e-4f; - } - - internal static bool IsCompatible(VarianceResult previous, VarianceResult current) { - return IsMetadataCompatible(previous, current) && - previous.totalFrames == current.totalFrames && - previous.headFrames == current.headFrames && - previous.tailFrames == current.tailFrames && - Math.Abs(previous.frameMs - current.frameMs) < 1e-4f && - SameLength(previous.energy, current.energy) && - SameLength(previous.breathiness, current.breathiness) && - SameLength(previous.voicing, current.voicing) && - SameLength(previous.tension, current.tension); - } - - static bool SameLength(float[]? a, float[]? b) { - return (a == null) == (b == null) && (a == null || a.Length == b!.Length); - } - - internal static VarianceResult CloneResult(VarianceResult result) { - return new VarianceResult { - energy = result.energy?.ToArray(), - breathiness = result.breathiness?.ToArray(), - voicing = result.voicing?.ToArray(), - tension = result.tension?.ToArray(), - frameMs = result.frameMs, - headFrames = result.headFrames, - tailFrames = result.tailFrames, - totalFrames = result.totalFrames, - }; - } - } -} +using System; +using System.Collections.Generic; +using System.Linq; + +namespace OpenUtau.Core.DiffSinger { + internal sealed class VariancePatchState { + public readonly float[] pitch; + public readonly float[]? speakerEmbed; + public readonly VarianceResult result; + + public VariancePatchState(float[] pitch, float[]? speakerEmbed, VarianceResult result) { + this.pitch = pitch.ToArray(); + this.speakerEmbed = speakerEmbed?.ToArray(); + this.result = DiffSingerVariancePatch.CloneResult(result); + } + } + + internal sealed class VariancePatchStateCache { + readonly int capacity; + readonly Dictionary> entries = new(); + readonly LinkedList<(ulong key, VariancePatchState state)> recency = new(); + + internal VariancePatchStateCache(int capacity) { + if (capacity <= 0) { + throw new ArgumentOutOfRangeException(nameof(capacity)); + } + this.capacity = capacity; + } + + internal int Count => entries.Count; + + internal bool TryGetValue(ulong key, out VariancePatchState state) { + if (!entries.TryGetValue(key, out var node)) { + state = null!; + return false; + } + recency.Remove(node); + recency.AddFirst(node); + state = node.Value.state; + return true; + } + + internal void Set(ulong key, VariancePatchState state) { + if (entries.TryGetValue(key, out var existing)) { + existing.Value = (key, state); + recency.Remove(existing); + recency.AddFirst(existing); + return; + } + var node = recency.AddFirst((key, state)); + entries.Add(key, node); + if (entries.Count <= capacity) { + return; + } + var oldest = recency.Last!; + recency.RemoveLast(); + entries.Remove(oldest.Value.key); + } + } + + internal static class DiffSingerVariancePatch { + public static ulong BuildStateKey(ulong baseHash, int phrasePosition, int phraseEnd) { + unchecked { + ulong hash = baseHash; + hash = (hash ^ (uint)phrasePosition) * 1099511628211UL; + hash = (hash ^ (uint)phraseEnd) * 1099511628211UL; + return hash; + } + } + + internal static bool[] BuildChangedFrameMask( + IReadOnlyList previous, + IReadOnlyList current, + float epsilon) { + int length = Math.Max(previous.Count, current.Count); + var mask = new bool[length]; + for (int i = 0; i < length; i++) { + mask[i] = i >= previous.Count || i >= current.Count || + Math.Abs(previous[i] - current[i]) > epsilon; + } + return mask; + } + + internal static bool[] BuildChangedFrameMask( + IReadOnlyList previous, + IReadOnlyList current, + int frameCount, + float epsilon) { + if (frameCount <= 0) { + return Array.Empty(); + } + if (previous.Count != current.Count || previous.Count % frameCount != 0) { + return Enumerable.Repeat(true, frameCount).ToArray(); + } + int valuesPerFrame = previous.Count / frameCount; + var mask = new bool[frameCount]; + for (int frame = 0; frame < frameCount; frame++) { + int offset = frame * valuesPerFrame; + for (int i = 0; i < valuesPerFrame; i++) { + if (Math.Abs(previous[offset + i] - current[offset + i]) > epsilon) { + mask[frame] = true; + break; + } + } + } + return mask; + } + + internal static bool[] ExpandToChannels( + IReadOnlyList frameMask, + int channelCount) { + if (channelCount < 0) { + throw new ArgumentOutOfRangeException(nameof(channelCount)); + } + var mask = new bool[frameMask.Count * channelCount]; + for (int frame = 0; frame < frameMask.Count; frame++) { + if (!frameMask[frame]) continue; + for (int channel = 0; channel < channelCount; channel++) { + mask[frame * channelCount + channel] = true; + } + } + return mask; + } + + internal static VarianceResult HardCompose( + VarianceResult previous, + VarianceResult predicted, + IReadOnlyList retakeMask, + int channelCount) { + if (!IsCompatible(previous, predicted) || + retakeMask.Count != previous.totalFrames * channelCount) { + return CloneResult(predicted); + } + int channel = 0; + var energy = ComposeEnabledChannel(previous.energy, predicted.energy, retakeMask, previous.totalFrames, ref channel, channelCount); + var breathiness = ComposeEnabledChannel(previous.breathiness, predicted.breathiness, retakeMask, previous.totalFrames, ref channel, channelCount); + var voicing = ComposeEnabledChannel(previous.voicing, predicted.voicing, retakeMask, previous.totalFrames, ref channel, channelCount); + var tension = ComposeEnabledChannel(previous.tension, predicted.tension, retakeMask, previous.totalFrames, ref channel, channelCount); + return new VarianceResult { + energy = energy, + breathiness = breathiness, + voicing = voicing, + tension = tension, + frameMs = predicted.frameMs, + headFrames = predicted.headFrames, + tailFrames = predicted.tailFrames, + totalFrames = predicted.totalFrames, + }; + } + + static float[]? ComposeEnabledChannel( + float[]? previous, + float[]? predicted, + IReadOnlyList mask, + int frameCount, + ref int channel, + int channelCount) { + if (previous == null && predicted == null) { + return null; + } + int currentChannel = channel++; + return ComposeChannel(previous, predicted, mask, frameCount, currentChannel, channelCount); + } + + static float[]? ComposeChannel( + float[]? previous, + float[]? predicted, + IReadOnlyList mask, + int frameCount, + int channel, + int channelCount) { + if (previous == null || predicted == null) { + return predicted?.ToArray(); + } + if (previous.Length != frameCount || predicted.Length != frameCount) { + return predicted.ToArray(); + } + var result = previous.ToArray(); + for (int frame = 0; frame < frameCount; frame++) { + if (mask[frame * channelCount + channel]) { + result[frame] = predicted[frame]; + } + } + return result; + } + + internal static bool IsMetadataCompatible(VarianceResult previous, VarianceResult current) { + return previous.totalFrames == current.totalFrames && + previous.headFrames == current.headFrames && + previous.tailFrames == current.tailFrames && + Math.Abs(previous.frameMs - current.frameMs) < 1e-4f; + } + + internal static bool IsChannelLayoutCompatible( + VarianceResult result, + int totalFrames, + bool predictEnergy, + bool predictBreathiness, + bool predictVoicing, + bool predictTension) { + return ChannelMatches(result.energy, predictEnergy, totalFrames) && + ChannelMatches(result.breathiness, predictBreathiness, totalFrames) && + ChannelMatches(result.voicing, predictVoicing, totalFrames) && + ChannelMatches(result.tension, predictTension, totalFrames); + } + + internal static bool IsCompatible(VarianceResult previous, VarianceResult current) { + return IsMetadataCompatible(previous, current) && + SameLength(previous.energy, current.energy) && + SameLength(previous.breathiness, current.breathiness) && + SameLength(previous.voicing, current.voicing) && + SameLength(previous.tension, current.tension); + } + + static bool ChannelMatches(float[]? values, bool enabled, int totalFrames) { + return enabled ? values?.Length == totalFrames : values == null; + } + + static bool SameLength(float[]? a, float[]? b) { + return (a == null) == (b == null) && (a == null || a.Length == b!.Length); + } + + internal static VarianceResult CloneResult(VarianceResult result) { + return new VarianceResult { + energy = result.energy?.ToArray(), + breathiness = result.breathiness?.ToArray(), + voicing = result.voicing?.ToArray(), + tension = result.tension?.ToArray(), + frameMs = result.frameMs, + headFrames = result.headFrames, + tailFrames = result.tailFrames, + totalFrames = result.totalFrames, + }; + } + } +} diff --git a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs index e4b171368..945e13caf 100644 --- a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs +++ b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs @@ -1,197 +1,235 @@ -using OpenUtau.Core.DiffSinger; -using Xunit; - -namespace OpenUtau.Core { - public class DiffSingerVariancePatchTest { - [Fact] - public void BuildChangedFrameMaskMarksOnlyChangedFrames() { - var mask = DiffSingerVariancePatch.BuildChangedFrameMask( - new[] { 1f, 1f, 1f, 2f }, - new[] { 1f, 2f, 1f, 2f }, - 1e-4f); - - Assert.Equal(new[] { false, true, false, false }, mask); - } - - [Fact] - public void BuildChangedFrameMaskGroupsSpeakerEmbeddingByFrame() { - var mask = DiffSingerVariancePatch.BuildChangedFrameMask( - new[] { 1f, 2f, 3f, 4f, 5f, 6f }, - new[] { 1f, 2f, 3f, 40f, 5f, 6f }, - 3, - 1e-4f); - - Assert.Equal(new[] { false, true, false }, mask); - } - - [Fact] - public void BuildChangedFrameMaskMarksAllFramesForIncompatibleEmbeddingShape() { - var mask = DiffSingerVariancePatch.BuildChangedFrameMask( - new[] { 1f, 2f, 3f, 4f }, - new[] { 1f, 2f, 3f }, - 2, - 1e-4f); - - Assert.Equal(new[] { true, true }, mask); - } - - [Fact] - public void ExpandToChannelsUsesSharedFrameMask() { - var mask = DiffSingerVariancePatch.ExpandToChannels( - new[] { false, true, false }, 3); - - Assert.Equal( - new[] { false, false, false, true, true, true, false, false, false }, - mask); - } - - [Fact] - public void HardComposePreservesUnmaskedFramesExactly() { - var previous = Result( - new[] { 1f, 2f, 3f, 4f }, - new[] { 5f, 6f, 7f, 8f }); - var predicted = Result( - new[] { 10f, 20f, 30f, 40f }, - new[] { 50f, 60f, 70f, 80f }); - var mask = DiffSingerVariancePatch.ExpandToChannels( - new[] { false, true, false, true }, 2); - - var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 2); - - Assert.Equal(new[] { 1f, 20f, 3f, 40f }, result.energy); - Assert.Equal(new[] { 5f, 60f, 7f, 80f }, result.breathiness); - } - - [Fact] - public void HardComposeDoesNotLeakModelChangesOutsideMask() { - var previous = Result(new[] { 1f, 2f, 3f }); - var predicted = Result(new[] { 100f, 200f, 300f }); - var mask = DiffSingerVariancePatch.ExpandToChannels( - new[] { false, true, false }, 1); - - var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 1); - - Assert.Equal(new[] { 1f, 200f, 3f }, result.energy); - } - - [Fact] - public void HardComposeHandlesNullMiddleChannel() { - var previous = Result( - new[] { 1f, 2f }, - voicing: new[] { 3f, 4f }); - var predicted = Result( - new[] { 10f, 20f }, - voicing: new[] { 30f, 40f }); - var mask = new[] { false, false, true, true }; - - var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 2); - - Assert.Equal(new[] { 1f, 20f }, result.energy); - Assert.Null(result.breathiness); - Assert.Equal(new[] { 3f, 40f }, result.voicing); - } - - [Fact] - public void HardComposeHandlesAllChannels() { - var previous = Result( - new[] { 1f, 2f }, - new[] { 3f, 4f }, - new[] { 5f, 6f }, - new[] { 7f, 8f }); - var predicted = Result( - new[] { 10f, 20f }, - new[] { 30f, 40f }, - new[] { 50f, 60f }, - new[] { 70f, 80f }); - var mask = DiffSingerVariancePatch.ExpandToChannels(new[] { false, true }, 4); - - var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 4); - - Assert.Equal(new[] { 1f, 20f }, result.energy); - Assert.Equal(new[] { 3f, 40f }, result.breathiness); - Assert.Equal(new[] { 5f, 60f }, result.voicing); - Assert.Equal(new[] { 7f, 80f }, result.tension); - } - - [Fact] - public void HardComposePreservesAllPreviousChannelsForFalseMask() { - var previous = Result( - new[] { 1f, 2f }, - new[] { 3f, 4f }, - new[] { 5f, 6f }, - new[] { 7f, 8f }); - var predicted = Result( - new[] { 10f, 20f }, - new[] { 30f, 40f }, - new[] { 50f, 60f }, - new[] { 70f, 80f }); - var mask = DiffSingerVariancePatch.ExpandToChannels(new[] { false, false }, 4); - - var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 4); - - Assert.Equal(previous.energy, result.energy); - Assert.Equal(previous.breathiness, result.breathiness); - Assert.Equal(previous.voicing, result.voicing); - Assert.Equal(previous.tension, result.tension); - } - - [Fact] - public void HardComposeFallsBackToPredictedForIncompatibleMetadata() { - var previous = Result(new[] { 1f, 2f, 3f }, frameMs: 50); - var predicted = Result(new[] { 10f, 20f, 30f }, frameMs: 60); - var mask = new[] { true, false, true }; - - var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 1); - - Assert.Equal(predicted.energy, result.energy); - } - - [Fact] - public void VariancePatchStateCacheEvictsLeastRecentlyUsedState() { - var cache = new VariancePatchStateCache(2); - cache.Set(1, State(1)); - cache.Set(2, State(2)); - Assert.True(cache.TryGetValue(1, out _)); - - cache.Set(3, State(3)); - - Assert.Equal(2, cache.Count); - Assert.True(cache.TryGetValue(1, out _)); - Assert.False(cache.TryGetValue(2, out _)); - Assert.True(cache.TryGetValue(3, out _)); - } - - [Fact] - public void IsMetadataCompatibleRejectsFrameLayoutChanges() { - var previous = Result(new[] { 1f, 2f, 3f }); - var changed = Result(new[] { 1f, 2f, 3f, 4f }); - - Assert.False(DiffSingerVariancePatch.IsMetadataCompatible(previous, changed)); - } - - static VariancePatchState State(float value) { - return new VariancePatchState( - new[] { value }, - null, - Result(new[] { value })); - } - - static VarianceResult Result( - float[] energy, - float[]? breathiness = null, - float[]? voicing = null, - float[]? tension = null, - float frameMs = 50) { - return new VarianceResult { - energy = energy, - breathiness = breathiness, - voicing = voicing, - tension = tension, - frameMs = frameMs, - headFrames = 1, - tailFrames = 1, - totalFrames = energy.Length, - }; - } - } -} +using OpenUtau.Core.DiffSinger; +using Xunit; + +namespace OpenUtau.Core { + public class DiffSingerVariancePatchTest { + [Fact] + public void BuildChangedFrameMaskMarksOnlyChangedFrames() { + var mask = DiffSingerVariancePatch.BuildChangedFrameMask( + new[] { 1f, 1f, 1f, 2f }, + new[] { 1f, 2f, 1f, 2f }, + 1e-4f); + + Assert.Equal(new[] { false, true, false, false }, mask); + } + + [Fact] + public void BuildChangedFrameMaskGroupsSpeakerEmbeddingByFrame() { + var mask = DiffSingerVariancePatch.BuildChangedFrameMask( + new[] { 1f, 2f, 3f, 4f, 5f, 6f }, + new[] { 1f, 2f, 3f, 40f, 5f, 6f }, + 3, + 1e-4f); + + Assert.Equal(new[] { false, true, false }, mask); + } + + [Fact] + public void BuildChangedFrameMaskMarksAllFramesForIncompatibleEmbeddingShape() { + var mask = DiffSingerVariancePatch.BuildChangedFrameMask( + new[] { 1f, 2f, 3f, 4f }, + new[] { 1f, 2f, 3f }, + 2, + 1e-4f); + + Assert.Equal(new[] { true, true }, mask); + } + + [Fact] + public void ExpandToChannelsUsesSharedFrameMask() { + var mask = DiffSingerVariancePatch.ExpandToChannels( + new[] { false, true, false }, 3); + + Assert.Equal( + new[] { false, false, false, true, true, true, false, false, false }, + mask); + } + + [Fact] + public void HardComposePreservesUnmaskedFramesExactly() { + var previous = Result( + new[] { 1f, 2f, 3f, 4f }, + new[] { 5f, 6f, 7f, 8f }); + var predicted = Result( + new[] { 10f, 20f, 30f, 40f }, + new[] { 50f, 60f, 70f, 80f }); + var mask = DiffSingerVariancePatch.ExpandToChannels( + new[] { false, true, false, true }, 2); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 2); + + Assert.Equal(new[] { 1f, 20f, 3f, 40f }, result.energy); + Assert.Equal(new[] { 5f, 60f, 7f, 80f }, result.breathiness); + } + + [Fact] + public void HardComposeDoesNotLeakModelChangesOutsideMask() { + var previous = Result(new[] { 1f, 2f, 3f }); + var predicted = Result(new[] { 100f, 200f, 300f }); + var mask = DiffSingerVariancePatch.ExpandToChannels( + new[] { false, true, false }, 1); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 1); + + Assert.Equal(new[] { 1f, 200f, 3f }, result.energy); + } + + [Fact] + public void HardComposeHandlesNullMiddleChannel() { + var previous = Result( + new[] { 1f, 2f }, + voicing: new[] { 3f, 4f }); + var predicted = Result( + new[] { 10f, 20f }, + voicing: new[] { 30f, 40f }); + var mask = new[] { false, false, true, true }; + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 2); + + Assert.Equal(new[] { 1f, 20f }, result.energy); + Assert.Null(result.breathiness); + Assert.Equal(new[] { 3f, 40f }, result.voicing); + } + + [Fact] + public void HardComposeHandlesAllChannels() { + var previous = Result( + new[] { 1f, 2f }, + new[] { 3f, 4f }, + new[] { 5f, 6f }, + new[] { 7f, 8f }); + var predicted = Result( + new[] { 10f, 20f }, + new[] { 30f, 40f }, + new[] { 50f, 60f }, + new[] { 70f, 80f }); + var mask = DiffSingerVariancePatch.ExpandToChannels(new[] { false, true }, 4); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 4); + + Assert.Equal(new[] { 1f, 20f }, result.energy); + Assert.Equal(new[] { 3f, 40f }, result.breathiness); + Assert.Equal(new[] { 5f, 60f }, result.voicing); + Assert.Equal(new[] { 7f, 80f }, result.tension); + } + + [Fact] + public void HardComposePreservesAllPreviousChannelsForFalseMask() { + var previous = Result( + new[] { 1f, 2f }, + new[] { 3f, 4f }, + new[] { 5f, 6f }, + new[] { 7f, 8f }); + var predicted = Result( + new[] { 10f, 20f }, + new[] { 30f, 40f }, + new[] { 50f, 60f }, + new[] { 70f, 80f }); + var mask = DiffSingerVariancePatch.ExpandToChannels(new[] { false, false }, 4); + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 4); + + Assert.Equal(previous.energy, result.energy); + Assert.Equal(previous.breathiness, result.breathiness); + Assert.Equal(previous.voicing, result.voicing); + Assert.Equal(previous.tension, result.tension); + } + + [Fact] + public void HardComposeFallsBackToPredictedForIncompatibleMetadata() { + var previous = Result(new[] { 1f, 2f, 3f }, frameMs: 50); + var predicted = Result(new[] { 10f, 20f, 30f }, frameMs: 60); + var mask = new[] { true, false, true }; + + var result = DiffSingerVariancePatch.HardCompose(previous, predicted, mask, 1); + + Assert.Equal(predicted.energy, result.energy); + } + + [Fact] + public void IsChannelLayoutCompatibleAcceptsExpectedChannels() { + var result = Result( + new[] { 1f, 2f }, + voicing: new[] { 3f, 4f }); + + Assert.True(DiffSingerVariancePatch.IsChannelLayoutCompatible( + result, 2, true, false, true, false)); + } + + [Fact] + public void IsChannelLayoutCompatibleRejectsMissingEnabledChannel() { + var result = Result(new[] { 1f, 2f }); + + Assert.False(DiffSingerVariancePatch.IsChannelLayoutCompatible( + result, 2, true, true, false, false)); + } + + [Fact] + public void IsChannelLayoutCompatibleRejectsWrongChannelLength() { + var result = Result( + new[] { 1f, 2f }, + new[] { 3f }); + + Assert.False(DiffSingerVariancePatch.IsChannelLayoutCompatible( + result, 2, true, true, false, false)); + } + + [Fact] + public void IsChannelLayoutCompatibleRejectsUnexpectedDisabledChannel() { + var result = Result( + new[] { 1f, 2f }, + tension: new[] { 3f, 4f }); + + Assert.False(DiffSingerVariancePatch.IsChannelLayoutCompatible( + result, 2, true, false, false, false)); + } + + [Fact] + public void VariancePatchStateCacheEvictsLeastRecentlyUsedState() { + var cache = new VariancePatchStateCache(2); + cache.Set(1, State(1)); + cache.Set(2, State(2)); + Assert.True(cache.TryGetValue(1, out _)); + + cache.Set(3, State(3)); + + Assert.Equal(2, cache.Count); + Assert.True(cache.TryGetValue(1, out _)); + Assert.False(cache.TryGetValue(2, out _)); + Assert.True(cache.TryGetValue(3, out _)); + } + + [Fact] + public void IsMetadataCompatibleRejectsFrameLayoutChanges() { + var previous = Result(new[] { 1f, 2f, 3f }); + var changed = Result(new[] { 1f, 2f, 3f, 4f }); + + Assert.False(DiffSingerVariancePatch.IsMetadataCompatible(previous, changed)); + } + + static VariancePatchState State(float value) { + return new VariancePatchState( + new[] { value }, + null, + Result(new[] { value })); + } + + static VarianceResult Result( + float[] energy, + float[]? breathiness = null, + float[]? voicing = null, + float[]? tension = null, + float frameMs = 50) { + return new VarianceResult { + energy = energy, + breathiness = breathiness, + voicing = voicing, + tension = tension, + frameMs = frameMs, + headFrames = 1, + tailFrames = 1, + totalFrames = energy.Length, + }; + } + } +}