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 .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}),
|
"width": ("INT", {"default": 1024, "min": 8, "step": 8}),
|
||||||
"height": ("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,
|
height,
|
||||||
width,
|
width,
|
||||||
regions,
|
regions,
|
||||||
ip_adapter_active,
|
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
base_mask = torch.zeros((height, width)).unsqueeze(0)
|
base_mask = torch.zeros((height, width)).unsqueeze(0)
|
||||||
global_mask = (torch.ones((height, width)) * global_prompt_weight).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()
|
new_model = model.clone()
|
||||||
|
|
||||||
if not isinstance(regions, list):
|
if not isinstance(regions, list):
|
||||||
regions = [regions]
|
regions = [regions]
|
||||||
|
|
||||||
num_conds = (
|
num_conds = len(regions) + 1
|
||||||
len(regions) + 1 + (1 if ip_adapter_active and len(regions) % 2 == 0 else 0)
|
|
||||||
)
|
|
||||||
|
|
||||||
mask = [base_mask] + [
|
mask = [base_mask] + [
|
||||||
global_mask
|
global_mask
|
||||||
if i == 0
|
if i == 0
|
||||||
else ip_even_mask
|
|
||||||
if ip_adapter_active and len(regions) % 2 == 0 and i == num_conds - 1
|
|
||||||
else F.interpolate(
|
else F.interpolate(
|
||||||
regions[i - 1]["mask"].unsqueeze(0),
|
regions[i - 1]["mask"].unsqueeze(0),
|
||||||
size=(height, width),
|
size=(height, width),
|
||||||
@@ -207,10 +197,7 @@ class AttentionCouple:
|
|||||||
self.mask = mask / mask.sum(dim=0, keepdim=True)
|
self.mask = mask / mask.sum(dim=0, keepdim=True)
|
||||||
|
|
||||||
self.conds = [
|
self.conds = [
|
||||||
base_prompt[0][0]
|
base_prompt[0][0] if i == 0 else regions[i - 1]["cond"][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]
|
|
||||||
for i in range(0, num_conds)
|
for i in range(0, num_conds)
|
||||||
]
|
]
|
||||||
num_tokens = [cond.shape[1] for cond in self.conds]
|
num_tokens = [cond.shape[1] for cond in self.conds]
|
||||||
@@ -253,6 +240,13 @@ class AttentionCouple:
|
|||||||
qs = torch.cat(qs, dim=0)
|
qs = torch.cat(qs, dim=0)
|
||||||
ks = torch.cat(ks, dim=0).to(k)
|
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
|
return qs, ks, ks
|
||||||
|
|
||||||
def attn2_output_patch(out, extra_options):
|
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