Add BigVGAN vocoder with CUDA support
Browse files- Auto-detects GPU at startup, loads BigVGAN from HuggingFace Hub
- Falls back to Griffin-LIM on CPU
- Loads fine-tuned BigVGAN weights from models/vocoder/ if present
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
- Dockerfile +1 -2
- main.py +96 -42
- requirements.txt +1 -0
Dockerfile
CHANGED
|
@@ -17,11 +17,10 @@ RUN useradd -m -u 1000 user
|
|
| 17 |
# Install uv for faster pip installs
|
| 18 |
RUN pip install --no-cache-dir uv
|
| 19 |
|
| 20 |
-
# Install
|
| 21 |
RUN uv pip install --system --no-cache-dir \
|
| 22 |
torch \
|
| 23 |
torchaudio \
|
| 24 |
-
--index-url https://download.pytorch.org/whl/cpu \
|
| 25 |
&& uv cache clean
|
| 26 |
|
| 27 |
# Install app dependencies
|
|
|
|
| 17 |
# Install uv for faster pip installs
|
| 18 |
RUN pip install --no-cache-dir uv
|
| 19 |
|
| 20 |
+
# Install PyTorch with CUDA support (GPU Space)
|
| 21 |
RUN uv pip install --system --no-cache-dir \
|
| 22 |
torch \
|
| 23 |
torchaudio \
|
|
|
|
| 24 |
&& uv cache clean
|
| 25 |
|
| 26 |
# Install app dependencies
|
main.py
CHANGED
|
@@ -1,9 +1,10 @@
|
|
| 1 |
"""Kicks API β VAE kick drum synthesizer deployed on Hugging Face Spaces.
|
| 2 |
|
| 3 |
-
Uses
|
| 4 |
PCA parameters are pre-computed and loaded from pca_params.npz.
|
| 5 |
"""
|
| 6 |
|
|
|
|
| 7 |
import io
|
| 8 |
import os
|
| 9 |
|
|
@@ -34,6 +35,9 @@ LOG_MEL_MAX = 2.5 # headroom above observed max
|
|
| 34 |
N_PCS = 5
|
| 35 |
LATENT_DIM = 128
|
| 36 |
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 39 |
# VAE model (identical architecture to kicks/model.py)
|
|
@@ -45,7 +49,6 @@ class VAE(nn.Module):
|
|
| 45 |
super().__init__()
|
| 46 |
self.latent_dim = latent_dim
|
| 47 |
|
| 48 |
-
# Encoder: (B, 1, 128, 256) β (B, 256, 8, 16)
|
| 49 |
self.encoder = nn.Sequential(
|
| 50 |
nn.Conv2d(1, 32, 3, stride=2, padding=1), nn.BatchNorm2d(32), nn.ReLU(),
|
| 51 |
nn.Conv2d(32, 64, 3, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(),
|
|
@@ -54,8 +57,6 @@ class VAE(nn.Module):
|
|
| 54 |
)
|
| 55 |
self.fc_mu = nn.Linear(256 * 8 * 16, latent_dim)
|
| 56 |
self.fc_logvar = nn.Linear(256 * 8 * 16, latent_dim)
|
| 57 |
-
|
| 58 |
-
# Decoder: z β (B, 1, 128, 256)
|
| 59 |
self.fc_decode = nn.Linear(latent_dim, 256 * 8 * 16)
|
| 60 |
self.decoder = nn.Sequential(
|
| 61 |
nn.ConvTranspose2d(256, 128, 3, stride=2, padding=1, output_padding=1),
|
|
@@ -75,10 +76,44 @@ class VAE(nn.Module):
|
|
| 75 |
|
| 76 |
|
| 77 |
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 78 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 79 |
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 80 |
class GriffinLimVocoder:
|
| 81 |
-
"""Mel spectrogram β audio via pseudo-inverse
|
| 82 |
|
| 83 |
def __init__(self, n_iter: int = 64):
|
| 84 |
n_stft = N_FFT // 2 + 1
|
|
@@ -89,22 +124,16 @@ class GriffinLimVocoder:
|
|
| 89 |
n_mels=N_MELS,
|
| 90 |
sample_rate=SAMPLE_RATE,
|
| 91 |
)
|
| 92 |
-
# fb shape: (n_stft, n_mels) β fb.T shape: (n_mels, n_stft)
|
| 93 |
-
# pinv(fb.T) shape: (n_stft, n_mels)
|
| 94 |
self.fb_pinv = torch.linalg.pinv(fb.T)
|
| 95 |
self.griffin_lim = torchaudio.transforms.GriffinLim(
|
| 96 |
-
n_fft=N_FFT,
|
| 97 |
-
|
| 98 |
-
hop_length=HOP_LENGTH,
|
| 99 |
-
n_iter=n_iter,
|
| 100 |
)
|
| 101 |
|
| 102 |
def __call__(self, log_mel: torch.Tensor) -> torch.Tensor:
|
| 103 |
-
"""Input: (B, n_mels, T). Output: (B, T)."""
|
| 104 |
mel_linear = torch.exp(log_mel).cpu()
|
| 105 |
-
# Linear spectrogram via pseudo-inverse
|
| 106 |
linear_spec = torch.clamp(self.fb_pinv @ mel_linear, min=0.0)
|
| 107 |
-
linear_spec = linear_spec + 1e-4
|
| 108 |
return self.griffin_lim(linear_spec)
|
| 109 |
|
| 110 |
|
|
@@ -122,9 +151,10 @@ def denormalize(spec: torch.Tensor) -> torch.Tensor:
|
|
| 122 |
class AppState:
|
| 123 |
device: torch.device
|
| 124 |
model: VAE
|
| 125 |
-
vocoder:
|
| 126 |
-
|
| 127 |
-
|
|
|
|
| 128 |
pc_names: list[str]
|
| 129 |
pc_mins: list[float]
|
| 130 |
pc_maxs: list[float]
|
|
@@ -147,13 +177,19 @@ app.add_middleware(
|
|
| 147 |
|
| 148 |
@app.on_event("startup")
|
| 149 |
async def startup():
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 153 |
pca_path = os.environ.get("PCA_PARAMS_PATH", "pca_params.npz")
|
| 154 |
pca_data = np.load(pca_path, allow_pickle=True)
|
| 155 |
-
state.pca_components = pca_data["components"]
|
| 156 |
-
state.pca_mean = pca_data["mean"]
|
| 157 |
state.pc_names = list(pca_data["pc_names"])
|
| 158 |
state.pc_mins = list(pca_data["pc_mins"])
|
| 159 |
state.pc_maxs = list(pca_data["pc_maxs"])
|
|
@@ -161,21 +197,32 @@ async def startup():
|
|
| 161 |
|
| 162 |
# Load VAE checkpoint
|
| 163 |
ckpt_path = os.environ.get("VAE_CHECKPOINT_PATH", "vae_best.pth")
|
| 164 |
-
checkpoint = torch.load(ckpt_path, map_location=
|
| 165 |
latent_dim = checkpoint.get("latent_dim", LATENT_DIM)
|
| 166 |
state.model = VAE(latent_dim=latent_dim)
|
| 167 |
state.model.load_state_dict(checkpoint["model"])
|
|
|
|
| 168 |
state.model.eval()
|
| 169 |
print(f"VAE loaded (latent_dim={latent_dim})")
|
| 170 |
|
| 171 |
-
#
|
| 172 |
-
state.
|
| 173 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 174 |
print("API ready")
|
| 175 |
|
| 176 |
|
| 177 |
def _parse_pc_values(request: Request) -> list[float]:
|
| 178 |
-
"""Map slider positions [0, 1] to PC-space values."""
|
| 179 |
values = []
|
| 180 |
for i in range(N_PCS):
|
| 181 |
raw = float(request.query_params.get(f"pc{i + 1}", "0.5"))
|
|
@@ -185,17 +232,26 @@ def _parse_pc_values(request: Request) -> list[float]:
|
|
| 185 |
|
| 186 |
|
| 187 |
def _pc_to_latent(pc_values: list[float]) -> torch.Tensor:
|
| 188 |
-
"""Transform 5D PC slider values β 128D latent vector via PCA inverse."""
|
| 189 |
z_np = (np.array(pc_values) @ state.pca_components) + state.pca_mean
|
| 190 |
-
return torch.tensor(z_np, dtype=torch.float32).unsqueeze(0)
|
| 191 |
|
| 192 |
|
| 193 |
-
def
|
| 194 |
-
"""
|
| 195 |
log_mel = denormalize(spec)
|
| 196 |
log_mel = log_mel.squeeze(1) # (1, 128, 256)
|
| 197 |
-
|
| 198 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 199 |
return waveform.squeeze(0).numpy()
|
| 200 |
|
| 201 |
|
|
@@ -210,14 +266,12 @@ async def get_config():
|
|
| 210 |
{
|
| 211 |
"id": i + 1,
|
| 212 |
"name": state.pc_names[i],
|
| 213 |
-
"min": 0.0,
|
| 214 |
-
"
|
| 215 |
-
"default": 0.5,
|
| 216 |
-
"step": 0.01,
|
| 217 |
}
|
| 218 |
for i in range(N_PCS)
|
| 219 |
],
|
| 220 |
-
"vocoder":
|
| 221 |
}
|
| 222 |
|
| 223 |
|
|
@@ -227,8 +281,8 @@ async def generate(request: Request):
|
|
| 227 |
z = _pc_to_latent(pc_values)
|
| 228 |
|
| 229 |
with torch.no_grad():
|
| 230 |
-
spec = state.model.decode(z)
|
| 231 |
-
audio =
|
| 232 |
|
| 233 |
buf = io.BytesIO()
|
| 234 |
sf.write(buf, audio, SAMPLE_RATE, format="WAV")
|
|
@@ -250,4 +304,4 @@ async def spectrogram_data(request: Request):
|
|
| 250 |
|
| 251 |
@app.get("/health")
|
| 252 |
async def health():
|
| 253 |
-
return {"status": "ok"}
|
|
|
|
| 1 |
"""Kicks API β VAE kick drum synthesizer deployed on Hugging Face Spaces.
|
| 2 |
|
| 3 |
+
Uses BigVGAN vocoder on GPU (falls back to Griffin-LIM on CPU).
|
| 4 |
PCA parameters are pre-computed and loaded from pca_params.npz.
|
| 5 |
"""
|
| 6 |
|
| 7 |
+
import glob
|
| 8 |
import io
|
| 9 |
import os
|
| 10 |
|
|
|
|
| 35 |
N_PCS = 5
|
| 36 |
LATENT_DIM = 128
|
| 37 |
|
| 38 |
+
VOCODER_DIR = "models/vocoder" # fine-tuned BigVGAN weights
|
| 39 |
+
BIGVGAN_MODEL = "nvidia/bigvgan_v2_44khz_128band_256x"
|
| 40 |
+
|
| 41 |
|
| 42 |
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 43 |
# VAE model (identical architecture to kicks/model.py)
|
|
|
|
| 49 |
super().__init__()
|
| 50 |
self.latent_dim = latent_dim
|
| 51 |
|
|
|
|
| 52 |
self.encoder = nn.Sequential(
|
| 53 |
nn.Conv2d(1, 32, 3, stride=2, padding=1), nn.BatchNorm2d(32), nn.ReLU(),
|
| 54 |
nn.Conv2d(32, 64, 3, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(),
|
|
|
|
| 57 |
)
|
| 58 |
self.fc_mu = nn.Linear(256 * 8 * 16, latent_dim)
|
| 59 |
self.fc_logvar = nn.Linear(256 * 8 * 16, latent_dim)
|
|
|
|
|
|
|
| 60 |
self.fc_decode = nn.Linear(latent_dim, 256 * 8 * 16)
|
| 61 |
self.decoder = nn.Sequential(
|
| 62 |
nn.ConvTranspose2d(256, 128, 3, stride=2, padding=1, output_padding=1),
|
|
|
|
| 76 |
|
| 77 |
|
| 78 |
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 79 |
+
# BigVGAN vocoder
|
| 80 |
+
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 81 |
+
def load_bigvgan(device: torch.device):
|
| 82 |
+
"""Load BigVGAN from HF Hub, optionally with fine-tuned weights."""
|
| 83 |
+
import bigvgan as _bigvgan
|
| 84 |
+
|
| 85 |
+
# Patch for huggingface_hub >= 1.0 compatibility
|
| 86 |
+
original = _bigvgan.BigVGAN._from_pretrained.__func__
|
| 87 |
+
|
| 88 |
+
@classmethod
|
| 89 |
+
def _patched(cls, **kwargs):
|
| 90 |
+
kwargs.setdefault("proxies", None)
|
| 91 |
+
kwargs.setdefault("resume_download", False)
|
| 92 |
+
return original(cls, **kwargs)
|
| 93 |
+
|
| 94 |
+
_bigvgan.BigVGAN._from_pretrained = _patched
|
| 95 |
+
|
| 96 |
+
model = _bigvgan.BigVGAN.from_pretrained(BIGVGAN_MODEL, use_cuda_kernel=False)
|
| 97 |
+
|
| 98 |
+
# Load fine-tuned weights if available
|
| 99 |
+
pth_files = sorted(glob.glob(os.path.join(VOCODER_DIR, "*.pth")))
|
| 100 |
+
if pth_files:
|
| 101 |
+
pth_path = pth_files[0]
|
| 102 |
+
checkpoint = torch.load(pth_path, map_location=device, weights_only=False)
|
| 103 |
+
model.load_state_dict(checkpoint["generator"])
|
| 104 |
+
print(f"Loaded fine-tuned BigVGAN from {pth_path}")
|
| 105 |
+
else:
|
| 106 |
+
print("Using pretrained BigVGAN from HuggingFace Hub")
|
| 107 |
+
|
| 108 |
+
model.remove_weight_norm()
|
| 109 |
+
return model.eval().to(device)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 113 |
+
# Griffin-LIM fallback vocoder
|
| 114 |
# ββββββββββββββββββββββββββββββββββββββββββββββ
|
| 115 |
class GriffinLimVocoder:
|
| 116 |
+
"""Mel spectrogram β audio via pseudo-inverse + Griffin-LIM."""
|
| 117 |
|
| 118 |
def __init__(self, n_iter: int = 64):
|
| 119 |
n_stft = N_FFT // 2 + 1
|
|
|
|
| 124 |
n_mels=N_MELS,
|
| 125 |
sample_rate=SAMPLE_RATE,
|
| 126 |
)
|
|
|
|
|
|
|
| 127 |
self.fb_pinv = torch.linalg.pinv(fb.T)
|
| 128 |
self.griffin_lim = torchaudio.transforms.GriffinLim(
|
| 129 |
+
n_fft=N_FFT, win_length=WIN_SIZE,
|
| 130 |
+
hop_length=HOP_LENGTH, n_iter=n_iter,
|
|
|
|
|
|
|
| 131 |
)
|
| 132 |
|
| 133 |
def __call__(self, log_mel: torch.Tensor) -> torch.Tensor:
|
|
|
|
| 134 |
mel_linear = torch.exp(log_mel).cpu()
|
|
|
|
| 135 |
linear_spec = torch.clamp(self.fb_pinv @ mel_linear, min=0.0)
|
| 136 |
+
linear_spec = linear_spec + 1e-4
|
| 137 |
return self.griffin_lim(linear_spec)
|
| 138 |
|
| 139 |
|
|
|
|
| 151 |
class AppState:
|
| 152 |
device: torch.device
|
| 153 |
model: VAE
|
| 154 |
+
vocoder: object
|
| 155 |
+
vocoder_type: str
|
| 156 |
+
pca_components: np.ndarray
|
| 157 |
+
pca_mean: np.ndarray
|
| 158 |
pc_names: list[str]
|
| 159 |
pc_mins: list[float]
|
| 160 |
pc_maxs: list[float]
|
|
|
|
| 177 |
|
| 178 |
@app.on_event("startup")
|
| 179 |
async def startup():
|
| 180 |
+
# Auto-detect device
|
| 181 |
+
if torch.cuda.is_available():
|
| 182 |
+
state.device = torch.device("cuda")
|
| 183 |
+
print("Using CUDA GPU")
|
| 184 |
+
else:
|
| 185 |
+
state.device = torch.device("cpu")
|
| 186 |
+
print("CUDA not available β falling back to CPU")
|
| 187 |
+
|
| 188 |
+
# Load PCA parameters
|
| 189 |
pca_path = os.environ.get("PCA_PARAMS_PATH", "pca_params.npz")
|
| 190 |
pca_data = np.load(pca_path, allow_pickle=True)
|
| 191 |
+
state.pca_components = pca_data["components"]
|
| 192 |
+
state.pca_mean = pca_data["mean"]
|
| 193 |
state.pc_names = list(pca_data["pc_names"])
|
| 194 |
state.pc_mins = list(pca_data["pc_mins"])
|
| 195 |
state.pc_maxs = list(pca_data["pc_maxs"])
|
|
|
|
| 197 |
|
| 198 |
# Load VAE checkpoint
|
| 199 |
ckpt_path = os.environ.get("VAE_CHECKPOINT_PATH", "vae_best.pth")
|
| 200 |
+
checkpoint = torch.load(ckpt_path, map_location=state.device, weights_only=False)
|
| 201 |
latent_dim = checkpoint.get("latent_dim", LATENT_DIM)
|
| 202 |
state.model = VAE(latent_dim=latent_dim)
|
| 203 |
state.model.load_state_dict(checkpoint["model"])
|
| 204 |
+
state.model.to(state.device)
|
| 205 |
state.model.eval()
|
| 206 |
print(f"VAE loaded (latent_dim={latent_dim})")
|
| 207 |
|
| 208 |
+
# Load vocoder β try BigVGAN with CUDA, fall back to Griffin-LIM
|
| 209 |
+
if state.device.type == "cuda":
|
| 210 |
+
try:
|
| 211 |
+
state.vocoder = load_bigvgan(state.device)
|
| 212 |
+
state.vocoder_type = "bigvgan"
|
| 213 |
+
print("BigVGAN vocoder ready")
|
| 214 |
+
except Exception as e:
|
| 215 |
+
print(f"BigVGAN loading failed ({e}), falling back to Griffin-LIM")
|
| 216 |
+
state.vocoder = GriffinLimVocoder()
|
| 217 |
+
state.vocoder_type = "griffinlim"
|
| 218 |
+
else:
|
| 219 |
+
state.vocoder = GriffinLimVocoder()
|
| 220 |
+
state.vocoder_type = "griffinlim"
|
| 221 |
+
|
| 222 |
print("API ready")
|
| 223 |
|
| 224 |
|
| 225 |
def _parse_pc_values(request: Request) -> list[float]:
|
|
|
|
| 226 |
values = []
|
| 227 |
for i in range(N_PCS):
|
| 228 |
raw = float(request.query_params.get(f"pc{i + 1}", "0.5"))
|
|
|
|
| 232 |
|
| 233 |
|
| 234 |
def _pc_to_latent(pc_values: list[float]) -> torch.Tensor:
|
|
|
|
| 235 |
z_np = (np.array(pc_values) @ state.pca_components) + state.pca_mean
|
| 236 |
+
return torch.tensor(z_np, dtype=torch.float32).to(state.device).unsqueeze(0)
|
| 237 |
|
| 238 |
|
| 239 |
+
def _spec_to_audio(spec: torch.Tensor) -> np.ndarray:
|
| 240 |
+
"""Convert normalized spectrogram to waveform using the loaded vocoder."""
|
| 241 |
log_mel = denormalize(spec)
|
| 242 |
log_mel = log_mel.squeeze(1) # (1, 128, 256)
|
| 243 |
+
|
| 244 |
+
if state.vocoder_type == "bigvgan":
|
| 245 |
+
with torch.no_grad():
|
| 246 |
+
waveform = state.vocoder(log_mel.to(state.device)) # (1, 1, T)
|
| 247 |
+
waveform = waveform.squeeze(1).cpu() # (1, T)
|
| 248 |
+
waveform = torchaudio.functional.highpass_biquad(waveform, SAMPLE_RATE, cutoff_freq=25.0)
|
| 249 |
+
waveform = torchaudio.functional.lowpass_biquad(waveform, SAMPLE_RATE, cutoff_freq=20000.0)
|
| 250 |
+
waveform = waveform / (waveform.abs().max() + 1e-8)
|
| 251 |
+
else:
|
| 252 |
+
waveform = state.vocoder(log_mel) # (1, T)
|
| 253 |
+
waveform = waveform / (waveform.abs().max() + 1e-8)
|
| 254 |
+
|
| 255 |
return waveform.squeeze(0).numpy()
|
| 256 |
|
| 257 |
|
|
|
|
| 266 |
{
|
| 267 |
"id": i + 1,
|
| 268 |
"name": state.pc_names[i],
|
| 269 |
+
"min": 0.0, "max": 1.0,
|
| 270 |
+
"default": 0.5, "step": 0.01,
|
|
|
|
|
|
|
| 271 |
}
|
| 272 |
for i in range(N_PCS)
|
| 273 |
],
|
| 274 |
+
"vocoder": state.vocoder_type,
|
| 275 |
}
|
| 276 |
|
| 277 |
|
|
|
|
| 281 |
z = _pc_to_latent(pc_values)
|
| 282 |
|
| 283 |
with torch.no_grad():
|
| 284 |
+
spec = state.model.decode(z)
|
| 285 |
+
audio = _spec_to_audio(spec)
|
| 286 |
|
| 287 |
buf = io.BytesIO()
|
| 288 |
sf.write(buf, audio, SAMPLE_RATE, format="WAV")
|
|
|
|
| 304 |
|
| 305 |
@app.get("/health")
|
| 306 |
async def health():
|
| 307 |
+
return {"status": "ok", "vocoder": state.vocoder_type, "device": str(state.device)}
|
requirements.txt
CHANGED
|
@@ -3,3 +3,4 @@ scikit-learn
|
|
| 3 |
soundfile>=0.12
|
| 4 |
fastapi>=0.115.0
|
| 5 |
uvicorn[standard]>=0.34.0
|
|
|
|
|
|
| 3 |
soundfile>=0.12
|
| 4 |
fastapi>=0.115.0
|
| 5 |
uvicorn[standard]>=0.34.0
|
| 6 |
+
bigvgan>=2.4.1
|