SymphonyGen / app.py
Xuzheng's picture
Cap files at 32 bars
2da2c95 verified
Raw
History Blame Contribute Delete
21.7 kB
"""
SymphonyGen demo Space.
Paper: https://arxiv.org/abs/2604.25498 (ISMIR 2026)
Code: https://github.com/symphonygen/symphonygen
Models: https://huggingface.co/SymphonyGen/SymphonyGen
The Space clones the code repository, downloads the released packed
checkpoints, and exposes a few inference options: skeleton source (sampled /
uploaded / analyzed), model variant, dissonance-averse sampling weights, and
an instrument mask. Audio previews are rendered with headless MuseScore 3
(as in the paper), falling back to FluidSynth + MuseScore_General.sf3.
"""
import os
import shutil
import subprocess
import sys
import tempfile
import uuid
from pathlib import Path
# ZeroGPU support (no-op elsewhere). Must be imported before torch so that
# the CUDA emulation layer can intercept module-level model placements.
try:
import spaces
gpu_decorator = spaces.GPU(duration=120)
except ImportError:
def gpu_decorator(fn):
return fn
APP_DIR = Path(__file__).resolve().parent
REPO_URL = "https://github.com/symphonygen/symphonygen"
REPO_DIR = APP_DIR / "symphonygen"
ASSET_DIR = APP_DIR / "asset"
OUT_ROOT = Path(tempfile.gettempdir()) / "symphonygen_demo"
# --- Code + checkpoints ----------------------------------------------------
if not REPO_DIR.exists():
subprocess.run(["git", "clone", "--depth", "1", REPO_URL, str(REPO_DIR)], check=True)
os.environ.setdefault("ASSET_DIR", str(ASSET_DIR))
os.environ.setdefault("WORK_DIR", str(OUT_ROOT / "work"))
os.makedirs(os.environ["WORK_DIR"], exist_ok=True)
sys.path.insert(0, str(REPO_DIR))
from huggingface_hub import snapshot_download
snapshot_download("SymphonyGen/SymphonyGen", local_dir=str(ASSET_DIR), allow_patterns=["*.pt"])
# --- Headless MuseScore 3.6.2 (the paper's renderer) -----------------------
# Same setup as utils/headless_musescore.sh in the code repo: extract the
# AppImage (no FUSE in containers) and run squashfs-root/AppRun under Xvfb.
MUSESCORE_APPIMAGE = "MuseScore-3.6.2.548021370-x86_64.AppImage"
MUSESCORE_URL = f"https://github.com/musescore/MuseScore/releases/download/v3.6.2/{MUSESCORE_APPIMAGE}"
MUSESCORE_APPRUN = ASSET_DIR / "squashfs-root" / "AppRun"
def setup_headless_musescore():
if MUSESCORE_APPRUN.exists():
return
appimage = ASSET_DIR / MUSESCORE_APPIMAGE
try:
if not appimage.exists():
from urllib.request import urlretrieve
urlretrieve(MUSESCORE_URL, appimage)
appimage.chmod(0o755)
subprocess.run(
[f"./{MUSESCORE_APPIMAGE}", "--appimage-extract"],
cwd=ASSET_DIR, check=True, capture_output=True, timeout=600,
)
subprocess.run(["chmod", "-R", "+x", str(ASSET_DIR / "squashfs-root")], check=True)
except Exception as e:
print(f"Headless MuseScore setup failed ({e}); falling back to FluidSynth.")
setup_headless_musescore()
import gradio as gr
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import mido
import torch
from miditoolkit import MidiFile
from arch.config import BAR_NUM
from arch.harmo.data_tensor import HarmonyTensorConverter
from arch.harmo.generator import HarmonyGenerator
from arch.harmo.model import HarmonyGPT
from arch.harmo.sampling import HarmPosConstrainer
from arch.symph import generate
from arch.symph.cond_gen import SymphonyGenerator3D
from arch.symph.data_tensor import MusicTensorConverter3D
from arch.symph.model import MusicModel3D
from arch.symph.sampling import DissonanceConstrainer
from data_prep.harmo_analysis import HarmonySkeletonAnalyzer
from data_prep.main import export_midi, preprocess_midi
from rl.reward.harmo_filter import filter_harmony_by_rule
from utils.remove_harmo_outline import POSSIBLE_HARMO_OUTLINE_TRACK_NAMES, midi_remove_harmo_outline
MODEL_VARIANTS = {
# label: (checkpoint, default register decay)
"reinforced+track (GRPO, CLaMP 3 + track density reward)": ("grpo_clamp+track_epoch_6.pt", True),
"reinforced (GRPO, pure CLaMP 3 reward)": ("grpo_clamp_epoch_10.pt", False),
"pretrained (no RL)": ("stage_two_pretrained.pt", False),
}
SKELETON_SOURCES = [
"Sample a new skeleton (harmony model)",
"Upload a harmony-skeleton MIDI",
"Upload any MIDI and analyze its harmony (re-orchestration)",
]
# Audio rendering: headless MuseScore 3 first (the paper's renderer), then
# FluidSynth with the MuseScore General soundfont, then FluidR3.
SOUNDFONTS = [
"/usr/share/sounds/sf3/MuseScore_General.sf3",
"/usr/share/sounds/sf3/MuseScore_General_Small.sf3",
"/usr/share/sounds/sf2/FluidR3_GM.sf2",
]
# Set PROVIDE_MIDI=0 in the Space variables to expose audio previews only.
PROVIDE_MIDI = os.environ.get("PROVIDE_MIDI", "1") == "1"
EXAMPLE_DIR = APP_DIR / "examples"
# Harmony beats must stay on quarters for the symphony model (see Results/batch_filter_gen.py)
HarmPosConstrainer.limit_at_quarter = True
_model_cache = {}
def load_symph_model(ckpt_name: str) -> MusicModel3D:
if ckpt_name not in _model_cache:
model = MusicModel3D.from_pretrained(str(ASSET_DIR / ckpt_name))
_model_cache[ckpt_name] = model.eval()
return _model_cache[ckpt_name]
def load_harmo_model() -> HarmonyGPT:
if "harmo" not in _model_cache:
model = HarmonyGPT.from_pretrained(str(ASSET_DIR / "stage_one_pretrained.pt"))
_model_cache["harmo"] = model.eval()
return _model_cache["harmo"]
# Preload every model at module level: ZeroGPU records the CUDA placements
# during startup (emulated CUDA), so @spaces.GPU calls skip transfer costs.
load_harmo_model()
for _ckpt_name, _ in MODEL_VARIANTS.values():
load_symph_model(_ckpt_name)
def sample_skeleton(apply_filters: bool):
""" Sample one harmony skeleton, optionally curated by the paper's filters. """
generator = HarmonyGenerator(HarmonyTensorConverter())
harmo_dict = generator.run(load_harmo_model(), num_batches=1, batch_size=8)
candidates = list(harmo_dict.values())
note = ""
if apply_filters:
good = filter_harmony_by_rule(list(candidates), num_requested=1)
good = generator.filter_by_log_prob(load_harmo_model(), good, num_requested=1)
if good:
candidates = good
else:
note = "No sampled skeleton passed the quality filters; using an unfiltered one. "
bars = candidates[0]
# Sampled skeletons are pruned into pure template chords (paper Sec. 4.2)
HarmonySkeletonAnalyzer(triad_match_strict=True).purify_harmony_(bars)
return bars, note
MAX_PREVIEW_BARS = BAR_NUM # the model's supported piece length (32 bars)
def _bar_starts(midi: MidiFile, max_bars: int) -> list[int]:
""" Bar boundary ticks honoring time-signature changes.
Same walk as BeatToTime.parse (music/lib/beat_to_time.py): each bar spans
whole_ticks * numerator / denominator; a signature change landing mid-bar
truncates that bar at the change.
"""
whole_ticks = 4 * midi.ticks_per_beat
changes = sorted(midi.time_signature_changes, key=lambda c: c.time)
starts = [0]
tick, i = 0, 0
ticks_per_bar = whole_ticks # 4/4 default
while len(starts) <= max_bars:
while i < len(changes) and changes[i].time <= tick:
ticks_per_bar = max(1, whole_ticks * changes[i].numerator // changes[i].denominator)
i += 1
next_tick = tick + ticks_per_bar
if i < len(changes) and tick < changes[i].time < next_tick:
next_tick = changes[i].time
starts.append(next_tick)
tick = next_tick
return starts
def piano_roll_image(midi_path: Path, title: str) -> str | None:
""" Toy single-color piano roll of the first MAX_PREVIEW_BARS bars. """
try:
midi = MidiFile(str(midi_path))
tpb = midi.ticks_per_beat
bar_starts = _bar_starts(midi, MAX_PREVIEW_BARS)
end_tick = bar_starts[-1]
spans_by_pitch: dict[int, list[tuple[float, float]]] = {}
for inst in midi.instruments:
if inst.is_drum:
continue
for n in inst.notes:
if n.start >= end_tick:
continue
start = n.start / tpb
dur = (min(n.end, end_tick) - n.start) / tpb
spans_by_pitch.setdefault(n.pitch, []).append((start, max(dur, 0.1)))
if not spans_by_pitch:
return None
fig, ax = plt.subplots(figsize=(10, 3.5), dpi=110)
for t in bar_starts:
ax.axvline(t / tpb, color="0.88", linewidth=0.6, zorder=0)
for pitch, spans in spans_by_pitch.items():
ax.broken_barh(spans, (pitch - 0.4, 0.8), color="#46608a", linewidth=0)
pitches = list(spans_by_pitch)
ax.set_xlim(0, end_tick / tpb)
ax.set_ylim(min(pitches) - 2, max(pitches) + 2)
label_bars = range(0, MAX_PREVIEW_BARS + 1, 4)
ax.set_xticks([bar_starts[b] / tpb for b in label_bars])
ax.set_xticklabels([str(b) for b in label_bars])
ax.set_xlabel("bar")
ax.set_ylabel("pitch")
ax.set_title(title, fontsize=9)
fig.tight_layout()
img_path = midi_path.with_suffix(".roll.png")
fig.savefig(img_path)
plt.close(fig)
return str(img_path)
except Exception:
return None
def _has_harmo_outline_track(midi_path: Path) -> bool:
""" Whether the MIDI is a harmony-skeleton file, identified by track name
as in utils/remove_harmo_outline.py. """
try:
midi_obj = mido.MidiFile(str(midi_path))
except Exception:
return False
return any(t.name in POSSIBLE_HARMO_OUTLINE_TRACK_NAMES for t in midi_obj.tracks)
def truncate_midi_to_bars(midi_path: Path, max_bars: int = BAR_NUM) -> Path:
""" Cap a user-uploaded MIDI at max_bars upfront (the model does not
support longer pieces). Returns the original path when already short. """
try:
midi = MidiFile(str(midi_path))
end_tick = _bar_starts(midi, max_bars)[-1]
if midi.max_tick <= end_tick:
return midi_path
for inst in midi.instruments:
inst.notes = [n for n in inst.notes if n.start < end_tick]
for n in inst.notes:
n.end = min(n.end, end_tick)
for attr in ("control_changes", "pitch_bends", "pedals"):
events = getattr(inst, attr, None)
if events is not None:
setattr(inst, attr, [e for e in events if getattr(e, "time", getattr(e, "start", 0)) < end_tick])
for attr in ("tempo_changes", "time_signature_changes", "key_signature_changes", "markers", "lyrics"):
events = getattr(midi, attr, None)
if events is not None:
setattr(midi, attr, [e for e in events if e.time < end_tick])
capped_path = midi_path.with_suffix(".cap.mid")
midi.dump(str(capped_path))
return capped_path
except Exception:
return midi_path
def export_skeleton_midi(bars, out_path: Path):
""" Export only the harmony skeleton of the given bars. """
harmo_bars = [bar.__class__(duration_q=bar.duration_q) for bar in bars]
for harmo_bar, bar in zip(harmo_bars, bars):
harmo_bar.harmony_track = bar.harmony_track
export_midi(harmo_bars, out_path)
def preview_uploaded_midi(skeleton_source: str, midi_file: str | None):
""" Preview the harmony skeleton of an upload (not the raw MIDI): the
existing outline track when present, otherwise the analyzed harmony. """
if not midi_file:
return gr.update(value=None)
path = truncate_midi_to_bars(Path(midi_file))
try:
bars = preprocess_midi(
path,
analyze_harmo=skeleton_source == SKELETON_SOURCES[2],
analyze_harmo_if_not_exist=True,
)
skeleton_path = path.with_suffix(".skeleton.mid")
export_skeleton_midi(bars, skeleton_path)
except Exception:
return gr.update(value=None)
return piano_roll_image(skeleton_path, f"Harmony skeleton (first {MAX_PREVIEW_BARS} bars)")
def _mscore_cmd() -> list[str] | None:
""" MuseScore invocation, preferring the paper's 3.6.2 AppImage (see
utils/midi2audio.py in the code repo). """
if not shutil.which("xvfb-run"):
return None
if MUSESCORE_APPRUN.exists():
return ["xvfb-run", "--auto-servernum", str(MUSESCORE_APPRUN)]
for name in ("mscore3", "musescore3", "mscore", "musescore"):
if shutil.which(name):
return ["xvfb-run", "--auto-servernum", name]
return None
def _render_midi_to_wav(render_src: Path, wav_path: Path) -> str | None:
""" Tries headless MuseScore 3 under Xvfb, then FluidSynth with the first
available soundfont. """
mscore = _mscore_cmd()
if mscore:
try:
subprocess.run(
mscore + ["-o", str(wav_path), str(render_src)],
check=True, capture_output=True, timeout=300,
)
if wav_path.exists():
return str(wav_path)
except Exception:
pass
soundfont = next((sf for sf in SOUNDFONTS if os.path.exists(sf)), None)
if shutil.which("fluidsynth") and soundfont:
try:
subprocess.run(
["fluidsynth", "-ni", soundfont, str(render_src), "-F", str(wav_path), "-r", "44100"],
check=True, capture_output=True, timeout=300,
)
except Exception:
return None
return str(wav_path) if wav_path.exists() else None
def render_audio(midi_path: Path) -> str | None:
""" Render the song preview WAV, muting the harmony outline track. """
no_harmo_path = midi_path.with_suffix(".no_harmo.mid")
try:
render_src = midi_remove_harmo_outline(midi_path, no_harmo_path)
except Exception:
render_src = midi_path
return _render_midi_to_wav(Path(render_src), midi_path.with_suffix(".wav"))
def render_skeleton_audio(skeleton_path: Path) -> str | None:
""" Render the harmony skeleton itself, with every track forced to piano
(no outline muting here — the skeleton is what we want to hear). """
try:
midi = MidiFile(str(skeleton_path))
for inst in midi.instruments:
inst.program = 0 # acoustic grand piano
inst.is_drum = False
piano_path = skeleton_path.with_suffix(".piano.mid")
midi.dump(str(piano_path))
except Exception:
return None
return _render_midi_to_wav(piano_path, piano_path.with_suffix(".wav"))
@gpu_decorator
def generate_song(
skeleton_source: str,
midi_file: str | None,
apply_filters: bool,
variant_label: str,
dissonance_averse: bool,
hn_weight: float,
nn_weight: float,
register_decay: bool,
forbid_piano: bool,
):
out_dir = OUT_ROOT / uuid.uuid4().hex[:8]
out_dir.mkdir(parents=True, exist_ok=True)
note = ""
# 1. The harmony skeleton condition
analyze_harmo = False
if skeleton_source == SKELETON_SOURCES[0]:
bars, note = sample_skeleton(apply_filters)
cond = bars
else:
if not midi_file:
raise gr.Error("Please upload a MIDI file for this skeleton source.")
cond = truncate_midi_to_bars(Path(midi_file))
if cond != Path(midi_file):
note += f"Input capped at the first {BAR_NUM} bars (model limit). "
analyze_harmo = skeleton_source == SKELETON_SOURCES[2]
if not analyze_harmo and not _has_harmo_outline_track(cond):
# Not a skeleton MIDI (no harmony outline track): silently fall
# back to analyzing its harmony instead.
analyze_harmo = True
# 2. Sampling configuration (module-level, as in arch/symph/generator.py)
ckpt_name, _ = MODEL_VARIANTS[variant_label]
generate.FORBID_PIANO = forbid_piano
DissonanceConstrainer.enabled = dissonance_averse
DissonanceConstrainer.HN_weight = hn_weight
DissonanceConstrainer.NN_weight = nn_weight
DissonanceConstrainer.register_decay = register_decay
# 3. Generate
generator = SymphonyGenerator3D(MusicTensorConverter3D())
model = load_symph_model(ckpt_name)
with torch.no_grad():
song_dict = generator.run(model, [cond], group_size=1, analyze_harmo=analyze_harmo)
if not song_dict:
raise gr.Error("Generation produced no decodable song; please try again.")
bars = next(iter(song_dict.values()))
song_path = out_dir / "symphonygen_song.mid"
export_midi(bars, song_path, harmo_is_cond=True)
skeleton_path = out_dir / "harmony_skeleton.mid"
export_skeleton_midi(bars, skeleton_path)
return song_path, skeleton_path, note
def run_generation(*args):
""" CPU wrapper: generation on GPU, then audio/piano-roll rendering
outside the @spaces.GPU window (saves ZeroGPU quota). """
song_path, skeleton_path, note = generate_song(*args)
audio_path = render_audio(song_path)
if audio_path is None:
note += "Audio preview unavailable (no MuseScore/FluidSynth renderer found)."
skeleton_audio = render_skeleton_audio(skeleton_path)
skeleton_img = piano_roll_image(skeleton_path, "Harmony skeleton (the condition)")
song_img = piano_roll_image(song_path, "Generated orchestral piece")
song_file = str(song_path) if PROVIDE_MIDI else None
skeleton_file = str(skeleton_path) if PROVIDE_MIDI else None
return (audio_path, song_file, skeleton_file, note or "Done.",
skeleton_img, song_img, skeleton_audio)
def variant_changed(variant_label: str):
""" Follow the paper's recommended register decay per variant. """
return gr.update(value=MODEL_VARIANTS[variant_label][1])
with gr.Blocks(title="SymphonyGen Demo") as demo:
gr.Markdown(
"# 🎻 SymphonyGen\n"
"**3D Hierarchical Orchestral Generation with Controllable Harmony Skeleton** (ISMIR 2026)\n\n"
"[📄 Paper](https://arxiv.org/abs/2604.25498) · "
"[💻 Code](https://github.com/symphonygen/symphonygen) · "
"[🤗 Checkpoints](https://huggingface.co/SymphonyGen/SymphonyGen) · "
"[🌐 Audio demo page](https://symphonygen.github.io)\n\n"
"Generates a 32-bar orchestral piece conditioned on a beat-quantized multi-voice "
"harmony skeleton. The skeleton can be sampled from the released harmony model, "
"uploaded directly, or analyzed from any MIDI (re-orchestration)."
)
with gr.Row():
with gr.Column():
skeleton_source = gr.Radio(SKELETON_SOURCES, value=SKELETON_SOURCES[0], label="Harmony skeleton source")
midi_file = gr.File(label=f"Conditioning MIDI (for the upload options; capped at the first {BAR_NUM} bars)", file_types=[".mid", ".midi"], type="filepath")
input_roll = gr.Image(label="Input piano roll (harmony skeleton)", interactive=False)
input_audio = gr.Audio(label="Skeleton audio preview (piano)")
apply_filters = gr.Checkbox(True, label="Curate sampled skeletons with the paper's quality filters (density / repetition / log-probability)")
variant = gr.Dropdown(list(MODEL_VARIANTS), value=list(MODEL_VARIANTS)[0], label="Model variant")
with gr.Accordion("Dissonance-averse sampling", open=False):
dissonance_averse = gr.Checkbox(True, label="Enable")
hn_weight = gr.Slider(0.0, 5.0, value=1.0, step=0.5, label="λ_hn (harmonic-vs-non-harmonic clash weight)")
nn_weight = gr.Slider(0.0, 20.0, value=10.0, step=1.0, label="λ_nn (non-harmonic-vs-non-harmonic clash weight)")
register_decay = gr.Checkbox(True, label="Register-dependent decay (Low Interval Limit)")
forbid_piano = gr.Checkbox(True, label="Forbid the piano instrument family")
run_btn = gr.Button("Generate", variant="primary")
with gr.Column():
audio_out = gr.Audio(label="Audio preview (headless MuseScore 3, as in the paper)")
output_roll = gr.Image(label="Output piano roll", interactive=False)
song_out = gr.File(label="Generated orchestral MIDI", visible=PROVIDE_MIDI)
skeleton_out = gr.File(label="Harmony skeleton MIDI (the condition)", visible=PROVIDE_MIDI)
status = gr.Textbox(label="Status", interactive=False)
example_midis = sorted(EXAMPLE_DIR.glob("*.mid")) if EXAMPLE_DIR.exists() else []
if example_midis:
gr.Examples(
examples=[[SKELETON_SOURCES[1], str(p)] for p in example_midis],
inputs=[skeleton_source, midi_file],
label="Example harmony skeletons (SymphonyNet validation set)",
)
# api_visibility="private": endpoints unavailable to gradio_client (UI-only)
variant.change(variant_changed, inputs=variant, outputs=register_decay,
api_visibility="private")
midi_file.change(preview_uploaded_midi, inputs=[skeleton_source, midi_file],
outputs=input_roll, api_visibility="private")
skeleton_source.change(preview_uploaded_midi, inputs=[skeleton_source, midi_file],
outputs=input_roll, api_visibility="private")
run_btn.click(
run_generation,
inputs=[skeleton_source, midi_file, apply_filters, variant,
dissonance_averse, hn_weight, nn_weight, register_decay, forbid_piano],
outputs=[audio_out, song_out, skeleton_out, status, input_roll, output_roll, input_audio],
api_visibility="private",
)
if __name__ == "__main__":
# No API footer link; one generation at a time (endpoints are private above).
demo.queue(max_size=20, default_concurrency_limit=1).launch(
footer_links=["gradio", "settings"]
)