From cd996c1a5bcd732e5227c0ac0d4ba5ebbbd0513e Mon Sep 17 00:00:00 2001 From: jax Date: Mon, 14 Apr 2025 20:40:37 +0800 Subject: [PATCH] style lora support from Civitai --- __init__.py | 4 ++- nodes/cai_utils.py | 69 ++++++++++++++++++++++++++++++++++++++++++++ nodes/comfy_nodes.py | 59 +++++++++++++++++++++++++++++++++++++ 3 files changed, 131 insertions(+), 1 deletion(-) create mode 100644 nodes/cai_utils.py diff --git a/__init__.py b/__init__.py index 1abefdf..34c1d70 100644 --- a/__init__.py +++ b/__init__.py @@ -1,5 +1,5 @@ -from .nodes.comfy_nodes import EasyControlLoadFlux, EasyControlLoadLora, EasyControlLoadMultiLora, EasyControlGenerate, EasyControlLoadStyleLora +from .nodes.comfy_nodes import EasyControlLoadFlux, EasyControlLoadLora, EasyControlLoadMultiLora, EasyControlGenerate, EasyControlLoadStyleLora, EasyControlLoadStyleLoraFromCivitai # 注册节点 @@ -9,6 +9,7 @@ NODE_CLASS_MAPPINGS = { "EasyControlLoadMultiLora": EasyControlLoadMultiLora, "EasyControlGenerate": EasyControlGenerate, "EasyControlLoadStyleLora": EasyControlLoadStyleLora, + "EasyControlLoadStyleLoraFromCivitai": EasyControlLoadStyleLoraFromCivitai, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -17,6 +18,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "EasyControlLoadMultiLora": "Load Multiple EasyControl Loras", "EasyControlGenerate": "EasyControl Generate", "EasyControlLoadStyleLora": "Load EasyControl Style Lora", + "EasyControlLoadStyleLoraFromCivitai": "Load EasyControl Style Lora from Civitai", } WEB_DIRECTORY = "./web" diff --git a/nodes/cai_utils.py b/nodes/cai_utils.py new file mode 100644 index 0000000..a343fe2 --- /dev/null +++ b/nodes/cai_utils.py @@ -0,0 +1,69 @@ +import requests +from requests.exceptions import HTTPError +import os +import re + +def download_file_with_token(url, params=None, save_path='.'): + try: + # Send a GET request to the URL + with requests.get(url, params=params, stream=True) as response: + response.raise_for_status() # Raise an error for bad responses + print(f"Downloading model successfully from {response.url}") + + # Create temporary file path + temp_path = save_path + '.download' + + # Write response content to temporary file first + with open(temp_path, 'wb') as file: + for chunk in response.iter_content(chunk_size=8192): + file.write(chunk) + + # After successful download, rename temp file to target path + os.replace(temp_path, save_path) + + print(f"File downloaded successfully: {save_path}") + return True + + except requests.HTTPError as http_err: + print(f'HTTP error occurred: {http_err}') + except Exception as err: + print(f'An error occurred: {err}') + # Clean up temp file if it exists + if os.path.exists(temp_path): + try: + os.remove(temp_path) + except: + pass + return False + + +def download_cai(model_id, token, lora_path): + """ + Download a LoRA model from CivitAI directly to the specified lora_path. + + :param model_id: The ID of the model to download. + :param token: The authentication token (optional). + :param lora_path: The full path (including filename) where the file will be saved. + :param full_url: Full URL for downloading the model (optional). + """ + # Ensure the directory of lora_path exists + directory_path = os.path.dirname(lora_path) + if not os.path.exists(directory_path): + os.makedirs(directory_path, exist_ok=True) + + # Determine the URL for the download + if not model_id: + print("Either model_id must be provided for model download.") + return False + + url = f'https://civitai.com/api/download/models/{model_id}' + params = {'token': token} if token else {} + + # Call the download function and specify the exact file path + download_success = download_file_with_token(url, params, lora_path) + if download_success: + print("File downloaded successfully.") + return True + else: + print("Failed to download the file.") + return False \ No newline at end of file diff --git a/nodes/comfy_nodes.py b/nodes/comfy_nodes.py index 1b0b138..6b440c9 100644 --- a/nodes/comfy_nodes.py +++ b/nodes/comfy_nodes.py @@ -5,6 +5,8 @@ import folder_paths from PIL import Image import numpy as np +from .cai_utils import download_cai + # Add the parent directory to the Python path so we can import from easycontrol sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) @@ -129,6 +131,63 @@ class EasyControlLoadStyleLora: return (pipe,) +class EasyControlLoadStyleLoraFromCivitai: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "pipe": ("EASYCONTROL_PIPE",), + "lora_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05}), + "civitai_model_id": ("STRING", {"default": "", "tooltip": "The ID of the model to download from CivitAI."}), + } + } + + RETURN_TYPES = ("EASYCONTROL_PIPE",) + FUNCTION = "load_lora" + CATEGORY = "EasyControl" + + def load_lora(self, pipe, lora_weight, civitai_model_id): + + civitai_token_id = os.getenv("CIVITAI_TOKEN", "").strip() + if not civitai_token_id: + raise RuntimeError("CIVITAI_TOKEN environment variable is not set or empty.") + + loras_dir = folder_paths.get_folder_paths("loras")[0] + + lora_filename = f"tmp_{civitai_model_id or 'downloaded_lora'}.safetensors" # 生成临时文件名 + lora_path = os.path.join(loras_dir, lora_filename) + + file_exists = os.path.exists(lora_path) + if not file_exists: + self.download_from_civitai(civitai_model_id, civitai_token_id, lora_path) + else: + print(f"LoRA file already exists at {lora_path}, skipping download") + + + # Load LoRA weights + print(f"Loading FLUX Style LoRA: {lora_filename}, Weight: {lora_weight}") + + weight_name = lora_filename + + # handle offload + device = next(pipe.transformer.parameters()).device + + # Load LoRA weights + pipe.load_lora_weights(lora_path, weight_name=weight_name, device=device) + + # Fuse LoRA + # pipe.fuse_lora(lora_weights=[lora_weight]) + return (pipe,) + + def download_from_civitai(self, model_id, token_id, lora_path): + print("Downloading LoRA from CivitAI") + print(f"\tModel ID: {model_id}") + print(f"\tToken ID: {token_id}") + print(f"\tSave path: {lora_path}") + # 实现下载逻辑 + download_cai(model_id, token_id, lora_path) + + class EasyControlLoadMultiLora: @classmethod def INPUT_TYPES(cls):