DExter / runtime.py
LING
Load legacy DExter checkpoint explicitly
3e52e55
Raw History Blame Contribute Delete
4.08 kB
"""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