Add lora with metadata and health endpoints & IPA fix instead of explicit flag

This commit is contained in:
ramyma
2024-12-09 18:06:18 +02:00
parent def82a1fe1
commit 8d617cea60
4 changed files with 83 additions and 17 deletions
+2 -2
View File
@@ -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"]
+9 -15
View File
@@ -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):
+1
View File
@@ -0,0 +1 @@
from .api import *
+71
View File
@@ -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("<Q", n_bytes)[
0
] # '<Q' means little-endian unsigned long long
# Read n bytes and decode JSON
metadata_bytes = io_device.read(n)
io_device.close()
metadata = json.loads(metadata_bytes.decode("utf-8"))
# Retrieve the value associated with "__metadata__"
return metadata.get("__metadata__")