File size: 3,783 Bytes
d6b7165
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
"""
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