From 69d1f2716d85a24bc5d2d8f724123777fed347d9 Mon Sep 17 00:00:00 2001 From: Tung Nguyen Date: Fri, 14 Jul 2023 16:02:31 +0700 Subject: [PATCH] chore: delay runner start up --- __init__.py | 2 ++ modules/workflow.py | 17 ++++++++++++++--- 2 files changed, 16 insertions(+), 3 deletions(-) diff --git a/__init__.py b/__init__.py index 7211c4b..9cd17c0 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,5 @@ import os +import time import shutil import threading from typing import Callable @@ -71,6 +72,7 @@ def init(): # update checkpoint hash in background def _update_checkpoints_hash(cb: Callable): update_checkpoints_hash() + time.sleep(10) # wait for server to start cb() def _cb(): diff --git a/modules/workflow.py b/modules/workflow.py index 64cdacb..78653a8 100644 --- a/modules/workflow.py +++ b/modules/workflow.py @@ -8,13 +8,14 @@ from datetime import datetime from typing import Dict, List from server import PromptServer -from nodes import NODE_CLASS_MAPPINGS from folder_paths import models_dir, get_filename_list from .log import logger from .nodes import NODE_CLASS_MAPPINGS as _NODE_CLASS_MAPPINGS -ALL_NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **_NODE_CLASS_MAPPINGS} +node_class_mappings_loaded = False +ALL_NODE_CLASS_MAPPINGS = {**_NODE_CLASS_MAPPINGS} + Graph = Dict[str, List[str]] root_dir = os.path.dirname(inspect.getfile(PromptServer)) @@ -32,6 +33,16 @@ promp_args = {"prompt", "negative_prompt"} seed_args = {"seed", "noise_seed"} +def get_node_class_mapping(): + if not node_class_mappings_loaded: + from nodes import NODE_CLASS_MAPPINGS + + ALL_NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS) + node_class_mappings_loaded = True + + return ALL_NODE_CLASS_MAPPINGS + + def __dfs_sort_helper( graph: Graph, v: str, n: int, visited: Dict[str, bool], topNums: Dict[str, int] ) -> int: @@ -260,7 +271,7 @@ def workflow_to_prompt(workflow, args: dict = {}): if node["type"] in virtual_nodes or node["type"] in input_nodes: continue - obj_class = ALL_NODE_CLASS_MAPPINGS.get(node["type"], None) + obj_class = get_node_class_mapping().get(node["type"], None) if obj_class is None: logger.error(f"Unknown node {node['type']}") continue