#!/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()