style lora support from Civitai
This commit is contained in:
+3
-1
@@ -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"
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user