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)