#!/usr/bin/env python3
"""Music Transcriber + Trap Mode for Hugging Face Spaces (ZeroGPU)"""
import gradio as gr
import spaces
from pathlib import Path
import tempfile
import os
os.environ["HF_HOME"] = "/tmp/hf_cache"
os.environ["TORCH_HOME"] = "/tmp/hf_cache"
os.environ["TRANSFORMERS_CACHE"] = "/tmp/hf_cache"
from trap_transformer import TrapTransformer
@spaces.GPU(duration=600)
def transcribe_audio(audio_file, model_size, instruments, trap_mode=False):
"""Transcribe audio ā MIDI, optionally apply Trap transformation."""
if audio_file is None:
return None, "ā Please upload an audio file first.", '
0
'
from muscriptor import TranscriptionModel
# --- Parse instrument filter ---
instrument_group = None
instrument_names = []
if instruments:
try:
from muscriptor.tokenizer.mt3 import MT3_FULL_PLUS_GROUP_NAMES
names = [n.strip().lower() for n in instruments.split(",")]
valid_names = [n for n in names if n in MT3_FULL_PLUS_GROUP_NAMES]
instrument_group = " ".join(str(MT3_FULL_PLUS_GROUP_NAMES[n]) for n in valid_names)
instrument_names = valid_names
except Exception:
pass
# --- Load model ---
try:
model = TranscriptionModel.load_model(model_size)
except Exception as e:
return None, f"ā Model load failed: {e}", '0
'
# --- Transcribe ---
try:
midi_bytes = model.transcribe_to_midi(str(audio_file), instrument_group=instrument_group)
# Save original MIDI
original_path = Path(tempfile.mktemp(suffix=".mid"))
original_path.write_bytes(midi_bytes)
# Count notes from streaming pass
note_events = []
for event in model.transcribe(str(audio_file), instrument_group=instrument_group):
if hasattr(event, "note_id"):
note_events.append(event)
msg = f"ā
Complete!\nš¼ {len(note_events)} notes detected"
if instrument_names:
msg += f"\nš Filtered: {', '.join(instrument_names)}"
# --- Trap transformation ---
if trap_mode:
trap_path = Path(tempfile.mktemp(suffix="_trap.mid"))
transformer = TrapTransformer(bpm=140)
transformer.transform_to_trap(original_path, str(trap_path))
original_path.unlink(missing_ok=True)
msg += "\nš„ Trap mode applied (808s, hi-hats, half-time drums)"
return str(trap_path), msg, f'{len(note_events)}
'
return str(original_path), msg, f'{len(note_events)}
'
except Exception as e:
return None, f"ā Error: {e}", '0
'
def list_instruments():
try:
from muscriptor.tokenizer.mt3 import MT3_FULL_PLUS_GROUP_NAMES
return "\n".join(f"⢠{name}" for name in sorted(MT3_FULL_PLUS_GROUP_NAMES.keys()))
except Exception:
return "Could not load instrument list."
# ---------- UI ----------
with gr.Blocks(title="šµ Music Transcriber") as demo:
gr.Markdown("# šµ Music Transcriber\nTurn audio into MIDI using AI. Optionally transform to Trap style! š„")
with gr.Row():
# LEFT COLUMN ā Inputs
with gr.Column(scale=1):
audio_input = gr.Audio(
type="filepath",
label="Upload Audio (MP3, WAV, FLAC, OGG, M4A)",
interactive=True,
)
model_select = gr.Radio(
choices=[
("Small (~100M) - Fastest ā”", "small"),
("Medium (~300M) - Balanced š", "medium"),
("Large (~1.3B) - Best Quality š", "large"),
],
value="medium",
label="Model Size",
)
instrument_input = gr.Textbox(
label="Filter Instruments (optional)",
placeholder="e.g., acoustic_piano, acoustic_guitar",
)
trap_checkbox = gr.Checkbox(
label="š„ Transform to Trap Style (808s, hi-hats, half-time drums)",
value=False,
)
transcribe_btn = gr.Button("š¹ Transcribe Audio", variant="primary", size="lg")
# RIGHT COLUMN ā Outputs
with gr.Column(scale=1):
status_output = gr.Textbox(label="Status", lines=5)
note_count_output = gr.HTML(
'0
'
)
midi_output = gr.File(label="Download MIDI", visible=False)
with gr.Accordion("š View Available Instruments", open=False):
gr.Markdown(list_instruments())
# --- Wire up the button ---
transcribe_btn.click(
fn=transcribe_audio,
inputs=[audio_input, model_select, instrument_input, trap_checkbox],
outputs=[midi_output, status_output, note_count_output],
).then(
lambda path, html: (gr.update(visible=bool(path)), html),
inputs=[midi_output, note_count_output],
outputs=[midi_output, note_count_output],
)
if __name__ == "__main__":
demo.launch(theme=gr.themes.Soft(primary_hue="purple"), show_error=True)