#!/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)