fix bug, update v3.0.0

This commit is contained in:
billwuhao
2025-02-23 02:11:44 +08:00
parent d6347e1424
commit ce0bc5b8f8
12 changed files with 247995 additions and 157 deletions
+1
View File
@@ -3,6 +3,7 @@ txtfiles/others.txt
txtfiles/poses.txt
txtfiles/styles.txt
txtfiles/test.txt
txtfiles/civit_nsfw.json
# Byte-compiled / optimized / DLL files
__pycache__/
+157
View File
@@ -0,0 +1,157 @@
import random
import re
import os
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
class DeepseekRun:
node_dir = os.path.dirname(os.path.abspath(__file__))
comfy_path = os.path.dirname(os.path.dirname(node_dir))
model_path = os.path.join(comfy_path, "models", "LLM")
ds_model_path = os.path.join(model_path, "DeepScaleR-1.5B-Preview")
model_paths = {
"DeepScaleR-1.5B-Preview": ds_model_path, # 可以添加更多模型和路径
}
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (list(cls.model_paths.keys()), {"tooltip": "models are expected to be in Comfyui/models/LLM folder"}),
"user_prompt": ("STRING", {"default": "", "multiline": True}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"max_tokens": ("INT", {"default": 1000, "min": 0, "max": 0xffffffffffffffff}),
"temperature": ("FLOAT", {"default": 1, "min": 0, "max": 2}),
"top_k": ("INT", {"default": 50, "min": 0, "max": 101}),
"top_p": ("FLOAT", {"default": 1, "min": 0, "max": 1}),
"unload_model": ("BOOLEAN", {
"default": False,
"tooltip": "If True, unload the model from memory after execution. Next execution will reload the model."}), # Added unload_model input
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("STRING",)
FUNCTION = "dsgen"
CATEGORY = "MW-OneButtonPrompt"
_model_cache = {}
def dsgen(self, model, user_prompt, seed=0, temperature=1.0, max_tokens=1000, top_k=25, top_p=1.0, unload_model=False):
if seed:
set_seed(self.hash_seed(seed))
match = re.search(r'{(.*?)}', user_prompt)
if match:
content_in_brackets = match.group(1)
# 按 "|" 拆分
options = content_in_brackets.split('|')
# 随机选择一个
chosen_option = random.choice(options).strip() # 使用 strip() 去除可能存在的首尾空格
# 替换原字符串 "{}" 中的内容
user_prompt = user_prompt.replace(match.group(0), chosen_option)
else:
user_prompt = user_prompt
messages = [
{"role": "user", "content": user_prompt},
]
model_dict = self.load_model(model)
if model_dict is None: # Check if model loading failed
return ("Error loading model. Check console.", )
device = model_dict["device"] # 获取设备信息
dsmodel = model_dict["dsmodel"]
tokenizer = model_dict["tokenizer"]
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
# 构造输入时直接使用 device
model_inputs = tokenizer([text], return_tensors="pt").to(device)
generated_ids = dsmodel.generate(
**model_inputs,
max_new_tokens=max_tokens,
temperature=temperature,
top_k=top_k,
top_p=top_p,
do_sample=True
)
generated_ids = [
output_ids[len(input_ids):] for input_ids, output_ids in zip(model_inputs.input_ids, generated_ids)
]
response = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
response = re.sub(r'<think>[\s\S]*?</think>', '', response).strip()
if unload_model: # Check if unload_model is True
self.unload_model_from_cache(model) # Unload the model if requested
print(f"DeepseekRun: Model '{model}' unloaded from cache.") # Inform user model is unloaded
return (response, )
def load_model(self, model):
model_path = self.model_paths.get(model)
if model in self._model_cache:
return self._model_cache[model]
try:
dsmodel = AutoModelForCausalLM.from_pretrained(
model_path,
device_map="auto", # 自动分配
torch_dtype="auto",
low_cpu_mem_usage=True,
trust_remote_code=True
)
tokenizer = AutoTokenizer.from_pretrained(model_path)
# 获取实际使用的设备
device = next(dsmodel.parameters()).device
model_info = {
"dsmodel": dsmodel,
"tokenizer": tokenizer,
"device": device
}
self._model_cache[model] = model_info
print(f"DeepseekRun: Model '{model}' loaded to cache.") # Inform user model is loaded
return model_info
except Exception as e:
print(f"DeepseekRun: Error loading model {model} from {model_path}: {e}")
return None
def unload_model_from_cache(self, model): # New method to unload model
if model in self._model_cache:
model_info = self._model_cache.pop(model) # Remove model from cache
del model_info # Optionally delete model_info dictionary to release references (may not be strictly needed in Python)
torch.cuda.empty_cache() # Clear CUDA cache to try and free VRAM
print(f"DeepseekRun: Model '{model}' removed from cache and CUDA cache cleared.") # Inform user of unload action
else:
print(f"DeepseekRun: Model '{model}' not found in cache, cannot unload.") # Inform user if model was not in cache
def hash_seed(self, seed):
import hashlib
# Convert the seed to a string and then to bytes
seed_bytes = str(seed).encode('utf-8')
# Create a SHA-256 hash of the seed bytes
hash_object = hashlib.sha256(seed_bytes)
# Convert the hash to an integer
hashed_seed = int(hash_object.hexdigest(), 16)
# Ensure the hashed seed is within the acceptable range for set_seed
return hashed_seed % (2**32)
NODE_CLASS_MAPPINGS = {
"DeepseekRun": DeepseekRun
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DeepseekRun": "Deepseek Run"
}
+182
View File
@@ -0,0 +1,182 @@
from PIL import Image, ImageSequence, ImageOps
import torch
import requests
from io import BytesIO
import os
import numpy as np
import json
import random
from transformers import set_seed
def get_image_data_from_url(url, proxies=None):
"""
Checks if a URL is likely an image and returns the image data if it is.
Args:
url (str): The URL to check.
Returns:
bytes or None:
- Image data (bytes) if the URL is likely an image.
- None if the URL is not likely an image or if there was an error.
"""
try:
response = requests.get(url, stream=True, allow_redirects=True, timeout=10, proxies=proxies)
response.raise_for_status() # Raise HTTPError for bad responses (4xx or 5xx)
content_type = response.headers.get('Content-Type', '').lower()
if content_type.startswith('image/'):
image_content = response.content # 显式读取 response.content 到内存
if image_content is None:
return None
return Image.open(BytesIO(image_content)) # Return the image data as bytes
else:
return None # Not an image content type
except requests.exceptions.RequestException as e:
print(f"Error fetching URL: {url}. Error: {e}")
return None # Error during request
def pil2tensor(img):
output_images = []
output_masks = []
for i in ImageSequence.Iterator(img):
i = ImageOps.exif_transpose(i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return (output_image, output_mask)
def get_random_imginfo(data: dict, img_type: str, proxies=None):
"""
Selects and returns a random key-value pair from a dictionary.
Args:
data (dict): The input dictionary.
img_type (str): The type of image to return.
Returns:
tuple or None: A tuple containing the (key, value) of a randomly
selected item from the dictionary. Returns None if
the dictionary is empty.
"""
if not data:
return None # Return None if the dictionary is empty
if img_type == "Prompt":
data = data["Prompt"]
else:
data = {**data["Prompt"], **data["NoPrompt"]}
keys = list(data)
for i in range(6):
random_url = random.choice(keys)
imageinfo = get_image_data_from_url(random_url, proxies=proxies)
if imageinfo:
return imageinfo, data[random_url][1]
else:
raise ValueError("Failed to find a valid image URL after 6 attempts.")
class LoadImageInfoFromCivitai:
node_dir = os.path.dirname(os.path.abspath(__file__))
jsonfile_path = os.path.join(node_dir, "txtfiles")
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"output_type": (["Img", "Img+Prompt"], {"default": "Prompt"}),
"nsfw": ("BOOLEAN", {
"default": False,
"tooltip": "When True, need `civit_nsfw.json` in the `txtfiles` folder, otherwise it is invalid"}),
"proxy": ("STRING", {
"multiline": False,
"default": "http://127.0.0.1:None",
"tooltip": "When load Img, if unable to access Civitai site, proxy needs to be filled in"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
CATEGORY = "MW-OneButtonPrompt"
RETURN_TYPES = ("IMAGE", "MASK", "STRING")
RETURN_NAMES = ("Image", "Mask", "Prompt")
FUNCTION = "load"
def hash_seed(self, seed):
import hashlib
# Convert the seed to a string and then to bytes
seed_bytes = str(seed).encode('utf-8')
# Create a SHA-256 hash of the seed bytes
hash_object = hashlib.sha256(seed_bytes)
# Convert the hash to an integer
hashed_seed = int(hash_object.hexdigest(), 16)
# Ensure the hashed seed is within the acceptable range for set_seed
return hashed_seed % (2**32)
def load(self, output_type, nsfw, proxy, seed=0):
if seed:
set_seed(self.hash_seed(seed))
proxy = None if proxy == "http://127.0.0.1:None" else proxy
img, prompt = self.load_json_file(output_type, nsfw, proxies={"https": proxy,"http": proxy,})
print(f"LoadImageInfoFromCivitai.load: load_json_file returned img type: {type(img)}") # Debug print
if img is None: # Check if get_image_data_from_url failed
print("LoadImageInfoFromCivitai.load: get_image_data_from_url returned None. Image download failed.") # Debug print
return (None, None, prompt)
img_out, mask_out = pil2tensor(img)
print(f"LoadImageInfoFromCivitai.load: pil2tensor returned img_out type: {type(img_out)}, mask_out type: {type(mask_out)}") # Debug print
if img_out is None or mask_out is None: # Check if pil2tensor failed
print("LoadImageInfoFromCivitai.load: pil2tensor returned None outputs. Image processing failed.") # Debug print
return (None, None, prompt) # Return None for IMAGE and MASK, and error message
return (img_out, mask_out, prompt)
def load_json_file(self, output_type, nsfw, proxies=None):
if nsfw:
file_path = self.jsonfile_path + "/civit_nsfw.json"
if not os.path.exists(self.jsonfile_path + "/civit_nsfw.json"):
file_path = self.jsonfile_path + "/civit_sfw.json"
try:
with open(file_path, "r", encoding="utf-8") as f:
data = json.load(f)
except Exception as e:
print(f"Error loading JSON file: {file_path}. Error: {e}")
return None, None
if output_type == "Img+Prompt":
img, prompt = get_random_imginfo(data, "Prompt", proxies=proxies)
else:
img, prompt = get_random_imginfo(data, "Img", proxies=proxies)
if not nsfw:
try:
with open(self.jsonfile_path + "/civit_sfw.json", "r", encoding="utf-8") as f:
data = json.load(f)
except Exception as e:
print(f"Error loading JSON file: {self.jsonfile_path + '/civit_sfw.json'}. Error: {e}")
return None, None
if output_type == "Img+Prompt":
img, prompt = get_random_imginfo(data, "Prompt", proxies=proxies)
else:
img, prompt = get_random_imginfo(data, "Img", proxies=proxies)
return img, prompt
+15 -156
View File
@@ -1,160 +1,12 @@
import random
import os
import re
from collections.abc import Callable
import torch
from pathlib import Path
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
import folder_paths
from comfy import model_management, model_patcher
def create_path_dict(paths: list[str], predicate: Callable[[Path], bool] = lambda _: True) -> dict[str, str]:
"""
Creates a flat dictionary of the contents of all given paths: ``{name: absolute_path}``.
Non-recursive. Optionally takes a predicate to filter items. Duplicate names overwrite (the last one wins).
Args:
paths (list[str]):
The paths to search for items.
predicate (Callable[[Path], bool]):
(Optional) If provided, each path is tested against this filter.
Returns ``True`` to include a path.
Default: Include everything
"""
flattened_paths = [item for path in paths for item in Path(path).iterdir() if predicate(item)]
return {item.name: str(item.absolute()) for item in flattened_paths}
class DeepseekRun:
@classmethod
def INPUT_TYPES(s):
all_llm_paths = folder_paths.get_folder_paths("LLM")
s.model_paths = create_path_dict(all_llm_paths, lambda x: x.is_dir())
return {
"required": {
"model": ([*s.model_paths], {"tooltip": "models are expected to be in Comfyui/models/LLM folder"}),
"user_prompt": ("STRING", {"default": "", "multiline": True}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"max_tokens": ("INT", {"default": 1000, "min": 0, "max": 0xffffffffffffffff}),
"temperature": ("FLOAT", {"default": 1, "min": 0, "max": 2}),
"top_k": ("INT", {"default": 20, "min": 0, "max": 101}),
"top_p": ("FLOAT", {"default": 1, "min": 0, "max": 1}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("STRING",)
FUNCTION = "dsgen"
CATEGORY = "MW-OneButtonPrompt"
_model_cache = {}
def load_model(self, model):
model_path = DeepseekRun.model_paths.get(model)
if model in self._model_cache:
return self._model_cache[model]
dsmodel = AutoModelForCausalLM.from_pretrained(
model_path,
device_map="auto", # 自动分配
torch_dtype="auto",
low_cpu_mem_usage=True,
trust_remote_code=True
)
tokenizer = AutoTokenizer.from_pretrained(model_path)
# 获取实际使用的设备
device = next(dsmodel.parameters()).device
self._model_cache[model] = {
"model": dsmodel,
"tokenizer": tokenizer,
"device": device
}
return {
"model": dsmodel,
"tokenizer": tokenizer,
"device": device
}
def hash_seed(self, seed):
import hashlib
# Convert the seed to a string and then to bytes
seed_bytes = str(seed).encode('utf-8')
# Create a SHA-256 hash of the seed bytes
hash_object = hashlib.sha256(seed_bytes)
# Convert the hash to an integer
hashed_seed = int(hash_object.hexdigest(), 16)
# Ensure the hashed seed is within the acceptable range for set_seed
return hashed_seed % (2**32)
def dsgen(self, model, user_prompt, seed=0, temperature=1.0, max_tokens=1000, top_k=25, top_p=1.0, **kwargs):
if seed:
set_seed(self.hash_seed(seed))
match = re.search(r'{(.*?)}', user_prompt)
if match:
content_in_brackets = match.group(1)
# 按 "|" 拆分
options = content_in_brackets.split('|')
# 随机选择一个
chosen_option = random.choice(options).strip() # 使用 strip() 去除可能存在的首尾空格
# 替换原字符串 "{}" 中的内容
user_prompt = user_prompt.replace(match.group(0), chosen_option)
else:
user_prompt = user_prompt
messages = [
{"role": "user", "content": user_prompt.format(**kwargs)},
]
# 单次加载模型
model_dict = self.load_model(model)
device = model_dict["device"] # 获取设备信息
dsmodel = model_dict["model"]
tokenizer = model_dict["tokenizer"]
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
# 构造输入时直接使用 device
model_inputs = tokenizer([text], return_tensors="pt").to(device)
generated_ids = dsmodel.generate(
**model_inputs,
max_new_tokens=max_tokens,
temperature=temperature,
top_k=top_k,
top_p=top_p,
do_sample=True
)
generated_ids = [
output_ids[len(input_ids):] for input_ids, output_ids in zip(model_inputs.input_ids, generated_ids)
]
response = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
response = re.sub(r'<think>[\s\S]*?</think>', '', response).strip()
return (response, )
# ------------------------------------
node_dir = os.path.dirname(os.path.abspath(__file__))
jsonfile_path = os.path.join(node_dir, "/txtfiles/")
def process_txt_file(txtfile: str):
script_dir = os.path.dirname(os.path.abspath(__file__)) # Script directory
file_path = os.path.join(script_dir, "./txtfiles/", txtfile)
file_path = os.path.join(node_dir, "./txtfiles/", txtfile)
if os.path.exists(file_path):
with open(file_path, 'r', encoding='utf-8') as file:
@@ -163,9 +15,9 @@ def process_txt_file(txtfile: str):
if processed_lines:
return processed_lines
else:
file_path = os.path.join(script_dir, "./txtfiles/", "example_" + txtfile)
file_path = os.path.join(node_dir, "./txtfiles/", "example_" + txtfile)
else:
file_path = os.path.join(script_dir, "./txtfiles/", "example_" + txtfile)
file_path = os.path.join(node_dir, "./txtfiles/", "example_" + txtfile)
with open(file_path, 'r', encoding='utf-8') as file:
lines = file.readlines()
@@ -240,7 +92,7 @@ class OneButtonPromptFlux:
FUNCTION = "fluxprompt"
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
@@ -273,12 +125,19 @@ class OneButtonPromptFlux:
return (generate_prompt(subject, pose, style, lora_trigger_or_prefix, refresh, test, seed),)
from .DeepSeekRone import DeepseekRun
from .LoadCivitai import LoadImageInfoFromCivitai
NODE_CLASS_MAPPINGS = {
"DeepseekRun": DeepseekRun,
"OneButtonPromptFlux": OneButtonPromptFlux,
"DeepseekRun": DeepseekRun
"LoadImageInfoFromCivitai": LoadImageInfoFromCivitai
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DeepseekRun": "Deepseek Run",
"OneButtonPromptFlux": "One Button Prompt Flux",
"DeepseekRun": "Deepseek Run"
"LoadImageInfoFromCivitai": "Load Image Info From Civitai"
}
+16
View File
@@ -12,6 +12,22 @@ This is a node for generating Flux prompts with one click in ComfyUI.
## 📣 Updates
[2025-02-20]⚒️: Support [C](https://civitai.com/images) Station images and prompts.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-23_00-40-23.png)
Default use [Civitai](https://civitai.com/images). If you need to use images from other websites, please modify the file `\ComfyUI_OneButtonPrompt_Flux\txtfiles\civit_sfw.json` yourself. I will periodically update the `civit_sfw.json` file
If you need nsfw images, please create a new `civit_nsfw.json` file in the `ComfyUI_OneButtonPrompt_Flux\txtfiles\` folder, which should be in the same format as the content of the `civit_sfw.json` file.
Collaborate with Deepseek r1 to optimize and enhance prompts.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-23_01-14-08.png)
Reverse inference prompts.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-23_01-37-50.png)
[2025-02-19] ⚒️: Support local DeepSeek R1.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-19_10-32-16.png)
+16
View File
@@ -8,6 +8,22 @@
## 📣 更新
[2025-02-20]⚒️: 支持 [C](https://civitai.com/images) 站图片及提示词.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-23_00-40-23.png)
默认使用 [Civitai](https://civitai.com/images) 的图片, 如果需要使用其他站图片, 请自行修改 `\ComfyUI_OneButtonPrompt_Flux\txtfiles\civit_sfw.json` 文件. 将不定期更新 `civit_sfw.json` 文件.
如果你需要 nsfw 图片, 请在 `ComfyUI_OneButtonPrompt_Flux\txtfiles\` 文件夹下新建 `civit_nsfw.json` 文件, 需与 `civit_sfw.json` 文件内容保持同样的格式.
配合 deepseek r1 优化增强提示词.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-23_01-14-08.png)
反推提示词.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-23_01-37-50.png)
[2025-02-19]⚒️: 支持本地 DeepSeek R1.
![](https://github.com/billwuhao/ComfyUI_OneButtonPrompt_Flux/blob/master/images/2025-02-19_10-32-16.png)
Binary file not shown.

After

Width:  |  Height:  |  Size: 477 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 290 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 256 KiB

+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "onebuttonprompt_flux"
description = "ComfyUI_OneButtonPrompt_Flux is a Flux prompt generation node. The subject can be \"human\", \"other\" or a combination of both. For human, pose settings can be enabled. Additionally, various styles can be applied. Finally, combine it with \"Prompt Enhancement\" to seamlessly automate image generation, eliminating the hassle of designing prompts."
version = "2.0.0"
version = "3.0.0"
license = {file = "LICENSE"}
[project.urls]
+247606
View File
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long