This commit is contained in:
billwuhao
2025-03-10 07:07:04 +08:00
parent 10fed8ca99
commit 9f4fcee10a
11 changed files with 3167 additions and 1 deletions
+22
View File
@@ -0,0 +1,22 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- master
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+370
View File
@@ -0,0 +1,370 @@
import os
import time
import torch
from .utils import *
from .config import nota_lx, nota_small, nota_medium
from transformers import GPT2Config
from abctoolkit.utils import Barline_regexPattern
# from abctoolkit.transpose import Note_list, Pitch_sign_list
from abctoolkit.duration import calculate_bartext_duration
node_dir = os.path.dirname(os.path.abspath(__file__))
comfy_path = os.path.dirname(os.path.dirname(node_dir))
output_path = os.path.join(comfy_path, "output")
# Path to weights for inference
nota_model_path = os.path.join(comfy_path, "models", "TTS", "NotaGen")
# Folder to save output files
ORIGINAL_OUTPUT_FOLDER = os.path.join(output_path, 'notagen_original')
INTERLEAVED_OUTPUT_FOLDER = os.path.join(output_path, 'notagen_interleaved')
os.makedirs(ORIGINAL_OUTPUT_FOLDER, exist_ok=True)
os.makedirs(INTERLEAVED_OUTPUT_FOLDER, exist_ok=True)
class NotaGenRun:
model_names = ["notagenx.pth", "notagen_small.pth", "notagen_medium.pth", "notagen_large.pth"]
periods = ["Baroque", "Classical", "Romantic"]
composers = ["Bach, Johann Sebastian", "Corelli, Arcangelo", "Handel, George Frideric", "Scarlatti, Domenico", "Vivaldi, Antonio", "Beethoven, Ludwig van",
"Haydn, Joseph", "Mozart, Wolfgang Amadeus", "Paradis, Maria Theresia von", "Reichardt, Louise", "Saint-Georges, Joseph Bologne", "Schroter, Corona",
"Bartok, Bela", "Berlioz, Hector", "Bizet, Georges", "Boulanger, Lili", "Boulton, Harold", "Brahms, Johannes", "Burgmuller, Friedrich",
"Butterworth, George", "Chaminade, Cecile", "Chausson, Ernest", "Chopin, Frederic", "Cornelius, Peter", "Debussy, Claude", "Dvorak, Antonin",
"Faisst, Clara", "Faure, Gabriel", "Franz, Robert", "Gonzaga, Chiquinha", "Grandval, Clemence de", "Grieg, Edvard", "Hensel, Fanny",
"Holmes, Augusta Mary Anne", "Jaell, Marie", "Kinkel, Johanna", "Kralik, Mathilde", "Lang, Josephine", "Lehmann, Liza", "Liszt, Franz",
"Mayer, Emilie", "Medtner, Nikolay", "Mendelssohn, Felix", "Munktell, Helena", "Parratt, Walter", "Prokofiev, Sergey", "Rachmaninoff, Sergei",
"Ravel, Maurice", "Saint-Saens, Camille", "Satie, Erik", "Schubert, Franz", "Schumann, Clara", "Schumann, Robert", "Scriabin, Aleksandr",
"Shostakovich, Dmitry", "Sibelius, Jean", "Smetana, Bedrich", "Tchaikovsky, Pyotr", "Viardot, Pauline", "Warlock, Peter", "Wolf, Hugo", "Zumsteeg, Emilie"]
instrumentations = ["Chamber", "Choral", "Keyboard", "Orchestral", "Vocal-Orchestral", "Art Song"]
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
nota_model_path = nota_model_path
node_dir = node_dir
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (s.model_names, {"default": "notagenx.pth"}),
"period": (s.periods, {"default": "Romantic", }),
"composer": (s.composers, {"default": "Bach, Johann Sebastian", }),
"instrumentation": (s.instrumentations, {"default": "Keyboard", }),
"num_samples": ("INT", {"default": 1, "min": 1}),
# "temperature": ("FLOAT", {"default": 0.8, "min": 0, "max": 1, "step": 0.1}),
# "top_k": ("INT", {"default": 50, "min": 0}),
# "top_p": ("FLOAT", {"default": 0.95, "min": 0, "max": 1, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
},
"optional": {
"abc2xml": ("BOOLEAN", {"default": False}),
"python_path": ("STRING", {"default": "", "multiline": False, "tooltip": "The absolute path of python.exe"}),
# "save_path": ("STRING", {"default": "", "tooltip": "(optional) Default Save to output/notagen_xxx"}),
# "custom_prompt": ("STRING", {"default": "", "multiline": True, "tooltip": "(optional) The format must be `period | composer | instrumentation`."}),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "inference_patch"
CATEGORY = "MW-NotaGen"
# Note_list = Note_list + ['z', 'x']
def rest_unreduce(self, abc_lines):
tunebody_index = None
for i in range(len(abc_lines)):
if '[V:' in abc_lines[i]:
tunebody_index = i
break
metadata_lines = abc_lines[: tunebody_index]
tunebody_lines = abc_lines[tunebody_index:]
part_symbol_list = []
voice_group_list = []
for line in metadata_lines:
if line.startswith('%%score'):
for round_bracket_match in re.findall(r'\((.*?)\)', line):
voice_group_list.append(round_bracket_match.split())
existed_voices = [item for sublist in voice_group_list for item in sublist]
if line.startswith('V:'):
symbol = line.split()[0]
part_symbol_list.append(symbol)
if symbol[2:] not in existed_voices:
voice_group_list.append([symbol[2:]])
z_symbol_list = [] # voices that use z as rest
x_symbol_list = [] # voices that use x as rest
for voice_group in voice_group_list:
z_symbol_list.append('V:' + voice_group[0])
for j in range(1, len(voice_group)):
x_symbol_list.append('V:' + voice_group[j])
part_symbol_list.sort(key=lambda x: int(x[2:]))
unreduced_tunebody_lines = []
for i, line in enumerate(tunebody_lines):
unreduced_line = ''
line = re.sub(r'^\[r:[^\]]*\]', '', line)
pattern = r'\[V:(\d+)\](.*?)(?=\[V:|$)'
matches = re.findall(pattern, line)
line_bar_dict = {}
for match in matches:
key = f'V:{match[0]}'
value = match[1]
line_bar_dict[key] = value
# calculate duration and collect barline
dur_dict = {}
for symbol, bartext in line_bar_dict.items():
right_barline = ''.join(re.split(Barline_regexPattern, bartext)[-2:])
bartext = bartext[:-len(right_barline)]
try:
bar_dur = calculate_bartext_duration(bartext)
except:
bar_dur = None
if bar_dur is not None:
if bar_dur not in dur_dict.keys():
dur_dict[bar_dur] = 1
else:
dur_dict[bar_dur] += 1
try:
ref_dur = max(dur_dict, key=dur_dict.get)
except:
pass # use last ref_dur
if i == 0:
prefix_left_barline = line.split('[V:')[0]
else:
prefix_left_barline = ''
for symbol in part_symbol_list:
if symbol in line_bar_dict.keys():
symbol_bartext = line_bar_dict[symbol]
else:
if symbol in z_symbol_list:
symbol_bartext = prefix_left_barline + 'z' + str(ref_dur) + right_barline
elif symbol in x_symbol_list:
symbol_bartext = prefix_left_barline + 'x' + str(ref_dur) + right_barline
unreduced_line += '[' + symbol + ']' + symbol_bartext
unreduced_tunebody_lines.append(unreduced_line + '\n')
unreduced_lines = metadata_lines + unreduced_tunebody_lines
return unreduced_lines
def inference_patch(self, model, period, composer, instrumentation, num_samples, abc2xml, python_path, seed):
if model == "notagenx.pth" or model == "notagen_large.pth":
cf = nota_lx
elif model == "notagen_small.pth":
cf = nota_small
elif model == "notagen_medium.pth":
cf = nota_medium
patch_size = cf["PATCH_SIZE"]
patch_length = cf["PATCH_LENGTH"]
char_num_layers = cf["CHAR_NUM_LAYERS"]
patch_num_layers = cf["PATCH_NUM_LAYERS"]
hidden_size = cf["HIDDEN_SIZE"]
patch_config = GPT2Config(num_hidden_layers=patch_num_layers,
max_length=patch_length,
max_position_embeddings=patch_length,
n_embd=hidden_size,
num_attention_heads=hidden_size // 64,
vocab_size=1)
byte_config = GPT2Config(num_hidden_layers=char_num_layers,
max_length=patch_size + 1,
max_position_embeddings=patch_size + 1,
hidden_size=hidden_size,
num_attention_heads=hidden_size // 64,
vocab_size=128)
nota_model = NotaGenLMHeadModel(encoder_config=patch_config, decoder_config=byte_config, model=model)
print("Parameter Number: " + str(sum(p.numel() for p in nota_model.parameters() if p.requires_grad)))
nota_model_path = os.path.join(self.nota_model_path, model)
checkpoint = torch.load(nota_model_path, map_location=torch.device(self.device))
nota_model.load_state_dict(checkpoint['model'])
nota_model = nota_model.to(self.device)
nota_model.eval()
prompt_lines=[
'%' + period + '\n',
'%' + composer + '\n',
'%' + instrumentation + '\n']
patchilizer = Patchilizer(model)
file_no = 1
bos_patch = [patchilizer.bos_token_id] * (patch_size - 1) + [patchilizer.eos_token_id]
while file_no <= num_samples:
start_time = time.time()
# start_time_format = time.strftime("%Y%m%d-%H%M%S")
prompt_patches = patchilizer.patchilize_metadata(prompt_lines)
byte_list = list(''.join(prompt_lines))
print(''.join(byte_list), end='')
prompt_patches = [[ord(c) for c in patch] + [patchilizer.special_token_id] * (patch_size - len(patch)) for patch
in prompt_patches]
prompt_patches.insert(0, bos_patch)
input_patches = torch.tensor(prompt_patches, device=self.device).reshape(1, -1)
failure_flag = False
end_flag = False
cut_index = None
tunebody_flag = False
while True:
predicted_patch = nota_model.generate(input_patches.unsqueeze(0),
top_k=9,
top_p=0.9,
temperature=1.2)
if not tunebody_flag and patchilizer.decode([predicted_patch]).startswith('[r:'): # start with [r:0/
tunebody_flag = True
r0_patch = torch.tensor([ord(c) for c in '[r:0/']).unsqueeze(0).to(self.device)
temp_input_patches = torch.concat([input_patches, r0_patch], axis=-1)
predicted_patch = nota_model.generate(temp_input_patches.unsqueeze(0),
top_k=9,
top_p=0.9,
temperature=1.2)
predicted_patch = [ord(c) for c in '[r:0/'] + predicted_patch
if predicted_patch[0] == patchilizer.bos_token_id and predicted_patch[1] == patchilizer.eos_token_id:
end_flag = True
break
next_patch = patchilizer.decode([predicted_patch])
for char in next_patch:
byte_list.append(char)
print(char, end='')
patch_end_flag = False
for j in range(len(predicted_patch)):
if patch_end_flag:
predicted_patch[j] = patchilizer.special_token_id
if predicted_patch[j] == patchilizer.eos_token_id:
patch_end_flag = True
predicted_patch = torch.tensor([predicted_patch], device=self.device) # (1, 16)
input_patches = torch.cat([input_patches, predicted_patch], dim=1) # (1, 16 * patch_len)
if len(byte_list) > 102400:
failure_flag = True
break
if time.time() - start_time > 20 * 60:
failure_flag = True
break
if input_patches.shape[1] >= patch_length * patch_size and not end_flag:
print('Stream generating...')
abc_code = ''.join(byte_list)
abc_lines = abc_code.split('\n')
tunebody_index = None
for i, line in enumerate(abc_lines):
if line.startswith('[r:') or line.startswith('[V:'):
tunebody_index = i
break
if tunebody_index is None or tunebody_index == len(abc_lines) - 1:
break
metadata_lines = abc_lines[:tunebody_index]
tunebody_lines = abc_lines[tunebody_index:]
metadata_lines = [line + '\n' for line in metadata_lines]
if not abc_code.endswith('\n'):
tunebody_lines = [tunebody_lines[i] + '\n' for i in range(len(tunebody_lines) - 1)] + [
tunebody_lines[-1]]
else:
tunebody_lines = [tunebody_lines[i] + '\n' for i in range(len(tunebody_lines))]
if cut_index is None:
cut_index = len(tunebody_lines) // 2
abc_code_slice = ''.join(metadata_lines + tunebody_lines[-cut_index:])
input_patches = patchilizer.encode_generate(abc_code_slice)
input_patches = [item for sublist in input_patches for item in sublist]
input_patches = torch.tensor([input_patches], device=self.device)
input_patches = input_patches.reshape(1, -1)
if not failure_flag:
generation_time_cost = time.time() - start_time
abc_text = ''.join(byte_list)
filename = time.strftime("%Y%m%d-%H%M%S") + \
"_" + format(generation_time_cost, '.2f') + '_' + str(file_no) + ".abc"
# unreduce
unreduced_output_path = os.path.join(INTERLEAVED_OUTPUT_FOLDER, filename)
abc_lines = abc_text.split('\n')
abc_lines = list(filter(None, abc_lines))
abc_lines = [line + '\n' for line in abc_lines]
try:
abc_lines = self.rest_unreduce(abc_lines)
with open(unreduced_output_path, 'w') as file:
file.writelines(abc_lines)
print(f"Saved to {unreduced_output_path}",)
if abc2xml:
import subprocess
xml_filename = f"{INTERLEAVED_OUTPUT_FOLDER}/{filename.rsplit(".", 1)[0]}.xml"
try:
subprocess.run(
[python_path, f"{self.node_dir}/abc2xml.py", '-o', INTERLEAVED_OUTPUT_FOLDER, unreduced_output_path, ],
check=True,
capture_output=True,
text=True
)
print(f"Conversion to {xml_filename}",)
except subprocess.CalledProcessError as e:
print(f"Conversion failed: {e.stderr}" if e.stderr else "Unknown error")
raise
except:
pass
else:
# original
original_output_path = os.path.join(ORIGINAL_OUTPUT_FOLDER, filename)
with open(original_output_path, 'w') as w:
w.write(abc_text)
print(f"Saved to {original_output_path}",)
if abc2xml:
import subprocess
xml_filename = f"{ORIGINAL_OUTPUT_FOLDER}/{filename.rsplit(".", 1)[0]}.xml"
try:
subprocess.run(
[python_path, f"{self.node_dir}/abc2xml.py", '-o', ORIGINAL_OUTPUT_FOLDER, original_output_path, ],
check=True,
capture_output=True,
text=True
)
print(f"Conversion to {xml_filename}",)
except subprocess.CalledProcessError as e:
print(f"Conversion failed: {e.stderr}" if e.stderr else "Unknown error")
raise
file_no += 1
else:
print('Generation failed.')
return (f"Saved to {INTERLEAVED_OUTPUT_FOLDER} and {ORIGINAL_OUTPUT_FOLDER}",)
NODE_CLASS_MAPPINGS = {
"NotaGenRun": NotaGenRun,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"NotaGenRun": "NotaGen Run",
}
+13
View File
@@ -0,0 +1,13 @@
[中文](README.md) | [English](README-en.md)
# Symbolic Music Generation, NotaGen node for ComfyUI.
![image](https://github.com/billwuhao/ComfyUI_KokoroTTS_MW/blob/master/images/2025-03-10_06-24-03.png)
Download the model to `ComfyUI\models\TTS\NotaGen` and rename it as required:
[NotaGen-X](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagenx_p_size_16_p_length_1024_p_layers_20_h_size_1280.pth) → `notagenx.pth`
[NotaGen-small](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagen_pretrain_p_size_16_p_length_2048_p_layers_12_c_layers_3_h_size_768_lr_0.0002_batch_8.pth) → `notagen_small.pth`
[NotaGen-medium](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagen_pretrain_p_size_16_p_length_2048_p_layers_16_c_layers_3_h_size_1024_lr_0.0001_batch_4.pth) → `notagen_medium.pth`
[NotaGen-large](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagen_pretrain_p_size_16_p_length_1024_p_layers_20_c_layers_6_h_size_1280_lr_0.0001_batch_4.pth) → `notagen_large.pth`
+13 -1
View File
@@ -1 +1,13 @@
# ComfyUI_NotaGen
[中文](README.md) | [English](README-en.md)
# 符号音乐生成. NotaGen 的 ComfyUI 节点.
![image](https://github.com/billwuhao/ComfyUI_KokoroTTS_MW/blob/master/images/2025-03-10_06-24-03.png)
将模型下载放到 `ComfyUI\models\TTS\NotaGen` 下, 并按要求重命名:
[NotaGen-X](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagenx_p_size_16_p_length_1024_p_layers_20_h_size_1280.pth) → `notagenx.pth`
[NotaGen-small](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagen_pretrain_p_size_16_p_length_2048_p_layers_12_c_layers_3_h_size_768_lr_0.0002_batch_8.pth) → `notagen_small.pth`
[NotaGen-medium](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagen_pretrain_p_size_16_p_length_2048_p_layers_16_c_layers_3_h_size_1024_lr_0.0001_batch_4.pth) → `notagen_medium.pth`
[NotaGen-large](https://huggingface.co/ElectricAlexis/NotaGen/blob/main/weights_notagen_pretrain_p_size_16_p_length_1024_p_layers_20_c_layers_6_h_size_1280_lr_0.0001_batch_4.pth) → `notagen_large.pth`
+3
View File
@@ -0,0 +1,3 @@
from .NotaGenNode import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+2252
View File
File diff suppressed because it is too large Load Diff
+30
View File
@@ -0,0 +1,30 @@
# Configurations for model
nota_lx = {
"PATCH_STREAM": True, # Stream training / inference
"PATCH_SIZE": 16, # Patch Size
"PATCH_LENGTH": 1024, # Patch Length
"CHAR_NUM_LAYERS": 6, # Number of layers in the decoder
"PATCH_NUM_LAYERS": 20, # Number of layers in the encoder
"HIDDEN_SIZE": 1280, # Hidden Size
"PATCH_SAMPLING_BATCH_SIZE": 0 # Batch size for patch during training, 0 for full conaudio
}
nota_small = {
"PATCH_STREAM": True, # Stream training / inference
"PATCH_SIZE": 16, # Patch Size
"PATCH_LENGTH": 2048, # Patch Length
"CHAR_NUM_LAYERS": 3, # Number of layers in the decoder
"PATCH_NUM_LAYERS": 12, # Number of layers in the encoder
"HIDDEN_SIZE": 768, # Hidden Size
"PATCH_SAMPLING_BATCH_SIZE": 0 # Batch size for patch during training, 0 for full conaudio
}
nota_medium = {
"PATCH_STREAM": True, # Stream training / inference
"PATCH_SIZE": 16, # Patch Size
"PATCH_LENGTH": 2048, # Patch Length
"CHAR_NUM_LAYERS": 3, # Number of layers in the decoder
"PATCH_NUM_LAYERS": 16, # Number of layers in the encoder
"HIDDEN_SIZE": 1024, # Hidden Size
"PATCH_SAMPLING_BATCH_SIZE": 0 # Batch size for patch during training, 0 for full conaudio
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 36 KiB

+14
View File
@@ -0,0 +1,14 @@
[project]
name = "notagen-mw"
description = "Symbolic Music Generation, NotaGen node for ComfyUI."
version = "1.0.0"
license = {file = "LICENSE"}
[project.urls]
Repository = "https://github.com/billwuhao/ComfyUI_NotaGen"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "mw"
DisplayName = "ComfyUI_NotaGen"
Icon = "https://github.com/billwuhao/aiart.website/blob/master/hb.png"
+7
View File
@@ -0,0 +1,7 @@
# transformers==4.40.0
# numpy==1.26.4
# gradio==5.17.1
wandb>=0.17.2
abctoolkit>=0.0.6
samplings>=0.1.7
pyparsing>=3.2.1
+443
View File
@@ -0,0 +1,443 @@
import torch
import random
import bisect
# import json
import re
from .config import nota_lx, nota_small, nota_medium
from transformers import GPT2Model, GPT2LMHeadModel, PreTrainedModel
from samplings import top_p_sampling, top_k_sampling, temperature_sampling
# from tokenizers import Tokenizer
class Patchilizer:
def __init__(self, model):
if model == "notagenx.pth" or model == "notagen_large.pth":
cf = nota_lx
elif model == "notagen_small.pth":
cf = nota_small
elif model == "notagen_medium.pth":
cf = nota_medium
self.stream = cf["PATCH_STREAM"]
self.patch_size = cf["PATCH_SIZE"]
self.patch_length = cf["PATCH_LENGTH"]
# self.char_num_layers = cf["CHAR_NUM_LAYERS"]
# self.patch_num_layers = cf["PATCH_NUM_LAYERS"]
# self.hidden_size = cf["HIDDEN_SIZE"]
# self.psbs = cf["PATCH_SAMPLING_BATCH_SIZE"]
self.delimiters = ["|:", "::", ":|", "[|", "||", "|]", "|"]
self.regexPattern = '(' + '|'.join(map(re.escape, self.delimiters)) + ')'
self.bos_token_id = 1
self.eos_token_id = 2
self.special_token_id = 0
def split_bars(self, body_lines):
"""
Split a body of music into individual bars.
"""
new_bars = []
try:
for line in body_lines:
line_bars = re.split(self.regexPattern, line)
line_bars = list(filter(None, line_bars))
new_line_bars = []
if len(line_bars) == 1:
new_line_bars = line_bars
else:
if line_bars[0] in self.delimiters:
new_line_bars = [line_bars[i] + line_bars[i + 1] for i in range(0, len(line_bars), 2)]
else:
new_line_bars = [line_bars[0]] + [line_bars[i] + line_bars[i + 1] for i in range(1, len(line_bars), 2)]
if 'V' not in new_line_bars[-1]:
new_line_bars[-2] += new_line_bars[-1] # 吸收最后一个 小节线+\n 的组合
new_line_bars = new_line_bars[:-1]
new_bars += new_line_bars
except:
pass
return new_bars
def split_patches(self, abc_text, generate_last=False):
if not generate_last and len(abc_text) % self.patch_size != 0:
abc_text += chr(self.eos_token_id)
patches = [abc_text[i : i + self.patch_size] for i in range(0, len(abc_text), self.patch_size)]
return patches
def patch2chars(self, patch):
"""
Convert a patch into a bar.
"""
bytes = ''
for idx in patch:
if idx == self.eos_token_id:
break
if idx < self.eos_token_id:
pass
bytes += chr(idx)
return bytes
def patchilize_metadata(self, metadata_lines):
metadata_patches = []
for line in metadata_lines:
metadata_patches += self.split_patches(line)
return metadata_patches
def patchilize_tunebody(self, tunebody_lines, encode_mode='train'):
tunebody_patches = []
bars = self.split_bars(tunebody_lines)
if encode_mode == 'train':
for bar in bars:
tunebody_patches += self.split_patches(bar)
elif encode_mode == 'generate':
for bar in bars[:-1]:
tunebody_patches += self.split_patches(bar)
tunebody_patches += self.split_patches(bars[-1], generate_last=True)
return tunebody_patches
def encode_train(self, abc_text, add_special_patches=True, cut=True):
lines = abc_text.split('\n')
lines = list(filter(None, lines))
lines = [line + '\n' for line in lines]
tunebody_index = -1
for i, line in enumerate(lines):
if '[V:' in line:
tunebody_index = i
break
metadata_lines = lines[ : tunebody_index]
tunebody_lines = lines[tunebody_index : ]
if self.stream:
tunebody_lines = ['[r:' + str(line_index) + '/' + str(len(tunebody_lines) - line_index - 1) + ']' + line for line_index, line in
enumerate(tunebody_lines)]
metadata_patches = self.patchilize_metadata(metadata_lines)
tunebody_patches = self.patchilize_tunebody(tunebody_lines, encode_mode='train')
if add_special_patches:
bos_patch = chr(self.bos_token_id) * (self.patch_size - 1) + chr(self.eos_token_id)
eos_patch = chr(self.bos_token_id) + chr(self.eos_token_id) * (self.patch_size - 1)
metadata_patches = [bos_patch] + metadata_patches
tunebody_patches = tunebody_patches + [eos_patch]
if self.stream:
if len(metadata_patches) + len(tunebody_patches) > self.patch_length:
available_cut_indexes = [0] + [index + 1 for index, patch in enumerate(tunebody_patches) if '\n' in patch]
line_index_for_cut_index = list(range(len(available_cut_indexes)))
end_index = len(metadata_patches) + len(tunebody_patches) - self.patch_length
biggest_index = bisect.bisect_left(available_cut_indexes, end_index)
available_cut_indexes = available_cut_indexes[:biggest_index + 1]
if len(available_cut_indexes) == 1:
choices = ['head']
elif len(available_cut_indexes) == 2:
choices = ['head', 'tail']
else:
choices = ['head', 'tail', 'middle']
choice = random.choice(choices)
if choice == 'head':
patches = metadata_patches + tunebody_patches[0:]
else:
if choice == 'tail':
cut_index = len(available_cut_indexes) - 1
else:
cut_index = random.choice(range(1, len(available_cut_indexes) - 1))
line_index = line_index_for_cut_index[cut_index]
stream_tunebody_lines = tunebody_lines[line_index : ]
stream_tunebody_patches = self.patchilize_tunebody(stream_tunebody_lines, encode_mode='train')
if add_special_patches:
stream_tunebody_patches = stream_tunebody_patches + [eos_patch]
patches = metadata_patches + stream_tunebody_patches
else:
patches = metadata_patches + tunebody_patches
else:
patches = metadata_patches + tunebody_patches
if cut:
patches = patches[ : self.patch_length]
else:
pass
# encode to ids
id_patches = []
for patch in patches:
id_patch = [ord(c) for c in patch] + [self.special_token_id] * (self.patch_size - len(patch))
id_patches.append(id_patch)
return id_patches
def encode_generate(self, abc_code, add_special_patches=True):
lines = abc_code.split('\n')
lines = list(filter(None, lines))
tunebody_index = None
for i, line in enumerate(lines):
if line.startswith('[V:') or line.startswith('[r:'):
tunebody_index = i
break
metadata_lines = lines[ : tunebody_index]
tunebody_lines = lines[tunebody_index : ]
metadata_lines = [line + '\n' for line in metadata_lines]
if self.stream:
if not abc_code.endswith('\n'):
tunebody_lines = [tunebody_lines[i] + '\n' for i in range(len(tunebody_lines) - 1)] + [tunebody_lines[-1]]
else:
tunebody_lines = [tunebody_lines[i] + '\n' for i in range(len(tunebody_lines))]
else:
tunebody_lines = [line + '\n' for line in tunebody_lines]
metadata_patches = self.patchilize_metadata(metadata_lines)
tunebody_patches = self.patchilize_tunebody(tunebody_lines, encode_mode='generate')
if add_special_patches:
bos_patch = chr(self.bos_token_id) * (self.patch_size - 1) + chr(self.eos_token_id)
metadata_patches = [bos_patch] + metadata_patches
patches = metadata_patches + tunebody_patches
patches = patches[ : self.patch_length]
# encode to ids
id_patches = []
for patch in patches:
if len(patch) < self.patch_size and patch[-1] != chr(self.eos_token_id):
id_patch = [ord(c) for c in patch]
else:
id_patch = [ord(c) for c in patch] + [self.special_token_id] * (self.patch_size - len(patch))
id_patches.append(id_patch)
return id_patches
def decode(self, patches):
"""
Decode patches into music.
"""
return ''.join(self.patch2chars(patch) for patch in patches)
class PatchLevelDecoder(PreTrainedModel):
"""
A Patch-level Decoder model for generating patch features in an auto-regressive manner.
It inherits PreTrainedModel from transformers.
"""
def __init__(self, config, model):
if model == "notagenx.pth" or model == "notagen_large.pth":
cf = nota_lx
elif model == "notagen_small.pth":
cf = nota_small
elif model == "notagen_medium.pth":
cf = nota_medium
self.patch_size = cf["PATCH_SIZE"]
super().__init__(config)
self.patch_embedding = torch.nn.Linear(self.patch_size * 128, config.n_embd)
torch.nn.init.normal_(self.patch_embedding.weight, std=0.02)
self.base = GPT2Model(config)
def forward(self,
patches: torch.Tensor,
masks=None) -> torch.Tensor:
"""
The forward pass of the patch-level decoder model.
:param patches: the patches to be encoded
:param masks: the masks for the patches
:return: the encoded patches
"""
patches = torch.nn.functional.one_hot(patches, num_classes=128).to(self.dtype)
patches = patches.reshape(len(patches), -1, self.patch_size * (128))
patches = self.patch_embedding(patches.to(self.device))
if masks==None:
return self.base(inputs_embeds=patches)
else:
return self.base(inputs_embeds=patches,
attention_mask=masks)
class CharLevelDecoder(PreTrainedModel):
"""
A Char-level Decoder model for generating the chars within each patch in an auto-regressive manner
based on the encoded patch features. It inherits PreTrainedModel from transformers.
"""
def __init__(self, config, model):
super().__init__(config)
self.special_token_id = 0
self.bos_token_id = 1
self.base = GPT2LMHeadModel(config)
if model == "notagenx.pth" or model == "notagen_large.pth":
cf = nota_lx
elif model == "notagen_small.pth":
cf = nota_small
elif model == "notagen_medium.pth":
cf = nota_medium
self.psbs = cf["PATCH_SAMPLING_BATCH_SIZE"]
def forward(self,
encoded_patches: torch.Tensor,
target_patches: torch.Tensor):
"""
The forward pass of the char-level decoder model.
:param encoded_patches: the encoded patches
:param target_patches: the target patches
:return: the output of the model
"""
# preparing the labels for model training
target_patches = torch.cat((torch.ones_like(target_patches[:,0:1])*self.bos_token_id, target_patches), dim=1)
# print('target_patches shape:', target_patches.shape)
target_masks = target_patches == self.special_token_id
labels = target_patches.clone().masked_fill_(target_masks, -100)
# masking the labels for model training
target_masks = torch.ones_like(labels)
target_masks = target_masks.masked_fill_(labels == -100, 0)
# select patches
if self.psbs != 0 and self.psbs < target_patches.shape[0]:
indices = list(range(len(target_patches)))
random.shuffle(indices)
selected_indices = sorted(indices[:self.psbs])
target_patches = target_patches[selected_indices,:]
target_masks = target_masks[selected_indices,:]
encoded_patches = encoded_patches[selected_indices,:]
# get input embeddings
inputs_embeds = torch.nn.functional.embedding(target_patches, self.base.transformer.wte.weight)
# concatenate the encoded patches with the input embeddings
inputs_embeds = torch.cat((encoded_patches.unsqueeze(1), inputs_embeds[:,1:,:]), dim=1)
output = self.base(inputs_embeds=inputs_embeds,
attention_mask=target_masks,
labels=labels)
# output_hidden_states=True=True)
return output
def generate(self,
encoded_patch: torch.Tensor, # [hidden_size]
tokens: torch.Tensor): # [1]
"""
The generate function for generating a patch based on the encoded patch and already generated tokens.
:param encoded_patch: the encoded patch
:param tokens: already generated tokens in the patch
:return: the probability distribution of next token
"""
encoded_patch = encoded_patch.reshape(1, 1, -1) # [1, 1, hidden_size]
tokens = tokens.reshape(1, -1)
# Get input embeddings
tokens = torch.nn.functional.embedding(tokens, self.base.transformer.wte.weight)
# Concatenate the encoded patch with the input embeddings
tokens = torch.cat((encoded_patch, tokens[:,1:,:]), dim=1)
# Get output from model
outputs = self.base(inputs_embeds=tokens)
# Get probabilities of next token
probs = torch.nn.functional.softmax(outputs.logits.squeeze(0)[-1], dim=-1)
return probs
class NotaGenLMHeadModel(PreTrainedModel):
"""
NotaGen is a language model with a hierarchical structure.
It includes a patch-level decoder and a char-level decoder.
The patch-level decoder is used to generate patch features in an auto-regressive manner.
The char-level decoder is used to generate the chars within each patch in an auto-regressive manner.
It inherits PreTrainedModel from transformers.
"""
def __init__(self, encoder_config, decoder_config, model):
super().__init__(encoder_config)
self.special_token_id = 0
self.bos_token_id = 1
self.eos_token_id = 2
self.patch_level_decoder = PatchLevelDecoder(encoder_config, model)
self.char_level_decoder = CharLevelDecoder(decoder_config, model)
if model == "notagenx.pth" or model == "notagen_large.pth":
cf = nota_lx
elif model == "notagen_small.pth":
cf = nota_small
elif model == "notagen_medium.pth":
cf = nota_medium
self.patch_size = cf["PATCH_SIZE"]
def forward(self,
patches: torch.Tensor,
masks: torch.Tensor):
"""
The forward pass of the bGPT model.
:param patches: the patches to be encoded
:param masks: the masks for the patches
:return: the decoded patches
"""
patches = patches.reshape(len(patches), -1, self.patch_size)
encoded_patches = self.patch_level_decoder(patches, masks)["last_hidden_state"]
left_shift_masks = masks * (masks.flip(1).cumsum(1).flip(1) > 1)
masks[:, 0] = 0
encoded_patches = encoded_patches[left_shift_masks == 1]
patches = patches[masks == 1]
return self.char_level_decoder(encoded_patches, patches)
def generate(self,
patches: torch.Tensor,
top_k=0,
top_p=1,
temperature=1.0):
"""
The generate function for generating patches based on patches.
:param patches: the patches to be encoded
:param top_k: the top k for sampling
:param top_p: the top p for sampling
:param temperature: the temperature for sampling
:return: the generated patches
"""
if patches.shape[-1] % self.patch_size != 0:
tokens = patches[:,:,-(patches.shape[-1]%self.patch_size):].squeeze(0, 1)
tokens = torch.cat((torch.tensor([self.bos_token_id], device=self.device), tokens), dim=-1)
patches = patches[:,:,:-(patches.shape[-1]%self.patch_size)]
else:
tokens = torch.tensor([self.bos_token_id], device=self.device)
patches = patches.reshape(len(patches), -1, self.patch_size) # [bs, seq, patch_size]
encoded_patches = self.patch_level_decoder(patches)["last_hidden_state"] # [bs, seq, hidden_size]
generated_patch = []
while True:
prob = self.char_level_decoder.generate(encoded_patches[0][-1], tokens).cpu().detach().numpy() # [128]
prob = top_k_sampling(prob, top_k=top_k, return_probs=True) # [128]
prob = top_p_sampling(prob, top_p=top_p, return_probs=True) # [128]
token = temperature_sampling(prob, temperature=temperature) # int
char = chr(token)
generated_patch.append(token)
if len(tokens) >= self.patch_size:# or token == self.eos_token_id:
break
else:
tokens = torch.cat((tokens, torch.tensor([token], device=self.device)), dim=0)
return generated_patch