From d2bec782fffb003aba2a8df7bd04ca5cae41ed8e Mon Sep 17 00:00:00 2001 From: billwuhao Date: Fri, 23 May 2025 17:43:20 +0800 Subject: [PATCH] Global cache, multi node sharing --- NotaGenNode.py | 60 ++++++++++++++++++++++++-------------------------- pyproject.toml | 2 +- 2 files changed, 30 insertions(+), 32 deletions(-) diff --git a/NotaGenNode.py b/NotaGenNode.py index c44d645..4741e4b 100644 --- a/NotaGenNode.py +++ b/NotaGenNode.py @@ -22,6 +22,8 @@ os.makedirs(ORIGINAL_OUTPUT_FOLDER, exist_ok=True) os.makedirs(INTERLEAVED_OUTPUT_FOLDER, exist_ok=True) +MODEL_CACHE = None +PATCHILIZER = None class NotaGenRun: model_names = ["notagenx.pth", "notagen_small.pth", "notagen_medium.pth", "notagen_large.pth"] periods = ["Baroque", "Classical", "Romantic"] @@ -38,11 +40,8 @@ class NotaGenRun: def __init__(self): self.device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") - self.nota_model_path = nota_model_path + nota_model_path = nota_model_path self.node_dir = node_dir - self.model_cache = None - self.nota_model = None - self.patchilizer = None @classmethod def INPUT_TYPES(s): @@ -110,15 +109,14 @@ class NotaGenRun: 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) - if self.model_cache is None: - self.model_cache = torch.load(nota_model_path, map_location=torch.device(self.device)) + nota_model_path = os.path.join(nota_model_path, model) - if self.nota_model is None: - self.nota_model = nota_model - self.nota_model.load_state_dict(self.model_cache['model']) - self.nota_model = self.nota_model.to(self.device) - self.nota_model.eval() + global MODEL_CACHE, PATCHILIZER + if MODEL_CACHE is None: + MODEL_CACHE = torch.load(nota_model_path, map_location=torch.device(self.device)) + nota_model.load_state_dict(MODEL_CACHE['model']) + nota_model = nota_model.to(self.device) + nota_model.eval() if custom_prompt.strip(): period, composer, instrumentation = [i.strip() for i in custom_prompt.split('|')] @@ -128,11 +126,11 @@ class NotaGenRun: '%' + composer + '\n', '%' + instrumentation + '\n'] - if self.patchilizer is None: - self.patchilizer = Patchilizer(model) + if PATCHILIZER is None: + PATCHILIZER = Patchilizer(model) # file_no = 1 - bos_patch = [self.patchilizer.bos_token_id] * (patch_size - 1) + [self.patchilizer.eos_token_id] + bos_patch = [PATCHILIZER.bos_token_id] * (patch_size - 1) + [PATCHILIZER.eos_token_id] num_gen = 0 unreduced_xml_path = None save_xml_original = False @@ -141,11 +139,11 @@ class NotaGenRun: start_time = time.time() # start_time_format = time.strftime("%Y%m%d-%H%M%S") - prompt_patches = self.patchilizer.patchilize_metadata(prompt_lines) + 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] + [self.patchilizer.special_token_id] * (patch_size - len(patch)) for patch + 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) @@ -157,23 +155,23 @@ class NotaGenRun: tunebody_flag = False while True: - predicted_patch = self.nota_model.generate(input_patches.unsqueeze(0), + predicted_patch = nota_model.generate(input_patches.unsqueeze(0), top_k=top_k, top_p=top_p, temperature=temperature) - if not tunebody_flag and self.patchilizer.decode([predicted_patch]).startswith('[r:'): # start with [r:0/ + 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 = self.nota_model.generate(temp_input_patches.unsqueeze(0), + predicted_patch = nota_model.generate(temp_input_patches.unsqueeze(0), top_k=top_k, top_p=top_p, temperature=temperature) predicted_patch = [ord(c) for c in '[r:0/'] + predicted_patch - if predicted_patch[0] == self.patchilizer.bos_token_id and predicted_patch[1] == self.patchilizer.eos_token_id: + if predicted_patch[0] == PATCHILIZER.bos_token_id and predicted_patch[1] == PATCHILIZER.eos_token_id: end_flag = True break - next_patch = self.patchilizer.decode([predicted_patch]) + next_patch = PATCHILIZER.decode([predicted_patch]) for char in next_patch: byte_list.append(char) @@ -182,8 +180,8 @@ class NotaGenRun: patch_end_flag = False for j in range(len(predicted_patch)): if patch_end_flag: - predicted_patch[j] = self.patchilizer.special_token_id - if predicted_patch[j] == self.patchilizer.eos_token_id: + 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) @@ -223,7 +221,7 @@ class NotaGenRun: cut_index = len(tunebody_lines) // 2 abc_code_slice = ''.join(metadata_lines + tunebody_lines[-cut_index:]) - input_patches = self.patchilizer.encode_generate(abc_code_slice) + 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) @@ -332,9 +330,9 @@ class NotaGenRun: if unload_model: import gc - self.patchilizer = None - self.nota_model = None - self.model_cache = None + PATCHILIZER = None + nota_model = None + MODEL_CACHE = None gc.collect() torch.cuda.empty_cache() @@ -347,9 +345,9 @@ class NotaGenRun: else: if unload_model: import gc - self.patchilizer = None - self.nota_model = None - self.model_cache = None + PATCHILIZER = None + nota_model = None + MODEL_CACHE = None gc.collect() torch.cuda.empty_cache() diff --git a/pyproject.toml b/pyproject.toml index f99aa70..9fa4e79 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "notagen-mw" description = "Symbolic Music Generation, NotaGen node for ComfyUI." -version = "2.2.5" +version = "2.2.6" license = {file = "LICENSE"} [project.urls]