"""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