diff --git a/NotaGenNode.py b/NotaGenNode.py index c2ae726..03996f2 100644 --- a/NotaGenNode.py +++ b/NotaGenNode.py @@ -38,44 +38,56 @@ class NotaGenRun: device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") nota_model_path = nota_model_path node_dir = node_dir + + model_cache = None @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", }), - "custom_prompt": ("STRING", {"default": "Romantic | Bach, Johann Sebastian | Keyboard", - "multiline": True, - "tooltip": "Custom prompt must ||."}), - # "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}), + "period": (s.periods, {"default": "Romantic"}), + "composer": (s.composers, {"default": "Bach, Johann Sebastian"}), + "instrumentation": (s.instrumentations, {"default": "Keyboard"}), + "custom_prompt": ("STRING", { + "default": "Romantic | Bach, Johann Sebastian | Keyboard", + "multiline": True, + "tooltip": "Custom prompt must follow format: ||" + }), + "unload_model":("BOOLEAN", {"default": False}), + "temperature": ("FLOAT", {"default": 0.8, "min": 0, "max": 5, "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}), - "comfy_python_path": ("STRING", {"default": "", "multiline": False, "tooltip": "The absolute path of python.exe in the Comfyui environment"}), - # "audio_sheet_music": ("BOOLEAN", {"default": True}), - "musescore4_path": ("STRING", {"default": "", "tooltip": r"The absolute path as `D:\APP\MuseScorePortable\App\MuseScore\bin\MuseScore4.exe`"}), - # "abc2xml": ("BOOLEAN", {"default": True}), + "comfy_python_path": ("STRING", { + "default": "", + "multiline": False, + "tooltip": "Absolute path of python.exe in ComfyUI environment" + }), + "musescore4_path": ("STRING", { + "default": "", + "tooltip": r"Absolute path e.g. D:\APP\MuseScorePortable\App\MuseScore\bin\MuseScore4.exe" + }), }, } RETURN_TYPES = ("AUDIO", "IMAGE", "STRING") RETURN_NAMES = ("audio", "score", "message") FUNCTION = "inference_patch" - CATEGORY = "MW-NotaGen" + CATEGORY = "MW/MW-NotaGen" - # Note_list = Note_list + ['z', 'x'] def inference_patch(self, model, period, composer, instrumentation, custom_prompt, - # num_samples, - # abc2xml, comfy_python_path, - # audio_sheet_music, musescore4_path, + unload_model, + temperature, + top_k, + top_p, seed): + if seed != 0: + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + if model == "notagenx.pth" or model == "notagen_large.pth": cf = nota_lx elif model == "notagen_small.pth": @@ -106,7 +118,11 @@ 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) - checkpoint = torch.load(nota_model_path, map_location=torch.device(self.device)) + if self.model_cache is None: + checkpoint = torch.load(nota_model_path, map_location=torch.device(self.device)) + self.model_cache = checkpoint + else: + checkpoint = self.model_cache nota_model.load_state_dict(checkpoint['model']) nota_model = nota_model.to(self.device) nota_model.eval() @@ -148,17 +164,17 @@ class NotaGenRun: tunebody_flag = False while True: predicted_patch = nota_model.generate(input_patches.unsqueeze(0), - top_k=9, - top_p=0.9, - temperature=1.2) + top_k=top_k, + top_p=top_p, + temperature=temperature) 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) + 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] == patchilizer.bos_token_id and predicted_patch[1] == patchilizer.eos_token_id: end_flag = True @@ -319,6 +335,13 @@ class NotaGenRun: else: image1 = self.get_empty_image() + if unload_model: + del patchilizer + del nota_model + del checkpoint + torch.cuda.empty_cache() + self.model_cache = None + return ( audio, image1, @@ -326,16 +349,22 @@ class NotaGenRun: ) else: + if unload_model: + del patchilizer + del nota_model + del checkpoint + torch.cuda.empty_cache() + self.model_cache = None print(f".abc and .xml was saved to {INTERLEAVED_OUTPUT_FOLDER} and {ORIGINAL_OUTPUT_FOLDER}") raise Exception("Conversion of .mp3 and .png failed, try again or check if MuseScore4 installation was successful.") def get_empty_audio(self): - """返回空音频""" + """Return empty audio""" return {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100} def get_empty_image(self): - """返回空图片""" + """Return empty image""" import numpy as np return torch.from_numpy(np.zeros((1, 64, 64, 3), dtype=np.float32)) @@ -429,7 +458,7 @@ class NotaGenRun: return unreduced_lines def wait_for_file(self, file_path, timeout=15, check_interval=0.3): - """等待文件生成完成""" + """Wait for file generation to complete""" start_time = time.time() while time.time() - start_time < timeout: @@ -446,7 +475,7 @@ class NotaGenRun: return False def wait_for_png_sequence(self, base_path, timeout=15, check_interval=0.3): - """等待PNG序列生成完成""" + """Wait for PNG sequence generation to complete""" import glob start_time = time.time() diff --git a/README-CN.md b/README-CN.md index 53b4b35..2d0a0c4 100644 --- a/README-CN.md +++ b/README-CN.md @@ -9,6 +9,8 @@ https://github.com/user-attachments/assets/0671657f-e66b-4000-a0aa-48520f15b782 ## 📣 更新 +[2025-03-21]⚒️: 增加更多可调参数, 更自由畅玩. 可选是否卸载模型. + [2025-03-15]⚒️: 支持 Linux Ubuntu/Debian 系列, 以及服务器, 其他未测试. 本地 Linux 电脑, 安装 `musescore` 等: diff --git a/README.md b/README.md index d9a1da0..18c4c23 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,8 @@ https://github.com/user-attachments/assets/0671657f-e66b-4000-a0aa-48520f15b782 ## 📣 Updates +[2025-03-21] ⚒️: Added more tunable parameters for more creative freedom. Optional model unloading. + [2025-03-15]⚒️: Supports Linux Ubuntu/Debian series, as well as servers, others untested, as well as servers. For local Linux computers, install `musescore` etc.: diff --git a/images/2025-03-10_06-24-03.png b/images/2025-03-10_06-24-03.png index aa0cac7..331ec0d 100644 Binary files a/images/2025-03-10_06-24-03.png and b/images/2025-03-10_06-24-03.png differ diff --git a/pyproject.toml b/pyproject.toml index eef1666..915614b 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.1.2" +version = "2.2.1" license = {file = "LICENSE"} [project.urls] diff --git a/requirements.txt b/requirements.txt index d0c27e0..df11b83 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,3 +2,4 @@ wandb>=0.17.2 abctoolkit>=0.0.6 samplings>=0.1.7 pyparsing>=3.2.1 +transformers>=4.40.0