projectlosangeles commited on
Commit
21e6556
·
verified ·
1 Parent(s): 9a1e8cd

Upload 6 files

Browse files
Files changed (6) hide show
  1. TMIDIX.py +0 -0
  2. app.py +1017 -0
  3. midi_to_colab_audio.py +0 -0
  4. packages.txt +1 -0
  5. requirements.txt +10 -0
  6. x_transformer_2_3_1.py +0 -0
TMIDIX.py ADDED
The diff for this file is too large to render. See raw diff
 
app.py ADDED
@@ -0,0 +1,1017 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #===================================================================
2
+ # https://huggingface.co/spaces/asigalov61/Orpheus-Music-Transformer
3
+ #===================================================================
4
+
5
+ """
6
+ Orpheus Music Transformer Gradio App
7
+ SOTA 8k multi-instrumental music transformer trained on 2.31M+ high-quality MIDIs
8
+ """
9
+
10
+ #===================================================================
11
+ # Environment requirements (fully cross platform and minimal)
12
+ #===================================================================
13
+ # pip requirements
14
+ #-------------------------------------------------------------------
15
+ # !pip install tqdm
16
+ # !pip install numpy
17
+ # !pip install matplotlib
18
+ # !pip install gradio
19
+ # !pip install hf-transfer
20
+ # !pip install huggingface_hub
21
+ # !pip install torch
22
+ # !pip install einops
23
+ # !pip install einx
24
+ # !pip install scikit-learn
25
+ #===================================================================
26
+ # apt requirements
27
+ #-------------------------------------------------------------------
28
+ # !sudo apt install fluidsynth -y
29
+ #===================================================================
30
+ # Required modules
31
+ #-------------------------------------------------------------------
32
+ # Download modules from https://github.com/asigalov61/tegridy-tools
33
+ #-------------------------------------------------------------------
34
+ # TMIDIX.py
35
+ # x_transformer_2_3_1.py
36
+ # midi_to_colab_audio.py
37
+ #===================================================================
38
+
39
+ # -----------------------------
40
+ # CONFIGURATION & GLOBALS
41
+ # -----------------------------
42
+ TIME_ZONE = 'US/Pacific'
43
+
44
+ SEQ_LEN = 8192
45
+ PAD_IDX = 18819
46
+
47
+ MODELS_CHECKPOINTS = [
48
+ {
49
+ 'checkpoint_tag': 'Medium Base Model',
50
+ 'checkpoint_name': 'Orpheus_Music_Transformer_Trained_Model_128497_steps_0.6934_loss_0.7927_acc.pth',
51
+ 'checkpoint_depth': 8,
52
+ 'checkpoint_heads': 32
53
+ },
54
+ {
55
+ 'checkpoint_tag': 'Large Base Model',
56
+ 'checkpoint_name': 'Orpheus_Music_Transformer_Large_Trained_Model_43860_steps_0.6682_loss_0.8054_acc.pth',
57
+ 'checkpoint_depth': 16,
58
+ 'checkpoint_heads': 16
59
+ },
60
+ {
61
+ 'checkpoint_tag': 'Large Fine-Tuned Model',
62
+ 'checkpoint_name': 'Orpheus_Music_Transformer_Large_Quality_Fine_Tuned_Model_2027_steps_1.2913_loss_0.6263_acc.pth',
63
+ 'checkpoint_depth': 16,
64
+ 'checkpoint_heads': 16
65
+ }
66
+ ]
67
+
68
+ MODEL_DEVICE = 'cuda'
69
+
70
+ SOUNDFONT_BANK = 'SGM-v2.01-YamahaGrand-Guit-Bass-v2.7.sf2'
71
+ AUDIO_SAMPLE_RATE = 16000
72
+ AUDIO_FORMAT = 'mp3'
73
+
74
+ NUM_OUT_BATCHES = 10
75
+ PREVIEW_LENGTH = 120 # in tokens
76
+
77
+ OUTPUT_MIDIS_DIR = 'output_midis'
78
+
79
+ # -----------------------------
80
+ # START-UP INFO FUNCTIONS
81
+ # -----------------------------
82
+ SEP = '=' * 70
83
+
84
+ def print_sep():
85
+ print(SEP)
86
+
87
+ print_sep()
88
+ print("Orpheus Music Transformer Gradio App")
89
+ print_sep()
90
+ print("Loading modules...")
91
+
92
+ # -----------------------------
93
+ # ENVIRONMENT & MODULES IMPORTS
94
+ # -----------------------------
95
+
96
+ import os
97
+
98
+ os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
99
+
100
+ RUNNING_IN_SPACE = (
101
+ os.environ.get("SYSTEM", "").lower() == "spaces"
102
+ or "SPACE_ID" in os.environ
103
+ or "HF_SPACE_ID" in os.environ
104
+ )
105
+
106
+ import argparse
107
+
108
+ from pathlib import Path
109
+ from io import BytesIO
110
+
111
+ import time as reqtime
112
+ import datetime
113
+ from pytz import timezone
114
+
115
+ PDT = timezone(TIME_ZONE)
116
+
117
+ import random
118
+
119
+ if RUNNING_IN_SPACE:
120
+ import spaces
121
+ GPU = spaces.GPU
122
+ else:
123
+ def GPU(*args, **kwargs):
124
+ def wrapper(fn):
125
+ return fn
126
+ return wrapper
127
+
128
+ import gradio as gr
129
+
130
+ import TMIDIX
131
+
132
+ from midi_to_colab_audio import midi_to_colab_audio
133
+
134
+ import matplotlib.pyplot as plt
135
+
136
+ from huggingface_hub import hf_hub_download
137
+
138
+ # -----------------------------
139
+ # PyTorch
140
+ # -----------------------------
141
+
142
+ import torch
143
+
144
+ os.environ['USE_FLASH_ATTENTION'] = '1'
145
+
146
+ torch.set_float32_matmul_precision('high')
147
+ torch.backends.cuda.matmul.allow_tf32 = True
148
+ torch.backends.cudnn.allow_tf32 = True
149
+ torch.backends.cuda.enable_mem_efficient_sdp(True)
150
+ torch.backends.cuda.enable_math_sdp(True)
151
+ torch.backends.cuda.enable_flash_sdp(True)
152
+ torch.backends.cuda.enable_cudnn_sdp(True)
153
+
154
+ MODEL_DTYPE = torch.bfloat16
155
+
156
+ # -----------------------------
157
+ # X-Transformer
158
+ # -----------------------------
159
+
160
+ from x_transformer_2_3_1 import TransformerWrapper, AutoregressiveWrapper, Decoder, top_p
161
+
162
+ print_sep()
163
+ print("PyTorch version:", torch.__version__)
164
+ print("Done loading modules!")
165
+ print_sep()
166
+
167
+ # -----------------------------
168
+ # SPACES AND LOCAL ARGS
169
+ # -----------------------------
170
+
171
+ def parse_local_args():
172
+ parser = argparse.ArgumentParser()
173
+ parser.add_argument("--soundfont-name", type=str, default="SGM-v2.01-YamahaGrand-Guit-Bass-v2.7.sf2")
174
+ return parser.parse_args()
175
+
176
+ args = parse_local_args() if not RUNNING_IN_SPACE else None
177
+
178
+ if args:
179
+ SOUNDFONT_BANK = args.soundfont_name
180
+
181
+ # -----------------------------
182
+ # MODELS INIT FUNCTIONS
183
+ # -----------------------------
184
+ print_sep()
185
+
186
+ #------------------------------------------------------------------------
187
+
188
+ def load_model(model_dic):
189
+
190
+ print('Instantiating model...')
191
+
192
+ model = TransformerWrapper(
193
+ num_tokens=PAD_IDX + 1,
194
+ max_seq_len=SEQ_LEN,
195
+ attn_layers=Decoder(
196
+ dim=2048,
197
+ depth=model_dic['checkpoint_depth'],
198
+ heads=model_dic['checkpoint_heads'],
199
+ rotary_pos_emb=True,
200
+ attn_flash=True
201
+ )
202
+ )
203
+ model = AutoregressiveWrapper(model,
204
+ ignore_index=PAD_IDX,
205
+ pad_value=PAD_IDX
206
+ )
207
+
208
+ print('Done!')
209
+ print_sep()
210
+ print("Model will use", MODEL_DTYPE.__repr__().split('.')[-1], "precision...")
211
+ print("Model will use", MODEL_DEVICE, "device...")
212
+ print_sep()
213
+ print("Loading model checkpoint...")
214
+ print('Checkpoint name:', model_dic['checkpoint_name'])
215
+ print_sep()
216
+
217
+ checkpoint = hf_hub_download(
218
+ repo_id='asigalov61/Orpheus-Music-Transformer',
219
+ filename=model_dic['checkpoint_name']
220
+ )
221
+
222
+ model.load_state_dict(torch.load(checkpoint, map_location='cpu'))
223
+
224
+ model.eval()
225
+
226
+ model.cpu()
227
+
228
+ model = torch.compile(model, mode='max-autotune')
229
+
230
+ print_sep()
231
+ print("Done!")
232
+ print_sep()
233
+
234
+ return model_dic['checkpoint_tag'], model
235
+
236
+ #------------------------------------------------------------------------
237
+
238
+ models_dict = {}
239
+
240
+ for model_dic in MODELS_CHECKPOINTS:
241
+ tag, model = load_model(model_dic)
242
+ models_dict[tag] = model
243
+
244
+ #------------------------------------------------------------------------
245
+
246
+ ctx = torch.amp.autocast(device_type=MODEL_DEVICE,
247
+ dtype=MODEL_DTYPE
248
+ )
249
+
250
+ print_sep()
251
+ print("Done!")
252
+ print_sep()
253
+
254
+ # -----------------------------
255
+ # SOUNDFONT LOADING FUNCTION
256
+ # -----------------------------
257
+ print('Loading SoundFont...')
258
+ print_sep()
259
+
260
+ SOUNDFONT_PATH = hf_hub_download(repo_id='projectlosangeles/soundfonts4u',
261
+ repo_type='dataset',
262
+ filename=SOUNDFONT_BANK
263
+ )
264
+
265
+ print_sep()
266
+ print('Done!')
267
+ print('=' * 70)
268
+
269
+ # -----------------------------
270
+ # MIDI PROCESSING FUNCTIONS
271
+ # -----------------------------
272
+ def load_midi(input_midi):
273
+
274
+ """Process the input MIDI file and create a token sequence."""
275
+
276
+ raw_score = TMIDIX.midi2single_track_ms_score(input_midi.name)
277
+
278
+ escore_notes = TMIDIX.advanced_score_processor(raw_score,
279
+ return_enhanced_score_notes=True,
280
+ apply_sustain=True
281
+ )
282
+
283
+ if escore_notes and escore_notes[0]:
284
+
285
+ escore_notes = TMIDIX.augment_enhanced_score_notes(escore_notes[0],
286
+ sort_drums_last=True
287
+ )
288
+
289
+ escore_notes = TMIDIX.remove_duplicate_pitches_from_escore_notes(escore_notes)
290
+
291
+ escore_notes = TMIDIX.fix_escore_notes_durations(escore_notes,
292
+ min_notes_gap=0
293
+ )
294
+
295
+ dscore = TMIDIX.delta_score_notes(escore_notes)
296
+
297
+ dcscore = TMIDIX.chordify_score([d[1:] for d in dscore])
298
+
299
+ melody_chords = [18816]
300
+
301
+ #=======================================================
302
+ # MAIN PROCESSING CYCLE
303
+ #=======================================================
304
+
305
+ for i, c in enumerate(dcscore):
306
+
307
+ delta_time = c[0][0]
308
+
309
+ melody_chords.append(delta_time)
310
+
311
+ for e in c:
312
+
313
+ #=======================================================
314
+
315
+ # Durations
316
+ dur = max(1, min(255, e[1]))
317
+
318
+ # Patches
319
+ pat = max(0, min(128, e[5]))
320
+
321
+ # Pitches
322
+ ptc = max(1, min(127, e[3]))
323
+
324
+ # Velocities
325
+ # Calculating octo-velocity
326
+ vel = max(8, min(127, e[4]))
327
+ velocity = round(vel / 15)-1
328
+
329
+ #=======================================================
330
+ # FINAL NOTE SEQ
331
+ #=======================================================
332
+
333
+ # Writing final note
334
+ pat_ptc = (128 * pat) + ptc
335
+ dur_vel = (8 * dur) + velocity
336
+
337
+ melody_chords.extend([pat_ptc+256, dur_vel+16768])
338
+
339
+ return melody_chords
340
+
341
+ else:
342
+ return [18816]
343
+
344
+ def save_midi(tokens):
345
+
346
+ """Convert token sequence back to a MIDI score and write it using TMIDIX.
347
+ """
348
+
349
+ time = 0
350
+ dur = 1
351
+ vel = 90
352
+ pitch = 60
353
+ channel = 0
354
+ patch = 0
355
+
356
+ patches = [-1] * 16
357
+
358
+ channels = [0] * 16
359
+ channels[9] = 1
360
+
361
+ song_f = []
362
+
363
+ for ss in tokens:
364
+
365
+ if 0 <= ss < 256:
366
+
367
+ time += ss * 16
368
+
369
+ if 256 <= ss < 16768:
370
+
371
+ patch = (ss-256) // 128
372
+
373
+ if patch < 128:
374
+
375
+ if patch not in patches:
376
+ if 0 in channels:
377
+ cha = channels.index(0)
378
+ channels[cha] = 1
379
+ else:
380
+ cha = 15
381
+
382
+ patches[cha] = patch
383
+ channel = patches.index(patch)
384
+ else:
385
+ channel = patches.index(patch)
386
+
387
+ if patch == 128:
388
+ channel = 9
389
+
390
+ pitch = (ss-256) % 128
391
+
392
+
393
+ if 16768 <= ss < 18816:
394
+
395
+ dur = ((ss-16768) // 8) * 16
396
+ vel = (((ss-16768) % 8)+1) * 15
397
+
398
+ song_f.append(['note', time, dur, channel, pitch, vel, patch])
399
+
400
+ if song_f is not None and song_f:
401
+
402
+ song_f = TMIDIX.remove_duplicate_pitches_from_escore_notes(song_f)
403
+
404
+ song_f = TMIDIX.fix_escore_notes_durations(song_f,
405
+ min_notes_gap=0
406
+ )
407
+
408
+ output_score, patches, overflow_patches = TMIDIX.patch_enhanced_score_notes(song_f)
409
+
410
+ now = datetime.datetime.now(PDT)
411
+ ms4 = now.strftime("%f")[:4] # first four digits of microseconds
412
+
413
+ fname = (
414
+ "Orpheus-Music-Transformer-Composition-"
415
+ + now.strftime(f"%Y-%m-%d-%H-%M-%S-{ms4}")
416
+ )
417
+
418
+ os.makedirs(OUTPUT_MIDIS_DIR, exist_ok=True)
419
+
420
+ output_fname = os.path.join(OUTPUT_MIDIS_DIR, fname)
421
+
422
+ TMIDIX.Tegridy_ms_SONG_to_MIDI_Converter(
423
+ output_score,
424
+ output_signature='Orpheus Music Transformer',
425
+ output_file_name=output_fname,
426
+ track_name='Project Los Angeles',
427
+ list_of_MIDI_patches=patches,
428
+ verbose=False
429
+ )
430
+ return output_fname, output_score
431
+
432
+ else:
433
+ return None, None
434
+
435
+ # -----------------------------
436
+ # TOKENS SANITIZER FUNCTIONS
437
+ # -----------------------------
438
+
439
+ def extract_pairs_and_prefix(lst):
440
+
441
+ RANGE1 = (0, 255)
442
+ RANGE2 = (256, 16767)
443
+ RANGE3 = (16768, 18815)
444
+ RANGE4 = (18816, 18819)
445
+
446
+ def in_range(x, r):
447
+ return r[0] <= x <= r[1]
448
+
449
+ prefix = []
450
+ started = False
451
+
452
+ for x in lst:
453
+ if in_range(x, RANGE2):
454
+ started = True
455
+ break
456
+
457
+ prefix.append(x)
458
+
459
+ pairs = []
460
+ pending = None
461
+
462
+ for x in lst:
463
+ if in_range(x, RANGE2):
464
+ pending = x
465
+
466
+ elif in_range(x, RANGE3):
467
+ if pending is not None:
468
+ pairs.append((pending, x))
469
+ pending = None
470
+
471
+ elif in_range(x, RANGE4):
472
+ pairs.append((x, x))
473
+
474
+ return prefix, pairs
475
+
476
+ def sanitize_tokens(tokens):
477
+
478
+ chords = []
479
+ cho = []
480
+
481
+ for t in tokens:
482
+ if t < 256:
483
+ if cho:
484
+ chords.append(cho)
485
+
486
+ cho = [t]
487
+
488
+ else:
489
+ cho.append(t)
490
+
491
+ if cho:
492
+ chords.append(cho)
493
+
494
+ san_tokens = []
495
+
496
+ for cho in chords:
497
+ pfx, ptcs_durs = extract_pairs_and_prefix(cho)
498
+
499
+ san_tokens.extend(pfx)
500
+
501
+ san_ptcs_durs = []
502
+ seen = []
503
+
504
+ for ptc, dur in ptcs_durs:
505
+ if 256 <= ptc < 16768:
506
+ if ptc not in seen:
507
+ san_tokens.append(ptc)
508
+ san_tokens.append(dur)
509
+ seen.append(ptc)
510
+
511
+ else:
512
+ san_tokens.append(ptc)
513
+
514
+ return san_tokens
515
+
516
+ # -----------------------------
517
+ # MUSIC GENERATION FUNCTIONS
518
+ # -----------------------------
519
+ @GPU
520
+ def generate_music(prime,
521
+ num_gen_tokens,
522
+ num_gen_batches,
523
+ model_temperature,
524
+ model_top_p,
525
+ model_selector
526
+ ):
527
+
528
+ """Generate music tokens given prime tokens and parameters."""
529
+
530
+ if len(prime) >= 6656:
531
+ prime = [18816] + prime[-6656:]
532
+
533
+ inputs = prime
534
+
535
+ print(f'Will use {model_selector[0]}...')
536
+
537
+ model = models_dict[model_selector[0]]
538
+
539
+ model.to(MODEL_DEVICE)
540
+
541
+ print("Generating...")
542
+ inp = torch.LongTensor([inputs] * num_gen_batches).to(MODEL_DEVICE)
543
+
544
+ if model_top_p < 1:
545
+ with ctx:
546
+ out = model.generate(
547
+ inp,
548
+ num_gen_tokens,
549
+ filter_logits_fn=top_p,
550
+ filter_kwargs={'thres': model_top_p},
551
+ temperature=model_temperature,
552
+ eos_token=18818,
553
+ return_prime=False,
554
+ verbose=False
555
+ )
556
+
557
+ else:
558
+ with ctx:
559
+ out = model.generate(
560
+ inp,
561
+ num_gen_tokens,
562
+ temperature=model_temperature,
563
+ eos_token=18818,
564
+ return_prime=False,
565
+ verbose=False
566
+ )
567
+
568
+ model.cpu()
569
+
570
+ print("Done!")
571
+ print_sep()
572
+ return out.tolist()
573
+
574
+ def generate_music_and_state(input_midi,
575
+ prime_instruments,
576
+ num_prime_tokens,
577
+ num_gen_tokens,
578
+ model_temperature,
579
+ model_top_p,
580
+ add_drums,
581
+ add_outro,
582
+ final_composition,
583
+ generated_batches,
584
+ block_lines,
585
+ model_selector
586
+ ):
587
+
588
+ """
589
+ Generate tokens using the model, update the composition state, and prepare outputs.
590
+ This function combines seed loading, token generation, and UI output packaging.
591
+ """
592
+
593
+ print_sep()
594
+ print("Request start time:", datetime.datetime.now(PDT).strftime("%Y-%m-%d %H:%M:%S"))
595
+ start_time = reqtime.time()
596
+
597
+ print_sep()
598
+ print('Requested model:', model_selector[0])
599
+
600
+ if input_midi is not None:
601
+ fn = os.path.basename(input_midi.name)
602
+ fn1 = fn.split('.')[0]
603
+ print('Input file name:', fn)
604
+
605
+ print('Prime instruments:', prime_instruments)
606
+ print('Num prime tokens:', num_prime_tokens)
607
+ print('Num gen tokens:', num_gen_tokens)
608
+
609
+ print('Model temp:', model_temperature)
610
+ print('Model top p:', model_top_p)
611
+
612
+ print('Add drums:', add_drums)
613
+ print('Add outro:', add_outro)
614
+
615
+ print_sep()
616
+
617
+ # Load seed from MIDI if there is no existing composition.
618
+ if not final_composition and input_midi is not None:
619
+ final_composition = load_midi(input_midi)
620
+
621
+ if num_prime_tokens < 6656:
622
+ final_composition = final_composition[:num_prime_tokens]
623
+
624
+ midi_fname, midi_score = save_midi(final_composition)
625
+ # Use the last note's time as a marker.
626
+ last_nd_note = [e for e in midi_score if e[3] != 9]
627
+ block_lines.append((last_nd_note[-1][1]+last_nd_note[-1][2]) // 1000 if final_composition else 0)
628
+
629
+ if not final_composition and input_midi is None and prime_instruments:
630
+ final_composition = [18816, 0]
631
+
632
+ if "Drums" in prime_instruments:
633
+ ci_num = random.choice([37, 42])
634
+
635
+ for _ in range(4):
636
+ final_composition.append((128*128)+ci_num+256)
637
+ final_composition.append((8*16)+7+16768)
638
+ final_composition.append(32)
639
+
640
+ nd_instruments = [i for i in prime_instruments[:4] if i != 'Drums']
641
+
642
+ if nd_instruments:
643
+ prime_chord = random.choice([c for c in TMIDIX.ALL_CHORDS_FULL if len(c) == len(nd_instruments)])
644
+
645
+ for i, instr in enumerate(nd_instruments):
646
+ instr_num = Patch2number[instr]
647
+ instr_oct = TMIDIX.Patch2octave[instr]
648
+
649
+ final_composition.append((128*instr_num)+(instr_oct+prime_chord[i])+256)
650
+ dur = random.randint(16, 32)
651
+ vel = random.randint(5, 7)
652
+ final_composition.append((8*dur)+vel+16768)
653
+
654
+ if 'Drums' in prime_instruments:
655
+ drum_pitch = random.choice([35, 36, 41, 43, 45, 47, 48, 50])
656
+ final_composition.append((128*128)+(drum_pitch)+256)
657
+ final_composition.append((8*16)+7+16768)
658
+
659
+ drum_seq = []
660
+ outro_seq = []
661
+
662
+ if final_composition:
663
+
664
+ if add_drums or add_outro:
665
+ final_composition = TMIDIX.trim_list_trail_range(final_composition, 16768, 18815)
666
+
667
+ if add_drums:
668
+ drum_pitches = random.sample([35, 36, 41, 43, 45], k=1)
669
+ for dp in drum_pitches:
670
+ drum_seq.append((128*128)+dp+256) # Drum patch/pitch token
671
+ drum_seq.append((8*16)+7+16768) # Dur/vel token
672
+
673
+ if add_outro:
674
+ outro_seq.append(18817) # Outro token
675
+
676
+ if not final_composition and input_midi is None and not prime_instruments:
677
+ final_composition = [18816, 0]
678
+
679
+ print_sep()
680
+ print('Composition has', len(final_composition+drum_seq+outro_seq), 'tokens')
681
+ print_sep()
682
+
683
+ batched_gen_tokens = generate_music(final_composition+drum_seq+outro_seq,
684
+ num_gen_tokens,
685
+ NUM_OUT_BATCHES,
686
+ model_temperature,
687
+ model_top_p,
688
+ model_selector
689
+ )
690
+
691
+ batched_gen_tokens_san = []
692
+
693
+ for tokens in batched_gen_tokens:
694
+ san_tokens = sanitize_tokens(tokens)
695
+ batched_gen_tokens_san.append(san_tokens)
696
+
697
+ batched_gen_tokens = batched_gen_tokens_san
698
+
699
+ batched_gen_tokens_ext = []
700
+
701
+ if drum_seq or outro_seq:
702
+ for tokens in batched_gen_tokens:
703
+ batched_gen_tokens_ext.append(drum_seq+outro_seq+tokens)
704
+
705
+ batched_gen_tokens = batched_gen_tokens_ext
706
+
707
+ output_batches = []
708
+ for i, tokens in enumerate(batched_gen_tokens):
709
+ preview_composition = final_composition+drum_seq+outro_seq
710
+ preview_tokens = preview_composition[-PREVIEW_LENGTH:]
711
+
712
+ plot_kwargs = {'plot_title': f'Batch # {i}', 'return_plt': True}
713
+
714
+ if len(preview_composition) > PREVIEW_LENGTH:
715
+ preview_score = save_midi(preview_tokens[:PREVIEW_LENGTH])[1]
716
+ plot_kwargs['block_lines_times_list'] = [(preview_score[-1][1]+preview_score[-1][2]) // 1000]
717
+
718
+ midi_fname, midi_score = save_midi(preview_tokens + tokens)
719
+ midi_plot = TMIDIX.plot_ms_SONG(midi_score,
720
+ **plot_kwargs
721
+ )
722
+
723
+ gradio_audio = midi_to_colab_audio(midi_fname + '.mid',
724
+ soundfont_path=SOUNDFONT_PATH,
725
+ sample_rate=AUDIO_SAMPLE_RATE,
726
+ output_for_gradio=True)
727
+
728
+ output_batches.append([(AUDIO_SAMPLE_RATE, gradio_audio), midi_plot, tokens, midi_fname + '.mid'])
729
+
730
+ # Update generated_batches (for use by add/remove functions)
731
+ generated_batches = batched_gen_tokens
732
+
733
+ # Flatten outputs: states then audio and plots for each batch.
734
+ outputs_flat = []
735
+ for batch in output_batches:
736
+ outputs_flat.extend([batch[0], batch[1], batch[3]])
737
+
738
+ print("Request end time:", datetime.datetime.now(PDT).strftime("%Y-%m-%d %H:%M:%S"))
739
+ print_sep()
740
+
741
+ end_time = reqtime.time()
742
+ execution_time = end_time - start_time
743
+
744
+ print(f"Request execution time: {execution_time} seconds")
745
+ print_sep()
746
+
747
+ return [final_composition, generated_batches, block_lines] + outputs_flat
748
+
749
+ # -----------------------------
750
+ # BATCH HANDLING FUNCTIONS
751
+ # -----------------------------
752
+ def add_batch(batch_number, final_composition, generated_batches, block_lines):
753
+ """Add tokens from the specified batch to the final composition and update outputs."""
754
+ if generated_batches:
755
+ final_composition.extend(generated_batches[batch_number])
756
+ midi_fname, midi_score = save_midi(final_composition)
757
+ last_nd_note = [e for e in midi_score if e[3] != 9]
758
+ block_lines.append((last_nd_note[-1][1]+last_nd_note[-1][2]) // 1000 if final_composition else 0)
759
+ midi_plot = TMIDIX.plot_ms_SONG(
760
+ midi_score,
761
+ plot_title='Orpheus Music Transformer Composition',
762
+ block_lines_times_list=block_lines[:-1],
763
+ return_plt=True
764
+ )
765
+ gradio_audio = midi_to_colab_audio(midi_fname + '.mid',
766
+ soundfont_path=SOUNDFONT_PATH,
767
+ sample_rate=AUDIO_SAMPLE_RATE,
768
+ output_for_gradio=True)
769
+ print("Added batch #", batch_number)
770
+ print_sep()
771
+ return (AUDIO_SAMPLE_RATE, gradio_audio), midi_plot, midi_fname + '.mid', final_composition, generated_batches, block_lines
772
+
773
+ else:
774
+ return None, None, None, [], [], []
775
+
776
+ def remove_batch(batch_number, num_tokens, final_composition, generated_batches, block_lines):
777
+ """Remove tokens from the final composition and update outputs."""
778
+ if final_composition and len(final_composition) > num_tokens:
779
+ final_composition = final_composition[:-num_tokens]
780
+ if block_lines:
781
+ block_lines.pop()
782
+ midi_fname, midi_score = save_midi(final_composition)
783
+
784
+ if midi_fname and midi_score:
785
+ midi_plot = TMIDIX.plot_ms_SONG(
786
+ midi_score,
787
+ plot_title='Orpheus Music Transformer Composition',
788
+ block_lines_times_list=block_lines[:-1],
789
+ return_plt=True
790
+ )
791
+ gradio_audio = midi_to_colab_audio(midi_fname + '.mid',
792
+ soundfont_path=SOUNDFONT_PATH,
793
+ sample_rate=AUDIO_SAMPLE_RATE,
794
+ output_for_gradio=True)
795
+ print("Removed batch #", batch_number)
796
+ print_sep()
797
+ return (AUDIO_SAMPLE_RATE, gradio_audio), midi_plot, midi_fname + '.mid', final_composition, generated_batches, block_lines
798
+
799
+ return None, None, None, [], [], []
800
+
801
+ # -----------------------------
802
+ # MISC FUNCTIONS
803
+ # -----------------------------
804
+
805
+ def clear():
806
+ """Clear outputs and reset state."""
807
+ print_sep()
808
+ print('Clear batch...')
809
+ print_sep()
810
+ return None, None, None, [], []
811
+
812
+ def reset(final_composition=[], generated_batches=[], block_lines=[]):
813
+ """Reset composition state."""
814
+ print_sep()
815
+ print('Reset composition...')
816
+ print_sep()
817
+ return [], [], []
818
+
819
+ def update_state_from_dropdown(choice, state):
820
+ """Store the dropdown value inside the global state list"""
821
+ print_sep()
822
+ print('Changed model from', state[0], 'to', choice)
823
+ print_sep()
824
+ state[0] = choice
825
+ return state
826
+
827
+ Patch2number = TMIDIX.reverse_dict(TMIDIX.Number2patch)
828
+ Patch2number['Drums'] = 128
829
+
830
+ # -----------------------------
831
+ # GRADIO INTERFACE SETUP
832
+ # -----------------------------
833
+ with gr.Blocks() as orpheus_app:
834
+
835
+ gr.Markdown("<h1 style='text-align: left; margin-bottom: 1rem'>Orpheus Music Transformer</h1>")
836
+ gr.Markdown("<h1 style='text-align: left; margin-bottom: 1rem'>SOTA 8k multi-instrumental music transformer trained on 2.31M+ high-quality MIDIs</h1>")
837
+ gr.Markdown("<h1 style='text-align: left; margin-bottom: 1rem'>🔥[2026]🔥 Now featuring large optimized model!</h1>")
838
+
839
+ with gr.Row(elem_classes="duplicate-row"):
840
+ gr.Button(
841
+ value="🪬 User Guide 🪬",
842
+ variant="huggingface",
843
+ size="md",
844
+ link="https://asigalov61.github.io/Orpheus-Music-Transformer-User-Guide/",
845
+ link_target="_blank"
846
+ )
847
+
848
+ gr.DuplicateButton(
849
+ value="🤗 Duplicate 🤗",
850
+ variant="huggingface",
851
+ size="md",
852
+ link="https://huggingface.co/spaces/asigalov61/Orpheus-Music-Transformer?duplicate=true",
853
+ link_target="_blank"
854
+ )
855
+
856
+ gr.Button(
857
+ value="❤️ Models ❤️",
858
+ variant="huggingface",
859
+ size="md",
860
+ link="https://huggingface.co/asigalov61/Orpheus-Music-Transformer",
861
+ link_target="_blank"
862
+ )
863
+
864
+ gr.Button(
865
+ value="🚀 Spaces 🚀",
866
+ variant="huggingface",
867
+ size="md",
868
+ link="https://huggingface.co/collections/asigalov61/orpheus-music-transformer",
869
+ link_target="_blank"
870
+ )
871
+
872
+ gr.Button(
873
+ value="🦖 Dataset 🦖",
874
+ variant="huggingface",
875
+ size="md",
876
+ link="https://huggingface.co/datasets/projectlosangeles/Godzilla-MIDI-Dataset",
877
+ link_target="_blank"
878
+ )
879
+
880
+ gr.HTML("""
881
+ <iframe width="100%" height="300" scrolling="no" frameborder="no" allow="autoplay" src="https://w.soundcloud.com/player/?url=https%3A//api.soundcloud.com/playlists/2042253855&color=%23ff5500&auto_play=false&hide_related=false&show_comments=true&show_user=true&show_reposts=false&show_teaser=true&visual=true"></iframe><div style="font-size: 10px; color: #cccccc;line-break: anywhere;word-break: normal;overflow: hidden;white-space: nowrap;text-overflow: ellipsis; font-family: Interstate,Lucida Grande,Lucida Sans Unicode,Lucida Sans,Garuda,Verdana,Tahoma,sans-serif;font-weight: 100;"><a href="https://soundcloud.com/aleksandr-sigalov-61" title="Project Los Angeles" target="_blank" style="color: #cccccc; text-decoration: none;">Project Los Angeles</a> · <a href="https://soundcloud.com/aleksandr-sigalov-61/sets/orpheus-music-transformer" title="Orpheus Music Transformer" target="_blank" style="color: #cccccc; text-decoration: none;">Orpheus Music Transformer</a></div>
882
+ """)
883
+
884
+ gr.Markdown("## Key Features")
885
+ gr.Markdown("""
886
+ - **Efficient Architecture with RoPE**: Large optimized 748M full attention autoregressive transformer with RoPE.
887
+ - **Extended Sequence Length**: 8k tokens that comfortably fit most music compositions and facilitate long-term music structure generation.
888
+ - **Premium Training Data**: Trained solely on the highest-quality MIDIs from the Godzilla MIDI dataset.
889
+ - **Optimized MIDI Encoding**: Extremely efficient MIDI representation using only 3 tokens per note and 7 tokens per tri-chord.
890
+ - **Distinct Encoding Order**: Features a unique duration/velocity last MIDI encoding order for refined musical expression.
891
+ - **Full-Range Instrumental Learning**: True full-range MIDI instruments encoding enabling the model to learn each instrument separately.
892
+ - **Natural Composition Endings**: Outro tokens that help generate smooth and natural musical conclusions.
893
+ """)
894
+
895
+ gr.Markdown("## Best Practices Tips")
896
+ gr.Markdown("""
897
+ - Good prime seed MIDI is everything!!!
898
+ - Trim the seed MIDI to exact (or at least - approximate) musical phrase.
899
+ - 30sec-1min (1024-1536 tokens) run time is ideal.
900
+ - Remove excessive instruments. 4-5 most pronounced instruments work best.
901
+ - Do not be discouraged by generated contunuations! Sometimes you need several tries to get it right!
902
+ - Do not strive for perfection! Instead, try to have fun and enjoy the music!
903
+ """)
904
+
905
+ # Global state variables for composition
906
+ final_composition = gr.State([])
907
+ generated_batches = gr.State([])
908
+ block_lines = gr.State([])
909
+ model_selector = gr.State([list(models_dict.keys())[0]])
910
+
911
+ gr.Markdown("## Upload seed MIDI or select prime instruments or simply click 'Generate' button for random output")
912
+
913
+ gr.Markdown("""
914
+ ### PLEASE NOTE:
915
+ - Orpheus Music Transformer is a primarily continuation/co-composition model!"
916
+ - The model works best if given some music context to work with
917
+ - Random generation from SOS token/embeddings may not always produce good results
918
+ """)
919
+
920
+ input_midi = gr.File(label="Input MIDI", file_types=[".midi", ".mid", ".kar"])
921
+ input_midi.upload(reset, [final_composition, generated_batches, block_lines],
922
+ [final_composition, generated_batches, block_lines])
923
+
924
+ gr.Markdown("## Generation options")
925
+ prime_instruments = gr.Dropdown(label="Prime instruments (select up to 5)", choices=list(Patch2number.keys()),
926
+ multiselect=True, max_choices=5, type="value",
927
+ info="NOTE: Custom MIDI overrides prime instruments"
928
+ )
929
+
930
+ prime_instruments.input(reset, [final_composition, generated_batches, block_lines],
931
+ [final_composition, generated_batches, block_lines])
932
+
933
+ num_prime_tokens = gr.Slider(16, 6656, value=6656, step=1, label="Number of prime tokens")
934
+ num_gen_tokens = gr.Slider(16, 1024, value=512, step=1, label="Number of tokens to generate")
935
+ requested_model = gr.Dropdown(label="Model to use",
936
+ choices=list(models_dict.keys()),
937
+ value=list(models_dict.keys())[0],
938
+ info="Use medium model when speed is important, use large models when quality is important"
939
+ )
940
+ model_temperature = gr.Slider(0.1, 1, value=0.9, step=0.01, label="Model temperature",
941
+ info="Increase for more creative output, decrease for more repetitive output"
942
+ )
943
+ model_top_p = gr.Slider(0.1, 1.0, value=0.96, step=0.01, label="Model sampling top p value",
944
+ info="1 == Disabled"
945
+ )
946
+ add_drums = gr.Checkbox(value=False, label="Add drums")
947
+ add_outro = gr.Checkbox(value=False, label="Add an outro")
948
+
949
+ generate_btn = gr.Button("Generate", variant="primary")
950
+
951
+ gr.Markdown("## Batch Previews")
952
+ outputs = [final_composition, generated_batches, block_lines]
953
+ # Two outputs (audio and plot) for each batch
954
+ for i in range(NUM_OUT_BATCHES):
955
+ with gr.Tab(f"Batch # {i}"):
956
+ audio_output = gr.Audio(label=f"Batch # {i} MIDI Audio", format=AUDIO_FORMAT)
957
+ plot_output = gr.Plot(label=f"Batch # {i} MIDI Plot")
958
+ midi_file = gr.File(label=f"Batch # {i} MIDI File")
959
+ outputs.extend([audio_output, plot_output, midi_file])
960
+
961
+ requested_model.change(
962
+ fn=update_state_from_dropdown,
963
+ inputs=[requested_model, model_selector],
964
+ outputs=model_selector
965
+ )
966
+
967
+ generate_btn.click(
968
+ generate_music_and_state,
969
+ [input_midi,
970
+ prime_instruments,
971
+ num_prime_tokens,
972
+ num_gen_tokens,
973
+ model_temperature,
974
+ model_top_p,
975
+ add_drums,
976
+ add_outro,
977
+ final_composition,
978
+ generated_batches,
979
+ block_lines,
980
+ model_selector
981
+ ],
982
+ outputs
983
+ )
984
+
985
+ gr.Markdown("## Add/Remove Batch")
986
+ batch_number = gr.Slider(0, NUM_OUT_BATCHES - 1, value=0, step=1, label="Batch number to add/remove")
987
+ add_btn = gr.Button("Add batch", variant="primary")
988
+ remove_btn = gr.Button("Remove batch", variant="stop")
989
+ clear_btn = gr.ClearButton()
990
+
991
+ final_audio_output = gr.Audio(label="Final MIDI audio", format=AUDIO_FORMAT)
992
+ final_plot_output = gr.Plot(label="Final MIDI plot")
993
+ final_file_output = gr.File(label="Final MIDI file")
994
+
995
+ add_btn.click(
996
+ add_batch,
997
+ [batch_number, final_composition, generated_batches, block_lines],
998
+ [final_audio_output, final_plot_output, final_file_output, final_composition, generated_batches, block_lines]
999
+ )
1000
+ remove_btn.click(
1001
+ remove_batch,
1002
+ [batch_number, num_gen_tokens, final_composition, generated_batches, block_lines],
1003
+ [final_audio_output, final_plot_output, final_file_output, final_composition, generated_batches, block_lines]
1004
+ )
1005
+ clear_btn.click(clear, inputs=None,
1006
+ outputs=[final_audio_output, final_plot_output, final_file_output, final_composition, block_lines])
1007
+
1008
+ # -----------------------------
1009
+ # APP LAUNCHER
1010
+ # -----------------------------
1011
+ if __name__ == "__main__":
1012
+ orpheus_app.launch(
1013
+ mcp_server=RUNNING_IN_SPACE, # MCP only on HF
1014
+ share=not RUNNING_IN_SPACE, # Share only locally
1015
+ server_name="0.0.0.0",
1016
+ server_port=7860
1017
+ )
midi_to_colab_audio.py ADDED
The diff for this file is too large to render. See raw diff
 
packages.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ fluidsynth
requirements.txt ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ tqdm
2
+ numpy
3
+ scikit-learn
4
+ matplotlib
5
+ gradio
6
+ hf-transfer
7
+ huggingface_hub
8
+ torch
9
+ einops
10
+ einx
x_transformer_2_3_1.py ADDED
The diff for this file is too large to render. See raw diff