diff --git a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs index 6d4193458..5cd8e7f21 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariance.cs @@ -1,353 +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; - readonly Dictionary variancePatchStates = new Dictionary(); - - 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 - if(dsConfig.speakers != null) { - var speakerEmbedManager = getSpeakerEmbedManager(); - var spkEmbedTensor = speakerEmbedManager.PhraseSpeakerEmbedByFrame(phrase, ph_dur, frameMs, totalFrames, headFrames, tailFrames); - AddVarianceInput(NamedOnnxValue.CreateFromTensor("spk_embed", spkEmbedTensor)); - } - 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); - } - Onnx.VerifyInputNames(varianceModel, varianceInputs); - var varianceCache = 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); - } - 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 (variancePatchKey.HasValue) { - result = ApplyVariancePatch(variancePatchKey.Value, pitch, 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; - } - } - - 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 bd74de2de..3ab87f4bf 100644 --- a/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs +++ b/OpenUtau.Core/DiffSinger/DiffSingerVariancePatch.cs @@ -1,143 +1,236 @@ -using System; -using System.Collections.Generic; -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 { - public readonly float[] pitch; - public readonly VarianceResult result; - - public VariancePatchState(float[] pitch, VarianceResult result) { - this.pitch = pitch.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; - hash = (hash ^ (uint)phrasePosition) * 1099511628211UL; - hash = (hash ^ (uint)phraseEnd) * 1099511628211UL; - return hash; - } - } - - 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); - } - 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, - }; - } - - internal static List FindChangedRanges( - IReadOnlyList previousPitch, - IReadOnlyList currentPitch, - 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 (start >= 0) { - ranges.Add(new VariancePatchRange(start, length)); - } - return ranges; - } - - 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); - } - } - return weights; - } - - 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(); - } - 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]; - } - return result; - } - - 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, - }; - } - - 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; - } - } -} +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 7b05b0b21..945e13caf 100644 --- a/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs +++ b/OpenUtau.Test/Core/DiffSinger/DiffSingerVariancePatchTest.cs @@ -1,69 +1,235 @@ -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 }; - - var ranges = DiffSingerVariancePatch.FindChangedRanges(previous, current, 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); - } - - [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); - - var merged = DiffSingerVariancePatch.Merge(previous, new[] { 60f, 61f, 62f }, currentResult); - - Assert.Equal(previousResult.energy!, merged.energy!); - } - - [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!); - } - - [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); - - var merged = DiffSingerVariancePatch.Merge(previous, new[] { 60f, 62f, 62f }, currentResult); - - Assert.Equal(currentResult.energy!, merged.energy!); - } - - static VarianceResult Result(float[] energy, float frameMs = 50) { - return new VarianceResult { - energy = energy, - 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, + }; + } + } +}