maestroai / scripts /train_classifier.py
gabrielpamplonapg
Fix mood loudness + electronic genre misclassification
c2dd476
Raw
History Blame Contribute Delete
5.43 kB
#!/usr/bin/env python3
"""
train_classifier.py — Train sklearn RandomForest classifier on labelled data.
Usage:
python3 scripts/train_classifier.py
Reads all JSON files from training_data/ directory.
Saves model to training_data/models/instrument_classifier.pkl
"""
import json
import os
import sys
from pathlib import Path
# Allow running from project root
PROJECT_ROOT = Path(__file__).parent.parent
sys.path.insert(0, str(PROJECT_ROOT))
# Data dir is overridable (MAESTRO_TRAINING_DIR) so tests can point at an
# isolated fixture dir instead of the repo's accumulated training_data/.
TRAINING_DIR = Path(os.environ.get("MAESTRO_TRAINING_DIR", PROJECT_ROOT / "training_data"))
MODEL_DIR = TRAINING_DIR / "models"
MODEL_PATH = MODEL_DIR / "instrument_classifier.pkl"
LABEL_ENCODER_PATH = MODEL_DIR / "label_encoder.pkl"
# Feature keys (must match ml_classifier.py)
FEATURE_KEYS = [
"spectral_centroid_hz",
"spectral_rolloff_hz",
"spectral_flatness",
"harmonic_ratio",
"zero_crossing_rate",
"percussive_ratio",
"dynamic_range_db",
"inharmonicity_score",
"tempo_bpm",
"pitch_std_hz",
"attack_sharpness",
]
def extract_feature_vector(features: dict) -> list:
"""Extract a flat feature vector from a features dict."""
vec = []
for key in FEATURE_KEYS:
vec.append(float(features.get(key, 0.0) or 0.0))
# mel_band_stats sub-keys
mel = features.get("mel_band_stats", {}) or {}
for sub in ["mel_sub", "mel_bass", "mel_low_mid", "mel_high_mid", "mel_high"]:
vec.append(float(mel.get(sub, 0.0) or 0.0))
# frequency_bands_db sub-keys
bands = features.get("frequency_bands_db", {}) or {}
for sub in ["sub_bass_db", "bass_db", "low_mid_db", "high_mid_db", "presence_db", "air_db"]:
vec.append(float(bands.get(sub, 0.0) or 0.0))
# mfcc_means — first 13 coefficients
mfcc = features.get("mfcc_means", []) or []
for i in range(13):
vec.append(float(mfcc[i]) if i < len(mfcc) else 0.0)
# spectral_contrast — first 7 bands
sc = features.get("spectral_contrast", []) or []
for i in range(7):
vec.append(float(sc[i]) if i < len(sc) else 0.0)
return vec
def load_training_data():
"""Load all JSON training records from training_data/ directory."""
X = []
y = []
skipped = 0
for json_file in TRAINING_DIR.glob("*.json"):
try:
with open(json_file) as f:
records = json.load(f)
# Support both list of records and single record
if isinstance(records, dict):
records = [records]
for record in records:
try:
features = record.get("features", {})
analysis = record.get("analysis", {})
instruments = analysis.get("instruments", [])
if not instruments:
skipped += 1
continue
label = instruments[0].get("model")
if not label:
skipped += 1
continue
vec = extract_feature_vector(features)
X.append(vec)
y.append(label)
except Exception as e:
skipped += 1
continue
except Exception as e:
print(f"Warning: Could not read {json_file}: {e}")
continue
if skipped > 0:
print(f"Skipped {skipped} records (missing features or labels)")
return X, y
def main():
print("Loading training data...")
X, y = load_training_data()
n_samples = len(X)
print(f"Found {n_samples} labelled samples")
if n_samples < 5:
print(f"Insufficient data ({n_samples} samples). Need at least 5 to train.")
sys.exit(1)
import numpy as np
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import cross_val_score
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import LabelEncoder, StandardScaler
import joblib
X_arr = np.array(X)
le = LabelEncoder()
y_enc = le.fit_transform(y)
n_classes = len(le.classes_)
print(f"Labels: {n_classes} unique instruments")
pipeline = Pipeline([
("scaler", StandardScaler()),
("clf", RandomForestClassifier(n_estimators=100, random_state=42)),
])
# Cross-validation (cv=3 if enough samples)
cv = min(3, n_samples // n_classes) if n_samples >= 3 else 2
cv = max(cv, 2) # at least 2
if n_samples >= cv * n_classes:
try:
scores = cross_val_score(pipeline, X_arr, y_enc, cv=cv, scoring="accuracy")
print(f"Cross-validation accuracy (cv={cv}): {scores.mean():.3f} ± {scores.std():.3f}")
except Exception as e:
print(f"Cross-validation skipped: {e}")
else:
print("Not enough samples for cross-validation, skipping.")
# Train on full data
print("Training final model on all data...")
pipeline.fit(X_arr, y_enc)
# Save model
MODEL_DIR.mkdir(parents=True, exist_ok=True)
joblib.dump(pipeline, MODEL_PATH)
joblib.dump(le, LABEL_ENCODER_PATH)
print(f"Model saved to: {MODEL_PATH}")
print(f"Label encoder saved to: {LABEL_ENCODER_PATH}")
print("Training complete.")
if __name__ == "__main__":
main()