Download handler.py from austinsd/midisim-endpoint: direct link, hf CLI and curl.
- Browser
- Download file 3.78 kB
-
https://huggingface.co/austinsd/midisim-endpoint/resolve/main/handler.py
- Command line
-
hf download hf://austinsd/midisim-endpoint/handler.py
-
curl -L -o handler.py https://huggingface.co/austinsd/midisim-endpoint/resolve/main/handler.py
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 | |