| """ |
| 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__) |
|
|
| |
| PAD_TOKEN = 0 |
| BOS_TOKEN = 1 |
| EOS_TOKEN = 2 |
| SEP_TOKEN = 3 |
|
|
| |
| SPECIAL_OFFSET = 4 |
|
|
| |
| |
| |
| 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 |
| TIMESHIFT_OFFSET = VELOCITY_OFFSET + VELOCITY_COUNT |
| TIMESHIFT_COUNT = 100 |
| TEMPO_OFFSET = TIMESHIFT_OFFSET + TIMESHIFT_COUNT |
| TEMPO_COUNT = 60 |
| POSITION_OFFSET = TEMPO_OFFSET + TEMPO_COUNT |
| POSITION_COUNT = 32 |
| 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: |
| |
| return VELOCITY_OFFSET + min(31, velocity // 4) |
|
|
| def timeshift_token(self, ms: float) -> int: |
| |
| idx = max(0, min(99, int(ms / 10))) |
| return TIMESHIFT_OFFSET + idx |
|
|
| def tempo_token(self, bpm: float) -> int: |
| |
| 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] |
|
|
| |
| 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 |
|
|
| |
| all_notes.sort(key=lambda n: (n.start, n.pitch)) |
|
|
| |
| 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)) |
|
|
| |
| bar_duration = 60.0 / current_tempo * 4 |
| current_bar = 0 |
| tokens.append(self.bar_token()) |
|
|
| prev_time = 0.0 |
| for note in all_notes: |
| |
| note_bar = int(note.start / bar_duration) |
| while current_bar < note_bar: |
| current_bar += 1 |
| tokens.append(self.bar_token()) |
|
|
| |
| dt = note.start - prev_time |
| if dt > 0: |
| |
| while dt > 1.0: |
| tokens.append(self.timeshift_token(1000)) |
| dt -= 1.0 |
| if dt > 0.005: |
| tokens.append(self.timeshift_token(dt * 1000)) |
|
|
| |
| pos_in_bar = (note.start % bar_duration) / bar_duration |
| pos_idx = int(pos_in_bar * 32) |
| tokens.append(self.position_token(pos_idx)) |
|
|
| |
| tokens.append(self.velocity_token(note.velocity)) |
| tokens.append(self.note_on_token(note.pitch)) |
|
|
| |
| 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 = {} |
|
|
| 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 |
| 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) |
|
|
| |
| 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 |
|
|