AriaLM / src /02_tokenizer.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
8.81 kB
"""
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