File size: 12,165 Bytes
2a545bd
c05966b
 
 
 
 
 
 
 
 
935db35
c05966b
 
 
 
 
 
 
 
 
 
 
 
 
2a545bd
 
 
 
 
 
c05966b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a545bd
 
 
b8cd695
 
2a545bd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c05966b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1194e24
c05966b
 
 
 
 
 
 
 
 
1194e24
 
c05966b
 
 
 
 
 
 
 
 
 
 
1194e24
 
 
 
 
c05966b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a545bd
c05966b
 
 
2a545bd
935db35
c05966b
 
935db35
 
 
c05966b
 
 
 
 
 
 
 
 
 
935db35
c05966b
2a545bd
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
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
"""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