Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
574d96369c |
@@ -48,6 +48,7 @@ class WorkflowPlanner:
|
||||
for idx, data, is_special_function in load_order:
|
||||
inputs, class_type = data["inputs"], data["class_type"]
|
||||
input_types = self.node_class_mappings[class_type].INPUT_TYPES()
|
||||
input_value_types = self.get_input_value_types(input_types)
|
||||
class_def = self.node_class_mappings[class_type]()
|
||||
|
||||
missing_required_variable = False
|
||||
@@ -106,7 +107,7 @@ class WorkflowPlanner:
|
||||
)
|
||||
inputs = self.update_inputs(inputs, executed_variables)
|
||||
seed_sync_code = self.create_prompt_seed_sync_code(
|
||||
idx, inputs, is_special_function
|
||||
idx, inputs, input_value_types, is_special_function
|
||||
)
|
||||
|
||||
target_lines = special_functions_code if is_special_function else code
|
||||
@@ -118,6 +119,7 @@ class WorkflowPlanner:
|
||||
class_def.FUNCTION,
|
||||
executed_variables[idx],
|
||||
is_special_function,
|
||||
input_value_types=input_value_types,
|
||||
**inputs,
|
||||
)
|
||||
)
|
||||
@@ -138,16 +140,24 @@ class WorkflowPlanner:
|
||||
func: str,
|
||||
variable_name: str,
|
||||
is_special_function: bool,
|
||||
input_value_types: dict[str, str] | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
args = ", ".join(self.format_arg(key, value) for key, value in kwargs.items())
|
||||
args = ", ".join(
|
||||
self.format_arg(key, value, (input_value_types or {}).get(key))
|
||||
for key, value in kwargs.items()
|
||||
)
|
||||
code = f"{variable_name} = {obj_name}.{func}({args})\n"
|
||||
if not is_special_function:
|
||||
code = f"\t{code}"
|
||||
return code
|
||||
|
||||
def create_prompt_seed_sync_code(
|
||||
self, node_id: str, inputs: dict, is_special_function: bool
|
||||
self,
|
||||
node_id: str,
|
||||
inputs: dict,
|
||||
input_value_types: dict[str, str],
|
||||
is_special_function: bool,
|
||||
) -> list[str]:
|
||||
seed_sync_lines = []
|
||||
for key in ("seed", "noise_seed"):
|
||||
@@ -156,8 +166,11 @@ class WorkflowPlanner:
|
||||
randomized_seed_variable = (
|
||||
f"node_{self.sanitize_node_id(str(node_id))}_{self.clean_variable_name(key)}"
|
||||
)
|
||||
randomized_seed_code = self.get_randomized_seed_code(
|
||||
input_value_types.get(key)
|
||||
)
|
||||
seed_sync_lines.append(
|
||||
f'{randomized_seed_variable} = prompt["{node_id}"]["inputs"]["{key}"] = random.randint(1, 2**64)'
|
||||
f'{randomized_seed_variable} = prompt["{node_id}"]["inputs"]["{key}"] = {randomized_seed_code}'
|
||||
)
|
||||
inputs[key] = {"variable_name": randomized_seed_variable}
|
||||
|
||||
@@ -167,22 +180,42 @@ class WorkflowPlanner:
|
||||
indentation = "" if is_special_function else "\t"
|
||||
return [f"{indentation}{line}\n" for line in seed_sync_lines]
|
||||
|
||||
def format_arg(self, key: str, value: Any) -> str:
|
||||
value_code = self.format_arg_value(key, value)
|
||||
def format_arg(self, key: str, value: Any, input_value_type: str | None = None) -> str:
|
||||
value_code = self.format_arg_value(key, value, input_value_type)
|
||||
if key.isidentifier() and not keyword.iskeyword(key):
|
||||
return f"{key}={value_code}"
|
||||
return f"**{{{json.dumps(key)}: {value_code}}}"
|
||||
|
||||
@staticmethod
|
||||
def format_arg_value(key: str, value: Any) -> str:
|
||||
def format_arg_value(
|
||||
key: str, value: Any, input_value_type: str | None = None
|
||||
) -> str:
|
||||
if isinstance(value, dict) and "variable_name" in value:
|
||||
return value["variable_name"]
|
||||
if key == "noise_seed" or key == "seed":
|
||||
return "random.randint(1, 2**64)"
|
||||
return WorkflowPlanner.get_randomized_seed_code(input_value_type)
|
||||
if isinstance(value, str):
|
||||
return json.dumps(value)
|
||||
return repr(value)
|
||||
|
||||
@staticmethod
|
||||
def get_input_value_types(input_types: dict) -> dict[str, str]:
|
||||
value_types = {}
|
||||
for section in ("required", "optional", "hidden"):
|
||||
for key, value in input_types.get(section, {}).items():
|
||||
if isinstance(value, tuple) and value:
|
||||
value_types[key] = value[0]
|
||||
elif isinstance(value, str):
|
||||
value_types[key] = value
|
||||
return value_types
|
||||
|
||||
@staticmethod
|
||||
def get_randomized_seed_code(input_value_type: str | None) -> str:
|
||||
randomized_seed_code = "random.randint(1, 2**64)"
|
||||
if input_value_type == "STRING":
|
||||
return f"str({randomized_seed_code})"
|
||||
return randomized_seed_code
|
||||
|
||||
def get_class_info(self, class_type: str) -> tuple[str, tuple[str, str], str]:
|
||||
class_obj = self.base_node_class_mappings.get(class_type)
|
||||
module_name = "nodes"
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "StringSeedNode",
|
||||
"inputs": {
|
||||
"seed": "seed-placeholder"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
{
|
||||
"33": {
|
||||
"class_type": "VaeDecode",
|
||||
"inputs": {
|
||||
"samples": "latent-placeholder"
|
||||
}
|
||||
},
|
||||
"42:0": {
|
||||
"class_type": "UpscaleModelLoader",
|
||||
"inputs": {
|
||||
"model_name": "4x-ultrasharp.safetensors"
|
||||
}
|
||||
},
|
||||
"42:1": {
|
||||
"class_type": "ImageUpscaleWithModel",
|
||||
"inputs": {
|
||||
"upscale_model": [
|
||||
"42:0",
|
||||
0
|
||||
],
|
||||
"image": [
|
||||
"33",
|
||||
0
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "TextConcatenateNode",
|
||||
"inputs": {
|
||||
"delimiter": "",
|
||||
"clean_whitespace": "true",
|
||||
"text_b": "\\"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
{
|
||||
"11": {
|
||||
"class_type": "DualClipLoader",
|
||||
"inputs": {
|
||||
"clip_name": "clip.safetensors"
|
||||
}
|
||||
},
|
||||
"633": {
|
||||
"class_type": "AnySwitchRgthree",
|
||||
"inputs": {
|
||||
"model": "model-placeholder"
|
||||
}
|
||||
},
|
||||
"631": {
|
||||
"class_type": "PowerLoraLoaderRgthree",
|
||||
"inputs": {
|
||||
"PowerLoraLoaderHeaderWidget": {
|
||||
"type": "PowerLoraLoaderHeaderWidget"
|
||||
},
|
||||
"lora_1": {
|
||||
"on": false,
|
||||
"lora": "lora.safetensors",
|
||||
"strength": 1.2
|
||||
},
|
||||
"lora_2": {
|
||||
"on": false,
|
||||
"lora": "lora2.safetensors",
|
||||
"strength": 0.7
|
||||
},
|
||||
"\u2795 Add Lora": "",
|
||||
"model": [
|
||||
"633",
|
||||
0
|
||||
],
|
||||
"clip": [
|
||||
"11",
|
||||
0
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "WindowsPathNode",
|
||||
"inputs": {
|
||||
"path": "C:\\ComfyUI\\models\\upscale_models\\RealESRGAN_x4plus.safetensors"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,312 @@
|
||||
import json
|
||||
import unittest
|
||||
from io import StringIO
|
||||
from pathlib import Path
|
||||
|
||||
from comfyui_to_python import ComfyUItoPython
|
||||
|
||||
|
||||
FIXTURE_DIR = Path(__file__).parent / "fixtures" / "unit" / "generator_codegen"
|
||||
|
||||
|
||||
class AnySwitchRgthree:
|
||||
CATEGORY = "utils"
|
||||
FUNCTION = "switch"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
}
|
||||
}
|
||||
|
||||
def switch(self, model):
|
||||
return (model,)
|
||||
|
||||
|
||||
class DualClipLoader:
|
||||
CATEGORY = "loaders"
|
||||
FUNCTION = "load_clip"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"clip_name": ("STRING",),
|
||||
}
|
||||
}
|
||||
|
||||
def load_clip(self, clip_name):
|
||||
return (clip_name,)
|
||||
|
||||
|
||||
class PowerLoraLoaderRgthree:
|
||||
CATEGORY = "loaders"
|
||||
FUNCTION = "load_loras"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"PowerLoraLoaderHeaderWidget": ("DICT",),
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP",),
|
||||
}
|
||||
}
|
||||
|
||||
def load_loras(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
|
||||
class UpscaleModelLoader:
|
||||
CATEGORY = "loaders"
|
||||
FUNCTION = "load_model"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model_name": ("STRING",),
|
||||
}
|
||||
}
|
||||
|
||||
def load_model(self, model_name):
|
||||
return (model_name,)
|
||||
|
||||
|
||||
class ImageUpscaleWithModel:
|
||||
CATEGORY = "image/upscaling"
|
||||
FUNCTION = "upscale"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"upscale_model": ("UPSCALE_MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
def upscale(self, upscale_model, image):
|
||||
return (image,)
|
||||
|
||||
|
||||
class VaeDecode:
|
||||
CATEGORY = "latent"
|
||||
FUNCTION = "decode"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"samples": ("LATENT",),
|
||||
}
|
||||
}
|
||||
|
||||
def decode(self, samples):
|
||||
return (samples,)
|
||||
|
||||
|
||||
class WindowsPathNode:
|
||||
CATEGORY = "paths"
|
||||
FUNCTION = "open_path"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"path": ("STRING",),
|
||||
}
|
||||
}
|
||||
|
||||
def open_path(self, path):
|
||||
return (path,)
|
||||
|
||||
|
||||
class TextConcatenateNode:
|
||||
CATEGORY = "text"
|
||||
FUNCTION = "text_concatenate"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"delimiter": ("STRING",),
|
||||
"clean_whitespace": ("STRING",),
|
||||
"text_b": ("STRING",),
|
||||
}
|
||||
}
|
||||
|
||||
def text_concatenate(self, delimiter, clean_whitespace, text_b):
|
||||
return (delimiter, clean_whitespace, text_b)
|
||||
|
||||
|
||||
class StringSeedNode:
|
||||
CATEGORY = "sampling"
|
||||
FUNCTION = "sample"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"seed": ("STRING",),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
},
|
||||
}
|
||||
|
||||
def sample(self, seed, prompt):
|
||||
return (seed, prompt)
|
||||
|
||||
|
||||
def load_fixture(name: str) -> dict:
|
||||
return json.loads((FIXTURE_DIR / name).read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
def export_workflow(workflow: dict, node_class_mappings: dict) -> str:
|
||||
output = StringIO()
|
||||
ComfyUItoPython(
|
||||
workflow=json.dumps(workflow),
|
||||
output_file=output,
|
||||
node_class_mappings=node_class_mappings,
|
||||
)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
class GeneratorCodegenIssueRegressionTest(unittest.TestCase):
|
||||
def test_export_uses_dictionary_expansion_for_rgthree_symbol_heavy_input_names(self):
|
||||
generated = export_workflow(
|
||||
load_fixture("unsafe-rgthree-kwargs.json"),
|
||||
{
|
||||
"AnySwitchRgthree": AnySwitchRgthree,
|
||||
"DualClipLoader": DualClipLoader,
|
||||
"PowerLoraLoaderRgthree": PowerLoraLoaderRgthree,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertIn(
|
||||
'powerloraloaderrgthree_631 = powerloraloaderrgthree.load_loras(',
|
||||
generated,
|
||||
)
|
||||
self.assertIn(
|
||||
'PowerLoraLoaderHeaderWidget={"type": "PowerLoraLoaderHeaderWidget"}',
|
||||
generated,
|
||||
)
|
||||
self.assertIn('**{"\\u2795 Add Lora": ""}', generated)
|
||||
self.assertNotIn('➕ Add Lora=""', generated)
|
||||
|
||||
def test_export_sanitizes_subgraph_identifiers_for_upscaler_workflows(self):
|
||||
generated = export_workflow(
|
||||
load_fixture("subgraph-upscaler-identifiers.json"),
|
||||
{
|
||||
"VaeDecode": VaeDecode,
|
||||
"UpscaleModelLoader": UpscaleModelLoader,
|
||||
"ImageUpscaleWithModel": ImageUpscaleWithModel,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertIn("upscalemodelloader_42_0 = upscalemodelloader.load_model(", generated)
|
||||
self.assertIn(
|
||||
"imageupscalewithmodel_42_1 = imageupscalewithmodel.upscale(",
|
||||
generated,
|
||||
)
|
||||
self.assertIn(
|
||||
"upscale_model=get_value_at_index(upscalemodelloader_42_0, 0)",
|
||||
generated,
|
||||
)
|
||||
self.assertNotIn("imageupscalewithmodel_42:1", generated)
|
||||
self.assertNotIn("upscalemodelloader_42:0", generated)
|
||||
|
||||
def test_export_preserves_windows_style_model_paths(self):
|
||||
generated = export_workflow(
|
||||
load_fixture("windows-path-string.json"),
|
||||
{
|
||||
"WindowsPathNode": WindowsPathNode,
|
||||
},
|
||||
)
|
||||
|
||||
globals_dict = {"__name__": "generated_workflow_module"}
|
||||
exec(generated, globals_dict)
|
||||
|
||||
self.assertEqual(
|
||||
globals_dict["build_workflow"]()["1"]["inputs"]["path"],
|
||||
r"C:\ComfyUI\models\upscale_models\RealESRGAN_x4plus.safetensors",
|
||||
)
|
||||
|
||||
def test_export_preserves_trailing_backslash_string_literals(self):
|
||||
generated = export_workflow(
|
||||
load_fixture("trailing-backslash-string.json"),
|
||||
{
|
||||
"TextConcatenateNode": TextConcatenateNode,
|
||||
},
|
||||
)
|
||||
|
||||
globals_dict = {"__name__": "generated_workflow_module"}
|
||||
exec(generated, globals_dict)
|
||||
|
||||
self.assertEqual(
|
||||
globals_dict["build_workflow"]()["1"]["inputs"]["text_b"],
|
||||
"\\",
|
||||
)
|
||||
|
||||
def test_export_randomizes_string_seed_inputs_as_strings(self):
|
||||
generated = export_workflow(
|
||||
load_fixture("string-seed-node.json"),
|
||||
{
|
||||
"StringSeedNode": StringSeedNode,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertIn(
|
||||
'node_1_seed = prompt["1"]["inputs"]["seed"] = str(random.randint(1, 2**64))',
|
||||
generated,
|
||||
)
|
||||
self.assertIn("seed=node_1_seed", generated)
|
||||
|
||||
def test_issue_cluster_regressions_render_parseable_python(self):
|
||||
workflows = [
|
||||
(
|
||||
load_fixture("unsafe-rgthree-kwargs.json"),
|
||||
{
|
||||
"AnySwitchRgthree": AnySwitchRgthree,
|
||||
"DualClipLoader": DualClipLoader,
|
||||
"PowerLoraLoaderRgthree": PowerLoraLoaderRgthree,
|
||||
},
|
||||
),
|
||||
(
|
||||
load_fixture("subgraph-upscaler-identifiers.json"),
|
||||
{
|
||||
"VaeDecode": VaeDecode,
|
||||
"UpscaleModelLoader": UpscaleModelLoader,
|
||||
"ImageUpscaleWithModel": ImageUpscaleWithModel,
|
||||
},
|
||||
),
|
||||
(
|
||||
load_fixture("trailing-backslash-string.json"),
|
||||
{
|
||||
"TextConcatenateNode": TextConcatenateNode,
|
||||
},
|
||||
),
|
||||
(
|
||||
load_fixture("windows-path-string.json"),
|
||||
{
|
||||
"WindowsPathNode": WindowsPathNode,
|
||||
},
|
||||
),
|
||||
(
|
||||
load_fixture("string-seed-node.json"),
|
||||
{
|
||||
"StringSeedNode": StringSeedNode,
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
for workflow, mapping in workflows:
|
||||
generated = export_workflow(workflow, mapping)
|
||||
compile(generated, "<generated_workflow>", "exec")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user