midisim-endpoint / handler.py
austinsd's picture
Add handler.py
d6b7165 verified
Raw History Blame Contribute Delete
3.78 kB
"""
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": "<base64-encoded MIDI bytes>" }
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