File size: 6,362 Bytes
69b220d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 | #!/usr/bin/env python3
"""Maurice — cold-start a romantic-piano opening from silence, conditioned on composer + key.
python generate.py --composer liszt --key "Eb major" --out opening.mid
python generate.py --composer debussy --key "D major" --best-of 8 --out opening.mid
Requires: torch, safetensors, aria-utils (the AbsTokenizer), pretty_midi.
The model prefixes generation with [instrument, composer, key, <S>]; those conditioning
tokens are stripped before detokenising back to MIDI.
Best composers (by ear on this checkpoint): liszt, chopin, debussy, scriabin.
Use composer="unknown" for a generic-romantic voice. Every call is fresh (no seeding).
"""
import argparse, io, json, math, os, random
import torch
from modeling_maurice import Maurice
HERE = os.path.dirname(os.path.abspath(__file__))
COND = json.load(open(os.path.join(HERE, 'conditioning.json')))
INSTR, BOS, EOS = COND['instr_prefix_id'], COND['bos'], COND['eos']
STRIP = set(COND['composers'].values()) | set(COND['keys'].values()) | {INSTR, BOS, EOS, COND['pad_id']}
_KEYPC = {'C': 0, 'C#': 1, 'DB': 1, 'D': 2, 'D#': 3, 'EB': 3, 'E': 4, 'F': 5, 'F#': 6, 'GB': 6,
'G': 7, 'G#': 8, 'AB': 8, 'A': 9, 'A#': 10, 'BB': 10, 'B': 11}
def key_to_id(k):
"""'Db', 'C# minor', 'Eb major' -> conditioning key id. None -> a random key."""
if not k:
return COND['keys'][random.choice(list(COND['keys']))]
s = k.strip().replace('♯', '#').replace('♭', 'b')
tonic = s[0].upper() + ('#' if len(s) > 1 and s[1] == '#' else 'B' if len(s) > 1 and s[1] in 'bB' else '')
pc = _KEYPC.get(tonic)
mode = 'min' if ('min' in s.lower() or s.lower().rstrip().endswith('m')) else 'maj'
return COND['keys'].get(f'{pc}_{mode}') if pc is not None else None
def load():
model = Maurice.from_pretrained(HERE)
from ariautils.tokenizer import AbsTokenizer
return model, AbsTokenizer()
def _sample(logits, temp, topp):
p = torch.softmax(logits.float() / temp, dim=-1)
sp, si = torch.sort(p, descending=True)
sp = sp * ((torch.cumsum(sp, 0) - sp) < topp) # top-p (nucleus)
return si[torch.multinomial(sp / sp.sum(), 1)].item()
def generate_tokens(model, composer, key_id, max_tokens=1024, temp=0.98, topp=0.96, guidance=1.0):
"""Cold-start note tokens. guidance>1 with a SPECIFIC composer enables classifier-free
guidance (CFG): decode cond (composer) vs uncond ('unknown') in lockstep and steer
logits = uncond + guidance*(cond - uncond) to amplify composer/key identity."""
comp = (composer or 'unknown').lower()
comp_id = COND['composers'].get(comp, COND['composers']['unknown'])
seed = [INSTR, comp_id, key_id, BOS]
cfg = guidance > 1.0 and comp != 'unknown'
seed_u = [INSTR, COND['composers']['unknown'], key_id, BOS] if cfg else None
toks, notes = list(seed), []
lc, cc = model.infer(torch.tensor([seed]), cache=None, pos=0)
lu, cu = model.infer(torch.tensor([seed_u]), cache=None, pos=0) if cfg else (None, None)
for _ in range(max_tokens):
logits = (lu[0].float() + guidance * (lc[0].float() - lu[0].float())) if cfg else lc[0]
nxt = _sample(logits, temp, topp)
if nxt == EOS or len(toks) >= model.ctx - 1:
break
toks.append(nxt)
if nxt not in STRIP: # keep note tokens; drop conditioning prefixes
notes.append(nxt)
x = torch.tensor([[nxt]])
lc, cc = model.infer(x, cache=cc, pos=len(toks) - 1)
if cfg:
lu, cu = model.infer(x, cache=cu, pos=len(toks) - 1)
return notes
def to_midi(tok, note_ids):
md = tok.detokenize(tok.decode([INSTR, BOS] + note_ids))
return md.to_midi()
def rhythm_score(midi):
"""Advisory musicality proxy (rhythmic aliveness dominates). Higher = better; used for best-of-N."""
import pretty_midi
buf = io.BytesIO(); midi.save(file=buf); pm = pretty_midi.PrettyMIDI(io.BytesIO(buf.getvalue()))
_, tempi = pm.get_tempo_changes(); spb = 60.0 / (tempi[0] if len(tempi) else 120.0)
ns = sorted((n.start / spb, n.pitch, n.end / spb) for i in pm.instruments for n in i.notes)
if len(ns) < 8:
return -99.0
ms = [p for _, p, _ in ns]
onsets = []; cur = None
for b, p, e in ns:
if cur is None or b - cur[0] > 0.08: cur = [b, [p]]; onsets.append(cur)
else: cur[1].append(p)
durs = len(set(round((e - b) * 8) / 8 for b, _, e in ns))
poly = sum(len(o[1]) >= 2 for o in onsets) / len(onsets)
bal = min(sum(m < 48 for m in ms), sum(m >= 60 for m in ms)) / len(ms)
cl = lambda x, a, b: max(a, min(b, x))
return round(2.4 * cl(durs / 26, 0, 1.25) - 3.2 * cl((7 - durs) / 6, 0, 1)
+ 1.6 * math.exp(-((poly - 0.45) / 0.42) ** 2) + 1.4 * cl(bal / 0.15, 0, 1), 2)
def cold_start(model, tok, composer='unknown', key=None, best_of=1, guidance=1.0, **kw):
key_id = key_to_id(key)
pool = [to_midi(tok, generate_tokens(model, composer, key_id, guidance=guidance, **kw)) for _ in range(max(1, best_of))]
return max(pool, key=rhythm_score) if best_of > 1 else pool[0]
if __name__ == '__main__':
ap = argparse.ArgumentParser()
ap.add_argument('--composer', default='unknown')
ap.add_argument('--key', default=None, help='e.g. "Db", "C# minor", "Eb major"; omit for random')
ap.add_argument('--best-of', type=int, default=1, help='generate N, keep the most musical (rhythm-scored)')
ap.add_argument('--tokens', type=int, default=1024, help='generation length in tokens (~3 per note); default 1024 ≈ a short full piece')
ap.add_argument('--guidance', type=float, default=1.0, help='classifier-free guidance (1-4); 1 = off (best/most musical). >1 sharpens a specific composer but tends toward frantic/over-dense')
ap.add_argument('--temp', type=float, default=0.98)
ap.add_argument('--topp', type=float, default=0.96)
ap.add_argument('--out', default='maurice.mid')
a = ap.parse_args()
model, tok = load()
midi = cold_start(model, tok, a.composer, a.key, best_of=a.best_of, guidance=a.guidance,
max_tokens=a.tokens, temp=a.temp, topp=a.topp)
midi.save(a.out)
n_notes = sum(1 for tr in midi.tracks for m in tr if m.type == 'note_on' and m.velocity > 0)
print(f'wrote {a.out} ({n_notes} notes, {a.composer}/{a.key or "random key"})')
|