Views
No views yet
1from miditok import PerTok, TokenizerConfig
2from transformers import AutoModelForCausalLM, GenerationConfig
3import torch
4import io
5import mido
6import rtmidi
7
8# Load tokenizer config and tokenizer
9config = TokenizerConfig(
10 num_velocities=8,
11 use_velocities=True,
12 use_chords=False,
13 use_rests=True,
14 use_tempos=True,
15 use_time_signatures=False,
16 use_sustain_pedals=False,
17 use_pitch_bends=False,
18 use_pitch_intervals=False,
19 use_programs=False,
20 use_pitchdrum_tokens=False,
21 ticks_per_quarter=320,
22 use_microtiming=False,
23 max_microtiming_shift=0.125
24)
25tokenizer = PerTok(config)
26tokenizer.from_pretrained("xingjianll/midi-tokenizer")
27
28# Load model and generation config
29model = AutoModelForCausalLM.from_pretrained("xingjianll/midi-gpt2")
30gen_config = GenerationConfig.from_pretrained("xingjianll/midi-gpt2")
31model.eval()
32
33# Generate token sequence
34with torch.no_grad():
35 input_ids = torch.tensor([[tokenizer["BOS_None"]]], dtype=torch.long)
36 output = model.generate(
37 input_ids=input_ids,
38 generation_config=gen_config
39 )
40
41# Decode tokens to MIDI bytes
42generated_ids = output[0].tolist()
43midi_bytes = tokenizer.decode([generated_ids]).dumps_midi()
44midi_file = io.BytesIO(midi_bytes)
45
46# Optionally save to disk
47# with open("output.mid", "wb") as f:
48# f.write(midi_bytes)
49
50# Setup MIDI output to GarageBand
51midiout = rtmidi.MidiOut()
52available_ports = midiout.get_ports()
53print("Available MIDI ports:", available_ports)
54
55# Choose GarageBand's virtual port
56port = mido.open_output('GarageBand Virtual In')
57
58# Play MIDI via GarageBand
59midi = mido.MidiFile(file=midi_file)
60for msg in midi.play():
61 print(msg)
62 port.send(msg)
63