moralec/MagicQuill
0
1import os2import time3import logging4from typing import Set, List, Dict, Tuple5 6supported_pt_extensions: Set[str] = set(['.ckpt', '.pt', '.bin', '.pth', '.safetensors', '.pkl'])7 8SupportedFileExtensionsType = Set[str]9ScanPathType = List[str]10folder_names_and_paths: Dict[str, Tuple[ScanPathType, SupportedFileExtensionsType]] = {}11 12base_path = os.path.dirname(os.path.realpath(__file__))13models_dir = os.path.join(base_path, "../models")14 15folder_names_and_paths["checkpoints"] = ([os.path.join(models_dir, "checkpoints")], supported_pt_extensions)16folder_names_and_paths["configs"] = ([os.path.join(models_dir, "configs")], [".yaml"])17 18folder_names_and_paths["loras"] = ([os.path.join(models_dir, "loras")], supported_pt_extensions)19folder_names_and_paths["vae"] = ([os.path.join(models_dir, "vae")], supported_pt_extensions)20folder_names_and_paths["clip"] = ([os.path.join(models_dir, "clip")], supported_pt_extensions)21folder_names_and_paths["unet"] = ([os.path.join(models_dir, "unet")], supported_pt_extensions)22folder_names_and_paths["clip_vision"] = ([os.path.join(models_dir, "clip_vision")], supported_pt_extensions)23folder_names_and_paths["style_models"] = ([os.path.join(models_dir, "style_models")], supported_pt_extensions)24folder_names_and_paths["embeddings"] = ([os.path.join(models_dir, "embeddings")], supported_pt_extensions)25folder_names_and_paths["diffusers"] = ([os.path.join(models_dir, "diffusers")], ["folder"])26folder_names_and_paths["vae_approx"] = ([os.path.join(models_dir, "vae_approx")], supported_pt_extensions)27 28folder_names_and_paths["controlnet"] = ([os.path.join(models_dir, "controlnet"), os.path.join(models_dir, "t2i_adapter")], supported_pt_extensions)29folder_names_and_paths["gligen"] = ([os.path.join(models_dir, "gligen")], supported_pt_extensions)30 31folder_names_and_paths["upscale_models"] = ([os.path.join(models_dir, "upscale_models")], supported_pt_extensions)32 33folder_names_and_paths["hypernetworks"] = ([os.path.join(models_dir, "hypernetworks")], supported_pt_extensions)34 35folder_names_and_paths["photomaker"] = ([os.path.join(models_dir, "photomaker")], supported_pt_extensions)36 37folder_names_and_paths["classifiers"] = ([os.path.join(models_dir, "classifiers")], {""})38 39output_directory = os.path.join(os.path.dirname(os.path.realpath(__file__)), "output")40temp_directory = os.path.join(os.path.dirname(os.path.realpath(__file__)), "temp")41input_directory = os.path.join(os.path.dirname(os.path.realpath(__file__)), "input")42user_directory = os.path.join(os.path.dirname(os.path.realpath(__file__)), "user")43 44filename_list_cache = {}45 46if not os.path.exists(input_directory):47 try:48 os.makedirs(input_directory)49 except:50 logging.error("Failed to create input directory")51 52def set_output_directory(output_dir):53 global output_directory54 output_directory = output_dir55 56def set_temp_directory(temp_dir):57 global temp_directory58 temp_directory = temp_dir59 60def set_input_directory(input_dir):61 global input_directory62 input_directory = input_dir63 64def get_output_directory():65 global output_directory66 return output_directory67 68def get_temp_directory():69 global temp_directory70 return temp_directory71 72def get_input_directory():73 global input_directory74 return input_directory75 76 77#NOTE: used in http server so don't put folders that should not be accessed remotely78def get_directory_by_type(type_name):79 if type_name == "output":80 return get_output_directory()81 if type_name == "temp":82 return get_temp_directory()83 if type_name == "input":84 return get_input_directory()85 return None86 87 88# determine base_dir rely on annotation if name is 'filename.ext [annotation]' format89# otherwise use default_path as base_dir90def annotated_filepath(name):91 if name.endswith("[output]"):92 base_dir = get_output_directory()93 name = name[:-9]94 elif name.endswith("[input]"):95 base_dir = get_input_directory()96 name = name[:-8]97 elif name.endswith("[temp]"):98 base_dir = get_temp_directory()99 name = name[:-7]100 else:101 return name, None102 103 return name, base_dir104 105 106def get_annotated_filepath(name, default_dir=None):107 name, base_dir = annotated_filepath(name)108 109 if base_dir is None:110 if default_dir is not None:111 base_dir = default_dir112 else:113 base_dir = get_input_directory() # fallback path114 115 return os.path.join(base_dir, name)116 117 118def exists_annotated_filepath(name):119 name, base_dir = annotated_filepath(name)120 121 if base_dir is None:122 base_dir = get_input_directory() # fallback path123 124 filepath = os.path.join(base_dir, name)125 return os.path.exists(filepath)126 127 128def add_model_folder_path(folder_name, full_folder_path):129 global folder_names_and_paths130 if folder_name in folder_names_and_paths:131 folder_names_and_paths[folder_name][0].append(full_folder_path)132 else:133 folder_names_and_paths[folder_name] = ([full_folder_path], set())134 135def get_folder_paths(folder_name):136 return folder_names_and_paths[folder_name][0][:]137 138def recursive_search(directory, excluded_dir_names=None):139 if not os.path.isdir(directory):140 return [], {}141 142 if excluded_dir_names is None:143 excluded_dir_names = []144 145 result = []146 dirs = {}147 148 # Attempt to add the initial directory to dirs with error handling149 try:150 dirs[directory] = os.path.getmtime(directory)151 except FileNotFoundError:152 logging.warning(f"Warning: Unable to access {directory}. Skipping this path.")153 154 logging.debug("recursive file list on directory {}".format(directory))155 for dirpath, subdirs, filenames in os.walk(directory, followlinks=True, topdown=True):156 subdirs[:] = [d for d in subdirs if d not in excluded_dir_names]157 for file_name in filenames:158 relative_path = os.path.relpath(os.path.join(dirpath, file_name), directory)159 result.append(relative_path)160 161 for d in subdirs:162 path = os.path.join(dirpath, d)163 try:164 dirs[path] = os.path.getmtime(path)165 except FileNotFoundError:166 logging.warning(f"Warning: Unable to access {path}. Skipping this path.")167 continue168 logging.debug("found {} files".format(len(result)))169 return result, dirs170 171def filter_files_extensions(files, extensions):172 return sorted(list(filter(lambda a: os.path.splitext(a)[-1].lower() in extensions or len(extensions) == 0, files)))173 174 175 176def get_full_path(folder_name, filename):177 global folder_names_and_paths178 if folder_name not in folder_names_and_paths:179 return None180 folders = folder_names_and_paths[folder_name]181 filename = os.path.relpath(os.path.join("/", filename), "/")182 for x in folders[0]:183 full_path = os.path.join(x, filename)184 if os.path.isfile(full_path):185 return full_path186 elif os.path.islink(full_path):187 logging.warning("WARNING path {} exists but doesn't link anywhere, skipping.".format(full_path))188 189 return None190 191def get_filename_list_(folder_name):192 global folder_names_and_paths193 output_list = set()194 folders = folder_names_and_paths[folder_name]195 output_folders = {}196 for x in folders[0]:197 files, folders_all = recursive_search(x, excluded_dir_names=[".git"])198 output_list.update(filter_files_extensions(files, folders[1]))199 output_folders = {**output_folders, **folders_all}200 201 return (sorted(list(output_list)), output_folders, time.perf_counter())202 203def cached_filename_list_(folder_name):204 global filename_list_cache205 global folder_names_and_paths206 if folder_name not in filename_list_cache:207 return None208 out = filename_list_cache[folder_name]209 210 for x in out[1]:211 time_modified = out[1][x]212 folder = x213 if os.path.getmtime(folder) != time_modified:214 return None215 216 folders = folder_names_and_paths[folder_name]217 for x in folders[0]:218 if os.path.isdir(x):219 if x not in out[1]:220 return None221 222 return out223 224def get_filename_list(folder_name):225 out = cached_filename_list_(folder_name)226 if out is None:227 out = get_filename_list_(folder_name)228 global filename_list_cache229 filename_list_cache[folder_name] = out230 return list(out[0])231 232def get_save_image_path(filename_prefix, output_dir, image_width=0, image_height=0):233 def map_filename(filename):234 prefix_len = len(os.path.basename(filename_prefix))235 prefix = filename[:prefix_len + 1]236 try:237 digits = int(filename[prefix_len + 1:].split('_')[0])238 except:239 digits = 0240 return (digits, prefix)241 242 def compute_vars(input, image_width, image_height):243 input = input.replace("%width%", str(image_width))244 input = input.replace("%height%", str(image_height))245 return input246 247 filename_prefix = compute_vars(filename_prefix, image_width, image_height)248 249 subfolder = os.path.dirname(os.path.normpath(filename_prefix))250 filename = os.path.basename(os.path.normpath(filename_prefix))251 252 full_output_folder = os.path.join(output_dir, subfolder)253 254 if os.path.commonpath((output_dir, os.path.abspath(full_output_folder))) != output_dir:255 err = "**** ERROR: Saving image outside the output folder is not allowed." + \256 "\n full_output_folder: " + os.path.abspath(full_output_folder) + \257 "\n output_dir: " + output_dir + \258 "\n commonpath: " + os.path.commonpath((output_dir, os.path.abspath(full_output_folder)))259 logging.error(err)260 raise Exception(err)261 262 try:263 counter = max(filter(lambda a: os.path.normcase(a[1][:-1]) == os.path.normcase(filename) and a[1][-1] == "_", map(map_filename, os.listdir(full_output_folder))))[0] + 1264 except ValueError:265 counter = 1266 except FileNotFoundError:267 os.makedirs(full_output_folder, exist_ok=True)268 counter = 1269 return full_output_folder, filename, counter, subfolder, filename_prefix270 