diff --git a/pyproject.toml b/pyproject.toml index c60fcae..412808f 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.7" +version = "2.2.8" license = {file = "LICENSE"} [project.urls] diff --git a/utils.py b/utils.py index 27fd96d..49d462f 100644 --- a/utils.py +++ b/utils.py @@ -1,7 +1,7 @@ import torch import random import bisect -# import json +import numpy as np import re from .config import nota_lx, nota_small, nota_medium @@ -358,6 +358,19 @@ class CharLevelDecoder(PreTrainedModel): return probs +def safe_normalize_probs(probs): + epsilon = 1e-12 + probs = np.array(probs, dtype=np.float64) + probs = np.where(np.isnan(probs) | (probs < 0), 0, probs) + probs = probs + epsilon + s = probs.sum() + if s > 0: + probs = probs / s + else: + probs = np.zeros_like(probs) + probs[0] = 1.0 + return probs + class NotaGenLMHeadModel(PreTrainedModel): """ NotaGen is a language model with a hierarchical structure. @@ -429,8 +442,11 @@ class NotaGenLMHeadModel(PreTrainedModel): while True: prob = self.char_level_decoder.generate(encoded_patches[0][-1], tokens).cpu().detach().numpy() # [128] + prob = safe_normalize_probs(prob) prob = top_k_sampling(prob, top_k=top_k, return_probs=True) # [128] + prob = safe_normalize_probs(prob) prob = top_p_sampling(prob, top_p=top_p, return_probs=True) # [128] + prob = safe_normalize_probs(prob) token = temperature_sampling(prob, temperature=temperature) # int char = chr(token) generated_patch.append(token) @@ -440,4 +456,4 @@ class NotaGenLMHeadModel(PreTrainedModel): else: tokens = torch.cat((tokens, torch.tensor([token], device=self.device)), dim=0) - return generated_patch + return generated_patch \ No newline at end of file