""" HuggingFace Inference Endpoints — Custom Handler for MIDISim ------------------------------------------------------------- Contract: __init__(path: str) — called once at container startup; path = repo root __call__(data: dict) — called per request Request payload (JSON): { "inputs": "" } Response: [float, float, ...] — 512-dim embedding vector (JSON array) The endpoint is registered as pipeline_tag: "feature-extraction" so the standard HF Inference Endpoints routing sends POST / with {"inputs": ...}. """ import base64 import logging import os import tempfile from pathlib import Path from typing import Any import numpy as np logger = logging.getLogger(__name__) MODEL_FILENAME = ( "midisim_small_pre_trained_model_2_epochs_43117_steps_0.3148_loss_0.9229_acc.pth" ) MAX_SEQ_LEN = int(os.getenv("MIDISIM_MAX_SEQ_LEN", "1024")) class EndpointHandler: def __init__(self, path: str = ""): """ Load the MIDISim model once at container startup. `path` is the local directory where the repo files were downloaded. """ import torch import midisim as ms self.ms = ms repo_dir = Path(path) if path else Path(".") # Model checkpoint is in the repo root model_path = repo_dir / MODEL_FILENAME if not model_path.exists(): logger.info(f"Downloading model checkpoint from HF Hub...") model_path = ms.download_model( filename=MODEL_FILENAME, local_dir=str(repo_dir), verbose=False, ) device = "cuda" if torch.cuda.is_available() else "cpu" logger.info(f"Loading MIDISim model on {device} ...") self.model, self.ctx, self.dtype = ms.load_model( model_path=str(model_path), depth=8, device=device, verbose=False, ) logger.info("MIDISim model ready.") def __call__(self, data: dict) -> Any: """ Parameters ---------- data : dict Must contain key "inputs" with a base64-encoded MIDI file string. Returns ------- list[float] — 512-dimensional embedding vector, or {"error": str} on failure. """ midi_b64 = data.get("inputs", "") if not midi_b64: return {"error": "No 'inputs' field in request body"} tmp_path = None try: midi_bytes = base64.b64decode(midi_b64) with tempfile.NamedTemporaryFile(suffix=".mid", delete=False) as f: f.write(midi_bytes) tmp_path = f.name # Tokenise — transpose_factor=0 → single sequence (no augmentation) token_seqs = self.ms.midi_to_tokens( tmp_path, max_seq_len=MAX_SEQ_LEN, transpose_factor=0, verbose=False, ) if not token_seqs: return {"error": "midi_to_tokens returned empty sequence"} # Encode → shape (1, 512) as numpy float32 emb = self.ms.get_embeddings_bf16( self.model, token_seqs, seq_len=MAX_SEQ_LEN, verbose=False, show_progress_bar=False, return_numpy=True, ) vec = emb[0].astype(np.float32).tolist() logger.info(f"Embedding: dim={len(vec)}, norm={np.linalg.norm(emb[0]):.4f}") return vec except Exception as e: logger.error(f"Handler error: {e}") return {"error": str(e)} finally: if tmp_path and os.path.exists(tmp_path): os.remove(tmp_path) # Made with Bob