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