File size: 4,079 Bytes
f55264b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3e52e55
f55264b
 
 
 
 
e5cbf19
f55264b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3e52e55
 
 
 
 
 
 
f55264b
e5cbf19
f55264b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e5cbf19
f55264b
e5cbf19
f55264b
e5cbf19
f55264b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Local inference wrapper around the official DExter source and checkpoint."""

from __future__ import annotations

import os
import sys
from functools import lru_cache
from pathlib import Path
from tempfile import gettempdir

ROOT = Path(__file__).resolve().parent
VENDOR = ROOT / "vendor"
os.environ.setdefault("MPLCONFIGDIR", str(Path(gettempdir()) / "dexter_matplotlib"))
sys.path.insert(0, str(VENDOR))

import model as dexter_model
import numpy as np
import numpy.lib.recfunctions as rfn
import partitura as pt
import torch
from hydra import compose, initialize_config_dir
from omegaconf import OmegaConf
from pytorch_lightning.core.saving import _load_state
from renderer import Renderer

CHECKPOINT = ROOT / "weights" / "last-v1.ckpt"
EXAMPLE = ROOT / "examples" / "Schubert_D783_no15.musicxml"
MAX_NOTES = 400
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")


@lru_cache(maxsize=1)
def load_model():
    if not CHECKPOINT.is_file():
        raise FileNotFoundError(
            f"Missing {CHECKPOINT}. Download Google Drive file "
            "1Rv_ba1wWexaFJ-KFOW8pdnFtssuYc3sy from the official DExter demo."
        )
    with initialize_config_dir(config_dir=str(VENDOR / "config"), version_base=None):
        cfg = compose(config_name="inference")
    args = OmegaConf.to_container(cfg.model.model.args, resolve=True)
    task = OmegaConf.to_container(cfg.task, resolve=True)
    for key in ("dataset_means", "dataset_stds"):
        if task.get(key) is None:
            task.pop(key)
    checkpoint = torch.load(str(CHECKPOINT), map_location="cpu", weights_only=False)
    model = _load_state(
        getattr(dexter_model, cfg.model.model.name),
        checkpoint,
        strict=True,
        **args,
        **task,
    )
    model.to(DEVICE).eval()
    return model, int(cfg.seg_len), int(cfg.overlap)


def render(score_path: str | Path, output_path: str | Path, seed: int = 13) -> Path:
    score_path = Path(score_path)
    output_path = Path(output_path)
    if score_path.suffix.lower() not in {".xml", ".musicxml"}:
        raise ValueError("Upload an uncompressed MusicXML file (.xml or .musicxml).")
    if score_path.stat().st_size > 2_000_000:
        raise ValueError("MusicXML is limited to 2 MB.")

    score = pt.load_musicxml(str(score_path))
    notes = score.note_array()
    if not 1 <= len(notes) <= MAX_NOTES:
        raise ValueError(f"Score must contain 1 to {MAX_NOTES} notes; found {len(notes)}.")

    model, seg_len, overlap = load_model()
    torch.manual_seed(int(seed))
    codec = rfn.structured_to_unstructured(
        notes[["onset_div", "duration_div", "pitch", "voice"]]
    )
    stride = seg_len - overlap
    starts = list(range(0, max(1, len(codec) - seg_len + 1), stride))
    if starts[-1] + seg_len < len(codec):
        starts.append(len(codec) - seg_len)
    segments = []
    for start in starts:
        window = codec[start : start + seg_len]
        padded = np.pad(window, ((0, seg_len - len(window)), (0, 0)))
        segments.append(padded)
    score_batch = torch.from_numpy(np.stack(segments)).to(DEVICE)
    batch = {
        "p_codec": torch.zeros(len(segments), seg_len, 5, device=DEVICE),
        "s_codec": score_batch,
        "c_codec": torch.zeros(len(segments), seg_len, 7, device=DEVICE),
    }

    with torch.inference_mode():
        prediction = model.inference_one(batch)
    combined = np.zeros((len(codec), 5), dtype=np.float64)
    counts = np.zeros((len(codec), 1), dtype=np.float64)
    for start, segment in zip(starts, prediction):
        valid = min(seg_len, len(codec) - start)
        combined[start : start + valid] += np.asarray(segment[0, :valid])
        counts[start : start + valid] += 1
    combined /= counts

    output_path.parent.mkdir(parents=True, exist_ok=True)
    Renderer(str(output_path.parent), combined).render_inference_sample(
        score, output_path=str(output_path)
    )
    if not output_path.is_file() or output_path.stat().st_size == 0:
        raise RuntimeError("DExter did not create a MIDI file.")
    return output_path