Add lora with metadata and health endpoints & IPA fix instead of explicit flag
This commit is contained in:
+2
-2
@@ -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
@@ -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):
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .api import *
|
||||
@@ -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__")
|
||||
Reference in New Issue
Block a user