diff --git a/Sana/nodes.py b/Sana/nodes.py index 8291c55..abce9ba 100644 --- a/Sana/nodes.py +++ b/Sana/nodes.py @@ -1,5 +1,6 @@ import torch import folder_paths +from nodes import EmptyLatentImage from .conf import sana_conf, sana_res from .loader import load_sana @@ -36,6 +37,15 @@ class SanaCheckpointLoader: return (model,) +class EmptySanaLatentImage(EmptyLatentImage): + CATEGORY = "ExtraModels/Sana" + TITLE = "Empty Sana Latent Image" + + def generate(self, width, height, batch_size=1): + latent = torch.zeros([batch_size, 32, height // 32, width // 32], device=self.device) + return ({"samples":latent}, ) + + class SanaResolutionSelect(): @classmethod def INPUT_TYPES(s): @@ -138,11 +148,5 @@ NODE_CLASS_MAPPINGS = { "SanaResolutionSelect" : SanaResolutionSelect, "SanaTextEncode" : SanaTextEncode, "SanaResolutionCond" : SanaResolutionCond, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "Sana Checkpoint Loader": "SanaCheckpointLoader", - "Sana Resolution Select": "SanaResolutionSelect", - "Sana Text Encoder": "SanaTextEncode", - "Sana Resolution Cond": "SanaResolutionCond", + "EmptySanaLatentImage": EmptySanaLatentImage, }