API to exchange workflows between multiple connected clients
Placeholder nodes to parametrize and run custom workflows from Krita
This commit is contained in:
+13
-1
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
|
||||
@@ -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,)
|
||||
Reference in New Issue
Block a user