"""Piano Performance Quality Evaluator - FastAPI Backend Upload audio (wav/mp3) or MIDI file and get a 19-dimension quality report. Audio is transcribed to MIDI via transkun before evaluation. """ import os import pickle import subprocess import sys import tempfile from pathlib import Path import numpy as np from fastapi import FastAPI, File, UploadFile, HTTPException from fastapi.responses import HTMLResponse, JSONResponse from fastapi.staticfiles import StaticFiles from fastapi.middleware.cors import CORSMiddleware from sklearn.linear_model import Ridge from sklearn.preprocessing import StandardScaler # -- Paths -- BASE_DIR = Path(__file__).parent MODELS_DIR = BASE_DIR / "saved_models" TEMPLATES_DIR = BASE_DIR / "templates" app = FastAPI(title="Piano Performance Quality Evaluator") # CORS for Vercel frontend app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) if (BASE_DIR / "static").exists(): app.mount("/static", StaticFiles(directory=str(BASE_DIR / "static")), name="static") # -- Load trained models on startup -- scaler: StandardScaler = None ridge: Ridge = None DIMENSION_NAMES = [ "timing_stability", "articulation_length", "articulation_hardness", "pedal_saturation", "pedal_clarity", "timbre_evenness", "timbre_richness", "timbre_brightness", "timbre_loudness", "dynamic_sophistication", "dynamic_range", "music_pace", "music_spaciousness", "music_balance", "music_expressiveness", "emotion_optimism", "emotion_energy", "emotion_imagination", "interpretation_quality", ] DIMENSION_CATEGORIES = { "Timing": ["timing_stability"], "Articulation": ["articulation_length", "articulation_hardness"], "Pedal": ["pedal_saturation", "pedal_clarity"], "Timbre": ["timbre_evenness", "timbre_richness", "timbre_brightness", "timbre_loudness"], "Dynamics": ["dynamic_sophistication", "dynamic_range"], "Musicality": ["music_pace", "music_spaciousness", "music_balance", "music_expressiveness"], "Emotion": ["emotion_optimism", "emotion_energy", "emotion_imagination"], "Interpretation": ["interpretation_quality"], } @app.on_event("startup") def load_models(): global scaler, ridge # Patch torch.load to always map to correct device import torch _orig_load = torch.load device = _get_device() def _safe_load(*args, **kwargs): if 'map_location' not in kwargs: kwargs['map_location'] = device return _orig_load(*args, **kwargs) torch.load = _safe_load scaler_path = MODELS_DIR / "scaler.pkl" ridge_path = MODELS_DIR / "ridge.pkl" if not scaler_path.exists() or not ridge_path.exists(): print("WARNING: Models not found. Run train_scoring_head.py first.") return with open(scaler_path, "rb") as f: scaler = pickle.load(f) with open(ridge_path, "rb") as f: ridge = pickle.load(f) print(f"Models loaded: scaler={scaler is not None}, ridge={ridge is not None}") def _get_device(): """Auto-detect best available device.""" import torch if os.environ.get("EVAL_FORCE_CPU") == "1": return "cpu" return "cuda" if torch.cuda.is_available() else "cpu" def _get_encoder(): """Lazy-load Aria model to avoid loading at import time.""" from evpmr.models.aria_model import AriaMidiModel if not hasattr(_get_encoder, "_model"): _get_encoder._model = AriaMidiModel(device=_get_device()) return _get_encoder._model @app.get("/api/health") async def health(): return {"status": "ok"} @app.get("/", response_class=HTMLResponse) async def index(): html_path = TEMPLATES_DIR / "index.html" if html_path.exists(): return HTMLResponse(content=html_path.read_text()) return HTMLResponse(content="

Piano Quality Evaluator API

Backend is running.

") AUDIO_EXTS = (".wav", ".mp3", ".flac", ".ogg", ".m4a") MIDI_EXTS = (".mid", ".midi") def _transcribe_audio(audio_path: str, midi_path: str): """Run transkun to transcribe audio to MIDI via wrapper script.""" device = _get_device() if device == "cpu": raise RuntimeError( "音频转写需要 GPU,当前环境无 GPU。请直接上传 MIDI 文件(.mid/.midi)。" ) wrapper = str(BASE_DIR / "transkun_wrapper.py") result = subprocess.run( [sys.executable, wrapper, audio_path, midi_path], capture_output=True, text=True, timeout=300, ) if result.returncode != 0: raise RuntimeError(f"transkun failed: {result.stderr[-500:]}") @app.post("/api/evaluate") async def evaluate(file: UploadFile = File(...)): """Evaluate an audio or MIDI file and return 19-dimension quality scores.""" ext = Path(file.filename).suffix.lower() if ext not in AUDIO_EXTS and ext not in MIDI_EXTS: raise HTTPException( status_code=400, detail=f"Unsupported format '{ext}'. Upload audio ({'/'.join(AUDIO_EXTS)}) or MIDI ({'/'.join(MIDI_EXTS)}).", ) if scaler is None or ridge is None: raise HTTPException(status_code=503, detail="Models not loaded. Run train_scoring_head.py first.") with tempfile.TemporaryDirectory(prefix="piano_eval_") as tmp_dir: tmp_dir = Path(tmp_dir) upload_path = tmp_dir / file.filename with open(upload_path, "wb") as f: f.write(await file.read()) midi_path = upload_path transcribed = False # If audio -> transcribe to MIDI first if ext in AUDIO_EXTS: midi_path = upload_path.with_suffix(".mid") try: _transcribe_audio(str(upload_path), str(midi_path)) transcribed = True except Exception as e: raise HTTPException(status_code=500, detail=f"Audio transcription failed: {str(e)}") try: # Encode MIDI -> 512-dim embedding model = _get_encoder() embedding = model.encode(str(midi_path)) # (512,) # Standardize + predict emb_scaled = scaler.transform(embedding.reshape(1, -1)) scores = ridge.predict(emb_scaled)[0] # (19,) scores = np.clip(scores, 0.0, 1.0) # Build response result = {} for i, name in enumerate(DIMENSION_NAMES): result[name] = round(float(scores[i]), 4) # Category averages categories = {} for cat, dims in DIMENSION_CATEGORIES.items(): vals = [result[d] for d in dims] categories[cat] = round(float(np.mean(vals)), 4) overall = round(float(np.mean(scores)), 4) return JSONResponse({ "status": "ok", "filename": file.filename, "input_type": "audio" if transcribed else "midi", "transcribed": transcribed, "overall_score": overall, "dimensions": result, "categories": categories, }) except HTTPException: raise except Exception as e: raise HTTPException(status_code=500, detail=f"Evaluation error: {str(e)}") if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=7860)