v1.0.0
This commit is contained in:
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
[中文](README.md) | [English](README-en.md)
|
||||
|
||||
# Symbolic Music Generation, NotaGen node for ComfyUI.
|
||||
|
||||

|
||||
|
||||
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`
|
||||
|
||||
@@ -1 +1,13 @@
|
||||
# ComfyUI_NotaGen
|
||||
[中文](README.md) | [English](README-en.md)
|
||||
|
||||
# 符号音乐生成. NotaGen 的 ComfyUI 节点.
|
||||
|
||||

|
||||
|
||||
将模型下载放到 `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`
|
||||
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from .NotaGenNode import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+2252
File diff suppressed because it is too large
Load Diff
@@ -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 |
@@ -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"
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user