799 lines
28 KiB
Python
799 lines
28 KiB
Python
#Utility functions for use in the nodes.
|
|
|
|
import pathlib
|
|
import hashlib
|
|
import datetime
|
|
import torch
|
|
import logging
|
|
|
|
import folder_paths
|
|
import comfy.utils
|
|
|
|
from .model_cache import cache
|
|
from .helpers_civitai import *
|
|
from .constants import MODEL_FILE_EXTENSIONS
|
|
|
|
def str_to_bool(value):
|
|
if isinstance(value, bool):
|
|
return value
|
|
if isinstance(value, str):
|
|
value = value.lower()
|
|
if value in {'true', '1', 'yes'}:
|
|
return True
|
|
if value in {'false', '0', 'no'}:
|
|
return False
|
|
raise ValueError(f"Cannot convert {value} to boolean.")
|
|
|
|
# Not currently used.
|
|
def bool_to_str(value):
|
|
if isinstance(value, bool):
|
|
return "true" if value else "false"
|
|
if isinstance(value, str):
|
|
value = value.lower()
|
|
if value in {'true', '1', 'yes'}:
|
|
return "true"
|
|
if value in {'false', '0', 'no'}:
|
|
return "false"
|
|
raise ValueError(f"Cannot convert {value} to string representation of boolean.")
|
|
|
|
def name_from_path(path):
|
|
return pathlib.Path(path).name
|
|
|
|
def get_path_without_base(folder_type:str, path:str) -> str:
|
|
"""Get the base path for a given folder type and path."""
|
|
for base in folder_paths.get_folder_paths(folder_type):
|
|
if path.startswith(base):
|
|
path = path[len(base):].lstrip("/\\")
|
|
break
|
|
return path
|
|
|
|
def get_file_extension(path: str) -> str:
|
|
"""Get the file extension from a path."""
|
|
return path.split(".")[-1] if "." in path else ""
|
|
|
|
def has_model_extension(path: str) -> bool:
|
|
"""Check if a file path has a model extension."""
|
|
if not path:
|
|
return False
|
|
extension = '.' + get_file_extension(path).lower()
|
|
return extension in {ext.lower() for ext in MODEL_FILE_EXTENSIONS}
|
|
|
|
def is_model_file(path: str) -> bool:
|
|
"""Check if a file is a valid model file based on its extension."""
|
|
return has_model_extension(path)
|
|
|
|
# fast_hash_threshold=100*1024*1024 # 100MB threshold
|
|
def get_file_sha256(path, fast_hash_threshold=0): # Full hashing by default
|
|
"""
|
|
Calculate SHA256 hash of a file with optimizations for large files.
|
|
|
|
Args:
|
|
path: File path to hash
|
|
fast_hash_threshold: File size threshold (in bytes) above which to use fast hashing.
|
|
Set to 0 to always use full hashing.
|
|
|
|
Returns:
|
|
First 10 characters of the SHA256 hash
|
|
"""
|
|
print(f"Calculating hash for {path}")
|
|
import hashlib
|
|
import os
|
|
|
|
# Use a larger buffer size for better I/O performance
|
|
BUFFER_SIZE = 65536 # 64KB chunks
|
|
|
|
try:
|
|
file_size = os.path.getsize(path)
|
|
print(f"File size: {file_size / (1024*1024):.1f} MB")
|
|
|
|
m = hashlib.sha256()
|
|
|
|
# For very large files, use a sampling approach for speed
|
|
if fast_hash_threshold > 0 and file_size > fast_hash_threshold:
|
|
print(f"Large file detected, using fast hashing strategy")
|
|
|
|
# Sample strategy: hash beginning, middle, and end of file + file size + filename
|
|
# This creates a unique signature that's very unlikely to collide for different files
|
|
sample_size = min(BUFFER_SIZE * 4, file_size // 10) # Sample up to 256KB or 10% of file
|
|
|
|
with open(path, 'rb') as f:
|
|
# Hash file metadata first (size and name)
|
|
m.update(str(file_size).encode())
|
|
m.update(os.path.basename(path).encode())
|
|
|
|
# Hash beginning of file
|
|
beginning = f.read(sample_size)
|
|
m.update(beginning)
|
|
|
|
# Hash middle of file
|
|
if file_size > sample_size * 2:
|
|
middle_pos = file_size // 2 - sample_size // 2
|
|
f.seek(middle_pos)
|
|
middle = f.read(sample_size)
|
|
m.update(middle)
|
|
|
|
# Hash end of file
|
|
if file_size > sample_size * 3:
|
|
f.seek(-sample_size, 2) # Seek to sample_size bytes from end
|
|
end = f.read(sample_size)
|
|
m.update(end)
|
|
else:
|
|
# Full file hashing for smaller files or when requested
|
|
with open(path, 'rb') as f:
|
|
while True:
|
|
chunk = f.read(BUFFER_SIZE)
|
|
if not chunk:
|
|
break
|
|
m.update(chunk)
|
|
|
|
except (OSError, IOError) as e:
|
|
print(f"Error reading file {path}: {e}")
|
|
raise
|
|
|
|
full_hash = m.hexdigest()
|
|
print(f"Got full hash {full_hash}")
|
|
result = full_hash[:10]
|
|
print(f"Got hash {result}")
|
|
return result
|
|
|
|
|
|
def get_files_in_dir(input_dirs=None, extensions=None):
|
|
if extensions is None:
|
|
extensions = ".*"
|
|
if input_dirs is None:
|
|
raise ValueError("input_dirs cannot be None")
|
|
|
|
input_files = []
|
|
# Check if input_dirs is a tuple or a list. If not, make it a list.
|
|
if not isinstance(input_dirs, (list, tuple)):
|
|
input_dirs = [input_dirs]
|
|
|
|
for dir in input_dirs:
|
|
if dir is None or dir == "":
|
|
continue
|
|
if pathlib.Path(dir).exists():
|
|
file_list = pathlib.Path(dir).rglob("*")
|
|
for file in file_list:
|
|
if file.exists() and not file.is_dir():
|
|
file_path = ""
|
|
if file.suffix.lower() in extensions:
|
|
# Get path relative to the input directory
|
|
try:
|
|
file_path = str(file.relative_to(dir))
|
|
except ValueError:
|
|
file_path = str(file)
|
|
input_files.append(file_path)
|
|
|
|
input_files = sorted(set(input_files))
|
|
return input_files
|
|
|
|
def last_used(file_path):
|
|
cache.load()
|
|
|
|
if file_path in cache.hash:
|
|
last_used = cache.by_path(file_path).get("lastUsed", None)
|
|
if last_used is not None:
|
|
return datetime.datetime.fromisoformat(last_used)
|
|
else:
|
|
return None
|
|
else:
|
|
return None
|
|
|
|
def days_since_last_used(file_path):
|
|
was_last_used = last_used(file_path)
|
|
if was_last_used is not None:
|
|
now = datetime.datetime.now()
|
|
delta = now - was_last_used
|
|
return delta.days
|
|
else:
|
|
return 365
|
|
|
|
def get_file_modification_date(file_path):
|
|
try:
|
|
file_path = pathlib.Path(file_path)
|
|
if file_path.exists():
|
|
return datetime.datetime.fromtimestamp(file_path.stat().st_mtime)
|
|
else:
|
|
print(f"File {file_path} does not exist.")
|
|
return datetime.datetime.now()
|
|
except Exception as e:
|
|
print(f"Error getting modification date for {file_path}: {e}")
|
|
return datetime.datetime.now()
|
|
|
|
# Called *after* verifying the json pulled successfully.
|
|
# This function updates the cache with the information from the CivitAI json.
|
|
def update_cache_from_civitai_json(file_path, json_data, timestamp=True):
|
|
the_files = json_data.get("files", [])
|
|
hashes = {}
|
|
|
|
if len(the_files) > 0:
|
|
hashes = the_files[0].get("hashes", {})
|
|
update_available = True
|
|
|
|
latest_model = None
|
|
if json_data.get("modelId", None) is not None:
|
|
latest_model = get_latest_model_version(json_data["modelId"])
|
|
if latest_model == json_data["id"] or latest_model is None:
|
|
update_available = False
|
|
|
|
if latest_model is None:
|
|
latest_model = ""
|
|
|
|
file_cache = cache.by_path(file_path)
|
|
file_cache.update({
|
|
'civitai': "True",
|
|
'civitai_failed_count': 0,
|
|
'model': json_data.get("model", {}),
|
|
'name': json_data.get("name", ""),
|
|
'baseModel': json_data.get("baseModel", ""),
|
|
'id': json_data.get("id", ""),
|
|
'modelId': json_data.get("modelId", ""),
|
|
'update_available': update_available,
|
|
'update_version_id': latest_model,
|
|
'trainedWords': json_data.get("trainedWords", []),
|
|
'downloadUrl': json_data.get("downloadUrl", ""),
|
|
'hashes': hashes
|
|
})
|
|
if timestamp:
|
|
print("Updating timestamp.")
|
|
cache.update_last_used_by_path(file_path)
|
|
|
|
print("Successfully pulled metadata.")
|
|
|
|
def update_cache_without_civitai_json(file_path, hash, timestamp=True):
|
|
file_cache = cache.by_path(file_path)
|
|
print(f"Unable to find on civitai.")
|
|
file_cache['civitai'] = "False"
|
|
file_cache['civitai_failed_count'] = file_cache.get('civitai_failed_count', 0) + 1
|
|
file_cache['hash'] = hash
|
|
cache.update_last_used_by_path(file_path)
|
|
|
|
def add_file_to_cache(file_path, hash=None):
|
|
file_path = str(file_path)
|
|
print(f"Adding {file_path} to cache.")
|
|
if hash is None:
|
|
hash = get_file_sha256(file_path)
|
|
|
|
if file_path not in cache.hash:
|
|
cache.hash[file_path] = hash
|
|
|
|
if cache.info.get(hash, None) is None:
|
|
cache.info[hash] = {
|
|
'civitai': "False",
|
|
'update_available': False,
|
|
'update_version_id': "",
|
|
'hash': hash,
|
|
'lastUsed': datetime.datetime.now().isoformat()
|
|
}
|
|
|
|
print(f"Adding {file_path} to cache with hash {hash}.")
|
|
return hash
|
|
|
|
def recheck_hash(file_path, hash):
|
|
new_hash = get_file_sha256(file_path)
|
|
if new_hash != hash:
|
|
print(f"Hash mismatch. Using new hash.")
|
|
if file_path in cache.hash:
|
|
print(f"Updating cache for {file_path} with new hash {new_hash}.")
|
|
if new_hash not in cache.info:
|
|
if hash in cache.info:
|
|
cache.info[new_hash] = cache.info[hash]
|
|
else:
|
|
print(f"File {file_path} not in cache. Adding with new hash {new_hash}.")
|
|
add_file_to_cache(file_path, new_hash)
|
|
hash = new_hash
|
|
return hash
|
|
|
|
# Model path is inconsistent. Sometimes it is a list, and sometimes a string.
|
|
# Also, different models vary on whether they use an absolute or relative path.
|
|
# Loras and VAEs use a relative path from the loras folder. They are converted to absolute paths
|
|
# when calling this function. get_full_path_or_raise will raise if passed an absolute path.
|
|
# Unfortunately, this is set in the main program, not this node set.
|
|
def pull_and_update_model_timestamp(file_paths, model_type):
|
|
if not isinstance(file_paths, (list, tuple)):
|
|
file_paths = [file_paths]
|
|
|
|
#if model_type == "lora":
|
|
# file_paths = [folder_paths.get_full_path_or_raise("loras", lora) for lora in file_paths]
|
|
|
|
for path in file_paths:
|
|
pull_metadata(path, model_type=model_type)
|
|
|
|
update_model_timestamp(file_paths)
|
|
|
|
def update_model_timestamp(file_paths):
|
|
cache.load()
|
|
# If the file_paths isn't a list, make it a list.
|
|
if not isinstance(file_paths, (list, tuple)):
|
|
file_paths = [file_paths]
|
|
for path in file_paths:
|
|
if path in cache.hash:
|
|
cache.update_last_used_by_path(path)
|
|
cache.save()
|
|
|
|
def pull_metadata(file_paths, timestamp = True, force_all = False, pbar = None, model_type = None):
|
|
pull_json = True
|
|
metadata_days_recheck = 7
|
|
|
|
cache.load()
|
|
cache.backup_counter += 1
|
|
if cache.backup_counter >= cache.num_of_backups_to_keep:
|
|
cache.prune_all_backups()
|
|
cache.backup_counter = 0
|
|
|
|
if isinstance(file_paths, str):
|
|
file_paths = [file_paths]
|
|
|
|
if not file_paths:
|
|
logging.warning("No file paths provided.")
|
|
return
|
|
|
|
num_not_pulled = 0
|
|
|
|
for file_path in file_paths:
|
|
force = force_all
|
|
hash = cache.hash.get(str(file_path), None)
|
|
if hash is None:
|
|
logging.debug(f"Hash not found in cache for {file_path}. Adding to cache.")
|
|
hash = add_file_to_cache(file_path)
|
|
|
|
file_cache = cache.by_path(file_path)
|
|
|
|
last_used_date = datetime.datetime.fromisoformat(file_cache['lastUsed']) if 'lastUsed' in file_cache else None
|
|
|
|
modified = get_file_modification_date(file_path)
|
|
# If file was modified after last used, force metadata pull
|
|
if last_used_date is not None and modified is not None and modified > last_used_date:
|
|
logging.info(f"File was modified after last used. Pulling metadata.")
|
|
force = True
|
|
|
|
# Only skip pull if not forced and civitai is True and recently pulled
|
|
civitai_val = False
|
|
try:
|
|
civitai_val = str_to_bool(file_cache.get('civitai', False))
|
|
except:
|
|
civitai_val = False
|
|
|
|
if not force and civitai_val == True:
|
|
if days_since_last_used(file_path) <= metadata_days_recheck:
|
|
num_not_pulled += 1
|
|
pull_json = False
|
|
|
|
if file_cache.get('blacklist'):
|
|
if not force:
|
|
logging.warning(f"File {file_path} is blacklisted (previously not found). Skipping metadata pull.")
|
|
pull_json = False
|
|
|
|
# If force, recalculate hash before any API call
|
|
if force:
|
|
logging.debug(f"Force flag is set. Recalculating hash for {file_path}.")
|
|
hash = recheck_hash(file_path, hash)
|
|
|
|
if pull_json or force:
|
|
logging.debug(f"Currently pulling metadata for {file_path}.")
|
|
json = get_civitai_model_version_json_by_hash(hash)
|
|
|
|
if 'error' in json:
|
|
retried = False
|
|
dead_model = False
|
|
if 'civitai_error' in json:
|
|
if 'Model not found' in json['civitai_error'] or 'No model with id' in json['civitai_error']:
|
|
dead_model = True
|
|
|
|
# Try fallback with modelId if available
|
|
if dead_model is False:
|
|
if 'modelId' in file_cache:
|
|
logging.debug(f"Using cached model id {file_cache.get('id', None)}")
|
|
json = get_civitai_model_version_json_by_id(file_cache['id'])
|
|
retried = True
|
|
else:
|
|
logging.debug(f"No cached model id.")
|
|
|
|
if 'error' in json:
|
|
if retried:
|
|
logging.error(f"Error: {json['error']}")
|
|
if dead_model:
|
|
file_cache['blacklist'] = True
|
|
print(f"Unable to find on civitai.")
|
|
file_cache['civitai'] = "False"
|
|
file_cache['civitai_failed_count'] = file_cache.get('civitai_failed_count', 0) + 1
|
|
update_cache_without_civitai_json(file_path, hash, timestamp=timestamp)
|
|
|
|
if 'error' not in json:
|
|
update_cache_from_civitai_json(file_path, json, timestamp=timestamp)
|
|
else:
|
|
retries = file_cache.get('civitai_failed_count', 0)
|
|
retries += 1
|
|
file_cache['civitai_failed_count'] = retries
|
|
|
|
cache.hash[file_path] = hash
|
|
cache.info[hash] = file_cache
|
|
if model_type is not None:
|
|
file_cache['model_type'] = model_type
|
|
|
|
if pbar is not None:
|
|
pbar.update(1)
|
|
|
|
if num_not_pulled > 0:
|
|
logging.info(f"Metadata pull complete. Skipped {num_not_pulled} files checked within the last {metadata_days_recheck} days.")
|
|
cache.save()
|
|
|
|
def lora_to_string(lora_name, model_weight, clip_weight):
|
|
lora_string = ' <lora:' + str(pathlib.Path(lora_name).name) + ":" + str(model_weight) + ">" # + ":" + str(clip_weight)
|
|
|
|
return lora_string
|
|
|
|
def lora_to_prompt(lora_stack = None):
|
|
lora_info = ''
|
|
if lora_stack is None:
|
|
return ""
|
|
else:
|
|
if isinstance(lora_stack, tuple) or isinstance(lora_stack, list):
|
|
if not isinstance(lora_stack[0], (list, tuple)):
|
|
lora_stack = [lora_stack]
|
|
for lora in lora_stack:
|
|
lora_info += lora_to_string(lora[0], lora[1], lora[2])
|
|
return lora_info
|
|
|
|
def get_lora_hash(lora_name):
|
|
lora_path = folder_paths.get_full_path_or_raise("loras", lora_name)
|
|
pull_metadata(lora_path)
|
|
|
|
return cache.hash[lora_path]
|
|
|
|
def model_scan(the_path, force = False):
|
|
the_paths = the_path
|
|
|
|
print(f"the_paths: {the_paths}")
|
|
|
|
model_list = []
|
|
for dir in the_paths:
|
|
print(f"dir: {dir}")
|
|
result = list(p.resolve() for p in pathlib.Path(dir).glob("**/*") if p.suffix in MODEL_FILE_EXTENSIONS)
|
|
model_list.extend(result)
|
|
|
|
model_list = list(set(model_list))
|
|
model_list = [str(x) for x in model_list]
|
|
print(f"Scanning {len(model_list)} models for metadata.")
|
|
pbar = comfy.utils.ProgressBar(len(model_list))
|
|
pull_metadata(model_list, force_all=force, pbar=pbar)
|
|
|
|
def grab_model_list(model_type: str, extra_models: list[str] | None = None) -> list[str]:
|
|
"""Get a list of model names based on the model type, including extra models."""
|
|
model_list = folder_paths.get_filename_list(model_type)
|
|
if extra_models is None:
|
|
return model_list
|
|
|
|
for extra_model in extra_models:
|
|
extra = []
|
|
try:
|
|
extra = [x for x in folder_paths.get_filename_list(extra_model)]
|
|
except Exception as e:
|
|
extra = []
|
|
#print(f"Extra models for {extra_model}: {extra}")
|
|
model_list += extra
|
|
model_list = list(set(model_list))
|
|
model_list.sort() # Remove duplicates
|
|
#print(f"Final model list for {model_type}: {model_list}")
|
|
return model_list
|
|
|
|
# Module-level cache for get_model_list
|
|
_model_list_cache = {}
|
|
|
|
def get_model_list(model_type: str) -> list[str]:
|
|
"""Get a list of model names based on the model type, with in-memory cache (1 min)."""
|
|
import time
|
|
global _model_list_cache
|
|
now = time.time()
|
|
cache_entry = _model_list_cache.get(model_type)
|
|
if cache_entry:
|
|
cached_time, cached_list = cache_entry
|
|
if now - cached_time < 60:
|
|
return cached_list
|
|
|
|
# Not cached or cache expired, fetch fresh
|
|
if model_type == "checkpoints":
|
|
result = grab_model_list("checkpoints")
|
|
elif model_type == "unet":
|
|
result = grab_model_list("unet", ["diffusion_models", "unet_gguf"])
|
|
elif model_type == "vae":
|
|
result = grab_model_list("vae")
|
|
elif model_type == "clip":
|
|
result = grab_model_list("clip", ["text_encoders", "clip_gguf"])
|
|
elif model_type == "loras":
|
|
result = grab_model_list("loras")
|
|
else:
|
|
result = []
|
|
|
|
_model_list_cache[model_type] = (now, result)
|
|
return result
|
|
|
|
def normalize_prompt_weights(text):
|
|
"""
|
|
Parse and normalize text prompts with weight formatting.
|
|
|
|
Converts prompts with implicit weights (parentheses) and explicit weights (text):weight
|
|
into a normalized format using only explicit weight notation.
|
|
|
|
Weight rules:
|
|
- Default weight: 1.0 (no parentheses in output)
|
|
- Each layer of parentheses adds 0.1 to the weight
|
|
- Explicit format: (text):weight stacks with surrounding parentheses
|
|
- Weight 1.1 outputs as (text) without weight notation
|
|
- All other weights output as (text:weight)
|
|
|
|
Example input: "The (cat) (jumped over (the dog):1.5)"
|
|
Example output: "The (cat) (jumped over (the dog:1.5))"
|
|
|
|
Cleanup rules applied:
|
|
- Multiple spaces reduced to 1
|
|
- Commas always followed by a space
|
|
- No space after ( or before )
|
|
- Space between ) and (
|
|
- Max one empty line in a row
|
|
- Unmatched parentheses are automatically balanced
|
|
"""
|
|
import re
|
|
|
|
# Balance parentheses
|
|
open_count = text.count('(')
|
|
close_count = text.count(')')
|
|
|
|
if open_count > close_count:
|
|
text = text + ')' * (open_count - close_count)
|
|
elif close_count > open_count:
|
|
text = '(' * (close_count - open_count) + text
|
|
|
|
def parse_weighted_text(text):
|
|
"""Parse text and return list of (content, weight, is_explicit) tuples."""
|
|
result = []
|
|
i = 0
|
|
|
|
while i < len(text):
|
|
# Look for opening parenthesis
|
|
paren_start = text.find('(', i)
|
|
|
|
if paren_start == -1:
|
|
# No more parentheses, add remaining text
|
|
if i < len(text):
|
|
result.append((text[i:], 1.0, False))
|
|
break
|
|
|
|
# Add text before parenthesis
|
|
if paren_start > i:
|
|
result.append((text[i:paren_start], 1.0, False))
|
|
|
|
# Find matching closing parenthesis and parse content
|
|
paren_depth = 0
|
|
j = paren_start
|
|
|
|
while j < len(text):
|
|
if text[j] == '(':
|
|
paren_depth += 1
|
|
elif text[j] == ')':
|
|
paren_depth -= 1
|
|
if paren_depth == 0:
|
|
break
|
|
j += 1
|
|
|
|
if paren_depth != 0:
|
|
# Unmatched parenthesis, treat as regular text
|
|
result.append((text[i:], 1.0, False))
|
|
break
|
|
|
|
# Extract content between parentheses
|
|
content = text[paren_start + 1:j]
|
|
|
|
# Check for explicit weight specification
|
|
weight_match = re.match(r'^(.*?):\s*([0-9]*\.?[0-9]+)\s*$', content)
|
|
|
|
if weight_match:
|
|
# Explicit weight specified
|
|
inner_content = weight_match.group(1).strip()
|
|
weight = float(weight_match.group(2))
|
|
result.append((inner_content, weight, True))
|
|
else:
|
|
# Implicit weight: recurse to calculate based on nested parentheses
|
|
nested_parts = parse_weighted_text(content)
|
|
|
|
# Adjust weights: add 0.1 for this level, for both implicit and explicit weights
|
|
for nested_content, nested_weight, is_explicit in nested_parts:
|
|
# Add 0.1 to both implicit and explicit weights for this parenthesis level
|
|
result.append((nested_content, nested_weight + 0.1, is_explicit))
|
|
|
|
i = j + 1
|
|
|
|
return result
|
|
|
|
def format_output(parts):
|
|
"""Format parsed parts into normalized output."""
|
|
output = []
|
|
|
|
for content, weight, is_explicit in parts:
|
|
# Round weight to avoid floating-point precision issues
|
|
weight = round(weight, 10)
|
|
|
|
# Recursively handle nested content
|
|
if '(' in content:
|
|
formatted = format_output(parse_weighted_text(content))
|
|
if weight == 1.0:
|
|
output.append(formatted)
|
|
elif weight == 1.1:
|
|
output.append(f"({formatted})")
|
|
else:
|
|
output.append(f"({formatted}:{weight})")
|
|
else:
|
|
content = content.strip()
|
|
if content:
|
|
if weight == 1.0:
|
|
output.append(content)
|
|
elif weight == 1.1:
|
|
output.append(f"({content})")
|
|
else:
|
|
output.append(f"({content}:{weight})")
|
|
|
|
return ' '.join(output)
|
|
|
|
# Parse the text
|
|
parts = parse_weighted_text(text)
|
|
|
|
# Format output
|
|
result = format_output(parts)
|
|
|
|
# Apply text cleanup
|
|
# Multiple spaces to single space
|
|
result = re.sub(r' +', ' ', result)
|
|
|
|
# Commas always followed by space
|
|
result = re.sub(r',\s*', ', ', result)
|
|
|
|
# Clean up spacing around parentheses
|
|
result = re.sub(r'\(\s+', '(', result) # No space after (
|
|
result = re.sub(r'\s+\)', ')', result) # No space before )
|
|
result = re.sub(r'\)\s+', ')', result) # No space after ) initially
|
|
|
|
# Add space after ) if followed by alphanumeric or (
|
|
result = re.sub(r'\)(?=[a-zA-Z0-9(])', ') ', result)
|
|
|
|
# Handle line breaks - max one empty line
|
|
result = re.sub(r'\n\s*\n\s*\n+', '\n\n', result)
|
|
|
|
# Clean leading/trailing whitespace
|
|
result = result.strip()
|
|
|
|
return result
|
|
|
|
|
|
def clean_keywords(keywords):
|
|
keywords = set(filter(None, (x.strip() for x in keywords)))
|
|
return ', '.join(keywords)
|
|
|
|
def clean_text(text):
|
|
ret = normalize_prompt_weights(text)
|
|
return ret
|
|
|
|
def clean_if_needed(text, clean):
|
|
return clean_text(text) if clean and text is not None else text
|
|
|
|
def condition_text(clip, text = None):
|
|
zero_text = text is None
|
|
text = text or ""
|
|
|
|
tokens = clip.tokenize(text)
|
|
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
|
|
cond = output.pop("cond")
|
|
|
|
if zero_text:
|
|
pooled_output = output.get("pooled_output")
|
|
if pooled_output is not None:
|
|
output["pooled_output"] = torch.zeros_like(pooled_output)
|
|
return [[torch.zeros_like(cond), output]]
|
|
|
|
return [[cond, output]]
|
|
|
|
|
|
def get_save_file_path(filename_prefix: str = "text", filename_ext: str = "txt") -> str:
|
|
"""
|
|
Generate a safe file path for saving files with automatic counter increment.
|
|
|
|
Args:
|
|
filename_prefix: Base filename, can include date/time variables like %year%, %month%, etc.
|
|
filename_ext: File extension (without dot)
|
|
|
|
Returns:
|
|
Complete file path including directory and filename with counter
|
|
"""
|
|
|
|
def _extract_counter_from_filename(filename: str) -> tuple[int, str]:
|
|
"""Extract counter from existing filename to determine next counter value."""
|
|
base_name = pathlib.Path(filename_prefix).name
|
|
prefix_len = len(base_name)
|
|
|
|
if len(filename) <= prefix_len + 1:
|
|
return 0, filename[:prefix_len + 1]
|
|
|
|
prefix = filename[:prefix_len + 1]
|
|
try:
|
|
# Remove file extension first, then extract counter
|
|
filename_no_ext = pathlib.Path(filename).stem
|
|
counter_part = filename_no_ext[prefix_len + 1:]
|
|
digits = int(counter_part)
|
|
except (ValueError, IndexError):
|
|
digits = 0
|
|
return digits, prefix
|
|
|
|
def _replace_date_variables(text: str) -> str:
|
|
"""Replace date/time variables in the filename prefix."""
|
|
now = datetime.datetime.now()
|
|
replacements = {
|
|
"%year%": str(now.year),
|
|
"%month%": str(now.month).zfill(2),
|
|
"%day%": str(now.day).zfill(2),
|
|
"%hour%": str(now.hour).zfill(2),
|
|
"%minute%": str(now.minute).zfill(2),
|
|
"%second%": str(now.second).zfill(2),
|
|
}
|
|
|
|
for placeholder, value in replacements.items():
|
|
text = text.replace(placeholder, value)
|
|
return text
|
|
|
|
# Get output directory and process filename prefix
|
|
output_dir = folder_paths.get_output_directory()
|
|
|
|
if "%" in filename_prefix:
|
|
filename_prefix = _replace_date_variables(filename_prefix)
|
|
|
|
# Parse the filename prefix path
|
|
filename_prefix_path = pathlib.Path(filename_prefix)
|
|
subfolder = filename_prefix_path.parent
|
|
base_filename = filename_prefix_path.name
|
|
|
|
# Construct full output path
|
|
output_path = pathlib.Path(output_dir)
|
|
full_output_folder = output_path / subfolder
|
|
|
|
# Security check: ensure we're not saving outside the output directory
|
|
try:
|
|
full_output_folder.resolve().relative_to(output_path.resolve())
|
|
except ValueError:
|
|
error_msg = (
|
|
"ERROR: Saving outside the output folder is not allowed.\n"
|
|
f" Target folder: {full_output_folder.resolve()}\n"
|
|
f" Output directory: {output_path.resolve()}"
|
|
)
|
|
logging.error(error_msg)
|
|
raise ValueError(error_msg)
|
|
|
|
# Ensure output directory exists
|
|
full_output_folder.mkdir(parents=True, exist_ok=True)
|
|
|
|
# Find the next available counter
|
|
counter = 1
|
|
try:
|
|
existing_files = [f.name for f in full_output_folder.iterdir() if f.is_file()]
|
|
matching_counters = []
|
|
|
|
for file in existing_files:
|
|
digits, prefix = _extract_counter_from_filename(file)
|
|
# Check if this file matches our pattern (same base name and ends with underscore)
|
|
if (prefix[:-1].lower() == base_filename.lower() and
|
|
len(prefix) > 0 and prefix[-1] == "_"):
|
|
matching_counters.append(digits)
|
|
|
|
if matching_counters:
|
|
counter = max(matching_counters) + 1
|
|
|
|
except Exception as e:
|
|
logging.warning(f"Error finding existing files, using counter=1: {e}")
|
|
counter = 1
|
|
|
|
# Generate final filename
|
|
final_filename = f"{base_filename}_{counter:05d}.{filename_ext}"
|
|
|
|
return str(full_output_folder / final_filename)
|
|
|
|
def unwrap_tuple(value):
|
|
"""Unwrap single-item tuples to their contained value."""
|
|
return value[0] if isinstance(value, tuple) and len(value) == 1 else value
|