kicks-api / export_pca.py
thepaul's picture
Initial deploy: VAE + Griffin-LIM API
7339046
Raw History Blame Contribute Delete
3.33 kB
#!/usr/bin/env python3
"""Export PCA parameters from trained kicks VAE for use in HF Space.
Run this from the kicks repo root:
python export_pca.py
Outputs: pca_params.npz
"""
import numpy as np
import torch
from sklearn.decomposition import PCA
from kicks import KickDataset, KickDataloader, VAE
from kicks.cluster import extract_latents, compute_descriptors
from kicks.config import load_vae_from_checkpoint, DATA_DIR, BEST_CHECKPOINT
N_PCS = 5
DESC_KEYS = ["sub", "punch", "click", "bright", "decay"]
DESC_NAME_MAP = {
"sub": "Sub", "punch": "Punch", "click": "Click",
"bright": "Bright", "decay": "Decay",
}
def main():
device = torch.device("cpu")
print(f"Loading dataset from {DATA_DIR}...")
dataset = KickDataset(DATA_DIR)
dataloader = KickDataloader(dataset, batch_size=32, shuffle=False)
print(f"Loading VAE from {BEST_CHECKPOINT}...")
model, _ = load_vae_from_checkpoint(BEST_CHECKPOINT, device)
print("Extracting latents and spectrograms...")
latents, spectrograms = extract_latents(model, dataloader, device)
print(f" {len(latents)} samples, latent dim {latents.shape[1]}")
print("Computing perceptual descriptors for PC naming...")
descriptors = [compute_descriptors(s) for s in spectrograms]
desc_arrays = {k: np.array([d[k] for d in descriptors]) for k in DESC_KEYS}
print(f"Fitting PCA ({latents.shape[1]}D -> {N_PCS}D)...")
pca = PCA(n_components=N_PCS)
pc_projected = pca.fit_transform(latents)
# Auto-name PCs from highest descriptor correlations, flip negative axes
pc_names = []
used = set()
for i in range(N_PCS):
pc_vals = pc_projected[:, i]
best_desc, best_corr = None, 0.0
for dk in DESC_KEYS:
if dk in used:
continue
dv = desc_arrays[dk]
if pc_vals.std() > 0 and dv.std() > 0:
r = float(np.corrcoef(pc_vals, dv)[0, 1])
else:
r = 0.0
if abs(r) > abs(best_corr):
best_corr = r
best_desc = dk
if best_desc and abs(best_corr) >= 0.15:
used.add(best_desc)
pc_names.append(DESC_NAME_MAP[best_desc])
if best_corr < 0:
pca.components_[i] *= -1
pc_projected[:, i] *= -1
print(f" PC{i+1} -> {pc_names[-1]} (r={best_corr:.2f}, flipped)")
else:
print(f" PC{i+1} -> {pc_names[-1]} (r={best_corr:.2f})")
else:
pc_names.append(f"PC{i+1}")
print(f" PC{i+1} -> PC{i+1} (no strong correlation)")
pc_mins = [float(np.percentile(pc_projected[:, i], 2)) for i in range(N_PCS)]
pc_maxs = [float(np.percentile(pc_projected[:, i], 98)) for i in range(N_PCS)]
# Save
np.savez(
"pca_params.npz",
components=pca.components_,
mean=pca.mean_,
explained_variance_ratio_=pca.explained_variance_ratio_,
pc_names=np.array(pc_names, dtype=object),
pc_mins=np.array(pc_mins),
pc_maxs=np.array(pc_maxs),
)
print(f"\nSaved pca_params.npz")
print(f"PC names: {pc_names}")
print(f"PC ranges: {list(zip(pc_mins, pc_maxs))}")
print(f"Explained variance: {pca.explained_variance_ratio_}")
if __name__ == "__main__":
main()