""" MIDI Tokenizer using REMI (REvamped MIDI-derived) representation. State-of-the-art tokenization for symbolic music generation. Handles: Note On/Off, Velocity, Time Shift, Tempo, Time Signature. """ import json import logging from pathlib import Path from typing import Optional import numpy as np logger = logging.getLogger(__name__) # Special token IDs PAD_TOKEN = 0 BOS_TOKEN = 1 EOS_TOKEN = 2 SEP_TOKEN = 3 # Event type offsets (after special tokens) SPECIAL_OFFSET = 4 # REMI vocabulary layout: # [PAD, BOS, EOS, SEP, NoteOn_0..127, NoteOff_0..127, Velocity_0..31, # TimeShift_0..99, Tempo_0..59, Position_0..31, Bar] NOTE_ON_OFFSET = SPECIAL_OFFSET NOTE_ON_COUNT = 128 NOTE_OFF_OFFSET = NOTE_ON_OFFSET + NOTE_ON_COUNT NOTE_OFF_COUNT = 128 VELOCITY_OFFSET = NOTE_OFF_OFFSET + NOTE_OFF_COUNT VELOCITY_COUNT = 32 # Quantized to 32 bins TIMESHIFT_OFFSET = VELOCITY_OFFSET + VELOCITY_COUNT TIMESHIFT_COUNT = 100 # 10ms to 1000ms in 10ms steps TEMPO_OFFSET = TIMESHIFT_OFFSET + TIMESHIFT_COUNT TEMPO_COUNT = 60 # 40-200 BPM quantized POSITION_OFFSET = TEMPO_OFFSET + TEMPO_COUNT POSITION_COUNT = 32 # 32 positions per bar (supports up to 32nd notes) BAR_OFFSET = POSITION_OFFSET + POSITION_COUNT BAR_COUNT = 1 VOCAB_SIZE = BAR_OFFSET + BAR_COUNT class MusicTokenizer: """Efficient REMI tokenizer for MIDI to token conversion.""" def __init__(self): self.vocab_size = VOCAB_SIZE self.pad_id = PAD_TOKEN self.bos_id = BOS_TOKEN self.eos_id = EOS_TOKEN def note_on_token(self, pitch: int) -> int: return NOTE_ON_OFFSET + max(0, min(127, pitch)) def note_off_token(self, pitch: int) -> int: return NOTE_OFF_OFFSET + max(0, min(127, pitch)) def velocity_token(self, velocity: int) -> int: # Quantize 0-127 to 0-31 bins return VELOCITY_OFFSET + min(31, velocity // 4) def timeshift_token(self, ms: float) -> int: # Quantize to 10ms steps, capped at 1000ms idx = max(0, min(99, int(ms / 10))) return TIMESHIFT_OFFSET + idx def tempo_token(self, bpm: float) -> int: # Map BPM range 40-200 to 0-59 idx = max(0, min(59, int((bpm - 40) / (160 / 59)))) return TEMPO_OFFSET + idx def position_token(self, pos: int) -> int: return POSITION_OFFSET + max(0, min(31, pos)) def bar_token(self) -> int: return BAR_OFFSET def decode_token(self, token_id: int) -> dict: """Decode a token ID back to its event type and value.""" if token_id == PAD_TOKEN: return {"type": "PAD", "value": 0} if token_id == BOS_TOKEN: return {"type": "BOS", "value": 0} if token_id == EOS_TOKEN: return {"type": "EOS", "value": 0} if token_id == SEP_TOKEN: return {"type": "SEP", "value": 0} if NOTE_ON_OFFSET <= token_id < NOTE_OFF_OFFSET: return {"type": "NoteOn", "value": token_id - NOTE_ON_OFFSET} if NOTE_OFF_OFFSET <= token_id < VELOCITY_OFFSET: return {"type": "NoteOff", "value": token_id - NOTE_OFF_OFFSET} if VELOCITY_OFFSET <= token_id < TIMESHIFT_OFFSET: return {"type": "Velocity", "value": (token_id - VELOCITY_OFFSET) * 4} if TIMESHIFT_OFFSET <= token_id < TEMPO_OFFSET: return {"type": "TimeShift", "value": (token_id - TIMESHIFT_OFFSET) * 10} if TEMPO_OFFSET <= token_id < POSITION_OFFSET: return {"type": "Tempo", "value": 40 + (token_id - TEMPO_OFFSET) * (160 / 59)} if POSITION_OFFSET <= token_id < BAR_OFFSET: return {"type": "Position", "value": token_id - POSITION_OFFSET} if token_id == BAR_OFFSET: return {"type": "Bar", "value": 0} return {"type": "Unknown", "value": token_id} def midi_to_tokens(self, midi_obj, max_len: Optional[int] = None) -> list[int]: """ Convert a pretty_midi.PrettyMIDI object to REMI token sequence. Uses note-level events sorted by onset time. """ tokens = [self.bos_id] # Collect all notes across instruments all_notes = [] for inst in midi_obj.instruments: if inst.is_drum: continue for note in inst.notes: all_notes.append(note) if not all_notes: tokens.append(self.eos_id) return tokens # Sort by start time, then by pitch all_notes.sort(key=lambda n: (n.start, n.pitch)) # Get tempo changes tempos = midi_obj.get_tempo_changes() current_tempo = 120.0 if len(tempos[1]) > 0: current_tempo = tempos[1][0] tokens.append(self.tempo_token(current_tempo)) # Compute bar duration bar_duration = 60.0 / current_tempo * 4 # Assume 4/4 current_bar = 0 tokens.append(self.bar_token()) prev_time = 0.0 for note in all_notes: # Bar tracking note_bar = int(note.start / bar_duration) while current_bar < note_bar: current_bar += 1 tokens.append(self.bar_token()) # Time shift from previous event dt = note.start - prev_time if dt > 0: # Break into chunks of max 1000ms while dt > 1.0: tokens.append(self.timeshift_token(1000)) dt -= 1.0 if dt > 0.005: # Ignore < 5ms tokens.append(self.timeshift_token(dt * 1000)) # Position within bar pos_in_bar = (note.start % bar_duration) / bar_duration pos_idx = int(pos_in_bar * 32) tokens.append(self.position_token(pos_idx)) # Velocity then NoteOn tokens.append(self.velocity_token(note.velocity)) tokens.append(self.note_on_token(note.pitch)) # Note duration as timeshift + NoteOff dur = note.end - note.start if dur > 0: while dur > 1.0: tokens.append(self.timeshift_token(1000)) dur -= 1.0 if dur > 0.005: tokens.append(self.timeshift_token(dur * 1000)) tokens.append(self.note_off_token(note.pitch)) prev_time = note.start if max_len and len(tokens) >= max_len - 1: break tokens.append(self.eos_id) if max_len: tokens = tokens[:max_len] return tokens def tokens_to_midi(self, tokens: list[int]): """Convert REMI tokens back to a PrettyMIDI object.""" import pretty_midi midi = pretty_midi.PrettyMIDI(initial_tempo=120.0) inst = pretty_midi.Instrument(program=0, name="Piano") current_time = 0.0 current_velocity = 80 active_notes = {} # pitch -> (start_time, velocity) for token_id in tokens: event = self.decode_token(token_id) etype = event["type"] val = event["value"] if etype in ("PAD", "BOS", "EOS", "SEP", "Bar", "Position"): continue elif etype == "Tempo": pass # Could adjust timing but simpler to ignore elif etype == "TimeShift": current_time += val / 1000.0 elif etype == "Velocity": current_velocity = max(1, min(127, val)) elif etype == "NoteOn": active_notes[val] = (current_time, current_velocity) elif etype == "NoteOff": if val in active_notes: start, vel = active_notes.pop(val) if current_time > start: note = pretty_midi.Note( velocity=vel, pitch=val, start=start, end=current_time, ) inst.notes.append(note) # Close any remaining active notes for pitch, (start, vel) in active_notes.items(): note = pretty_midi.Note( velocity=vel, pitch=pitch, start=start, end=current_time + 0.5 ) inst.notes.append(note) midi.instruments.append(inst) return midi def save(self, path: Path): data = {"vocab_size": self.vocab_size} path.parent.mkdir(parents=True, exist_ok=True) with open(path, "w") as f: json.dump(data, f) @classmethod def load(cls, path: Path) -> "MusicTokenizer": tok = cls() if path.exists(): with open(path) as f: data = json.load(f) tok.vocab_size = data.get("vocab_size", VOCAB_SIZE) return tok