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