Fornervegear's picture
Fix transcribe_and_postprocess API call for latest muscriptor
9a56ff7 verified
Raw History Blame Contribute Delete
21.6 kB
import spaces # MUST come before any torch / CUDA-touching import
import torch
import gradio as gr
import tempfile
import os
import time
import sys
import io
import json
import re
import base64
from pathlib import Path
import pretty_midi
from safetensors.torch import load_file
from huggingface_hub import hf_hub_download
from muscriptor.models.lm import LMModel, TorchAutocast
from muscriptor.modules.conditioners import (
MelSpectrogramConditioner,
ClassConditioner,
ConditioningProvider,
)
from muscriptor.tokenizer.mt3 import (
MT3Tokenizer,
MT3_FULL_PLUS_GROUP_NAMES,
get_group_program_map,
)
from muscriptor.transcription_model import (
TranscriptionModel,
_resolve_source,
_resolve_config,
_remap_single_codebook_keys,
_build_model,
)
_SAMPLE_RATE = 16000
def load_model_zerogpu():
"""Load the MuScriptor model in a ZeroGPU-compatible way.
On ZeroGPU, safetensors.load_file(device="cuda") fails because there's no
real GPU at module scope. However, .to("cuda") is intercepted by the
spaces hijack. So we:
1. Build the model with device="cuda" (conditioners store this as
self.device and use it at runtime to route tensors).
2. Load safetensors weights on CPU.
3. Load_state_dict into the model (weights land on CPU because the
model's tensors are still fake-CUDA via the hijack).
4. Call model.to("cuda") — ZeroGPU intercepts this and packs weights
to disk, streaming them into VRAM on the first @spaces.GPU call.
"""
device = torch.device("cuda") # ZeroGPU intercepts this
source = _resolve_source("large")
weights_path = Path(hf_hub_download(
repo_id="MuScriptor/muscriptor-large",
filename="model.safetensors",
))
cfg = _resolve_config(source, weights_path)
# Build the model with device="cuda" — conditioners store this and use it
# at runtime to route tensors. ZeroGPU intercepts .to("cuda") calls inside.
model = _build_model(device, cfg)
model.eval()
# Load weights on CPU (safetensors can't load to fake CUDA)
state_dict = load_file(str(weights_path), device="cpu")
state_dict = _remap_single_codebook_keys(state_dict)
# Load state dict — the model's parameters are fake-CUDA (ZeroGPU),
# but load_state_dict copies CPU data into them, which is fine.
model.load_state_dict(state_dict)
# Move everything to "cuda" — ZeroGPU intercepts this and packs to disk
model.to("cuda")
tokenizer = MT3Tokenizer(
instrument_vocabulary="MT3_FULL_PLUS",
max_shift_steps=1001,
)
return TranscriptionModel(model=model, tokenizer=tokenizer, device=device)
print("[muscriptor-space] Loading model...", file=sys.stderr, flush=True)
t0 = time.perf_counter()
model = load_model_zerogpu()
print(f"[muscriptor-space] Model loaded in {time.perf_counter() - t0:.1f}s", file=sys.stderr, flush=True)
# Build the instrument choices list (sorted by group ID for stable ordering)
INSTRUMENT_CHOICES = sorted(MT3_FULL_PLUS_GROUP_NAMES.keys(), key=lambda k: MT3_FULL_PLUS_GROUP_NAMES[k])
# MIDI program -> human-readable label using the model's own MT3_FULL_PLUS
# instrument groups (the model emits the representative program of a group,
# so this names lanes with the taxonomy the model actually decodes into).
# misc_programs="OMIT" keeps only the named groups 0-35; drums are handled
# via is_drum. Falls back to General MIDI names for unmapped programs.
_PROGRAM_TO_GROUP_LABEL: dict[int, str] = {}
try:
_gid_to_name = {v: k for k, v in MT3_FULL_PLUS_GROUP_NAMES.items()}
for _gid, _programs in get_group_program_map(
"MT3_FULL_PLUS", misc_programs="OMIT", is_mt3=True
).items():
_name = _gid_to_name.get(_gid)
if _name is None:
continue
_label = _name.replace("_", " ").title()
for _p in _programs:
_PROGRAM_TO_GROUP_LABEL[_p] = _label
except Exception:
pass
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
midi-player { width: 100%; display: block; margin-bottom: 8px; }
midi-visualizer { width: 100%; display: block; overflow: auto; }
.midi-roll-container { margin-top: 10px; }
.midi-roll-canvas { width: 100%; display: block; }
"""
# The html-midi-player web components (<midi-player> / <midi-visualizer>) rely
# on Magenta.js + Tone.js. Following Gradio's official custom-HTML-components
# guide (https://gradio.app/guides/custom-HTML-components), the library is
# loaded via the gr.HTML `head` parameter instead of embedding <script> tags in
# the component value — Gradio sanitizes value HTML and strips <script> tags,
# which is why the raw-HTML approach never mounted the web components.
_MIDI_PLAYER_HEAD = (
'<script src="https://cdn.jsdelivr.net/combine/'
"npm/tone@14.7.58,"
"npm/@magenta/music@1.23.1/es6/core.js,"
'npm/html-midi-player@1.5.0"></script>'
)
# Static markup for the player/visualizer. js_on_load wires the transcribed
# MIDI (a base64 data URI passed as the component value) into the <midi-player>
# and links it to the <midi-visualizer> so both mount and render.
_MIDI_PLAYER_TEMPLATE = (
'<div class="midi-player-container">'
'<midi-player sound-font visualizer=".midi-visualizer"></midi-player>'
'<midi-visualizer class="midi-visualizer" type="piano-roll"></midi-visualizer>'
"</div>"
)
# Runs once when the component loads. `element` is the component's DOM node,
# `props.value` is the base64 data URI (or "" when there's nothing yet).
# `watch('value', ...)` re-runs the sync every time the server pushes a new
# value (i.e. after each transcription), so the web components actually update.
_MIDI_PLAYER_JS = """
function syncPlayer() {
const player = element.querySelector('midi-player');
const container = element.querySelector('.midi-player-container');
if (!player || !container) { return; }
if (props.value) {
container.style.display = 'block';
player.src = props.value;
} else {
container.style.display = 'none';
player.removeAttribute('src');
}
}
syncPlayer();
watch('value', syncPlayer);
"""
# --- Piano-roll canvas (visual design ported from the space-demo-kit MIDI roll) ---
# One lane per instrument track: rounded lane background, label chip (colored
# dot + instrument name + note count), gridlines, note bars placed by
# onset/offset on x and pitch on y (per-lane pitch range), and a playhead
# synced to the <midi-player> above. Notes glow while they are being played.
_MIDI_ROLL_TEMPLATE = (
'<div class="midi-roll-container" style="display:none">'
'<canvas class="midi-roll-canvas"></canvas>'
"</div>"
)
_MIDI_ROLL_JS = r"""
const LANE_COLORS = ['#f97316', '#3b82f6', '#22c55e', '#a855f7', '#ec4899', '#eab308'];
const container = element.querySelector('.midi-roll-container');
const canvas = element.querySelector('.midi-roll-canvas');
let mt = null; // {span, tracks: [{label, drum, notes: [[t0, t1, pitch], ...]}]}
let lastT = 0; // last known playback time (seconds)
let live = false; // player currently playing
let rafId = null;
let wired = null; // the <midi-player> element we attached listeners to
function isDark() {
return document.body.classList.contains('dark') ||
document.documentElement.classList.contains('dark');
}
const PAD_L = 20, PAD_R = 20, PAD_T = 14, PAD_B = 18, LANE_H = 84, LANE_GAP = 12;
function layout() {
if (!mt) return;
const dpr = window.devicePixelRatio || 1;
const w = container.clientWidth || 600;
const nT = mt.tracks.length;
const h = PAD_T + PAD_B + nT * LANE_H + (nT - 1) * LANE_GAP;
canvas.width = Math.round(w * dpr);
canvas.height = Math.round(h * dpr);
canvas.style.width = w + 'px';
canvas.style.height = h + 'px';
}
function draw() {
if (!mt) return;
const dpr = window.devicePixelRatio || 1;
const cx = canvas.getContext('2d');
const W = canvas.width / dpr, H = canvas.height / dpr;
cx.setTransform(dpr, 0, 0, dpr, 0, 0);
cx.clearRect(0, 0, W, H);
const dark = isDark();
const span = mt.span;
const tNow = Math.min(Math.max(lastT, 0), span);
const innerW = W - PAD_L - PAD_R;
const tx = s => PAD_L + (s / span) * innerW;
// gridlines across the full roll (step adapts to the clip length)
const step = span > 240 ? 30 : span > 120 ? 15 : span > 60 ? 10 : span > 30 ? 5 : 1;
cx.strokeStyle = dark ? 'rgba(255,255,255,.07)' : 'rgba(17,24,39,.07)';
cx.lineWidth = 1;
for (let s = step; s < span; s += step) {
cx.beginPath(); cx.moveTo(tx(s), PAD_T); cx.lineTo(tx(s), H - PAD_B); cx.stroke();
}
mt.tracks.forEach((tr, i) => {
const y0 = PAD_T + i * (LANE_H + LANE_GAP);
const color = LANE_COLORS[i % LANE_COLORS.length];
// lane background
cx.fillStyle = dark ? 'rgba(255,255,255,.035)' : 'rgba(17,24,39,.04)';
cx.beginPath(); cx.roundRect(PAD_L - 10, y0, innerW + 20, LANE_H, 12); cx.fill();
// note area (below the label band)
const labelBand = 34;
const nTop = y0 + labelBand, nBot = y0 + LANE_H - 8;
let pMin = 127, pMax = 0;
for (const nn of tr.notes) { if (nn[2] < pMin) pMin = nn[2]; if (nn[2] > pMax) pMax = nn[2]; }
const pSpan = Math.max(1, pMax - pMin);
const nH = Math.max(4, Math.min(12, (nBot - nTop) / (pSpan + 1)));
for (const [t0, t1, p] of tr.notes) {
const x0 = tx(t0);
const w = Math.max(4, tx(t1) - x0);
const y = pMax === pMin
? (nTop + nBot - nH) / 2
: nTop + ((pMax - p) / pSpan) * (nBot - nTop - nH);
const active = live && tNow >= t0 && tNow <= t1;
const past = tNow > t1;
cx.save();
if (active) { cx.shadowColor = color; cx.shadowBlur = 12; cx.globalAlpha = 1; }
else cx.globalAlpha = past ? .85 : .3;
cx.fillStyle = color;
cx.beginPath();
cx.roundRect(x0, active ? y - 1 : y, w, active ? nH + 2 : nH, 2);
cx.fill();
cx.restore();
}
// label chip: colored dot + instrument name + note count
const label = tr.label.toUpperCase();
const count = tr.notes.length + (tr.notes.length === 1 ? ' NOTE' : ' NOTES');
cx.font = '700 13px system-ui, -apple-system, sans-serif';
const lw = cx.measureText(label).width;
cx.font = '500 12px ui-monospace, SFMono-Regular, Menlo, monospace';
const cw = cx.measureText(count).width;
const chipW = 12 + 9 + 7 + lw + 10 + cw + 12, chipH = 24;
cx.fillStyle = dark ? 'rgba(22,26,34,.92)' : 'rgba(255,255,255,.94)';
cx.shadowColor = 'rgba(0,0,0,.10)'; cx.shadowBlur = 8;
cx.beginPath(); cx.roundRect(PAD_L, y0 + 5, chipW, chipH, 999); cx.fill();
cx.shadowBlur = 0;
cx.fillStyle = color;
cx.beginPath(); cx.arc(PAD_L + 12 + 4, y0 + 5 + chipH / 2, 4, 0, 7); cx.fill();
cx.textBaseline = 'middle'; cx.textAlign = 'left';
cx.fillStyle = dark ? '#e5e7eb' : '#1f2430';
cx.font = '700 13px system-ui, -apple-system, sans-serif';
cx.fillText(label, PAD_L + 12 + 9 + 7, y0 + 6 + chipH / 2);
cx.fillStyle = dark ? '#9aa2ae' : '#6b7280';
cx.font = '500 12px ui-monospace, SFMono-Regular, Menlo, monospace';
cx.fillText(count, PAD_L + 12 + 9 + 7 + lw + 10, y0 + 6 + chipH / 2);
});
// playhead
const x = tx(tNow);
cx.fillStyle = dark ? '#e5e7eb' : '#111827';
cx.globalAlpha = live ? .85 : (lastT > 0 ? .3 : 0);
cx.beginPath(); cx.roundRect(x - 1.5, PAD_T - 4, 3, H - PAD_T - PAD_B + 8, 2); cx.fill();
cx.globalAlpha = 1;
}
// --- playhead sync with the <midi-player> web component (lives in the
// sibling gr.HTML component, so we look it up on document) ---
function tick() {
if (!live) { rafId = null; return; }
if (wired && typeof wired.currentTime === 'number') lastT = wired.currentTime;
draw();
rafId = requestAnimationFrame(tick);
}
function wirePlayer() {
const player = document.querySelector('midi-player');
if (!player || wired === player) return;
wired = player;
player.addEventListener('start', () => {
live = true;
if (rafId === null) rafId = requestAnimationFrame(tick);
});
player.addEventListener('stop', () => { // fires on both pause and finish
live = false;
if (rafId !== null) { cancelAnimationFrame(rafId); rafId = null; }
if (wired && typeof wired.currentTime === 'number') lastT = wired.currentTime;
draw();
});
}
// Low-rate poll: wires the player once it mounts, and catches scrubbing /
// theme flips while paused. Deliberately NOT a ResizeObserver.
setInterval(() => {
wirePlayer();
if (!live && wired && mt) {
const t = wired.currentTime;
if (typeof t === 'number' && Math.abs(t - lastT) > 0.05) { lastT = t; draw(); }
}
}, 400);
let resizeTimer = null;
window.addEventListener('resize', () => {
clearTimeout(resizeTimer);
resizeTimer = setTimeout(() => { layout(); draw(); }, 150);
});
function syncRoll() {
let data = null;
try { data = props.value ? JSON.parse(props.value) : null; } catch (e) { data = null; }
if (data && data.tracks && data.tracks.length) {
mt = data; lastT = 0; live = false;
container.style.display = 'block';
layout(); draw();
} else {
mt = null;
container.style.display = 'none';
}
}
syncRoll();
watch('value', syncRoll);
wirePlayer();
"""
def _midi_roll_data(midi_bytes: bytes | None) -> dict:
"""Per-track note data for the piano-roll canvas (and the summary stats).
Returns {"span": seconds, "tracks": [{"label", "drum", "notes":
[[onset, offset, pitch], ...]}, ...]} computed from the real returned
MIDI bytes with pretty_midi. All tracks are included (vocals too).
"""
if not midi_bytes:
return {}
try:
pm = pretty_midi.PrettyMIDI(io.BytesIO(midi_bytes))
except Exception:
return {}
tracks = []
for inst in pm.instruments:
if not inst.notes:
continue
if inst.is_drum:
label = "Drums"
elif inst.name and inst.name.strip():
label = inst.name.strip()
else:
label = _PROGRAM_TO_GROUP_LABEL.get(
inst.program
) or pretty_midi.program_to_instrument_name(inst.program)
notes = sorted(
[round(float(n.start), 4), round(float(n.end), 4), int(n.pitch)]
for n in inst.notes
)
tracks.append({"label": label, "drum": bool(inst.is_drum), "notes": notes})
if not tracks:
return {}
span = max(float(pm.get_end_time()), 0.001)
return {"span": round(span, 4), "tracks": tracks}
def _midi_player_src(midi_bytes: bytes | None) -> str:
"""Return the transcribed MIDI as a base64 data URI for the player.
The MIDI data is embedded directly as a base64 data URI so it needs no
separate file-serving route. The <midi-player> / <midi-visualizer> web
components consume this via the gr.HTML custom component's js_on_load.
"""
if not midi_bytes:
return ""
b64 = base64.b64encode(midi_bytes).decode("ascii")
return f"data:audio/midi;base64,{b64}"
@spaces.GPU(duration=120)
def transcribe_audio(
audio_path: str,
instruments: list[str] | None,
use_sampling: bool,
temperature: float,
progress=gr.Progress(track_tqdm=True),
):
"""Transcribe an audio recording into a downloadable MIDI file.
Upload any music recording (multi-instrument, any genre) and MuScriptor
will convert it into a MIDI file with per-note onset, offset, pitch, and
instrument information.
Args:
audio_path: Path to the uploaded audio file (wav, mp3, flac, etc.).
instruments: Optional list of instrument group names to condition the
transcription (improves coherence when you know which instruments
are present). Leave empty for automatic (unconditioned) transcription.
use_sampling: If True, use stochastic sampling instead of greedy decoding.
temperature: Sampling temperature (only used when use_sampling is True).
"""
if audio_path is None:
return None, "Please upload an audio file first.", "", ""
t0 = time.perf_counter()
# Run the transcription — returns MIDI bytes
try:
if hasattr(model, 'transcribe_to_midi'):
midi_bytes = model.transcribe_to_midi(
audio_path,
instruments=instruments if instruments else None,
use_sampling=use_sampling,
temperature=temperature,
)
else:
midi_bytes, _ = model.transcribe_and_postprocess(
audio_path,
instruments=instruments if instruments else None,
use_sampling=use_sampling,
temperature=temperature,
)
except Exception as e:
return None, f"Transcription failed: {e}", "", ""
elapsed = time.perf_counter() - t0
# Write to a temporary file for download
tmp = tempfile.NamedTemporaryFile(suffix=".mid", delete=False)
tmp.write(midi_bytes)
tmp.close()
# Per-track note data drives both the piano roll and the summary stats.
# (The previous mido-based stats loop hit an AttributeError on the first
# note — note_on messages have no .program — so it always reported 1 note.)
roll = _midi_roll_data(midi_bytes)
tracks = roll.get("tracks", [])
total_notes = sum(len(t["notes"]) for t in tracks)
if tracks:
per_track = ", ".join(f"{t['label']} {len(t['notes'])}" for t in tracks)
summary = (
f"Transcription complete in {elapsed:.1f}s. "
f"Found {total_notes} note{'s' if total_notes != 1 else ''} "
f"across {len(tracks)} track(s): {per_track}. "
f"Play it back below or download the MIDI file."
)
else:
summary = (
f"Transcription complete in {elapsed:.1f}s, but no notes were "
f"detected in the result."
)
player_src = _midi_player_src(midi_bytes)
roll_json = json.dumps(roll) if tracks else ""
return tmp.name, summary, player_src, roll_json
with gr.Blocks() as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""
# MuScriptor — Music Transcription (Audio → MIDI)
Upload a music recording and get a downloadable MIDI file with transcribed notes.
MuScriptor is a ~1.3B parameter multi-instrument automatic music transcription model
developed by [Mirelo](https://www.mirelo.ai/) x [Kyutai](https://kyutai.org/).
[Model card](https://huggingface.co/MuScriptor/muscriptor-large) · [Code](https://github.com/muscriptor/muscriptor) · [Audio samples](https://muscriptor.github.io)
"""
)
with gr.Row():
with gr.Column():
audio_input = gr.Audio(
label="Upload or record audio",
type="filepath",
sources=["upload", "microphone"],
)
with gr.Accordion("Advanced settings", open=False):
instrument_checkbox = gr.CheckboxGroup(
choices=INSTRUMENT_CHOICES,
value=[],
label="Instrument conditioning (optional)",
info="Select instruments present in the audio to improve transcription accuracy. Leave empty for automatic detection.",
)
use_sampling = gr.Checkbox(
label="Use sampling (stochastic decoding)",
value=False,
info="If enabled, uses temperature-based sampling instead of greedy decoding.",
)
temperature = gr.Slider(
label="Temperature",
minimum=0.1,
maximum=2.0,
value=1.0,
step=0.1,
info="Sampling temperature (only used when sampling is enabled).",
)
transcribe_btn = gr.Button("Transcribe", variant="primary")
with gr.Column():
midi_player = gr.HTML(
value="",
label="MIDI playback & visualization",
head=_MIDI_PLAYER_HEAD,
html_template=_MIDI_PLAYER_TEMPLATE,
js_on_load=_MIDI_PLAYER_JS,
)
midi_roll = gr.HTML(
value="",
label="Piano roll (per instrument)",
html_template=_MIDI_ROLL_TEMPLATE,
js_on_load=_MIDI_ROLL_JS,
)
midi_output = gr.File(label="Download MIDI file")
summary_output = gr.Textbox(label="Summary", interactive=False, visible=False)
transcribe_btn.click(
fn=transcribe_audio,
inputs=[audio_input, instrument_checkbox, use_sampling, temperature],
outputs=[midi_output, summary_output, midi_player, midi_roll],
api_name="transcribe",
)
gr.Examples(
examples=[
["example_medicine.mp3", [], False, 1.0],
["example_organic_flow.mp3", [], False, 1.0],
["example_water_afro_pop.mp3", [], False, 1.0],
],
inputs=[audio_input, instrument_checkbox, use_sampling, temperature],
outputs=[midi_output, summary_output, midi_player, midi_roll],
fn=transcribe_audio,
cache_examples=True,
cache_mode="lazy",
)
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)