Download engine.py from cpyang/motigen: direct link, hf CLI and curl.
- Browser
- Download file 8.74 kB
-
https://huggingface.co/spaces/cpyang/motigen/resolve/main/engine.py
- Command line
-
hf download hf://spaces/cpyang/motigen/engine.py
-
curl -L -o engine.py https://huggingface.co/spaces/cpyang/motigen/resolve/main/engine.py
8.74 kB
| """Cached patch/character decoding with the original trained motif attention bias.""" | |
| from dataclasses import dataclass | |
| import gzip | |
| import json | |
| import os | |
| import re | |
| from pathlib import Path | |
| import shutil | |
| import sys | |
| import tempfile | |
| import time | |
| import numpy as np | |
| import torch | |
| from safetensors.torch import load_file | |
| from samplings import top_k_sampling, top_p_sampling, temperature_sampling | |
| ROOT = Path(__file__).resolve().parent | |
| sys.path.insert(0, str(ROOT / "vendor")) | |
| from notagen_core import NotaGenLMHeadModel, build_notagen_configs, safe_normalize_probs | |
| PATCH_SIZE = 16 | |
| MAX_CONTEXT = 1024 | |
| def realization_parts(line): | |
| """Match finetune/utils.py Patchilizer.patchilize_metadata/split_bars. | |
| Realization tags occupy their own patches; notes use training's bar splits. | |
| In particular, the final non-voice segment is joined to the preceding bar. | |
| """ | |
| tag = re.match(r"^(%motif:abc(?::[a-z_]+)?: )", line) | |
| if not tag: | |
| return [line] | |
| prefix = tag.group(1) | |
| delimiters = ("|:", "::", ":|", "[|", "||", "|]", "|") | |
| pieces = [part for part in re.split("(" + "|".join(map(re.escape, delimiters)) + ")", line[len(prefix):]) if part] | |
| if len(pieces) <= 1: | |
| return [prefix, *pieces] | |
| start = 0 if pieces[0] in delimiters else 1 | |
| bars = pieces[:start] + ["".join(pieces[i:i + 2]) for i in range(start, len(pieces), 2)] | |
| if len(bars) > 1 and "V" not in bars[-1]: | |
| bars[-2:] = [bars[-2] + bars[-1]] | |
| return [prefix, *bars] | |
| def encode_prompt(text): | |
| patches, flags = [[1] * 15 + [2]], [False] | |
| for line in text.splitlines(keepends=True): | |
| for part in realization_parts(line): | |
| ids = list(part.encode("ascii")) | |
| if len(ids) % PATCH_SIZE: | |
| ids.append(2) | |
| for start in range(0, len(ids), PATCH_SIZE): | |
| patch = ids[start:start + PATCH_SIZE] | |
| patches.append(patch + [0] * (PATCH_SIZE - len(patch))) | |
| flags.append(line.startswith("%motif:")) | |
| return patches, flags | |
| def decode_patch(ids): | |
| result = [] | |
| for token in ids: | |
| if token == 2: | |
| break | |
| if token >= 32 or token in (9, 10, 13): | |
| result.append(chr(token)) | |
| return "".join(result) | |
| def load_model(): | |
| manifest = json.loads((ROOT / "checkpoint.json").read_text()) | |
| weights = Path(os.environ.get("ACCOMPGEN_CHECKPOINT", ROOT / manifest.get("export_filename", "model.safetensors"))) | |
| if not weights.exists(): | |
| raise FileNotFoundError("No checkpoint found. Run prepare.py before launching the demo.") | |
| configs = build_notagen_configs(**manifest["architecture"]) | |
| if weights.name.endswith(".safetensors.gz"): | |
| with tempfile.TemporaryDirectory(prefix="motigen-model-") as temporary: | |
| unpacked = Path(temporary) / "model.safetensors" | |
| with gzip.open(weights, "rb") as source, unpacked.open("wb") as target: | |
| shutil.copyfileobj(source, target, length=8 * 1024 * 1024) | |
| state = load_file(str(unpacked)) | |
| elif weights.suffix == ".safetensors": | |
| state = load_file(str(weights)) | |
| else: | |
| state = torch.load(weights, map_location="cpu", weights_only=True, mmap=True)["model"] | |
| # Meta initialization avoids allocating a second 2 GB set of parameters. | |
| with torch.device("meta"): | |
| model = NotaGenLMHeadModel(encoder_config=configs[0], decoder_config=configs[1]) | |
| model.load_state_dict(state, strict=True, assign=True) | |
| # GPT2 causal masks are non-persistent buffers, absent from the state dict. | |
| for module in model.modules(): | |
| if hasattr(module, "bias") and isinstance(module.bias, torch.Tensor) and module.bias.is_meta: | |
| shape = module.bias.shape | |
| module.bias = torch.tril(torch.ones(shape[-2:], dtype=torch.bool)).view(shape) | |
| if hasattr(module, "masked_bias") and module.masked_bias.is_meta: | |
| module.masked_bias = torch.tensor(-1e4) | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = model.to(device=device, dtype=torch.float16 if device == "cuda" else torch.float32).eval() | |
| torch.set_num_threads(min(4, os.cpu_count() or 1)) | |
| return model | |
| class CachedDecoder: | |
| """Per-request KV state; never shared between visitors.""" | |
| def __init__(self, model): | |
| self.model = model | |
| self.past = None | |
| def encode(self, patches, motif_flags): | |
| m = self.model | |
| encoder = m.patch_level_decoder | |
| ids = torch.tensor([patches], dtype=torch.long, device=m.device) | |
| embeds = torch.nn.functional.one_hot(ids, num_classes=128).to(m.dtype) | |
| embeds = encoder.patch_embedding(embeds.reshape(1, -1, 16 * 128)) | |
| bias = torch.tensor(motif_flags, device=m.device, dtype=m.dtype) * m.motif_attention_bias | |
| encoder._motif_bias = bias[None, None, None, :] | |
| try: | |
| output = encoder.base(inputs_embeds=embeds, attention_mask=torch.ones(1, len(motif_flags), device=m.device), | |
| past_key_values=self.past, use_cache=True) | |
| finally: | |
| encoder._motif_bias = None | |
| self.past = output.past_key_values | |
| return m.patch_proj(output.last_hidden_state[0, -1]) | |
| def patch(self, encoded, rng, temperature=1.2, top_k=9, top_p=0.9, prefix=()): | |
| base = self.model.char_level_decoder.base | |
| embeds = encoded.reshape(1, 1, -1) | |
| if prefix: | |
| tokens = torch.tensor([prefix], device=self.model.device) | |
| embeds = torch.cat((embeds, base.get_input_embeddings()(tokens)), dim=1) | |
| past, result = None, list(prefix) | |
| while len(result) < PATCH_SIZE: | |
| output = base(inputs_embeds=embeds, past_key_values=past, use_cache=True) | |
| past = output.past_key_values | |
| probs = safe_normalize_probs(torch.softmax(output.logits[0, -1].float(), dim=-1).cpu().numpy()) | |
| probs = safe_normalize_probs(top_k_sampling(probs, top_k=top_k, return_probs=True)) | |
| probs = safe_normalize_probs(top_p_sampling(probs, top_p=top_p, return_probs=True)) | |
| # samplings' temperature transform, sampled with a request-local RNG. | |
| probs = safe_normalize_probs(temperature_sampling(probs, temperature=temperature, return_probs=True)) | |
| token = int(rng.choice(len(probs), p=probs)) | |
| result.append(token) | |
| if token == 2: | |
| result.extend([0] * (PATCH_SIZE - len(result))) | |
| break | |
| embeds = base.get_input_embeddings()(torch.tensor([[token]], device=self.model.device)) | |
| return result | |
| class GenerationUpdate: | |
| text: str | |
| patches: int | |
| elapsed: float | |
| reason: str = "" | |
| def generate(model, prompt, *, temperature=1.2, seed=0, max_bars=16, max_seconds=100, max_patches=800): | |
| patches, flags = encode_prompt(prompt) | |
| rng = np.random.default_rng(seed) | |
| decoder = CachedDecoder(model) | |
| text, line, body = prompt, "", False | |
| started, last_yield = time.monotonic(), 0.0 | |
| reason = "Patch limit reached; showing the completed measures." | |
| with torch.inference_mode(): | |
| encoded = decoder.encode(patches, flags) | |
| for index in range(max_patches): | |
| ids = decoder.patch(encoded, rng, temperature) | |
| fragment = decode_patch(ids) | |
| if not body and fragment.startswith("[r:"): | |
| ids = decoder.patch(encoded, rng, temperature, prefix=tuple(b"[r:0/")) | |
| fragment = decode_patch(ids) | |
| body = True | |
| if ids[:2] == [1, 2]: | |
| reason = "Complete." | |
| break | |
| text += fragment | |
| flags.append((line + fragment).lstrip().startswith("%motif:")) | |
| line = (line + fragment).rsplit("\n", 1)[-1] | |
| patches.append(ids) | |
| elapsed = time.monotonic() - started | |
| if elapsed - last_yield >= 0.4: | |
| yield GenerationUpdate(text, index + 1, elapsed) | |
| last_yield = elapsed | |
| # Stream-format body lines correspond to complete score measures. | |
| body_lines = [s for s in text.splitlines(keepends=True) if s.startswith("[r:") and s.endswith("\n")] | |
| if len(body_lines) >= max_bars: | |
| reason = f"Finished {max_bars} measures." | |
| break | |
| if elapsed >= max_seconds: | |
| reason = "Time limit reached; showing the completed measures." | |
| break | |
| if len(patches) >= MAX_CONTEXT: | |
| reason = "Context limit reached; showing the completed measures." | |
| break | |
| encoded = decoder.encode([ids], flags) | |
| yield GenerationUpdate(text, index + 1, time.monotonic() - started, reason) | |