Spaces:
Running on Zero
Running on Zero
File size: 5,748 Bytes
2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 9e06a06 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 9e06a06 2697437 9e06a06 2697437 0fbada6 9e06a06 0fbada6 2697437 0fbada6 2697437 0fbada6 9e06a06 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 9e06a06 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 0fbada6 2697437 9e06a06 | 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 | #!/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.", '<div style="font-size: 2rem; font-weight: bold; color: #8b5cf6; text-align: center; padding: 1rem;">0</div>'
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}", '<div style="font-size: 2rem; font-weight: bold; color: #8b5cf6; text-align: center; padding: 1rem;">0</div>'
# --- 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'<div style="font-size: 2rem; font-weight: bold; color: #8b5cf6; text-align: center; padding: 1rem;">{len(note_events)}</div>'
return str(original_path), msg, f'<div style="font-size: 2rem; font-weight: bold; color: #8b5cf6; text-align: center; padding: 1rem;">{len(note_events)}</div>'
except Exception as e:
return None, f"β Error: {e}", '<div style="font-size: 2rem; font-weight: bold; color: #8b5cf6; text-align: center; padding: 1rem;">0</div>'
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(
'<div style="font-size: 2rem; font-weight: bold; color: #8b5cf6; text-align: center; padding: 1rem;">0</div>'
)
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) |