PMI_Calculator_GPU / pmi_core.py
luna0805's picture
Upload 6 files
abde881 verified
Raw
History Blame Contribute Delete
23.7 kB
"""
pmi_core.py — the PMI main line for HUMAN-REVISED MIDI.
revised MIDI -> notes (s,e,p) -> K-S key estimate (round-1 LOCKED) + snapping
-> transpose to C -> symbol string
-> Needleman-Wunsch (affine gap) alignment -> PMI %
Key estimation uses estimate_key_locked(): the FIRST K-S estimate (made on the
original, information-complete notes) is locked and used to drive conservative
snapping of out-of-key slips. Later rounds clean the absolute pitch but do NOT
re-estimate the key — except in the rare 'wildly-off' case where a re-estimate is wildly
different (tonic far away AND the locked key's correlation has collapsed), when the
new key is adopted once. This matches the observation that round-1 is the most
trustworthy estimate and re-estimating after snapping tends to drift.
Ported (same constants/profiles/scoring) from the batch pipeline so numbers match.
Self-contained: depends only on numpy + pretty_midi.
"""
import numpy as np
import pretty_midi
# ---- Krumhansl-Kessler profiles + alphabet (verbatim from batch pipeline) ----
_MAJOR = np.array([6.35,2.23,3.48,2.33,4.38,4.09,2.52,5.19,2.39,3.66,2.29,2.88])
_MINOR = np.array([6.33,2.68,3.52,5.38,2.60,3.53,2.54,4.75,3.98,2.69,3.34,3.17])
NOTE_NAMES = ['C','C#','D','D#','E','F','F#','G','G#','A','A#','B']
PC_TO_LETTER = ['C','d','D','e','E','F','g','G','a','A','b','B']
_MAJOR_INTERVALS = (0,2,4,5,7,9,11)
_MINOR_INTERVALS = (0,2,3,5,7,8,10)
MIN_DUR = 0.08
MARGIN_THRESH = 0.05
# scale-degree symbols (used only if encoding="scale_degree")
_DEG_MAJOR = {0:'1', 2:'2', 4:'3', 5:'4', 7:'5', 9:'6', 11:'7'}
_DEG_MINOR = {0:'1', 2:'2', 3:'3', 5:'4', 7:'5', 8:'6', 10:'7'}
_CHROMA_MAJOR = {1:'a', 3:'b', 6:'c', 8:'d', 10:'e'}
_CHROMA_MINOR = {1:'a', 4:'f', 6:'c', 9:'h', 11:'g'}
# ===========================================================================
# key estimation (plain K-S, NO snapping)
# ===========================================================================
def _pc(note): return int(note[2]) % 12
def scale_for(tonic, mode):
iv = _MAJOR_INTERVALS if mode == 'major' else _MINOR_INTERVALS
return {(tonic + i) % 12 for i in iv}
def note_weight(s, e, conf=None, min_dur=MIN_DUR):
dur = max(e - s, 1e-3)
w_dur = min(dur / min_dur, 1.0)
w_conf = float(np.clip(conf, 0.0, 1.0)) if conf is not None else 1.0
return dur * w_dur * w_conf
def pitch_class_histogram(notes, weighted=True, confs=None):
hist = np.zeros(12)
for i, note in enumerate(notes):
conf = (float(confs[i]) if confs is not None
else (float(note[3]) if len(note) > 3 else None))
w = note_weight(note[0], note[1], conf) if weighted else max(note[1]-note[0], 1e-3)
hist[_pc(note)] += w
return hist
def _correlate_all_keys(hist):
corrs = []
for mode, prof in (('major', _MAJOR), ('minor', _MINOR)):
for tonic in range(12):
corrs.append((np.corrcoef(hist, np.roll(prof, tonic))[0, 1], tonic, mode))
corrs.sort(reverse=True)
return corrs
def estimate_key(notes, margin_thresh=MARGIN_THRESH, use_out_of_key_tiebreak=True,
weighted=True, confs=None):
"""Plain K-S key estimate (no snapping). Accepts (s,e,p[,conf]) notes.
Returns dict(tonic, mode, scale, margin, confidence, out_of_key, hist, corrs) or None."""
hist = pitch_class_histogram(notes, weighted=weighted, confs=confs)
if hist.sum() == 0:
return None
corrs = _correlate_all_keys(hist)
top1, top2 = corrs[0], corrs[1]
margin = top1[0] - top2[0]
def ook(tonic, mode):
sc = scale_for(tonic, mode)
return sum(hist[i] for i in range(12) if i not in sc)
if use_out_of_key_tiebreak and margin < margin_thresh:
_, c, tonic, mode = sorted((ook(t, m), c, t, m) for c, t, m in (top1, top2))[0]
else:
c, tonic, mode = top1
return dict(tonic=tonic, mode=mode, scale=scale_for(tonic, mode), margin=margin,
confidence=c, out_of_key=ook(tonic, mode), hist=hist, corrs=corrs)
# ===========================================================================
# snapping (verbatim from batch pipeline) + iterative estimator
# ===========================================================================
SNAP_THRESH = 0.60
KEEP_THRESH = 0.30
MAX_ROUNDS = 4
def _conf(note, confs, idx):
if confs is not None:
return float(confs[idx])
return float(note[3]) if len(note) > 3 else None
def slip_likelihood(idx, notes, key, confs=None, min_dur=MIN_DUR):
if _pc(notes[idx]) in key['scale']:
return 0.0
s, e, p = notes[idx][0], notes[idx][1], _pc(notes[idx])
dur = max(e - s, 1e-3)
conf = _conf(notes[idx], confs, idx)
e_conf = (1.0 - float(np.clip(conf, 0.0, 1.0))) if conf is not None else 0.0
e_short = 1.0 - min(dur / min_dur, 1.0)
pcs = [_pc(n) for n in notes]
recur = pcs.count(p)
left_in = (idx == 0) or (pcs[idx-1] in key['scale'])
right_in = (idx == len(notes)-1) or (pcs[idx+1] in key['scale'])
e_iso = 1.0 if (recur == 1 and left_in and right_in) else 0.0
e_recur = 1.0 - 1.0 / recur
score = (0.5*e_conf + 0.25*e_short + 0.25*e_iso) * (1.0 - 0.5*e_recur)
return float(np.clip(score, 0.0, 1.0))
def is_minor_mode_member(p, key, notes):
if key['mode'] != 'minor':
return False
raised = {(key['tonic']+11) % 12, (key['tonic']+9) % 12}
if p % 12 not in raised:
return False
return [_pc(n) for n in notes].count(p % 12) >= 2
def _snap_target(p, scale):
for cand in ((p-1) % 12, (p+1) % 12):
if cand in scale:
return cand
return p
def classify_notes(notes, key, confs=None, snap_thresh=SNAP_THRESH, keep_thresh=KEEP_THRESH):
scale = key['scale']
out = []
for i, note in enumerate(notes):
p = _pc(note)
if p in scale:
out.append(dict(idx=i, label='diatonic', out_pc=p, confidence=1.0)); continue
if is_minor_mode_member(p, key, notes):
out.append(dict(idx=i, label='intentional_chromatic', out_pc=p, confidence=1.0)); continue
sc = slip_likelihood(i, notes, key, confs)
if sc >= snap_thresh:
out.append(dict(idx=i, label='corrected', out_pc=_snap_target(p, scale), confidence=sc))
elif sc <= keep_thresh:
out.append(dict(idx=i, label='intentional_chromatic', out_pc=p, confidence=1.0-sc))
else:
out.append(dict(idx=i, label='uncertain', out_pc=p, confidence=sc))
return out
def corrected_pitches(notes, classification):
out = [_pc(n) for n in notes]
for rec in classification:
out[rec['idx']] = rec['out_pc']
return out
# tonic distance on the circle of pitch classes (0..6 semitones)
def _tonic_dist(a, b):
d = abs(int(a) - int(b)) % 12
return min(d, 12 - d)
def estimate_key_locked(notes, max_rounds=MAX_ROUNDS, snap_thresh=SNAP_THRESH,
keep_thresh=KEEP_THRESH, confs=None,
outlier_tonic_dist=3, outlier_corr_drop=0.15, **est_kw):
"""Iterative key + snapping, but the ROUND-1 key is LOCKED and used to drive
all snapping. Subsequent rounds snap notes (cleaning absolute pitch) but do
NOT overwrite the key — UNLESS a re-estimate is wildly different from round 1
(the 'wildly-off' guard): tonic more than `outlier_tonic_dist` semitones away AND the
round-1 key's correlation has dropped by more than `outlier_corr_drop`. In that
rare case the new key is adopted once and re-locked.
Rationale: round-1 is estimated on the original, information-complete notes and
is the most trustworthy; re-estimating after snapping tends to drift. Returns
(key, classification). Caller's notes are never mutated.
"""
work = [list(n) for n in notes]
orig_pc = [_pc(n) for n in notes]
corrected = {}; frozen = set()
key = estimate_key(work, confs=confs, **est_kw) # ROUND-1 key
if key is None:
return None, []
locked = key # <-- lock it
cls = classify_notes(work, locked, confs, snap_thresh, keep_thresh)
for _ in range(max_rounds):
changed = False
for rec in cls:
i = rec['idx']
if rec['label'] == 'corrected' and i not in frozen and work[i][2] % 12 != rec['out_pc']:
work[i][2] = rec['out_pc']
corrected[i] = (rec['out_pc'], rec['confidence'])
frozen.add(i)
changed = True
# re-estimate ONLY to check for the 'wildly-off' (wildly-off) case; do not adopt by default
check = estimate_key(work, confs=confs, **est_kw)
if check is not None:
far_tonic = _tonic_dist(check['tonic'], locked['tonic']) > outlier_tonic_dist
# correlation the locked key now gets on the (snapped) histogram
locked_corr = next((c for c, t, m in check['corrs']
if t == locked['tonic'] and m == locked['mode']), None)
corr_drop = (locked_corr is not None) and (check['confidence'] - locked_corr > outlier_corr_drop)
if far_tonic and corr_drop:
locked = check # adopt + re-lock (rare)
# keep snapping against the LOCKED key
cls = classify_notes(work, locked, confs, snap_thresh, keep_thresh)
if not changed:
break
for rec in cls:
i = rec['idx']
rec['orig_pc'] = orig_pc[i]
if i in corrected:
rec['label'] = 'corrected'
rec['out_pc'], rec['confidence'] = corrected[i]
return locked, cls
# ===========================================================================
# transpose to C + encode
# ===========================================================================
def transpose_to_C(notes, tonic):
"""Every note's pitch class shifted so the estimated tonic maps to C (pc 0).
Input notes are (s,e,p[,...]); output is a time-ordered pitch-class list."""
return [(int(p) - int(tonic)) % 12 for _, _, p, *_ in sorted(notes)]
def encode(pcs):
return "".join(PC_TO_LETTER[p % 12] for p in pcs)
def encode_scale_degree(pcs, mode):
diatonic, chroma = (_DEG_MAJOR, _CHROMA_MAJOR) if mode == 'major' else (_DEG_MINOR, _CHROMA_MINOR)
return ''.join(diatonic.get(pc % 12) or chroma.get(pc % 12, 'x') for pc in pcs)
# ===========================================================================
# alignment + PMI (Needleman-Wunsch affine gap; Savage GOP=12, GEP=6)
# ===========================================================================
def nw_align(a, b, match=2, mismatch=-2, gop=12, gep=6):
n, m = len(a), len(b); NEG = -1e9
M = np.full((n+1, m+1), NEG); Ix = np.full((n+1, m+1), NEG); Iy = np.full((n+1, m+1), NEG); M[0,0] = 0
for i in range(1, n+1): Ix[i,0] = -gop - (i-1)*gep
for j in range(1, m+1): Iy[0,j] = -gop - (j-1)*gep
for i in range(1, n+1):
ai = a[i-1]
for j in range(1, m+1):
sc = match if ai == b[j-1] else mismatch
M[i,j] = max(M[i-1,j-1], Ix[i-1,j-1], Iy[i-1,j-1]) + sc
Ix[i,j] = max(M[i-1,j] - gop, Ix[i-1,j] - gep)
Iy[i,j] = max(M[i,j-1] - gop, Iy[i,j-1] - gep)
i, j = n, m; al = []; bl = []
st = max((M[n,m], 'M'), (Ix[n,m], 'X'), (Iy[n,m], 'Y'))[1]
while i > 0 or j > 0:
if i == 0: st = 'Y'
elif j == 0: st = 'X'
if st == 'M':
al.append(a[i-1]); bl.append(b[j-1])
st = max((M[i-1,j-1], 'M'), (Ix[i-1,j-1], 'X'), (Iy[i-1,j-1], 'Y'))[1]; i -= 1; j -= 1
elif st == 'X':
al.append(a[i-1]); bl.append('-')
st = 'X' if (i > 1 and abs(Ix[i,j] - (Ix[i-1,j] - gep)) < 1e-6) else 'M'; i -= 1
else:
al.append('-'); bl.append(b[j-1])
st = 'Y' if (j > 1 and abs(Iy[i,j] - (Iy[i,j-1] - gep)) < 1e-6) else 'M'; j -= 1
return "".join(reversed(al)), "".join(reversed(bl))
def pmi(seqA, seqB, **kw):
if not seqA or not seqB:
return dict(PMI=float('nan'), ID=0, alnA='', alnB='')
A, B = nw_align(seqA, seqB, **kw)
ID = sum(1 for x, y in zip(A, B) if x == y and x != '-')
return dict(PMI=100.0 * ID / ((len(seqA) + len(seqB)) / 2.0), ID=ID, alnA=A, alnB=B)
# ===========================================================================
# MIDI -> notes, and the per-song + pairwise drivers
# ===========================================================================
def notes_from_midi(midi_path, instrument_index=None):
"""Revised MIDI -> time-sorted [(start, end, pitch, conf), ...].
The editor encodes confidence as velocity (conf = velocity/127); hand-edited /
verified notes are written at velocity 127 (conf=1). Reading it back lets the
key estimator treat verified notes as high-confidence (never snapped) and only
let low-confidence auto-transcribed leftovers be eligible for snapping.
"""
pm = pretty_midi.PrettyMIDI(midi_path)
insts = [i for i in pm.instruments if not i.is_drum]
if not insts:
raise ValueError("No non-drum instrument tracks in MIDI.")
if instrument_index is not None:
insts = [pm.instruments[instrument_index]]
notes = [(float(n.start), float(n.end), int(n.pitch), float(n.velocity) / 127.0)
for inst in insts for n in inst.notes]
if not notes:
raise ValueError("MIDI contains no notes.")
notes.sort(key=lambda t: (t[0], t[1], t[2]))
return notes
def midi_to_symbols(midi_path, encoding="pitch_class", weighted=True,
instrument_index=None, snap=True):
"""revised MIDI -> dict(notes, key, transposed, symbols).
snap=True (default): use estimate_key_locked — the round-1 K-S key is locked
and drives conservative snapping of out-of-key slips; the cleaned (snapped)
absolute pitches are what get transposed to C. snap=False: plain single K-S.
"""
notes = notes_from_midi(midi_path, instrument_index=instrument_index)
if snap:
key, classification = estimate_key_locked(notes, confs=None, weighted=weighted)
if key is None:
raise ValueError("Key estimation returned None (empty pitch-class histogram).")
# transpose the FINAL (snapped) pitches, in note order
final_pcs = corrected_pitches([list(n) for n in notes], classification)
transposed = [(pc - key["tonic"]) % 12 for pc in final_pcs]
else:
key = estimate_key(notes, weighted=weighted)
if key is None:
raise ValueError("Key estimation returned None (empty pitch-class histogram).")
classification = []
transposed = transpose_to_C(notes, key["tonic"])
if encoding == "pitch_class":
symbols = encode(transposed)
elif encoding == "scale_degree":
symbols = encode_scale_degree(transposed, key["mode"])
else:
raise ValueError(f"Unknown encoding: {encoding}")
n_corr = sum(1 for r in classification if r.get('label') == 'corrected')
return dict(notes=notes, key=key, transposed=transposed, symbols=symbols,
n_corrected=n_corr)
def key_name(key):
return f"{NOTE_NAMES[key['tonic']]} {key['mode']}"
def key_name_tm(tonic, mode):
return f"{NOTE_NAMES[int(tonic) % 12]} {mode}"
# ---------------------------------------------------------------------------
# transpose a song to C using a GIVEN key (tonic+mode), with snapping vs that key.
# Used when the user picks a candidate / overrides the key.
# ---------------------------------------------------------------------------
def symbols_for_given_key(notes, tonic, mode, encoding="pitch_class", weighted=True):
"""Snap notes against the given (tonic,mode), transpose to C, encode.
Verified (high-confidence) notes stay protected; only low-confidence
out-of-key slips are snapped."""
key = dict(tonic=int(tonic) % 12, mode=mode, scale=scale_for(int(tonic) % 12, mode))
cls = classify_notes([list(n) for n in notes], key)
final_pcs = corrected_pitches([list(n) for n in notes], cls)
transposed = [(pc - key["tonic"]) % 12 for pc in final_pcs]
sym = encode(transposed) if encoding == "pitch_class" else encode_scale_degree(transposed, mode)
n_oot = sum(1 for pc in transposed if pc not in _MAJOR_INTERVALS) if mode == 'major' \
else sum(1 for pc in transposed if pc not in _MINOR_INTERVALS)
return dict(tonic=key["tonic"], mode=mode, transposed=transposed, symbols=sym,
n_out_of_key=n_oot)
def top_candidate_keys(notes, k=2, weighted=True):
"""Round-1 K-S top-k (tonic,mode) candidates, most likely first.
The first is the locked estimate; the second is the runner-up used for the
'listen and pick' A/B comparison."""
hist = pitch_class_histogram(notes, weighted=weighted)
if hist.sum() == 0:
return []
corrs = _correlate_all_keys(hist) # already sorted, best first
return [(t, m, float(c)) for c, t, m in corrs[:k]]
# ---------------------------------------------------------------------------
# absolute-pitch similarity (NO transposition) — used only to decide whether to
# OFFER soft key-unification. Never part of the main PMI.
# ---------------------------------------------------------------------------
def absolute_symbols(notes):
"""Encode notes by raw pitch class (no transposition)."""
return encode([_pc(n) for n in sorted(notes, key=lambda t: (t[0], t[1], t[2]))])
def absolute_similarity(notesA, notesB):
return pmi(absolute_symbols(notesA), absolute_symbols(notesB))["PMI"]
# ---------------------------------------------------------------------------
# per-song analysis for the /analyze endpoint
# ---------------------------------------------------------------------------
def analyze_one_midi(midi_path, encoding="pitch_class"):
"""One revised MIDI -> locked key estimate + top-2 candidates, each candidate
transposed-to-C so the front-end can play them back for the listen-and-pick.
Also returns the absolute (untransposed) pitch-class sequence for soft-unify."""
notes = notes_from_midi(midi_path)
base = midi_to_symbols(midi_path, encoding=encoding) # locked estimate + snap
cands = top_candidate_keys(notes, k=2)
cand_out = []
for t, m, corr in cands:
s = symbols_for_given_key(notes, t, m, encoding=encoding)
cand_out.append(dict(tonic=t, mode=m, corr=corr, name=key_name_tm(t, m),
symbols=s["symbols"], transposed=s["transposed"],
n_out_of_key=s["n_out_of_key"]))
return dict(
n_notes=len(notes),
key=dict(tonic=base["key"]["tonic"], mode=base["key"]["mode"],
name=key_name(base["key"]), margin=float(base["key"]["margin"]),
confidence=float(base["key"]["confidence"])),
symbols=base["symbols"],
transposed=base["transposed"],
abs_pcs=[_pc(n) for n in notes], # for soft-unify similarity
candidates=cand_out,
encoding=encoding,
)
# ---------------------------------------------------------------------------
# pairwise PMI from two analyses (+ optional user-forced keys + soft-unify)
# ---------------------------------------------------------------------------
def pmi_pair(notesA, notesB, encoding="pitch_class",
forced_A=None, forced_B=None,
soft_unify_threshold=70.0,
labelA="Song A", labelB="Song B"):
"""Compute the main PMI (each song its own key, locked estimate or user-forced),
plus a SEPARATE soft-unified PMI when the two melodies are already similar in
ABSOLUTE pitch (>= threshold). forced_A/forced_B are optional (tonic, mode)."""
# --- per-song key + symbols (main line) ---
if forced_A is not None:
A = symbols_for_given_key(notesA, forced_A[0], forced_A[1], encoding=encoding)
keyA = dict(tonic=A["tonic"], mode=A["mode"], margin=None,
confidence=None, scale=scale_for(A["tonic"], A["mode"]))
else:
kA, clsA = estimate_key_locked(notesA)
keyA = kA
pcsA = corrected_pitches([list(n) for n in notesA], clsA)
A = dict(tonic=kA["tonic"], mode=kA["mode"],
transposed=[(pc - kA["tonic"]) % 12 for pc in pcsA])
A["symbols"] = encode(A["transposed"]) if encoding == "pitch_class" \
else encode_scale_degree(A["transposed"], kA["mode"])
if forced_B is not None:
B = symbols_for_given_key(notesB, forced_B[0], forced_B[1], encoding=encoding)
keyB = dict(tonic=B["tonic"], mode=B["mode"], margin=None,
confidence=None, scale=scale_for(B["tonic"], B["mode"]))
else:
kB, clsB = estimate_key_locked(notesB)
keyB = kB
pcsB = corrected_pitches([list(n) for n in notesB], clsB)
B = dict(tonic=kB["tonic"], mode=kB["mode"],
transposed=[(pc - kB["tonic"]) % 12 for pc in pcsB])
B["symbols"] = encode(B["transposed"]) if encoding == "pitch_class" \
else encode_scale_degree(B["transposed"], kB["mode"])
main = pmi(A["symbols"], B["symbols"])
result = {
"PMI": main["PMI"], "ID": main["ID"], "alnA": main["alnA"], "alnB": main["alnB"],
"symbolsA": A["symbols"], "symbolsB": B["symbols"],
"keyA": key_name_tm(A["tonic"], A["mode"]),
"keyB": key_name_tm(B["tonic"], B["mode"]),
"n_notes_A": len(notesA), "n_notes_B": len(notesB),
"lenA": len(A["symbols"]), "lenB": len(B["symbols"]),
"encoding": encoding, "labelA": labelA, "labelB": labelB,
"forcedA": forced_A is not None, "forcedB": forced_B is not None,
}
# --- soft key-unification (SEPARATE, only OFFERED when abs-similarity high) ---
abs_sim = absolute_similarity(notesA, notesB)
result["abs_similarity"] = abs_sim
result["soft_unify_threshold"] = soft_unify_threshold
result["soft_unify_applicable"] = abs_sim >= soft_unify_threshold
if result["soft_unify_applicable"]:
# weight by margin (how decisively each song's own key beat the runner-up);
# the more decisive key wins as the shared key. forced keys count as fully decisive.
mA = 1.0 if forced_A is not None else float(keyA.get("margin") or 0.0)
mB = 1.0 if forced_B is not None else float(keyB.get("margin") or 0.0)
if mA >= mB:
uni_t, uni_m, src = A["tonic"], A["mode"], labelA
else:
uni_t, uni_m, src = B["tonic"], B["mode"], labelB
UA = symbols_for_given_key(notesA, uni_t, uni_m, encoding=encoding)
UB = symbols_for_given_key(notesB, uni_t, uni_m, encoding=encoding)
uni = pmi(UA["symbols"], UB["symbols"])
result["soft_unify"] = {
"shared_key": key_name_tm(uni_t, uni_m),
"shared_key_source": src,
"PMI": uni["PMI"], "ID": uni["ID"],
"alnA": uni["alnA"], "alnB": uni["alnB"],
"symbolsA": UA["symbols"], "symbolsB": UB["symbols"],
"weight_margin_A": mA, "weight_margin_B": mB,
}
return result
def pmi_from_two_midis(midi_path_A, midi_path_B, encoding="pitch_class",
forced_A=None, forced_B=None, soft_unify_threshold=70.0,
labelA="Song A", labelB="Song B"):
"""Convenience: read two MIDIs and run pmi_pair."""
notesA = notes_from_midi(midi_path_A)
notesB = notes_from_midi(midi_path_B)
return pmi_pair(notesA, notesB, encoding=encoding,
forced_A=forced_A, forced_B=forced_B,
soft_unify_threshold=soft_unify_threshold,
labelA=labelA, labelB=labelB)