File size: 13,166 Bytes
abde881
 
 
c91d158
 
 
 
 
 
 
0317c4b
abde881
520a389
abde881
 
 
244db97
abde881
 
 
 
 
 
 
 
75f7e7a
abde881
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dd5e91d
 
801f796
 
 
 
 
 
 
 
 
 
 
dd5e91d
801f796
 
 
520a389
 
 
801f796
 
 
 
 
 
 
 
 
dd5e91d
4ce751d
dd5e91d
 
 
4ce751d
dd5e91d
 
 
801f796
 
520a389
 
 
dd5e91d
 
28079df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4ce751d
 
801f796
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
abde881
dd5e91d
 
 
 
 
abde881
dd5e91d
 
 
 
 
 
 
 
 
 
abde881
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75f7e7a
244db97
75f7e7a
 
 
 
 
801f796
75f7e7a
 
28079df
 
520a389
28079df
520a389
28079df
 
 
 
 
 
 
 
 
 
 
abde881
 
75f7e7a
edbfb74
 
c91d158
edbfb74
 
 
 
abde881
 
 
801f796
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
abde881
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
801f796
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
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
# ============================================================================
#  app.py — ZeroGPU deployment via gradio.Server
#
#  gradio.Server extends FastAPI, so:
#    - @app.get("/")           serves the custom HTML editor (unchanged)
#    - @app.api()              wraps the GPU transcription with Gradio's queue
#                              + ZeroGPU allocation (@spaces.GPU)
#    - @app.post("/analyze")   and @app.post("/pmi") are plain FastAPI routes
#                              (CPU only — key estimation / PMI, no GPU)
#

# ============================================================================
import os, base64, tempfile, time

import spaces                      # ZeroGPU
from gradio import Server
from gradio.data_classes import FileData
from fastapi import UploadFile, File, Form
from fastapi.responses import FileResponse, JSONResponse
import pretty_midi

import pmi_core

# NOTE: heavy imports (torch, demucs, YourMT3+) happen lazily inside pipeline.py
# so the web server starts fast and the model is only touched inside the GPU call.
from pipeline import chroma_from_audio, demucs_stem, estimate_bpm, notes_from_audio


# ---------------------------------------------------------------------------
# helpers (identical to the FastAPI version)
# ---------------------------------------------------------------------------
def notes_to_midi_bytes(notes):
    pm = pretty_midi.PrettyMIDI()
    inst = pretty_midi.Instrument(program=0)
    for n in notes:
        s, e, p = n[0], n[1], n[2]
        conf = n[3] if len(n) > 3 else 1.0
        vel = max(1, min(127, int(round(conf * 127))))
        inst.notes.append(pretty_midi.Note(velocity=vel, pitch=int(p),
                                            start=float(s), end=float(e)))
    pm.instruments.append(inst)
    t = tempfile.NamedTemporaryFile(delete=False, suffix=".mid"); t.close()
    pm.write(t.name)
    data = open(t.name, "rb").read()
    os.unlink(t.name)
    return data


def _save_temp_midi(raw_bytes):
    t = tempfile.NamedTemporaryFile(delete=False, suffix=".mid")
    t.write(raw_bytes); t.flush(); t.close()
    return t.name


def _parse_forced(tonic, mode):
    if tonic is None or tonic == "" or mode is None or mode == "":
        return None
    try:
        return (int(tonic) % 12, "minor" if str(mode).lower().startswith("min") else "major")
    except (ValueError, TypeError):
        return None


# ---------------------------------------------------------------------------
# the GPU work: demucs + YourMT3+ + post-processing, ONE @spaces.GPU call.
#
# Why one call and not two (separation / transcription split): each entry into
# a @spaces.GPU function pays a fixed toll -- a process fork, CUDA re-attach,
# a pickle round-trip, and a fresh pass through the node queue -- and each
# entry runs the caller through the quota gate again. Splitting doubles all of
# that per clip, which is exactly what burned through visitors' free quota.
# Quota is charged on ACTUAL seconds used, not on the declared duration; the
# declaration is only (a) an admission check against the caller's remaining
# quota and (b) a queue-priority hint (shorter ranks higher). So the right
# declaration is a realistic estimate with modest headroom -- never a fat
# safety margin, which locks callers out long before their quota is spent.
# ---------------------------------------------------------------------------
# Set after every successful GPU call and read by _gpu_duration to decide
# whether the next request can reasonably expect a warm worker. Lives in the
# main process (the web server), which is where the duration callable runs.
_last_gpu_ok = 0.0
WARM_WINDOW = 120.0        # seconds; conservative -- covers back-to-back testing

def _resolve(audio_path):
    """FileData arrives as an object with .path, a dict, or a plain string."""
    if isinstance(audio_path, dict):
        return audio_path.get("path", "")
    if hasattr(audio_path, "path"):
        return audio_path.path
    return audio_path


def _audio_seconds(path):
    try:
        import soundfile as _sf
        with _sf.SoundFile(path) as f:
            return len(f) / float(f.samplerate)
    except Exception:
        return 30.0                      # unknown -> assume a modest clip

def _gpu_duration(audio_path, instrument="vocals"):
 
    sec = _audio_seconds(_resolve(audio_path))
    if (time.monotonic() - _last_gpu_ok) < WARM_WINDOW:
        return int(max(5, min(120, 0.2 * sec)))   
    return int(max(15, min(120, 0.25 * sec)))     
@spaces.GPU(duration=_gpu_duration)
def _transcribe_gpu(audio_path, instrument="vocals"):
    """Only the steps that actually need the GPU: demucs separation and YourMT3+
    inference. Chroma extraction and MIDI encoding are pure CPU work and run
    outside this window, where they cost the visitor no quota.

    Beat tracking stays inside despite being CPU work: notes_from_audio needs
    the tempo, and a second GPU entry to hand it back in would cost far more
    (fork, CUDA re-attach, queue, quota gate) than the ~0.5 s it takes."""
    stem_path = demucs_stem(audio_path, instrument)   # GPU: separation
    tempo = estimate_bpm(stem_path)                   # CPU, but needed inline:
                                                      # notes_from_audio needs it
    notes = notes_from_audio(stem_path, tempo=tempo)  # GPU: YourMT3+
    # Return the stem as BYTES, not a path: the GPU worker is a separate,
    # short-lived process and its return value crosses a pickle boundary, so a
    # path would rely on the file still being readable from the main process.
    # Bytes are self-contained.
    return open(stem_path, "rb").read(), notes


# ---------------------------------------------------------------------------
# Both excerpts in ONE GPU call.
## ---------------------------------------------------------------------------
def _pair_duration(audio_a, audio_b, instrument_a="vocals", instrument_b="vocals"):
    total = _audio_seconds(_resolve(audio_a)) + _audio_seconds(_resolve(audio_b))
    if (time.monotonic() - _last_gpu_ok) < WARM_WINDOW:
        return int(max(20, min(120, 8 + 0.2 * total)))
    return int(max(28, min(120, 15 + 0.25 * total)))


@spaces.GPU(duration=_pair_duration)
def _transcribe_pair_gpu(audio_a, audio_b, instrument_a="vocals", instrument_b="vocals"):
    """Separate + transcribe two excerpts in one GPU window. The model is moved
    to the device once by the first call into yourmt3_transcribe and stays there
    for the second, so the second excerpt costs only its own compute."""
    out = []
    for path, inst in ((audio_a, instrument_a), (audio_b, instrument_b)):
        stem_path = demucs_stem(_resolve(path), inst)
        tempo = estimate_bpm(stem_path)
        notes = notes_from_audio(stem_path, tempo=tempo)
        out.append((open(stem_path, "rb").read(), notes))
    return out[0], out[1]


# ---------------------------------------------------------------------------
# Warm-up at startup, on CPU, where time is free. Loading the YourMT3+
# checkpoint takes ~10-15 s; done lazily it lands inside the FIRST visitor's
# GPU window and is billed to THEIR quota. Doing it here means the GPU window
# only ever pays for the cheap cpu->cuda move, not the disk load. demucs
# weights are fetched here too, so no download ever happens on GPU time.
# ---------------------------------------------------------------------------
try:
    from pipeline import _load_model
    _load_model()                                  # YourMT3+ -> CPU, cached
except Exception as _e:
    print(f"[warmup] YourMT3+ preload failed (will retry lazily): {_e}", flush=True)
try:
    import demucs.pretrained as _dp
    _dp.get_model(name="htdemucs")                 # default (vocals) weights
except Exception as _e:
    print(f"[warmup] demucs preload failed (will retry lazily): {_e}", flush=True)


# ---------------------------------------------------------------------------
# gradio.Server (a FastAPI app with Gradio's API engine on top)
# ---------------------------------------------------------------------------
app = Server()


@app.get("/")
def index():
    return FileResponse("transcription-editor.html")


# transcription: wrapped by Gradio's queue + ZeroGPU. concurrency_limit=1
# because a single ZeroGPU slice serves one transcription at a time.
# Returns base64 strings so the existing front-end decoding still works.
@app.api(name="transcribe", concurrency_limit=1)
def transcribe(audio_path: FileData, instrument: str = "vocals") -> dict:
    """audio_path: a file uploaded via @gradio/client handle_file().
    instrument: which stem to isolate -- "vocals" (default, the validated path),
    "guitar" or "piano". Only single-line melodies are meaningful downstream:
    PMI aligns one melodic line, and its thresholds were calibrated on vocals.
    Returns {'midi': b64, 'vocal': b64, 'chroma': {...}}; the 'vocal' key holds
    whichever stem was requested."""
    path = _resolve(audio_path)
    if instrument not in ("vocals", "guitar", "piano"):
        instrument = "vocals"

    # --- GPU window: separation + transcription only --------------------
    global _last_gpu_ok
    stem_bytes, notes = _transcribe_gpu(path, instrument)
    _last_gpu_ok = time.monotonic()      # a worker was warm as of now

    # --- CPU (no quota): encoding + chroma for the completion hints ------
    midi_bytes = notes_to_midi_bytes(notes)
    _st = tempfile.NamedTemporaryFile(delete=False, suffix=".wav")
    _st.write(stem_bytes); _st.flush(); _st.close()
    try:
        chroma = chroma_from_audio(_st.name)
    finally:
        try: os.unlink(_st.name)
        except OSError: pass

    return {
        "midi":  base64.b64encode(midi_bytes).decode(),
        "vocal": base64.b64encode(stem_bytes).decode(),
        "chroma": {
            "data":     base64.b64encode(chroma["data"]).decode(),
            "energy":   base64.b64encode(chroma["energy"]).decode(),
            "n_frames": chroma["n_frames"],
            "sr":       chroma["sr"],
            "hop":      chroma["hop"],
        },
    }


@app.api(name="transcribe_pair", concurrency_limit=1)
def transcribe_pair(audio_a: FileData, audio_b: FileData,
                    instrument_a: str = "vocals",
                    instrument_b: str = "vocals") -> dict:
    """Both excerpts in one call, so ZeroGPU counts one run instead of two.
    Returns {'a': {...}, 'b': {...}} with each side shaped exactly like the
    single-excerpt response, so the front end can reuse the same decoding."""
    global _last_gpu_ok
    if instrument_a not in ("vocals", "guitar", "piano"): instrument_a = "vocals"
    if instrument_b not in ("vocals", "guitar", "piano"): instrument_b = "vocals"

    # --- the single GPU window -------------------------------------------
    (stem_a, notes_a), (stem_b, notes_b) = _transcribe_pair_gpu(
        audio_a, audio_b, instrument_a, instrument_b)
    _last_gpu_ok = time.monotonic()

    # --- CPU (no quota): encoding + chroma -------------------------------
    def _pack(stem_bytes, notes):
        t = tempfile.NamedTemporaryFile(delete=False, suffix=".wav")
        t.write(stem_bytes); t.flush(); t.close()
        try:
            chroma = chroma_from_audio(t.name)
        finally:
            try: os.unlink(t.name)
            except OSError: pass
        return {
            "midi":  base64.b64encode(notes_to_midi_bytes(notes)).decode(),
            "vocal": base64.b64encode(stem_bytes).decode(),
            "chroma": {
                "data":     base64.b64encode(chroma["data"]).decode(),
                "energy":   base64.b64encode(chroma["energy"]).decode(),
                "n_frames": chroma["n_frames"],
                "sr":       chroma["sr"],
                "hop":      chroma["hop"],
            },
        }

    return {"a": _pack(stem_a, notes_a), "b": _pack(stem_b, notes_b)}


# key analysis — CPU only, plain FastAPI route (no GPU, no gradio client needed)
@app.post("/analyze")
async def analyze_api(midi: UploadFile = File(...)):
    p = _save_temp_midi(await midi.read())
    try:
        return JSONResponse(pmi_core.analyze_one_midi(p))
    except Exception as e:
        return JSONResponse({"error": str(e)}, status_code=400)
    finally:
        try: os.unlink(p)
        except OSError: pass


# PMI — CPU only, plain FastAPI route
@app.post("/pmi")
async def pmi_api(midi_a: UploadFile = File(...), midi_b: UploadFile = File(...),
                  tonic_a: str = Form(None), mode_a: str = Form(None),
                  tonic_b: str = Form(None), mode_b: str = Form(None)):
    pa = _save_temp_midi(await midi_a.read())
    pb = _save_temp_midi(await midi_b.read())
    try:
        result = pmi_core.pmi_from_two_midis(
            pa, pb,
            forced_A=_parse_forced(tonic_a, mode_a),
            forced_B=_parse_forced(tonic_b, mode_b),
        )
        return JSONResponse(result)
    except Exception as e:
        return JSONResponse({"error": str(e)}, status_code=400)
    finally:
        for p in (pa, pb):
            try: os.unlink(p)
            except OSError: pass


if __name__ == "__main__":
    app.launch()