From 81c8781b293248c072e5ecda9ddcb4eb6de1fb51 Mon Sep 17 00:00:00 2001 From: Fillip Date: Mon, 19 Aug 2024 06:28:53 -0700 Subject: [PATCH] Add CSV Save + BF16 upscale --- __init__.py | 3 +++ nodes/FL_SaveCSV.py | 43 ++++++++++++++++++++++++++++++++++++++++ nodes/FL_UpscaleModel.py | 19 +++++++++++++----- 3 files changed, 60 insertions(+), 5 deletions(-) create mode 100644 nodes/FL_SaveCSV.py diff --git a/__init__.py b/__init__.py index 8f73b64..004d6b9 100644 --- a/__init__.py +++ b/__init__.py @@ -50,6 +50,7 @@ from .nodes.FL_KsamplerPlus import FL_KsamplerPlus from .nodes.FL_KsamplerBasic import FL_KsamplerBasic from .nodes.FL_KsamplerFractals import FL_FractalKSampler from .nodes.FL_UpscaleModel import FL_UpscaleModel +from .nodes.FL_SaveCSV import FL_SaveCSV @@ -107,6 +108,7 @@ NODE_CLASS_MAPPINGS = { "FL_KsamplerBasic": FL_KsamplerBasic, "FL_FractalKSampler": FL_FractalKSampler, "FL_UpscaleModel": FL_UpscaleModel, + "FL_SaveCSV": FL_SaveCSV, } @@ -163,6 +165,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FL_KsamplerBasic": "FL KSampler Basic", "FL_FractalKSampler": "FL Fractal KSampler", "FL_UpscaleModel": "FL Upscale Model", + "FL_SaveCSV": "FL Save CSV", } diff --git a/nodes/FL_SaveCSV.py b/nodes/FL_SaveCSV.py new file mode 100644 index 0000000..13f164e --- /dev/null +++ b/nodes/FL_SaveCSV.py @@ -0,0 +1,43 @@ +import os +import comfy.utils + +class FL_SaveCSV: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "csv_data": ("CSV",), + "output_directory": ("STRING", {"default": ""}), + "filename": ("STRING", {"default": "captions.csv"}), + }, + } + + RETURN_TYPES = () + FUNCTION = "save_csv" + CATEGORY = "🏵️Fill Nodes/Captioning" + OUTPUT_NODE = True + + def save_csv(self, csv_data, output_directory, filename): + # Ensure the output directory exists + os.makedirs(output_directory, exist_ok=True) + + # Construct the full file path + file_path = os.path.join(output_directory, filename) + + # Ensure the filename ends with .csv + if not file_path.lower().endswith('.csv'): + file_path += '.csv' + + # Write the CSV data to the file + try: + with open(file_path, 'wb') as f: + f.write(csv_data) + print(f"CSV file saved successfully: {file_path}") + except Exception as e: + print(f"Error saving CSV file: {str(e)}") + + return () + + @classmethod + def IS_CHANGED(cls, csv_data, output_directory, filename): + return float("NaN") \ No newline at end of file diff --git a/nodes/FL_UpscaleModel.py b/nodes/FL_UpscaleModel.py index 6d3a88b..6fad024 100644 --- a/nodes/FL_UpscaleModel.py +++ b/nodes/FL_UpscaleModel.py @@ -3,10 +3,9 @@ import comfy from comfy_extras.nodes_upscale_model import ImageUpscaleWithModel from tqdm import tqdm - class FL_UpscaleModel: rescale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] - precision_options = ["16", "32"] + precision_options = ["auto", "32", "16", "bfloat16"] RETURN_TYPES = ("IMAGE",) FUNCTION = "upscale" @@ -42,11 +41,21 @@ class FL_UpscaleModel: original_device = image.device original_dtype = image.dtype - if precision == "16": - dtype = torch.float16 + # Determine the appropriate dtype based on precision and device + if precision == "auto": + dtype = torch.float16 if original_device.type == "cuda" else torch.float32 + elif precision == "16": + dtype = torch.float16 if original_device.type == "cuda" else torch.bfloat16 + elif precision == "bfloat16": + dtype = torch.bfloat16 else: dtype = torch.float32 + # Ensure the chosen dtype is supported on the current device + if dtype == torch.float16 and original_device.type != "cuda": + print("Warning: float16 is not supported on CPU. Falling back to bfloat16.") + dtype = torch.bfloat16 + upscale_model = upscale_model.to(dtype).to(original_device) # Split the input batch into a list of individual images @@ -62,7 +71,7 @@ class FL_UpscaleModel: batch = torch.cat(image_list[i:i + batch_size]).to(dtype) with torch.no_grad(): - if dtype == torch.float16: + if dtype in [torch.float16, torch.bfloat16]: with torch.autocast(device_type=original_device.type, dtype=dtype): upscaled_batch = self.__imageScaler.upscale(upscale_model, batch)[0] else: