MuseMorphose / runtime.py
harp-dev's picture
Accept only .remi uploads
b8cd695 verified
Raw History Blame Contribute Delete
12.2 kB
"""Inference adapter for the official checkpoint and compatible REMI events."""
from __future__ import annotations
import pickle
import random
import sys
from functools import lru_cache
from pathlib import Path
import miditoolkit
import numpy as np
import torch
ROOT = Path(__file__).resolve().parent
sys.path.insert(0, str(ROOT / "model")) # Official model uses absolute sibling imports.
from model.musemorphose import MuseMorphose
from remi2midi import remi2midi
WEIGHTS = ROOT / "weights" / "musemorphose_pretrained_weights.pt"
VOCAB = ROOT / "pickles" / "remi_vocab.pkl"
RHYTHM_BOUNDS = [0.2, 0.25, 0.32, 0.38, 0.44, 0.5, 0.63]
POLYPHONY_BOUNDS = [2.63, 3.06, 3.50, 4.00, 4.63, 5.44, 6.44]
MAX_REMI_BYTES = 1_000_000
MAX_REMI_EVENTS = 10_000
MAX_REMI_BARS = 256
EVENT_NAMES = {
"Bar", "Beat", "Chord", "Tempo", "Note_Pitch", "Note_Velocity", "Note_Duration"
}
@lru_cache(maxsize=1)
def load_model():
if not WEIGHTS.is_file():
raise FileNotFoundError(f"Missing checkpoint: {WEIGHTS}")
event2idx, idx2event = pickle.loads(VOCAB.read_bytes())
model = MuseMorphose(
12,
8,
512,
2048,
12,
8,
512,
2048,
128,
512,
len(event2idx) + 1,
d_polyph_emb=64,
d_rfreq_emb=64,
cond_mode="in-attn",
)
try:
state = torch.load(WEIGHTS, map_location="cpu", weights_only=True)
except TypeError: # Older torch versions supported by the upstream project.
state = torch.load(WEIGHTS, map_location="cpu")
model.load_state_dict(state, strict=True)
model.eval()
return model, event2idx, idx2event
def _event_name(event):
return f"{event['name']}_{event['value']}"
def load_remi_text(path, event2idx):
"""Read vocabulary tokens from UTF-8 text without accepting pickle uploads."""
path = Path(path)
if path.suffix.lower() != ".remi":
raise ValueError("Upload a UTF-8 .remi event file.")
if path.stat().st_size > MAX_REMI_BYTES:
raise ValueError("REMI files must be at most 1 MB.")
try:
tokens = [line.strip() for line in path.read_text(encoding="utf-8-sig").splitlines()]
except UnicodeError as exc:
raise ValueError("The REMI file must use UTF-8 text.") from exc
tokens = [token for token in tokens if token]
if tokens and tokens[-1] == "EOS_None":
tokens.pop()
if not tokens or len(tokens) > MAX_REMI_EVENTS:
raise ValueError("The REMI file must contain 1–10,000 events.")
if tokens[0] != "Bar_None":
raise ValueError("The first REMI event must be Bar_None.")
positions = []
events = []
beat = -1
for index, token in enumerate(tokens):
if token not in event2idx:
raise ValueError(f"Unknown REMI token at line {index + 1}: {token}")
name = next(
(candidate for candidate in EVENT_NAMES if token.startswith(candidate + "_")),
None,
)
if name is None:
raise ValueError(f"Unsupported REMI event at line {index + 1}: {token}")
value = token[len(name) + 1 :]
if name == "Bar":
if token != "Bar_None":
raise ValueError("Only Bar_None is supported as a bar marker.")
positions.append(index)
beat = -1
if len(positions) > MAX_REMI_BARS:
raise ValueError("REMI files must have at most 256 bars.")
elif name == "Beat":
next_beat = int(value)
if next_beat < beat:
raise ValueError(f"Beat positions must increase within bar {len(positions)}.")
beat = next_beat
elif beat < 0:
raise ValueError(f"Bar {len(positions)} needs a Beat event before musical events.")
events.append({"name": name, "value": None if name == "Bar" else value})
for bar_number, (left, right) in enumerate(
zip(positions, positions[1:] + [len(tokens)]), start=1
):
if right - left > 128:
raise ValueError(f"Bar {bar_number} exceeds the 128-event encoder limit.")
for index in range(left + 1, right):
token = tokens[index]
if token.startswith("Note_Pitch_") and (
index + 2 >= right
or not tokens[index + 1].startswith("Note_Velocity_")
or not tokens[index + 2].startswith("Note_Duration_")
):
raise ValueError(f"Bar {bar_number} has an incomplete note event.")
if token.startswith("Note_Velocity_") and (
index == left or not tokens[index - 1].startswith("Note_Pitch_")
):
raise ValueError(f"Bar {bar_number} has an orphan velocity event.")
if token.startswith("Note_Duration_") and (
index < left + 2
or not tokens[index - 1].startswith("Note_Velocity_")
or not tokens[index - 2].startswith("Note_Pitch_")
):
raise ValueError(f"Bar {bar_number} has an orphan duration event.")
return positions, events
def _classes(events, n_bars):
poly = np.zeros((n_bars * 16,), dtype=np.float32)
rhythm = np.zeros_like(poly)
bar, beat = -1, 0
for event in events:
name, value = event["name"], event["value"]
if name == "Bar":
bar += 1
beat = 0
elif name == "Beat":
beat = int(value)
elif 0 <= bar < n_bars and name == "Note_Pitch":
rhythm[bar * 16 + beat] = 1
elif 0 <= bar < n_bars and name == "Note_Duration":
start = bar * 16 + beat
poly[start : min(start + int(value) // 120, len(poly))] += 1
return (
np.searchsorted(RHYTHM_BOUNDS, rhythm.reshape(n_bars, 16).mean(axis=1)),
np.searchsorted(POLYPHONY_BOUNDS, poly.reshape(n_bars, 16).mean(axis=1)),
)
def _piece(bar_positions, all_events, start_bar, n_bars, event2idx):
if start_bar < 0 or n_bars < 1 or start_bar + n_bars > len(bar_positions):
raise ValueError("The requested bar range is outside this piece.")
start = bar_positions[start_bar]
end = (
bar_positions[start_bar + n_bars]
if start_bar + n_bars < len(bar_positions)
else len(all_events)
)
events = all_events[start:end]
bar_starts = [bar_positions[start_bar + i] - start for i in range(n_bars)]
bar_ends = bar_starts[1:] + [len(events)]
bars = []
for left, right in zip(bar_starts, bar_ends):
ids = [event2idx[_event_name(event)] for event in events[left:right]]
bars.append(ids[:128])
if not any(event["name"] == "Note_Pitch" for event in events):
raise ValueError("The selected bars contain no piano notes.")
# Attribute classes are computed over the full piece in the official
# preprocessing; notes held across a bar line affect the next bar.
full_rhythm, full_poly = _classes(all_events, len(bar_positions))
rhythm = full_rhythm[start_bar : start_bar + n_bars]
poly = full_poly[start_bar : start_bar + n_bars]
return events, bars, rhythm, poly
def _sample_token(logits, temperature, top_p):
probs = torch.softmax(logits.float() / temperature, dim=-1)
sorted_probs, sorted_idx = probs.sort(descending=True)
keep = (sorted_probs.cumsum(0) - sorted_probs) < top_p
candidate_probs = sorted_probs[keep]
candidate_idx = sorted_idx[keep]
sampled = torch.multinomial(candidate_probs / candidate_probs.sum(), 1)
return int(candidate_idx[sampled].item())
def generate(
remi_path,
start_bar,
n_bars,
rhythm_shift,
poly_shift,
temperature,
top_p,
seed,
output_dir,
):
if not remi_path:
raise ValueError("Upload compatible REMI events.")
n_bars, start_bar, seed = int(n_bars), int(start_bar), int(seed)
if n_bars > 4:
raise ValueError("Choose at most four bars per request.")
if not 0.5 <= temperature <= 2 or not 0 < top_p <= 1:
raise ValueError("Temperature or top-p is out of range.")
torch.set_num_threads(4)
model, event2idx, idx2event = load_model()
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
positions, all_events = load_remi_text(remi_path, event2idx)
events, bars, rhythm, poly = _piece(
positions, all_events, start_bar, n_bars, event2idx
)
source = Path(remi_path).name
lengths = [len(bar) for bar in bars]
encoder = torch.full((128, n_bars), len(event2idx), dtype=torch.long)
mask = torch.ones((n_bars, 128), dtype=torch.bool)
for i, bar in enumerate(bars):
encoder[: len(bar), i] = torch.tensor(bar)
mask[i, : len(bar)] = False
with torch.inference_mode():
latents = model.get_sampled_latent(encoder, padding_mask=mask)
rhythm = np.clip(rhythm + int(rhythm_shift), 0, 7).astype(int).tolist()
poly = np.clip(poly + int(poly_shift), 0, 7).astype(int).tolist()
tokens = [event2idx["Bar_None"]]
context = tokens.copy()
segment = [0]
beat = 0
max_tokens = min(800, 200 * n_bars)
completed = False
with torch.inference_mode():
for _ in range(max_tokens):
current_bar = min(segment[-1], n_bars - 1)
inp = torch.tensor(context, dtype=torch.long)[:, None]
latent = torch.stack([latents[min(i, n_bars - 1)] for i in segment])[
:, None, :
]
rhythm_tensor = torch.tensor([rhythm[min(i, n_bars - 1)] for i in segment])[
:, None
]
poly_tensor = torch.tensor([poly[min(i, n_bars - 1)] for i in segment])[
:, None
]
logits = model.generate(inp, latent, rhythm_tensor, poly_tensor)[0]
token = _sample_token(logits, float(temperature), float(top_p))
word = idx2event.get(token, "PAD_None")
if word.startswith("Beat_"):
next_beat = int(word.split("_")[1])
if next_beat < beat:
continue
beat = next_beat
if word == "PAD_None":
continue
if word in {"Bar_None", "EOS_None"}:
if current_bar + 1 >= n_bars:
completed = True
break
beat = 0
segment.append(current_bar + 1)
context.append(event2idx["Bar_None"])
tokens.append(event2idx["Bar_None"])
else:
context.append(token)
segment.append(current_bar)
tokens.append(token)
if len(context) >= 1024:
context = context[-512:]
segment = segment[-512:]
if not completed:
raise RuntimeError(
"Generation reached the event limit. Try a different seed or fewer bars."
)
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
reference_path = output_dir / "reference.mid"
output_path = output_dir / "musemorphose.mid"
remi_output_path = output_dir / "musemorphose.remi"
source_events = [_event_name(event) for event in events]
_, tempos = remi2midi(source_events, str(reference_path), return_first_tempo=True)
result_events = [idx2event[token] for token in tokens]
remi_output_path.write_text("\n".join(result_events) + "\n", encoding="utf-8")
remi2midi(
result_events, str(output_path), enforce_tempo=True, enforce_tempo_val=tempos
)
saved_midi = miditoolkit.MidiFile(str(output_path))
saved_notes = sum(len(instrument.notes) for instrument in saved_midi.instruments)
if not saved_notes:
raise RuntimeError(
"The model generated no notes for this seed. Try a different seed."
)
details = {
"source": source,
"source_bars": n_bars,
"source_bar_lengths": lengths,
"rhythm_classes": rhythm,
"polyphony_classes": poly,
"generated_events": len(tokens),
"generated_notes": saved_notes,
}
return str(reference_path), str(output_path), str(remi_output_path), details