Spaces:
Running
Running
Download runtime.py from teamup-tech/MuseMorphose: direct link, hf CLI and curl.
- Browser
- Download file 12.2 kB
-
https://huggingface.co/spaces/teamup-tech/MuseMorphose/resolve/main/runtime.py
- Command line
-
hf download hf://spaces/teamup-tech/MuseMorphose/runtime.py
-
curl -L -o runtime.py https://huggingface.co/spaces/teamup-tech/MuseMorphose/resolve/main/runtime.py
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" | |
| } | |
| 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 | |