add fp16 support for svs and svc inference

This commit is contained in:
jlqian98
2026-03-16 13:40:30 +08:00
parent acd498ad9e
commit 6bc423f3af
9 changed files with 143 additions and 57 deletions
+6
View File
@@ -425,8 +425,14 @@ class ISTFTHead(FourierHead):
# phase = torch.atan2(y, x)
# S = mag * torch.exp(phase * 1j)
# better directly produce the complex value
# Always compute complex values in float32 to avoid ComplexHalf warning (then cast audio back if needed)
orig_dtype = mag.dtype
mag, x, y = mag.float(), x.float(), y.float()
S = mag * (x + 1j * y)
audio = self.istft(S)
if orig_dtype != torch.float32:
audio = audio.to(orig_dtype)
return audio
+31 -21
View File
@@ -4,6 +4,7 @@ import torch.nn.functional as F
import math
import numpy as np
from typing import Optional, Dict, Any, List
from contextlib import nullcontext
from soulxsinger.models.modules.vocoder import Vocoder
from soulxsinger.models.modules.decoder import CFMDecoder
@@ -11,6 +12,10 @@ from soulxsinger.models.modules.convnext import ConvNeXtV2Block
from soulxsinger.models.modules.mel_transform import MelSpectrogramEncoder
def _autocast_if(enabled: bool):
"""Return autocast(context) if enabled else no-op context. Use: with _autocast_if(use_amp): ..."""
return torch.amp.autocast(device_type="cuda", enabled=True) if enabled else nullcontext()
class SoulXSinger(nn.Module):
"""
SoulXSinger model.
@@ -102,7 +107,7 @@ class SoulXSinger(nn.Module):
return f0_coarse
def infer(self, meta: dict, auto_shift=False, pitch_shift=0, n_steps=32, cfg=3, control="melody"):
def infer(self, meta: dict, auto_shift=False, pitch_shift=0, n_steps=32, cfg=3, control="melody", use_fp16=False):
gt_note_text = meta['target']['phoneme']
gt_mel2note = meta['target']['mel2note']
@@ -147,8 +152,13 @@ class SoulXSinger(nn.Module):
if gt_note_pitch is None or pt_note_pitch is None:
gt_note_pitch, pt_note_pitch = torch.zeros_like(gt_note_type).int().to(gt_note_type.device), torch.zeros_like(pt_note_type).int().to(pt_note_type.device)
# convert prompt waveform to mel spectrogram
pt_mel = self.mel(pt_wav)
use_fp16 = use_fp16 and pt_wav.is_cuda
# mel is kept in fp32 (see build_model: model.mel.float() after model.half())
pt_mel = self.mel(pt_wav.float() if pt_wav.dtype != torch.float32 else pt_wav)
if use_fp16:
pt_mel = pt_mel.half()
pt_f0 = pt_f0.half()
gt_f0 = gt_f0.half()
len_prompt = pt_note_pitch.shape[1]
len_prompt_mel = pt_f0.shape[1]
@@ -165,23 +175,23 @@ class SoulXSinger(nn.Module):
note_pitch[note_pitch > 0] = note_pitch[note_pitch > 0] + f0_shift
note_pitch = torch.clamp(note_pitch, 0, 255)
features = self.note_pitch_encoder(note_pitch) + self.note_type_encoder(note_type) + self.note_text_encoder(note_text)
features = self.preflow(features)
features = self.expand_states(features, mel2note)
features = features + self.f0_encoder(f0_course)
gt_decoder_inp = features[:, len_prompt_mel:, :]
pt_decoder_inp = features[:, :len_prompt_mel, :]
with _autocast_if(use_fp16):
features = self.note_pitch_encoder(note_pitch) + self.note_type_encoder(note_type) + self.note_text_encoder(note_text)
features = self.preflow(features)
features = self.expand_states(features, mel2note)
features = features + self.f0_encoder(f0_course)
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, ...]).float()
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, ...])
return generated_audio
+41 -21
View File
@@ -4,6 +4,7 @@ import torch.nn.functional as F
import numpy as np
from tqdm import tqdm
from typing import Optional, Dict, Any, List, Tuple
from contextlib import nullcontext
from soulxsinger.models.modules.vocoder import Vocoder
from soulxsinger.models.modules.decoder import CFMDecoder
@@ -11,6 +12,10 @@ from soulxsinger.models.modules.mel_transform import MelSpectrogramEncoder
from soulxsinger.models.modules.whisper_encoder import WhisperEncoder
def _autocast_if(enabled: bool):
"""Return autocast(context) if enabled else no-op context. Use: with _autocast_if(use_amp): ..."""
return torch.amp.autocast(device_type="cuda", enabled=True) if enabled else nullcontext()
class SoulXSingerSVC(nn.Module):
"""
SoulXSinger SVC model.
@@ -186,6 +191,7 @@ class SoulXSingerSVC(nn.Module):
pitch_shift=0,
n_steps=32,
cfg=3,
use_fp16=False,
):
"""
SVC inference pipeline. First build vocal segments based on F0 contour, then run inference for each segment and merge results.
@@ -198,6 +204,7 @@ class SoulXSingerSVC(nn.Module):
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
use_fp16: if True, run in FP16 except mel extraction to save memory and speed.
"""
# calculate auto pitch shift
@@ -212,17 +219,29 @@ class SoulXSingerSVC(nn.Module):
else:
pitch_shift = pitch_shift
use_fp16 = use_fp16 and pt_wav.is_cuda
# mel is kept in fp32 (see build_model: model.mel.float() after model.half())
pt_mel = self.mel(pt_wav.float() if pt_wav.dtype != torch.float32 else pt_wav)
if use_fp16:
pt_mel = pt_mel.half()
pt_wav = pt_wav.half()
gt_wav = gt_wav.half()
pt_f0 = pt_f0.half()
gt_f0 = gt_f0.half()
# if target audio is less than 30 seconds, infer the whole audio
if gt_wav.shape[-1] < 30 * self.audio_cfg.sample_rate:
generated_audio = self.infer_segment(
pt_wav=pt_wav,
gt_wav=gt_wav,
pt_f0=pt_f0,
gt_f0=gt_f0,
pitch_shift=pitch_shift,
n_steps=n_steps,
cfg=cfg,
)
with _autocast_if(use_fp16):
generated_audio = self.infer_segment(
pt_mel=pt_mel,
pt_wav=pt_wav,
gt_wav=gt_wav,
pt_f0=pt_f0,
gt_f0=gt_f0,
pitch_shift=pitch_shift,
n_steps=n_steps,
cfg=cfg,
)
return generated_audio, pitch_shift
# if target audio is longer than 30 seconds, build vocal segments and infer each segment
@@ -258,15 +277,17 @@ class SoulXSingerSVC(nn.Module):
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,
)
with _autocast_if(use_fp16):
segment_generated_audio = self.infer_segment(
pt_mel=pt_mel,
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))
@@ -276,8 +297,7 @@ class SoulXSingerSVC(nn.Module):
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)
def infer_segment(self, pt_mel, pt_wav, gt_wav, pt_f0, gt_f0, pitch_shift=0, n_steps=32, cfg=3):
len_prompt_mel = pt_mel.shape[1]
pt_f0 = F.pad(pt_f0, (0, 0, 0, max(0, len_prompt_mel - pt_f0.shape[1])))[:, :len_prompt_mel]
@@ -308,7 +328,7 @@ class SoulXSingerSVC(nn.Module):
)
generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...])
generated_audio = generated_audio.squeeze()
generated_audio = generated_audio.squeeze().float()
# cut or pad to match gt_wav length
if generated_audio.shape[-1] > gt_wav.shape[-1]: