diff --git a/__init__.py b/__init__.py index 4d4c4f6..df1b688 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,4 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .server import * - -__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/attention_couple.py b/attention_couple.py index bcdb252..94773dc 100644 --- a/attention_couple.py +++ b/attention_couple.py @@ -152,10 +152,6 @@ class AttentionCouple: ), "width": ("INT", {"default": 1024, "min": 8, "step": 8}), "height": ("INT", {"default": 1024, "min": 8, "step": 8}), - "ip_adapter_active": ( - "BOOLEAN", - {"default": False, "tooltip": "Set to true if using IPA."}, - ), }, } @@ -173,27 +169,21 @@ class AttentionCouple: height, width, regions, - ip_adapter_active, **kwargs, ): base_mask = torch.zeros((height, width)).unsqueeze(0) global_mask = (torch.ones((height, width)) * global_prompt_weight).unsqueeze(0) - ip_even_mask = (torch.ones((height, width)) * 0.1).unsqueeze(0) new_model = model.clone() if not isinstance(regions, list): regions = [regions] - num_conds = ( - len(regions) + 1 + (1 if ip_adapter_active and len(regions) % 2 == 0 else 0) - ) + num_conds = len(regions) + 1 mask = [base_mask] + [ global_mask if i == 0 - else ip_even_mask - if ip_adapter_active and len(regions) % 2 == 0 and i == num_conds - 1 else F.interpolate( regions[i - 1]["mask"].unsqueeze(0), size=(height, width), @@ -207,10 +197,7 @@ class AttentionCouple: self.mask = mask / mask.sum(dim=0, keepdim=True) self.conds = [ - base_prompt[0][0] - if i == 0 - or (ip_adapter_active and len(regions) % 2 == 0 and i == num_conds - 1) - else regions[i - 1]["cond"][0][0] + base_prompt[0][0] if i == 0 else regions[i - 1]["cond"][0][0] for i in range(0, num_conds) ] num_tokens = [cond.shape[1] for cond in self.conds] @@ -253,6 +240,13 @@ class AttentionCouple: qs = torch.cat(qs, dim=0) ks = torch.cat(ks, dim=0).to(k) + if qs.size(0) % 2 == 1: + empty = torch.zeros_like(qs[0]).unsqueeze(0) + qs = torch.cat((qs, empty), dim=0) + + empty2 = torch.zeros_like(ks[0]).unsqueeze(0) + ks = torch.cat((ks, empty2), dim=0) + return qs, ks, ks def attn2_output_patch(out, extra_options): diff --git a/server/__init__.py b/server/__init__.py new file mode 100644 index 0000000..0a0e47b --- /dev/null +++ b/server/__init__.py @@ -0,0 +1 @@ +from .api import * diff --git a/server/api.py b/server/api.py new file mode 100644 index 0000000..1a3742a --- /dev/null +++ b/server/api.py @@ -0,0 +1,71 @@ +from server import PromptServer +from aiohttp import web +import folder_paths +import os +import logging + +import struct +import json +import pathlib + + +@PromptServer.instance.routes.get("/a8r8/loras") +async def loras(request): + try: + loras = [ + { + "path": lora, + "name": pathlib.Path(lora).stem, + "metadata": get_lora_metadata(lora), + } + for lora in folder_paths.get_filename_list("loras") + ] + + return web.json_response(loras) + except Exception as e: + logging.error(e) + return web.Response(status=400, text=str(e)) + + +@PromptServer.instance.routes.get("/a8r8/health") +async def health(_request): + try: + return web.json_response( + {}, + status=200, + ) + except Exception as e: + logging.error(e) + return web.Response(status=400, text=str(e)) + + +def get_lora_metadata(lora): + try: + if ".safetensors" in lora: + base_path = folder_paths.folder_names_and_paths["loras"][0][0] + + lora_path = os.path.join(base_path, lora) + + return read_metadata(lora_path) + return None + except FileNotFoundError as _e: + return None + + +def read_metadata(file_path): + # Open the file in binary mode + with open(file_path, "rb") as io_device: + # Read 8 bytes and unpack as little-endian unsigned integer + n_bytes = io_device.read(8) + + n = struct.unpack("