API to exchange workflows between multiple connected clients

Placeholder nodes to parametrize and run custom workflows from Krita
This commit is contained in:
Acly
2024-10-02 11:07:02 +02:00
parent 29e24ec52c
commit 5a45172d02
3 changed files with 151 additions and 36 deletions
+13 -1
View File
@@ -1,4 +1,4 @@
from . import api, nodes, tile, region, nsfw, translation
from . import api, nodes, tile, region, nsfw, translation, krita
NODE_CLASS_MAPPINGS = {
"ETN_LoadImageBase64": nodes.LoadImageBase64,
@@ -17,6 +17,12 @@ NODE_CLASS_MAPPINGS = {
"ETN_AttentionMask": region.AttentionMask,
"ETN_NSFWFilter": nsfw.NSFWFilter,
"ETN_Translate": translation.Translate,
"ETN_KritaOutput": krita.KritaOutput,
"ETN_KritaCanvas": krita.KritaCanvas,
"ETN_KritaSelection": krita.KritaSelection,
"ETN_KritaImageLayer": krita.KritaImageLayer,
"ETN_KritaMaskLayer": krita.KritaMaskLayer,
"ETN_IntParameter": krita.IntParameter,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_LoadImageBase64": "Load Image (Base64)",
@@ -35,5 +41,11 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_AttentionMask": "Regions Attention Mask",
"ETN_NSFWFilter": "NSFW Filter",
"ETN_Translate": "Translate Text",
"ETN_KritaOutput": "Krita Output",
"ETN_KritaCanvas": "Krita Canvas",
"ETN_KritaSelection": "Krita Selection",
"ETN_KritaImageLayer": "Krita Image Layer",
"ETN_KritaMaskLayer": "Krita Mask Layer",
"ETN_IntParameter": "Integer Parameter",
}
WEB_DIRECTORY = "./js"
+1 -35
View File
@@ -13,6 +13,7 @@ import folder_paths
import server
from .translation import available_languages, translate
from .krita import WorkflowExchange
input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
@@ -127,41 +128,6 @@ def has_invalid_filename(filename: str):
return None
class Publisher(NamedTuple):
name: str
id: str
workflow: dict
class WorkflowExchange:
def __init__(self, server: server.PromptServer):
self._server = server
self._publishers: dict[str, Publisher] = {}
self._subscribers: list[str] = []
async def publish(self, publisher_name: str, publisher_id: str, workflow: dict):
name = f"{publisher_name} ({publisher_id})"
publisher = Publisher(name, publisher_id, workflow)
for client_id in self._subscribers:
await self._notify(client_id, publisher)
self._publishers[publisher_id] = publisher
print(f"Published workflow from {name}: {workflow}")
async def subscribe(self, client_id: str):
if client_id in self._subscribers:
raise KeyError("Already subscribed")
self._subscribers.append(client_id)
for publisher in self._publishers.values():
await self._notify(client_id, publisher)
def unsubscribe(self, client_id: str):
self._subscribers.remove(client_id)
async def _notify(self, client_id: str, publisher: Publisher):
data = {"publisher": publisher.name, "workflow": publisher.workflow}
await self._server.send_json("etn_workflow_changed", data, client_id)
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
if _server is not None:
+137
View File
@@ -0,0 +1,137 @@
import torch
from typing import NamedTuple
import server
from .nodes import SendImageWebSocket
class Publisher(NamedTuple):
name: str
id: str
workflow: dict
class WorkflowExchange:
def __init__(self, server: server.PromptServer):
self._server = server
self._publishers: dict[str, Publisher] = {}
self._subscribers: list[str] = []
async def publish(self, publisher_name: str, publisher_id: str, workflow: dict):
publisher = Publisher(publisher_name, publisher_id, workflow)
for client_id in self._subscribers:
await self._notify(client_id, publisher)
self._publishers[publisher_id] = publisher
async def subscribe(self, client_id: str):
if client_id in self._subscribers:
raise KeyError("Already subscribed")
self._subscribers.append(client_id)
for publisher in self._publishers.values():
await self._notify(client_id, publisher)
def unsubscribe(self, client_id: str):
self._subscribers.remove(client_id)
async def _notify(self, client_id: str, publisher: Publisher):
data = {
"publisher": {"name": publisher.name, "id": publisher.id},
"workflow": publisher.workflow,
}
await self._server.send_json("etn_workflow_published", data, client_id)
class KritaOutput(SendImageWebSocket):
RETURN_TYPES = ()
FUNCTION = "send_images"
OUTPUT_NODE = True
CATEGORY = "krita"
class KritaCanvas:
@classmethod
def INPUT_TYPES(cls):
return {}
RETURN_TYPES = ("IMAGE", "INT", "INT", "INT")
RETURN_NAMES = ("image", "width", "height", "seed")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self):
empty = torch.zeroes(1, 512, 512, 3)
return (empty, 512, 512, 0)
class KritaSelection:
@classmethod
def INPUT_TYPES(cls):
return {}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self):
empty = torch.ones(1, 512, 512)
return (empty,)
class KritaImageLayer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Image"}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str):
empty = torch.zeros(1, 512, 512, 3)
return (empty,)
class KritaMaskLayer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Mask"}),
}
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str):
empty = torch.ones(1, 512, 512)
return (empty,)
class IntParameter:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Parameter"}),
"min": ("INT", {"default": 0}),
"max": ("INT", {"default": 100}),
"default": ("INT", {"default": 50}),
}
}
RETURN_TYPES = ("INT",)
RETURN_NAMES = ("value",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, min: int, max: int, default: int):
return (default,)