add svc inference code

This commit is contained in:
jlqian98
2026-03-03 09:35:06 +08:00
parent f8fce96f9c
commit edf4e11bf9
7 changed files with 986 additions and 22 deletions
+105
View File
@@ -0,0 +1,105 @@
import os
import torch
import json
import argparse
from tqdm import tqdm
import numpy as np
import soundfile as sf
from collections import OrderedDict
from omegaconf import DictConfig
from soulxsinger.utils.file_utils import load_config
from soulxsinger.models.soulxsinger_svc import SoulXSingerSVC
from soulxsinger.utils.audio_utils import load_wav
def build_model(
model_path: str,
config: DictConfig,
device: str = "cuda",
):
"""
Build the model from the pre-trained model path and model configuration.
Args:
model_path (str): Path to the checkpoint file.
config (DictConfig): Model configuration.
device (str, optional): Device to use. Defaults to "cuda".
Returns:
Tuple[torch.nn.Module, torch.nn.Module]: The initialized model and vocoder.
"""
if not os.path.isfile(model_path):
raise FileNotFoundError(
f"Model checkpoint not found: {model_path}. "
"Please download the pretrained model and place it at the path, or set --model_path."
)
model = SoulXSingerSVC(config).to(device)
print("Model initialized.")
print("Model parameters:", sum(p.numel() for p in model.parameters()) / 1e6, "M")
checkpoint = torch.load(model_path, weights_only=False, map_location=device)
if "state_dict" not in checkpoint:
raise KeyError(
f"Checkpoint at {model_path} has no 'state_dict' key. "
"Expected a checkpoint saved with model.state_dict()."
)
model.load_state_dict(checkpoint["state_dict"], strict=True)
model.eval()
model.to(device)
print("Model checkpoint loaded.")
return model
def process(args, config, model: torch.nn.Module):
"""Run the full inference pipeline given a data_processor and model.
"""
os.makedirs(args.save_dir, exist_ok=True)
pt_wav = load_wav(args.prompt_wav_path, config.audio.sample_rate).to(args.device)
gt_wav = load_wav(args.target_wav_path, config.audio.sample_rate).to(args.device)
pt_f0 = torch.from_numpy(np.load(args.prompt_f0_path)).unsqueeze(0).to(args.device)
gt_f0 = torch.from_numpy(np.load(args.target_f0_path)).unsqueeze(0).to(args.device)
n_step = args.n_steps if hasattr(args, "n_steps") else config.infer.n_steps
cfg = args.cfg if hasattr(args, "cfg") else config.infer.cfg
generated_audio, generated_shift = model.infer(pt_wav, gt_wav, pt_f0, gt_f0, auto_shift=args.auto_shift, pitch_shift=args.pitch_shift, n_steps=n_step, cfg=cfg)
generated_audio = generated_audio.squeeze().cpu().numpy()
if args.pitch_shift != generated_shift:
args.pitch_shift = generated_shift
# print(f"Applied pitch shift of {generated_shift} semitones to match GT F0 contour.")
sf.write(os.path.join(args.save_dir, "generated.wav"), generated_audio, config.audio.sample_rate)
print(f"Generated audio saved to {os.path.join(args.save_dir, 'generated.wav')}")
def main(args, config):
model = build_model(
model_path=args.model_path,
config=config,
device=args.device,
)
process(args, config, model)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--device", type=str, default="cuda")
parser.add_argument("--model_path", type=str, default='pretrained_models/soulx-singer/model.pt')
parser.add_argument("--config", type=str, default='soulxsinger/config/soulxsinger.yaml')
parser.add_argument("--prompt_wav_path", type=str, default='example/audio/zh_prompt.wav')
parser.add_argument("--target_wav_path", type=str, default='example/audio/zh_target.wav')
parser.add_argument("--prompt_f0_path", type=str, default='example/audio/zh_prompt_f0.npy')
parser.add_argument("--target_f0_path", type=str, default='example/audio/zh_target_f0.npy')
parser.add_argument("--save_dir", type=str, default='outputs')
parser.add_argument("--auto_shift", action="store_true")
parser.add_argument("--pitch_shift", type=int, default=0)
parser.add_argument("--n_steps", type=int, default=32)
parser.add_argument("--cfg", type=float, default=3.0)
args = parser.parse_args()
config = load_config(args.config)
main(args, config)
+27
View File
@@ -0,0 +1,27 @@
#!/bin/bash
script_dir=$(dirname "$(realpath "$0")")
root_dir=$(dirname "$script_dir")
cd $root_dir || exit
export PYTHONPATH=$root_dir:$PYTHONPATH
model_path=pretrained_models/SoulX-Singer-SVC/model.pt
config=soulxsinger/config/soulxsinger.yaml
prompt_wav_path=example/audio/zh_prompt.mp3
target_wav_path=example/audio/music.mp3
prompt_f0_path=example/audio/zh_prompt_f0.npy
target_f0_path=example/audio/music_f0.npy
save_dir=example/generated/music_svc
python -m cli.inference_svc \
--device cuda \
--model_path $model_path \
--config $config \
--prompt_wav_path $prompt_wav_path \
--target_wav_path $target_wav_path \
--prompt_f0_path $prompt_f0_path \
--target_f0_path $target_f0_path \
--save_dir $save_dir \
--auto_shift \
--pitch_shift 0
+6 -2
View File
@@ -15,6 +15,7 @@ save_dir=example/transcriptions/zh_prompt
language=Mandarin language=Mandarin
vocal_sep=False vocal_sep=False
max_merge_duration=30000 max_merge_duration=30000
midi_transcribe=True # Whether to transcribe vocal midi, set True for singing voice synthesis, False for singing voice conversion
python -m preprocess.pipeline \ python -m preprocess.pipeline \
--audio_path $audio_path \ --audio_path $audio_path \
@@ -22,7 +23,8 @@ python -m preprocess.pipeline \
--language $language \ --language $language \
--device $device \ --device $device \
--vocal_sep $vocal_sep \ --vocal_sep $vocal_sep \
--max_merge_duration $max_merge_duration --max_merge_duration $max_merge_duration \
--midi_transcribe $midi_transcribe
####### Run Target Annotation ####### ####### Run Target Annotation #######
@@ -31,6 +33,7 @@ save_dir=example/transcriptions/music
language=Mandarin language=Mandarin
vocal_sep=True vocal_sep=True
max_merge_duration=60000 max_merge_duration=60000
midi_transcribe=True # Whether to transcribe vocal midi, set True for singing voice synthesis, False for singing voice conversion
python -m preprocess.pipeline \ python -m preprocess.pipeline \
--audio_path $audio_path \ --audio_path $audio_path \
@@ -38,4 +41,5 @@ python -m preprocess.pipeline \
--language $language \ --language $language \
--device $device \ --device $device \
--vocal_sep $vocal_sep \ --vocal_sep $vocal_sep \
--max_merge_duration $max_merge_duration --max_merge_duration $max_merge_duration \
--midi_transcribe $midi_transcribe
+35 -20
View File
@@ -16,12 +16,13 @@ from preprocess.tools import (
class PreprocessPipeline: class PreprocessPipeline:
def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000): def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000, midi_transcribe: bool = True):
self.device = device self.device = device
self.language = language self.language = language
self.save_dir = save_dir self.save_dir = save_dir
self.vocal_sep = vocal_sep self.vocal_sep = vocal_sep
self.max_merge_duration = max_merge_duration self.max_merge_duration = max_merge_duration
self.midi_transcribe = midi_transcribe
if vocal_sep: if vocal_sep:
self.vocal_separator = VocalSeparator( self.vocal_separator = VocalSeparator(
@@ -37,26 +38,31 @@ class PreprocessPipeline:
model_path="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt", model_path="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt",
device=device, device=device,
) )
self.vocal_detector = VocalDetector( if self.midi_transcribe:
cut_wavs_output_dir= f"{save_dir}/cut_wavs", self.vocal_detector = VocalDetector(
) cut_wavs_output_dir= f"{save_dir}/cut_wavs",
self.lyric_transcriber = LyricTranscriber( )
zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch", self.lyric_transcriber = LyricTranscriber(
en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo", zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
device=device en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo",
) device=device
self.note_transcriber = NoteTranscriber( )
rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt", self.note_transcriber = NoteTranscriber(
rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt", rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt",
device=device rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt",
) device=device
)
else:
self.vocal_detector = None
self.lyric_transcriber = None
self.note_transcriber = None
def run( def run(
self, self,
audio_path: str, audio_path: str,
vocal_sep: bool = True, vocal_sep: bool = None,
max_merge_duration: int = 60000, max_merge_duration: int = None,
language: str = "Mandarin" language: str = None,
) -> None: ) -> None:
vocal_sep = self.vocal_sep if vocal_sep is None else vocal_sep vocal_sep = self.vocal_sep if vocal_sep is None else vocal_sep
max_merge_duration = self.max_merge_duration if max_merge_duration is None else max_merge_duration max_merge_duration = self.max_merge_duration if max_merge_duration is None else max_merge_duration
@@ -81,7 +87,11 @@ class PreprocessPipeline:
vocal_path = output_dir / "vocal.wav" vocal_path = output_dir / "vocal.wav"
sf.write(vocal_path, vocal, sample_rate) sf.write(vocal_path, vocal, sample_rate)
vocal_f0 = self.f0_extractor.process(str(vocal_path)) vocal_f0 = self.f0_extractor.process(str(vocal_path), f0_path=str(vocal_path).replace(".wav", "_f0.npy"))
if not self.midi_transcribe or self.vocal_detector is None or self.lyric_transcriber is None or self.note_transcriber is None:
return
segments = self.vocal_detector.process(str(vocal_path), f0=vocal_f0) segments = self.vocal_detector.process(str(vocal_path), f0=vocal_f0)
metadata = [] metadata = []
@@ -124,10 +134,11 @@ def main(args):
save_dir=args.save_dir, save_dir=args.save_dir,
vocal_sep=args.vocal_sep, vocal_sep=args.vocal_sep,
max_merge_duration=args.max_merge_duration, max_merge_duration=args.max_merge_duration,
midi_transcribe=args.midi_transcribe,
) )
pipeline.run( pipeline.run(
audio_path=args.audio_path, audio_path=args.audio_path,
language=args.language language=args.language,
) )
@@ -139,8 +150,12 @@ if __name__ == "__main__":
parser.add_argument("--save_dir", type=str, required=True, help="Directory to save the output files") parser.add_argument("--save_dir", type=str, required=True, help="Directory to save the output files")
parser.add_argument("--language", type=str, default="Mandarin", help="Language of the audio") parser.add_argument("--language", type=str, default="Mandarin", help="Language of the audio")
parser.add_argument("--device", type=str, default="cuda:0", help="Device to run the models on") parser.add_argument("--device", type=str, default="cuda:0", help="Device to run the models on")
parser.add_argument("--vocal_sep", type=bool, default=True, help="Whether to perform vocal separation") parser.add_argument("--vocal_sep", type=str, default="True", help="Whether to perform vocal separation")
parser.add_argument("--max_merge_duration", type=int, default=60000, help="Maximum merged segment duration in milliseconds") parser.add_argument("--max_merge_duration", type=int, default=60000, help="Maximum merged segment duration in milliseconds")
parser.add_argument("--midi_transcribe", type=str, default="True", help="Whether to do MIDI transcription")
args = parser.parse_args() args = parser.parse_args()
args.vocal_sep = args.vocal_sep.lower() == "true"
args.midi_transcribe = args.midi_transcribe.lower() == "true"
main(args) main(args)
@@ -0,0 +1,74 @@
"""Frozen Whisper encoder wrapper (wav -> encoder embeddings)."""
from __future__ import annotations
from typing import Optional
import torch
import torch.nn as nn
import torchaudio
from transformers import WhisperFeatureExtractor, WhisperModel
WHISPER_MEL_FRAMES = 3000 # 3000 frames at 16000 Hz
class WhisperEncoder():
def __init__(
self,
device: Optional[str] = None,
) -> None:
self.fe = WhisperFeatureExtractor.from_pretrained("openai/whisper-base")
self.model = WhisperModel.from_pretrained("openai/whisper-base")
self.model = self.model.to(device or ("cuda" if torch.cuda.is_available() else "cpu"))
def encode(
self,
wav: torch.Tensor,
sr: int,
) -> torch.Tensor:
wav = torchaudio.functional.resample(wav, orig_freq=sr, new_freq=self.fe.sampling_rate) if sr != self.fe.sampling_rate else wav
wav_np = wav.cpu().detach().numpy().astype("float32", copy=False)
inputs = self.fe(
wav_np,
sampling_rate=self.fe.sampling_rate,
return_tensors="pt",
padding=False,
truncation=False,
return_attention_mask=True,
)
input_features = inputs.input_features
num_frames = input_features.shape[-1]
if num_frames < WHISPER_MEL_FRAMES:
pad = WHISPER_MEL_FRAMES - num_frames
input_features = torch.nn.functional.pad(input_features, (0, pad))
else:
input_features = input_features[..., :WHISPER_MEL_FRAMES]
input_features = input_features.to(wav.device)
if self.model.device != wav.device:
self.model = self.model.to(wav.device)
attention_mask = inputs.attention_mask.to(wav.device) if inputs.attention_mask is not None else None
encoder_out = self.model.encoder(input_features).last_hidden_state
if attention_mask is not None:
valid_mel_frames = attention_mask.sum(dim=1)
valid_enc_frames = (valid_mel_frames + 1) // 2
max_valid_enc_frames = min(int(valid_enc_frames.max().item()), encoder_out.shape[1])
encoder_out = encoder_out[:, :max_valid_enc_frames, :]
valid_len = min(int(valid_enc_frames[0].item()), max_valid_enc_frames)
if valid_len < max_valid_enc_frames:
encoder_out[0, valid_len:, :] = 0
return encoder_out
if __name__ == "__main__":
torch.manual_seed(0)
audio = torch.randn(1, 24000 * 25).float().to("cuda")
encoder = WhisperEncoder()
whisper_encoder_out = encoder.encode(audio, sr=24000)
print(whisper_encoder_out.shape)
+303
View File
@@ -0,0 +1,303 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from tqdm import tqdm
from typing import Optional, Dict, Any, List, Tuple
from soulxsinger.models.modules.vocoder import Vocoder
from soulxsinger.models.modules.decoder import CFMDecoder
from soulxsinger.models.modules.mel_transform import MelSpectrogramEncoder
from soulxsinger.models.modules.whisper_encoder import WhisperEncoder
class SoulXSingerSVC(nn.Module):
"""
SoulXSinger SVC model.
"""
def __init__(self, config: Dict):
super(SoulXSingerSVC, self).__init__()
self.audio_cfg = config.audio
enc_cfg = config.model.encoder
cfm_cfg = config.model.flow_matching
self.whisper_encoder = WhisperEncoder()
self.f0_encoder = nn.Embedding(enc_cfg["f0_bin"], enc_cfg["f0_dim"])
self.cfm_decoder = CFMDecoder(cfm_cfg)
self.mel = MelSpectrogramEncoder(self.audio_cfg)
self.vocoder = Vocoder()
@staticmethod
def f0_to_coarse(f0, f0_bin=361, f0_min=32.7031956625, f0_shift=0):
"""
Convert continuous F0 values to discrete F0 bins (SIL and C1 - B6, 361 bins).
args:
f0: continuous F0 values
f0_bin: number of F0 bins
f0_min: minimum F0 value
f0_shift: shift value for F0 bins
returns:
f0_coarse: discrete F0 bins
"""
is_torch = isinstance(f0, torch.Tensor)
uv_mask = f0 <= 0
if is_torch:
f0_safe = torch.maximum(f0, torch.tensor(f0_min))
f0_cents = 1200 * torch.log2(f0_safe / f0_min)
else:
f0_safe = np.maximum(f0, f0_min)
f0_cents = 1200 * np.log2(f0_safe / f0_min)
f0_coarse = (f0_cents / 20) + 1
if is_torch:
f0_coarse = torch.round(f0_coarse).long()
f0_coarse = torch.clamp(f0_coarse, min=1, max=f0_bin - 1)
else:
f0_coarse = np.rint(f0_coarse).astype(int)
f0_coarse = np.clip(f0_coarse, 1, f0_bin - 1)
f0_coarse[uv_mask] = 0
if f0_shift != 0:
if is_torch:
voiced = f0_coarse > 0
if voiced.any():
shifted = f0_coarse[voiced] + f0_shift
f0_coarse[voiced] = torch.clamp(shifted, 1, f0_bin - 1)
else:
voiced = f0_coarse > 0
if np.any(voiced):
shifted = f0_coarse[voiced] + f0_shift
f0_coarse[voiced] = np.clip(shifted, 1, f0_bin - 1)
return f0_coarse
@staticmethod
def build_vocal_segments(
f0,
f0_rate: int = 50,
uv_frames_th: int = 5,
min_duration_sec: float = 5.0,
max_duration_sec: float = 30.0,
num_overlaps: int = 1,
ignore_silent_segments: bool = True,
) -> Tuple[List[Tuple[float, float]], List[Tuple[float, float]]]:
"""Build vocal segments based on F0 contour. First split by long silent runs, then merge into segments based on min and max duration constraints.
args:
f0: F0 contour of the audio, 1D array or tensor with shape (T,)
f0_rate: F0 sampling rate in Hz (e.g., 50 for 20ms hop size)
uv_frames_th: number of consecutive zero F0 frames to consider as a split point
min_duration_sec: minimum duration of each segment in seconds
max_duration_sec: maximum duration of each segment in seconds
num_overlaps: number of overlapping segments to create for each non-overlapping segment (for smooth inference)
ignore_silent_segments: whether to ignore segments that are mostly silent (e.g., > 95% zero F0)
returns:
overlap_segments: list of (overlap_start_sec, overlap_end_sec) for each segment, which may overlap with adjacent segments for smooth inference
segments: list of (seg_start_sec, seg_end_sec) for each segment, which are non-overlapping and used for final merging
"""
if isinstance(f0, torch.Tensor):
f0_np = f0.detach().float().cpu().numpy()
else:
f0_np = np.asarray(f0, dtype=np.float32)
f0_np = np.squeeze(f0_np)
total_frames = int(f0_np.shape[0])
if total_frames == 0:
return [], []
min_frames = max(1, int(round(min_duration_sec * f0_rate)))
max_frames = max(1, int(round(max_duration_sec * f0_rate)))
split_points = [0] # silence split points in frame indices, starting with 0 and ending with total_frames
def append_split_point(point: int):
# Ensure split points are within valid range and respect max_frames constraint
point = int(max(0, min(point, total_frames)))
while point - split_points[-1] > max_frames:
split_points.append(split_points[-1] + max_frames)
if point > split_points[-1]:
split_points.append(point)
idx = 0
while idx < total_frames:
if f0_np[idx] == 0:
run_start = idx
while idx < total_frames and f0_np[idx] == 0:
idx += 1
run_end = idx
if (run_end - run_start) >= uv_frames_th:
split_point = max(run_end - 5, (run_start + run_end) // 2)
append_split_point(split_point)
else:
idx += 1
append_split_point(total_frames)
# print(f"Initial split points (in seconds): {[round(p / f0_rate, 2) for p in split_points]}")
segments: List[Tuple[int, int]] = []
overlap_segments: List[Tuple[int, int]] = []
def append_segment(start_idx: int, end_idx: int, num_overlaps: int = num_overlaps):
segments.append((split_points[start_idx] / f0_rate, split_points[end_idx] / f0_rate))
overlap_start_idx = start_idx
if start_idx > 0 and (split_points[end_idx] - split_points[start_idx - num_overlaps]) <= max_frames:
overlap_start_idx = start_idx - num_overlaps
overlap_segments.append((split_points[overlap_start_idx] / f0_rate, split_points[end_idx] / f0_rate))
segment_start, segment_end = 0, 1
while segment_start < len(split_points) - 1:
while segment_end < len(split_points) and (split_points[segment_end] - split_points[segment_start]) < min_frames:
segment_end += 1
if segment_end >= len(split_points):
append_segment(segment_start, len(split_points) - 1, num_overlaps=num_overlaps)
break
append_segment(segment_start, segment_end, num_overlaps=num_overlaps)
segment_start = segment_end
segment_end = segment_start + 1
# print(f"Final segments (overlap_start, overlap_end, seg_start_time, seg_end_time) in seconds: {overlap_segments}")
if ignore_silent_segments:
filtered_idx = []
for i, seg in enumerate(overlap_segments):
start_frame = int(seg[0] * f0_rate)
end_frame = int(seg[1] * f0_rate)
total_frames = end_frame - start_frame
voice_frames = np.sum(f0_np[start_frame:end_frame] > 0)
if voice_frames / total_frames > 0.05 and voice_frames >= 10: # at least 10 voiced frames and >5% voiced frames
filtered_idx.append(i)
overlap_segments = [overlap_segments[i] for i in filtered_idx]
segments = [segments[i] for i in filtered_idx]
# print(f"Filtered segments with mostly silence removed: {overlap_segments}")
return overlap_segments, segments
def infer(
self,
pt_wav: str|torch.Tensor,
gt_wav: str|torch.Tensor,
pt_f0: str|torch.Tensor,
gt_f0: str|torch.Tensor,
auto_shift=False,
pitch_shift=0,
n_steps=32,
cfg=3,
):
"""
SVC inference pipeline. First build vocal segments based on F0 contour, then run inference for each segment and merge results.
args:
pt_wav: prompt waveform path or tensor
gt_wav: target waveform path or tensor
pt_f0: prompt F0 path or tensor
gt_f0: target F0 path or tensor
auto_shift: whether to automatically calculate pitch shift based on median F0 of prompt and target
pitch_shift: manual pitch shift in semitones (overrides auto_shift if > 0)
n_steps: number of diffusion steps for inference
cfg: classifier-free guidance scale for inference
"""
# calculate auto pitch shift
if auto_shift and pitch_shift == 0:
if gt_f0 is not None and pt_f0 is not None:
gt_f0_median = torch.median(gt_f0[gt_f0 > 0])
pt_f0_median = torch.median(pt_f0[pt_f0 > 0])
pitch_shift = torch.round(torch.log2(pt_f0_median / gt_f0_median) * 1200 / 100).int().item()
else:
print("Warning: pitch_shift is True but note_pitch or f0 is None. Set f0_shift to 0.")
pitch_shift = 0
else:
pitch_shift = pitch_shift
generated_audio = []
f0_rate = self.audio_cfg.sample_rate // self.audio_cfg.hop_size
overlap_segments, segments = self.build_vocal_segments(
gt_f0,
f0_rate=f0_rate,
uv_frames_th=10,
min_duration_sec=15.0,
max_duration_sec=30.0,
)
if len(segments) == 0:
segments = [(0.0, gt_wav.shape[-1] / self.audio_cfg.sample_rate)]
overlap_segments = [(0.0, gt_wav.shape[-1] / self.audio_cfg.sample_rate)]
generated_audio = torch.zeros_like(gt_wav)
for idx in tqdm(range(len(segments)), total=len(segments), desc="Inferring segments (SVC)", dynamic_ncols=True):
overlap_start_sec, overlap_end_sec = overlap_segments[idx]
seg_start_sec, seg_end_sec = segments[idx]
wav_start = int(round(overlap_start_sec * self.audio_cfg.sample_rate))
wav_end = int(round(overlap_end_sec * self.audio_cfg.sample_rate))
f0_start = int(round(overlap_start_sec * f0_rate))
f0_end = int(round(overlap_end_sec * f0_rate))
wav_start = max(0, min(wav_start, gt_wav.shape[-1]))
wav_end = max(wav_start, min(wav_end, gt_wav.shape[-1]))
f0_start = max(0, min(f0_start, gt_f0.shape[-1]))
f0_end = max(f0_start, min(f0_end, gt_f0.shape[-1]))
segment_gt_wav = gt_wav[:, wav_start:wav_end]
segment_gt_f0 = gt_f0[:, f0_start:f0_end]
segment_generated_audio = self.infer_segment(
pt_wav=pt_wav,
gt_wav=segment_gt_wav,
pt_f0=pt_f0,
gt_f0=segment_gt_f0,
pitch_shift=pitch_shift,
n_steps=n_steps,
cfg=cfg,
)
segment_start = int(round(seg_start_sec * self.audio_cfg.sample_rate))
segment_end = int(round(seg_end_sec * self.audio_cfg.sample_rate))
segment_generated_audio = segment_generated_audio[segment_start - wav_start: segment_end - wav_start]
generated_audio[:, segment_start:segment_end] = segment_generated_audio
return generated_audio, pitch_shift
def infer_segment(self, pt_wav, gt_wav, pt_f0, gt_f0, pitch_shift=0, n_steps=32, cfg=3):
pt_mel = self.mel(pt_wav)
len_prompt_mel = pt_f0.shape[1]
f0_course_pt = self.f0_to_coarse(pt_f0)
f0_course_gt = self.f0_to_coarse(gt_f0, f0_shift=pitch_shift * 5)
f0_course = torch.cat([f0_course_pt, f0_course_gt], 1)
pt_content_feat = self.whisper_encoder.encode(pt_wav, sr=self.audio_cfg.sample_rate)
gt_content_feat = self.whisper_encoder.encode(gt_wav, sr=self.audio_cfg.sample_rate)
content_feat = torch.cat([pt_content_feat, gt_content_feat], 1)
f0_feat = self.f0_encoder(f0_course)
min_len = min(content_feat.shape[1], f0_feat.shape[1])
content_feat = content_feat[:, :min_len, :]
f0_feat = f0_feat[:, :min_len, :]
features = content_feat + f0_feat
gt_decoder_inp = features[:, len_prompt_mel:, :]
pt_decoder_inp = features[:, :len_prompt_mel, :]
generated_mel = self.cfm_decoder.reverse_diffusion(
pt_mel,
pt_decoder_inp,
gt_decoder_inp,
n_timesteps=n_steps,
cfg=cfg
)
generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...])
generated_audio = generated_audio.squeeze()
# cut or pad to match gt_wav length
if generated_audio.shape[-1] > gt_wav.shape[-1]:
generated_audio = generated_audio[:gt_wav.shape[-1]]
elif generated_audio.shape[-1] < gt_wav.shape[-1]:
generated_audio = F.pad(generated_audio, (0, gt_wav.shape[-1] - generated_audio.shape[-1]))
return generated_audio
+436
View File
@@ -0,0 +1,436 @@
import random
import sys
import traceback
from datetime import datetime
from pathlib import Path
from typing import Literal
import gradio as gr
import librosa
import numpy as np
import soundfile as sf
import torch
from preprocess.pipeline import PreprocessPipeline
from soulxsinger.utils.file_utils import load_config
from cli.inference_svc import build_model as build_svc_model, process as svc_process
ROOT = Path(__file__).parent
SAMPLE_RATE = 44100
PROMPT_MAX_SEC_DEFAULT = 30
TARGET_MAX_SEC_DEFAULT = 600
SVC_EXAMPLE_PROMPT_AUDIO = "example/audio/svc_prompt_demo.mp3"
SVC_EXAMPLE_TARGET_AUDIO = "example/audio/svc_target_demo.mp3"
EXAMPLE_LIST = [[
str(ROOT / SVC_EXAMPLE_PROMPT_AUDIO),
str(ROOT / SVC_EXAMPLE_TARGET_AUDIO),
False,
True,
True,
True,
0,
32,
1.0,
42,
]]
_I18N = dict(
display_lang_label=dict(en="Display Language", zh="显示语言"),
title=dict(en="## SoulX-Singer SVC", zh="## SoulX-Singer SVC"),
prompt_audio_label=dict(en=f"Prompt audio", zh=f"Prompt 音频"),
target_audio_label=dict(en=f"Target audio", zh=f"Target 音频"),
prompt_vocal_sep_label=dict(en="Prompt vocal separation", zh="Prompt 人声分离"),
target_vocal_sep_label=dict(en="Target vocal separation", zh="Target 人声分离"),
auto_shift_label=dict(en="Auto pitch shift", zh="自动变调"),
auto_mix_acc_label=dict(en="Auto mix accompaniment", zh="自动混合伴奏"),
pitch_shift_label=dict(en="Pitch shift (semitones)", zh="指定变调(半音)"),
n_step_label=dict(en="n_step", zh="采样步数"),
cfg_label=dict(en="cfg scale", zh="cfg系数"),
seed_label=dict(en="Seed", zh="种子"),
examples_label=dict(en="Reference example (click to load)", zh="参考样例(点击加载)"),
run_btn=dict(en="🎤Singing Voice Conversion", zh="🎤歌声转换"),
output_audio_label=dict(en="Generated audio", zh="合成结果音频"),
warn_missing_audio=dict(en="Please provide both prompt audio and target audio.", zh="请同时上传 Prompt 与 Target 音频。"),
instruction_title=dict(en="Usage", zh="使用说明"),
instruction_p1=dict(
en="Upload the Prompt and Target audio, and configure the parameters",
zh="上传 Prompt 与 Target 音频,并配置相关参数",
),
instruction_p2=dict(
en="Click the button to start singing voice conversion.",
zh="点击「🎤歌声转换」开始最终生成。",
),
tips_title=dict(en="Tips", zh="提示"),
tip_p1=dict(
en="Input: The Prompt audio is recommended to be a clean and clear singing voice, while the Target audio can be either a pure vocal or a mixture with accompaniment. If the audio contains accompaniment, please check the vocal separation option.",
zh="输入:Prompt 音频建议是干净清晰的歌声,Target 音频可以是纯歌声或伴奏,这两者若带伴奏需要勾选分离选项",
),
tip_p2=dict(
en="Pitch shift: When there is a large pitch range difference between the Prompt and Target audio, you can try enabling auto pitch shift or manually adjusting the pitch shift in semitones. When a non-zero pitch shift is specified, auto pitch shift will not take effect. The accompaniment of auto mix will be pitch-shifted together with the vocal (keeping the same octave).",
zh="变调:Prompt 音频的音域和 Target 音频的音域差距较大的时候,可以尝试开启自动变调或手动调整变调半音数,指定非0的变调半音数时,自动变调不生效,自动混音的伴奏会配合歌声进行升降调(保持同一个八度)",
),
tip_p3=dict(
en="Model parameters: Generally, a larger number of sampling steps will yield better generation quality but also longer generation time; a larger cfg scale will increase timbre similarity and melody fidelity, but may cause more distortion, it is recommended to take a value between 1 and 3.",
zh="模型参数:一般采样步数越大,生成质量越好,但生成时间也越长;一般cfg系数越大,音色相似度和旋律保真度越高,但是会造成更多的失真,建议取1~3之间的值",
),
tip_p4=dict(
en="If you want to convert a long audio or a whole song with large pitch range, there may be instability in the generated voice. You can try converting in segments.",
zh="长音频或完整歌曲中,音域变化较大的情况有可能出现音色不稳定,可以尝试分段转换",
)
)
_GLOBAL_LANG: Literal["zh", "en"] = "zh"
def _i18n(key: str) -> str:
return _I18N[key][_GLOBAL_LANG]
def _print_exception(context: str) -> None:
print(f"[{context}]\n{traceback.format_exc()}", file=sys.stderr, flush=True)
def _get_device() -> str:
return "cuda:0" if torch.cuda.is_available() else "cpu"
def _session_dir() -> Path:
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f")
return ROOT / "outputs" / "gradio" / "svc" / timestamp
def _normalize_audio_input(audio):
return audio[0] if isinstance(audio, tuple) else audio
def _trim_and_save_audio(src_audio_path: str, dst_wav_path: Path, max_sec: int, sr: int = SAMPLE_RATE) -> None:
audio_data, _ = librosa.load(src_audio_path, sr=sr, mono=True)
audio_data = audio_data[: max_sec * sr]
dst_wav_path.parent.mkdir(parents=True, exist_ok=True)
sf.write(dst_wav_path, audio_data, sr)
def _usage_md() -> str:
return "\n\n".join([
f"### {_i18n('instruction_title')}",
f"**1.** {_i18n('instruction_p1')}",
f"**2.** {_i18n('instruction_p2')}",
])
def _tips_md() -> str:
return "\n\n".join([
f"### {_i18n('tips_title')}",
f"- {_i18n('tip_p1')}",
f"- {_i18n('tip_p2')}",
f"- {_i18n('tip_p3')}",
f"- {_i18n('tip_p4')}",
])
class AppState:
def __init__(self) -> None:
self.device = _get_device()
self.preprocess_pipeline = PreprocessPipeline(
device=self.device,
language="Mandarin",
save_dir=str(ROOT / "outputs" / "gradio" / "_placeholder" / "svc"),
vocal_sep=True,
max_merge_duration=60000,
midi_transcribe=False,
)
self.svc_config = load_config("soulxsinger/config/soulxsinger.yaml")
self.svc_model = build_svc_model(
model_path="pretrained_models/SoulX-Singer-SVC/model.pt",
config=self.svc_config,
device=self.device,
)
def run_preprocess(self, audio_path: Path, save_path: Path, vocal_sep: bool) -> tuple[bool, str, Path | None, Path | None]:
try:
self.preprocess_pipeline.save_dir = str(save_path)
self.preprocess_pipeline.run(
audio_path=str(audio_path),
vocal_sep=vocal_sep,
max_merge_duration=60000,
language="Mandarin",
)
vocal_wav = save_path / "vocal.wav"
vocal_f0 = save_path / "vocal_f0.npy"
if not vocal_wav.exists() or not vocal_f0.exists():
return False, f"preprocess output missing: {vocal_wav} or {vocal_f0}", None, None
return True, "ok", vocal_wav, vocal_f0
except Exception as e:
return False, f"preprocess failed: {e}", None, None
def run_svc(
self,
prompt_wav_path: Path,
target_wav_path: Path,
prompt_f0_path: Path,
target_f0_path: Path,
session_base: Path,
auto_shift: bool,
auto_mix_acc: bool,
pitch_shift: int,
n_step: int,
cfg: float,
seed: int,
) -> tuple[bool, str, Path | None]:
try:
torch.manual_seed(seed)
np.random.seed(seed)
random.seed(seed)
save_dir = session_base / "generated"
save_dir.mkdir(parents=True, exist_ok=True)
class Args:
pass
args = Args()
args.device = self.device
args.model_path = "soulx-singer-svc.pt"
args.config = "soulxsinger/config/soulxsinger.yaml"
args.prompt_wav_path = str(prompt_wav_path)
args.target_wav_path = str(target_wav_path)
args.prompt_f0_path = str(prompt_f0_path)
args.target_f0_path = str(target_f0_path)
args.save_dir = str(save_dir)
args.auto_shift = auto_shift
args.pitch_shift = int(pitch_shift)
args.n_steps = int(n_step)
args.cfg = float(cfg)
svc_process(args, self.svc_config, self.svc_model)
generated = save_dir / "generated.wav"
if not generated.exists():
return False, f"inference finished but output not found: {generated}", None
if auto_mix_acc:
acc_path = session_base / "transcriptions" / "target" / "acc.wav"
if acc_path.exists():
vocal_shift = args.pitch_shift
mul = -1 if vocal_shift < 0 else 1
acc_shift = abs(vocal_shift) % 12
acc_shift = mul * acc_shift
if acc_shift > 6:
acc_shift -= 12
if acc_shift < -6:
acc_shift += 12
mix_sr = self.svc_config.audio.sample_rate
vocal, _ = librosa.load(str(generated), sr=mix_sr, mono=True)
acc, _ = librosa.load(str(acc_path), sr=mix_sr, mono=True)
if acc_shift != 0:
acc = librosa.effects.pitch_shift(acc, sr=mix_sr, n_steps=acc_shift)
print(f"Applied pitch shift of {acc_shift} semitones to accompaniment to match vocal shift of {vocal_shift} semitones.")
mix_len = min(len(vocal), len(acc))
if mix_len > 0:
mixed = vocal[:mix_len] + acc[:mix_len]
peak = float(np.max(np.abs(mixed))) if mixed.size > 0 else 1.0
if peak > 1.0:
mixed = mixed / peak
mixed_path = save_dir / "generated_mixed.wav"
sf.write(str(mixed_path), mixed, mix_sr)
generated = mixed_path
return True, "svc inference done", generated
except Exception as e:
return False, f"svc inference failed: {e}", None
APP_STATE = AppState()
def _start_svc(prompt_audio, target_audio, prompt_vocal_sep, target_vocal_sep, auto_shift, auto_mix_acc, pitch_shift, n_step, cfg, seed):
try:
prompt_audio = _normalize_audio_input(prompt_audio)
target_audio = _normalize_audio_input(target_audio)
if not prompt_audio or not target_audio:
gr.Warning(_i18n("warn_missing_audio"))
return None
session_base = _session_dir()
audio_dir = session_base / "audio"
prompt_raw = audio_dir / "prompt.wav"
target_raw = audio_dir / "target.wav"
_trim_and_save_audio(prompt_audio, prompt_raw, PROMPT_MAX_SEC_DEFAULT)
_trim_and_save_audio(target_audio, target_raw, TARGET_MAX_SEC_DEFAULT)
prompt_ok, prompt_msg, prompt_wav, prompt_f0 = APP_STATE.run_preprocess(
audio_path=prompt_raw,
save_path=session_base / "transcriptions" / "prompt",
vocal_sep=bool(prompt_vocal_sep),
)
if not prompt_ok or prompt_wav is None or prompt_f0 is None:
print(prompt_msg, file=sys.stderr, flush=True)
return None
target_ok, target_msg, target_wav, target_f0 = APP_STATE.run_preprocess(
audio_path=target_raw,
save_path=session_base / "transcriptions" / "target",
vocal_sep=bool(target_vocal_sep),
)
if not target_ok or target_wav is None or target_f0 is None:
print(target_msg, file=sys.stderr, flush=True)
return None
ok, msg, generated = APP_STATE.run_svc(
prompt_wav_path=prompt_wav,
target_wav_path=target_wav,
prompt_f0_path=prompt_f0,
target_f0_path=target_f0,
session_base=session_base,
auto_shift=bool(auto_shift),
auto_mix_acc=bool(auto_mix_acc),
pitch_shift=int(pitch_shift),
n_step=int(n_step),
cfg=float(cfg),
seed=int(seed),
)
if not ok or generated is None:
print(msg, file=sys.stderr, flush=True)
return None
return str(generated)
except Exception:
_print_exception("_start_svc")
return None
def render_interface() -> gr.Blocks:
with gr.Blocks(title="SoulX-Singer SVC Demo", theme=gr.themes.Default()) as page:
with gr.Row(equal_height=True):
lang_choice = gr.Radio(
choices=["中文", "English"],
value="中文",
label=_i18n("display_lang_label"),
type="index",
interactive=True,
)
title_md = gr.Markdown(_i18n("title"))
usage_md = gr.Markdown(_usage_md())
with gr.Row(equal_height=True):
prompt_audio = gr.Audio(
label=_i18n("prompt_audio_label"),
type="filepath",
editable=False,
interactive=True,
)
target_audio = gr.Audio(
label=_i18n("target_audio_label"),
type="filepath",
editable=False,
interactive=True,
)
with gr.Row(equal_height=True):
prompt_vocal_sep = gr.Checkbox(label=_i18n("prompt_vocal_sep_label"), value=False, scale=1)
target_vocal_sep = gr.Checkbox(label=_i18n("target_vocal_sep_label"), value=True, scale=1)
auto_shift = gr.Checkbox(label=_i18n("auto_shift_label"), value=True, scale=1)
auto_mix_acc = gr.Checkbox(label=_i18n("auto_mix_acc_label"), value=True, scale=1)
with gr.Row(equal_height=True):
pitch_shift = gr.Slider(label=_i18n("pitch_shift_label"), value=0, minimum=-36, maximum=36, step=1, scale=1)
n_step = gr.Slider(label=_i18n("n_step_label"), value=32, minimum=1, maximum=200, step=1, scale=1)
cfg = gr.Slider(label=_i18n("cfg_label"), value=1.0, minimum=0.0, maximum=10.0, step=0.1, scale=1)
seed_input = gr.Slider(label=_i18n("seed_label"), value=42, minimum=0, maximum=10000, step=1, scale=1)
with gr.Row():
run_btn = gr.Button(value=_i18n("run_btn"), variant="primary", size="lg")
with gr.Row():
output_audio = gr.Audio(label=_i18n("output_audio_label"), type="filepath", interactive=False)
gr.Examples(
examples=EXAMPLE_LIST,
inputs=[prompt_audio, target_audio],
label=_i18n("examples_label"),
)
tips_md = gr.Markdown(_tips_md())
run_btn.click(
fn=_start_svc,
inputs=[
prompt_audio,
target_audio,
prompt_vocal_sep,
target_vocal_sep,
auto_shift,
auto_mix_acc,
pitch_shift,
n_step,
cfg,
seed_input,
],
outputs=[output_audio],
)
def _change_language(lang):
global _GLOBAL_LANG
_GLOBAL_LANG = ["zh", "en"][lang]
return [
gr.update(label=_i18n("display_lang_label")),
gr.update(value=_i18n("title")),
gr.update(value=_usage_md()),
gr.update(label=_i18n("prompt_audio_label")),
gr.update(label=_i18n("target_audio_label")),
gr.update(label=_i18n("prompt_vocal_sep_label")),
gr.update(label=_i18n("target_vocal_sep_label")),
gr.update(label=_i18n("auto_shift_label")),
gr.update(label=_i18n("auto_mix_acc_label")),
gr.update(label=_i18n("pitch_shift_label")),
gr.update(label=_i18n("n_step_label")),
gr.update(label=_i18n("cfg_label")),
gr.update(label=_i18n("seed_label")),
gr.update(value=_i18n("run_btn")),
gr.update(label=_i18n("output_audio_label")),
gr.update(value=_tips_md()),
]
lang_choice.change(
fn=_change_language,
inputs=[lang_choice],
outputs=[
lang_choice,
title_md,
usage_md,
prompt_audio,
target_audio,
prompt_vocal_sep,
target_vocal_sep,
auto_shift,
auto_mix_acc,
pitch_shift,
n_step,
cfg,
seed_input,
run_btn,
output_audio,
tips_md,
],
)
return page
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--port", type=int, default=7861, help="Gradio server port")
parser.add_argument("--share", action="store_true", help="Create public link")
args = parser.parse_args()
page = render_interface()
page.queue()
page.launch(share=args.share, server_name="0.0.0.0", server_port=args.port)