thepaul Claude Opus 4.6 commited on
Commit
ba9111a
Β·
1 Parent(s): 7339046

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>

Files changed (3) hide show
  1. Dockerfile +1 -2
  2. main.py +96 -42
  3. 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 CPU-only PyTorch (avoids ~3GB of CUDA bloat)
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 CPU-only PyTorch with Griffin-LIM vocoder (no GPU needed).
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
- # Griffin-LIM vocoder (runs on CPU, no GPU needed)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79
  # ──────────────────────────────────────────────
80
  class GriffinLimVocoder:
81
- """Mel spectrogram β†’ audio via pseudo-inverse mel filterbank + Griffin-LIM."""
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
- win_length=WIN_SIZE,
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 # noise floor to avoid phase artifacts
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: GriffinLimVocoder
126
- pca_components: np.ndarray # shape (N_PCS, LATENT_DIM)
127
- pca_mean: np.ndarray # shape (LATENT_DIM,)
 
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
- state.device = torch.device("cpu")
151
-
152
- # Load PCA parameters pre-computed from the training dataset
 
 
 
 
 
 
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"] # (N_PCS, LATENT_DIM)
156
- state.pca_mean = pca_data["mean"] # (LATENT_DIM,)
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="cpu", weights_only=False)
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
- # Initialize Griffin-LIM vocoder
172
- state.vocoder = GriffinLimVocoder()
173
- print("Griffin-LIM vocoder ready")
 
 
 
 
 
 
 
 
 
 
 
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 _decode_to_audio(spec: torch.Tensor) -> np.ndarray:
194
- """Decode normalized spectrogram β†’ waveform via Griffin-LIM."""
195
  log_mel = denormalize(spec)
196
  log_mel = log_mel.squeeze(1) # (1, 128, 256)
197
- waveform = state.vocoder(log_mel) # (1, T)
198
- waveform = waveform / (waveform.abs().max() + 1e-8)
 
 
 
 
 
 
 
 
 
 
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
- "max": 1.0,
215
- "default": 0.5,
216
- "step": 0.01,
217
  }
218
  for i in range(N_PCS)
219
  ],
220
- "vocoder": "griffinlim",
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) # (1, 1, 128, 256)
231
- audio = _decode_to_audio(spec)
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