rooroo79 commited on
Commit
d8d28a6
·
verified ·
1 Parent(s): 524c586

Add custom Inference Endpoint handler (SpectrogramDiffusionPipeline)

Browse files
Files changed (1) hide show
  1. handler.py +141 -0
handler.py ADDED
@@ -0,0 +1,141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Custom Hugging Face Inference Endpoint handler.
2
+
3
+ Loads Google's Multi-Instrument Spectrogram Diffusion
4
+ (`google/music-spectrogram-diffusion`) from the endpoint's local snapshot
5
+ and turns a MIDI file into a 16 kHz WAV.
6
+
7
+ This file must live at the *root* of the Hub repo the endpoint points at,
8
+ alongside `requirements.txt`. `google/music-spectrogram-diffusion` has no
9
+ `handler.py` and no Transformers `config.json`, so a Custom task against
10
+ that repo falls back to `pipeline()` and crashes. Fork the model, add these
11
+ two files, then set the endpoint to that fork with task=custom.
12
+
13
+ Request body (Band's `/midi-to-audio` client):
14
+
15
+ { "inputs": "<base64 .mid>", "parameters": { "prompt": "optional extra text" } }
16
+
17
+ Response:
18
+
19
+ { "audio_base64": "<base64 wav>", "sample_rate": 16000, "format": "wav" }
20
+
21
+ The prompt is accepted so Band can send clip text; this pipeline does not
22
+ condition on it. Keep the field so a later handler can without a client change.
23
+ """
24
+ from __future__ import annotations
25
+
26
+ import base64
27
+ import io
28
+ import os
29
+ import wave
30
+ from typing import Any
31
+
32
+ import numpy as np
33
+ import torch
34
+
35
+ try:
36
+ from diffusers import MidiProcessor, SpectrogramDiffusionPipeline
37
+ except ImportError: # moved under deprecated/ in newer Diffusers
38
+ from diffusers.pipelines.deprecated.spectrogram_diffusion import (
39
+ MidiProcessor,
40
+ SpectrogramDiffusionPipeline,
41
+ )
42
+
43
+
44
+ SAMPLE_RATE = 16000
45
+ FALLBACK_MODEL_ID = os.environ.get("HF_MSD_MODEL_ID", "google/music-spectrogram-diffusion")
46
+
47
+
48
+ class EndpointHandler:
49
+ def __init__(self, path: str = "") -> None:
50
+ # Inference Endpoints pass the local snapshot (`/repository`), which
51
+ # has `model_index.json` + decoder/encoder/melgan weights. Load that,
52
+ # not the Hub id — otherwise the replica re-downloads and ignores the
53
+ # fork this handler was deployed from.
54
+ model_path = path if _has_diffusers_index(path) else FALLBACK_MODEL_ID
55
+ device = "cuda" if torch.cuda.is_available() else "cpu"
56
+ dtype = torch.float16 if device == "cuda" else torch.float32
57
+ self.pipe = SpectrogramDiffusionPipeline.from_pretrained(
58
+ model_path,
59
+ torch_dtype=dtype,
60
+ )
61
+ try:
62
+ self.pipe.to(device)
63
+ except Exception:
64
+ # MelGAN is ONNX and may refuse `.to()`. Move the PyTorch modules.
65
+ for name in ("notes_encoder", "continuous_encoder", "decoder"):
66
+ module = getattr(self.pipe, name, None)
67
+ if module is not None:
68
+ module.to(device)
69
+ self.processor = MidiProcessor()
70
+ self.device = device
71
+
72
+ def __call__(self, data: Any) -> dict[str, Any]:
73
+ payload = data if isinstance(data, dict) else {}
74
+ midi_bytes = _midi_bytes(payload)
75
+ params = payload.get("parameters") if isinstance(payload.get("parameters"), dict) else {}
76
+ steps = params.get("num_inference_steps")
77
+ kwargs: dict[str, Any] = {}
78
+ if steps is not None:
79
+ kwargs["num_inference_steps"] = max(1, int(steps))
80
+
81
+ tokens = self.processor(midi_bytes)
82
+ audio = self.pipe(tokens, **kwargs)
83
+ samples = _to_mono_float(audio)
84
+ wav = _wav_bytes(samples, SAMPLE_RATE)
85
+ return {
86
+ "audio_base64": base64.b64encode(wav).decode("ascii"),
87
+ "sample_rate": SAMPLE_RATE,
88
+ "format": "wav",
89
+ }
90
+
91
+
92
+ def _has_diffusers_index(path: str) -> bool:
93
+ return bool(path) and os.path.isfile(os.path.join(path, "model_index.json"))
94
+
95
+
96
+ def _midi_bytes(payload: dict[str, Any]) -> bytes:
97
+ raw = payload.get("inputs", payload.get("midi_base64", ""))
98
+ if isinstance(raw, list) and raw:
99
+ raw = raw[0]
100
+ if isinstance(raw, dict):
101
+ raw = raw.get("midi_base64") or raw.get("inputs") or raw.get("data") or ""
102
+ if isinstance(raw, (bytes, bytearray)):
103
+ midi = bytes(raw)
104
+ elif isinstance(raw, str) and raw:
105
+ text = raw.strip()
106
+ if text.lower().startswith("data:") and "," in text:
107
+ text = text.split(",", 1)[1]
108
+ midi = base64.b64decode(text)
109
+ else:
110
+ raise ValueError("missing base64 MIDI in inputs")
111
+ if len(midi) < 8 or midi[:4] != b"MThd":
112
+ raise ValueError("inputs is not a MIDI file")
113
+ return midi
114
+
115
+
116
+ def _to_mono_float(audio: Any) -> np.ndarray:
117
+ if hasattr(audio, "audios") and audio.audios is not None:
118
+ arr = np.asarray(audio.audios, dtype=np.float32)
119
+ elif isinstance(audio, (list, tuple)) and audio:
120
+ arr = np.asarray(audio[0], dtype=np.float32)
121
+ else:
122
+ arr = np.asarray(audio, dtype=np.float32)
123
+ arr = np.squeeze(arr)
124
+ if arr.ndim > 1:
125
+ arr = arr.reshape(-1)
126
+ peak = float(np.max(np.abs(arr))) if arr.size else 0.0
127
+ if peak > 1.0:
128
+ arr = arr / peak
129
+ return arr
130
+
131
+
132
+ def _wav_bytes(samples: np.ndarray, rate: int) -> bytes:
133
+ pcm = np.clip(samples, -1.0, 1.0)
134
+ pcm = (pcm * 32767.0).astype(np.int16)
135
+ buf = io.BytesIO()
136
+ with wave.open(buf, "wb") as wf:
137
+ wf.setnchannels(1)
138
+ wf.setsampwidth(2)
139
+ wf.setframerate(rate)
140
+ wf.writeframes(pcm.tobytes())
141
+ return buf.getvalue()