motigen / engine.py
cpyang's picture
Use VexFlow staff engraving and match training tokenization for motif realizations
983cc69 verified
Raw History Blame Contribute Delete
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
@dataclass
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)