From 4f68674ae73b9a5673173aa4fa9ed3c1e84e43a9 Mon Sep 17 00:00:00 2001 From: qwaezrx Date: Sat, 2 Dec 2023 03:46:47 +0900 Subject: [PATCH] . --- .gitignore | 161 +++++++++++++++++++++++++++ __init__.py | 5 + lora.py | 128 ++++++++++++++++++++++ requirements.txt | 3 + s3_utils.py | 277 +++++++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 574 insertions(+) create mode 100644 .gitignore create mode 100644 requirements.txt create mode 100644 s3_utils.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..2bc1403 --- /dev/null +++ b/.gitignore @@ -0,0 +1,161 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/#use-with-ide +.pdm.toml + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +#.idea/ +runpod.toml diff --git a/__init__.py b/__init__.py index dd38eea..9ec126b 100644 --- a/__init__.py +++ b/__init__.py @@ -1,9 +1,14 @@ from .lora import * +from dotenv import load_dotenv + +load_dotenv() LIVE_NODE_CLASS_MAPPINGS = { "XL DreamBooth LoRA": XLDB_LoRA, + "S3 Bucket LoRA": S3Bucket_Load_LoRA, } LIVE_NODE_DISPLAY_NAME_MAPPINGS = { "XL DreamBooth LoRA": "XL DreamBooth LoRA", + "S3 Bucket LoRA": "S3 Bucket LoRA" } \ No newline at end of file diff --git a/lora.py b/lora.py index cfa1680..75611dd 100644 --- a/lora.py +++ b/lora.py @@ -6,6 +6,7 @@ import comfy.sd import comfy.utils import folder_paths from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from s3_utils import download_file sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy")) @@ -146,13 +147,140 @@ class XLDB_LoRA: lora_weights = torch.load(lora_path) return lora_weights + +class S3Bucket_Load_LoRA: + def __init__(self): + self.loaded_lora = None + + @classmethod + def INPUT_TYPES(s): + """ + Return a dictionary which contains config for all input fields. + Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT". + Input types "INT", "STRING" or "FLOAT" are special values for fields on the node. + The type can be a list for selection. + + Returns: `dict`: + - Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required` + - Value input_fields (`dict`): Contains input fields config: + * Key field_name (`string`): Name of a entry-point method's argument + * Value field_config (`tuple`): + + First value is a string indicate the type of field or a list for selection. + + Secound value is a config for type "INT", "STRING" or "FLOAT". + """ + # file_list = folder_paths.get_filename_list("loras") + # file_list.insert(0, "None") + + return { + "required": { + "model": ("MODEL",), + "clip": ("CLIP",), + "lora_name": ("STRING", { + "multiline": False, + }), + "strength_model": ("FLOAT", { + "default": 1.0, + "min": 0.0, + "max": 2.0, + "step": 0.01, + }), + "strength_clip": ("FLOAT", { + "default": 1.0, + "min": 0.0, + "max": 2.0, + "step": 0.01, + }), + } + } + return { + "required": { + "image": ("IMAGE",), + "int_field": ("INT", { + "default": 0, + "min": 0, #Minimum value + "max": 4096, #Maximum value + "step": 64, #Slider's step + "display": "number" # Cosmetic only: display as "number" or "slider" + }), + "float_field": ("FLOAT", { + "default": 1.0, + "min": 0.0, + "max": 10.0, + "step": 0.01, + "round": 0.001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding. + "display": "number"}), + "print_to_screen": (["enable", "disable"],), + "string_field": ("STRING", { + "multiline": False, #True if you want the field to look like the one on the ClipTextEncode node + "default": "Hello World!" + }), + }, + } + + RETURN_TYPES = ("MODEL", "CLIP") + FUNCTION = "load_lora" + + + def load_lora(self, model, clip, lora_name, strength_model, strength_clip): + """ + The entry point method. The name of this method must be the same as the value of property `FUNCTION`. + For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. + + Arguments: + - model (`MODEL`): The model object + - clip (`CLIP`): The clip object + - lora_name (`string`): The name of the lora file + + Returns: `tuple`: + - First value is a `MODEL` object + - Secound value is a `CLIP` object + """ + + if lora_name == "None": + return model, clip + lora_path = folder_paths.get_full_path("loras", lora_name) + lora = None + if not os.path.exists(lora_path): + download_file(bucket_path=lora_name, file_path=lora_path) + + if self.loaded_lora is not None: + if self.loaded_lora[0] == lora_path: + lora = self.loaded_lora[1] + else: + del self.loaded_lora + + if lora is None: + if lora_path and "checkpoint" in lora_path: + + lora = self.load_checkpoint_lora(model, clip, lora_path) + self.loaded_lora = (lora_path, lora) + else: + lora = comfy.utils.load_torch_file(lora_path, safe_load=True) + self.loaded_lora = (lora_path, lora) + + model_lora, clip_lora = comfy.sd.load_lora_for_models( + model, clip, lora, strength_model, strength_clip) + + return model_lora, clip_lora + + + def load_checkpoint_lora(self, model, clip, lora_path): + lora_path = Path(lora_path).parent + lora_path = os.path.join(lora_path, "pytorch_lora_weights.bin") + lora_path = str(lora_path) + lora_weights = torch.load(lora_path) + return lora_weights + + # A dictionary that contains all nodes you want to export with their names # NOTE: names should be globally unique NODE_CLASS_MAPPINGS = { "XLDB_LoRA": XLDB_LoRA, + "S3Bucket_Load_LoRA": S3Bucket_Load_LoRA, } # A dictionary that contains the friendly/humanly readable titles for the nodes NODE_DISPLAY_NAME_MAPPINGS = { "XLDB_LoRA": "XLDB LoRA", + "S3Bucket_Load_LoRA": "S3 Bucket Load LoRA", } diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..06c8e7b --- /dev/null +++ b/requirements.txt @@ -0,0 +1,3 @@ + +boto3 +python-dotenv \ No newline at end of file diff --git a/s3_utils.py b/s3_utils.py new file mode 100644 index 0000000..1c8c877 --- /dev/null +++ b/s3_utils.py @@ -0,0 +1,277 @@ +''' PodWorker | modules | upload.py ''' +# pylint: disable=too-many-arguments + +import os +import io +import time +import uuid +import shutil +import logging +import threading +import multiprocessing +from urllib.parse import urlparse +from typing import Optional, Tuple + +import boto3 +from boto3 import session +from boto3.s3.transfer import TransferConfig +from botocore.config import Config + + +logger = logging.getLogger("runpod upload utility") +FMT = "%(filename)-20s:%(lineno)-4d %(asctime)s %(message)s" +logging.basicConfig(level=logging.INFO, format=FMT, handlers=[logging.StreamHandler()]) + + +def extract_region_from_url(endpoint_url): + """ + Extracts the region from the endpoint URL. + """ + parsed_url = urlparse(endpoint_url) + # AWS/backblaze S3-like URL + if '.s3.' in endpoint_url: + return endpoint_url.split('.s3.')[1].split('.')[0] + + # DigitalOcean Spaces-like URL + if parsed_url.netloc.endswith('.digitaloceanspaces.com'): + return endpoint_url.split('.')[1].split('.digitaloceanspaces.com')[0] + + return None + +def extract_bucket_name_from_url(endpoint_url): + parsed_url = urlparse(endpoint_url) + if '.s3.' in endpoint_url: + return endpoint_url.split('.s3.')[0].split('/')[-1] + + return None + +# --------------------------- S3 Bucket Connection --------------------------- # +def get_boto_client( + bucket_creds: Optional[dict] = None) -> Tuple[boto3.client, TransferConfig]: # pragma: no cover # pylint: disable=line-too-long + ''' + Returns a boto3 client and transfer config for the bucket. + ''' + bucket_session = session.Session() + + boto_config = Config( + signature_version='s3v4', + retries={ + 'max_attempts': 3, + 'mode': 'standard' + } + ) + + transfer_config = TransferConfig( + multipart_threshold=1024 * 25, + max_concurrency=multiprocessing.cpu_count(), + multipart_chunksize=1024 * 25, + use_threads=True + ) + + if bucket_creds: + endpoint_url = bucket_creds['endpointUrl'] + access_key_id = bucket_creds['accessId'] + secret_access_key = bucket_creds['accessSecret'] + else: + endpoint_url = os.environ.get('BUCKET_ENDPOINT_URL', None) + access_key_id = os.environ.get('BUCKET_ACCESS_KEY_ID', None) + secret_access_key = os.environ.get('BUCKET_SECRET_ACCESS_KEY', None) + + if endpoint_url and access_key_id and secret_access_key: + # Extract region from the endpoint URL + region = extract_region_from_url(endpoint_url) + + boto_client = bucket_session.client( + 's3', + endpoint_url=endpoint_url, + aws_access_key_id=access_key_id, + aws_secret_access_key=secret_access_key, + config=boto_config, + region_name=region + ) + else: + boto_client = None + + return boto_client, transfer_config + + +def download_file(bucket_path, file_path): + boto_client, _ = get_boto_client() + bucket_name = extract_bucket_name_from_url(os.environ.get('BUCKET_ENDPOINT_URL', None)) + boto_client.download_file( + bucket_name, + bucket_path, + file_path + ) + +# ---------------------------------------------------------------------------- # +# Upload Image # +# ---------------------------------------------------------------------------- # +def upload_image(job_id, image_location, result_index=0, results_list=None): # pragma: no cover + ''' + Upload a single file to bucket storage. + ''' + image_name = str(uuid.uuid4())[:8] + boto_client, _ = get_boto_client() + file_extension = os.path.splitext(image_location)[1] + content_type = "image/" + file_extension.lstrip(".") + + with open(image_location, "rb") as input_file: + output = input_file.read() + + if boto_client is None: + # Save the output to a file + print("No bucket endpoint set, saving to disk folder 'simulated_uploaded'") + print("If this is a live endpoint, please reference the following:") + print("https://github.com/runpod/runpod-python/blob/main/docs/serverless/utils/rp_upload.md") # pylint: disable=line-too-long + + os.makedirs("simulated_uploaded", exist_ok=True) + sim_upload_location = f"simulated_uploaded/{image_name}{file_extension}" + + with open(sim_upload_location, "wb") as file_output: + file_output.write(output) + + if results_list is not None: + results_list[result_index] = sim_upload_location + + return sim_upload_location + + bucket = time.strftime('%m-%y') + boto_client.put_object( + Bucket=f'{bucket}', + Key=f'{job_id}/{image_name}{file_extension}', + Body=output, + ContentType=content_type + ) + + presigned_url = boto_client.generate_presigned_url( + 'get_object', + Params={ + 'Bucket': f'{bucket}', + 'Key': f'{job_id}/{image_name}{file_extension}' + }, ExpiresIn=604800) + + if results_list is not None: + results_list[result_index] = presigned_url + + return presigned_url + + +# ---------------------------------------------------------------------------- # +# Files To Upload # +# ---------------------------------------------------------------------------- # +def files(job_id, file_list): # pragma: no cover + ''' + Uploads a list of files in parallel. + Once all files are uploaded, the function returns the presigned URLs list. + ''' + upload_progress = [] # List of threads + file_urls = [None] * len(file_list) # Resulting list of URLs for each file + + for index, selected_file in enumerate(file_list): + new_upload = threading.Thread( + target=upload_image, + args=(job_id, selected_file, index, file_urls) + ) + + new_upload.start() + upload_progress.append(new_upload) + + # Wait for all uploads to finish + for upload in upload_progress: + upload.join() + + return file_urls + + +# --------------------------- Custom Bucket Upload --------------------------- # +def bucket_upload(job_id, file_list, bucket_creds): # pragma: no cover + ''' + Uploads files to bucket storage. + ''' + temp_bucket_session = session.Session() + + temp_boto_config = Config( + signature_version='s3v4', + retries={ + 'max_attempts': 3, + 'mode': 'standard' + } + ) + + temp_boto_client = temp_bucket_session.client( + 's3', + endpoint_url=bucket_creds['endpointUrl'], + aws_access_key_id=bucket_creds['accessId'], + aws_secret_access_key=bucket_creds['accessSecret'], + config=temp_boto_config + ) + + bucket_urls = [] + + for selected_file in file_list: + with open(selected_file, 'rb') as file_data: + temp_boto_client.put_object( + Bucket=str(bucket_creds['bucketName']), + Key=f'{job_id}/{selected_file}', + Body=file_data, + ) + + bucket_urls.append( + f"{bucket_creds['endpointUrl']}/{bucket_creds['bucketName']}/{job_id}/{selected_file}") + + return bucket_urls + + +# ------------------------- Single File Bucket Upload ------------------------ # +def upload_file_to_bucket( + file_name: str, file_location: str, + bucket_creds: Optional[dict] = None, + bucket_name: Optional[str] = None, + prefix: Optional[str] = None, + extra_args: Optional[dict] = None +) -> str: # pragma: no cover + ''' + Uploads a single file to bucket storage and returns a presigned URL. + ''' + boto_client, transfer_config = get_boto_client(bucket_creds) + + if not bucket_name: + bucket_name = time.strftime('%m-%y') + + key = f"{prefix}/{file_name}" if prefix else file_name + + if boto_client is None: + print("No bucket endpoint set, saving to disk folder 'local_upload'") + print("If this is a live endpoint, please reference the following:") + print("https://github.com/runpod/runpod-python/blob/main/docs/serverless/utils/rp_upload.md") # pylint: disable=line-too-long + + os.makedirs("local_upload", exist_ok=True) + local_upload_location = f"local_upload/{file_name}" + shutil.copyfile(file_location, local_upload_location) + + return local_upload_location + + file_size = os.path.getsize(file_location) + + upload_file_args = { + "Filename": file_location, + "Bucket": bucket_name, + "Key": key, + "Config": transfer_config, + } + + if extra_args: + upload_file_args["ExtraArgs"] = extra_args + + boto_client.upload_file(**upload_file_args) + + presigned_url = boto_client.generate_presigned_url( + 'get_object', + Params={ + 'Bucket': bucket_name, + 'Key': key + }, ExpiresIn=604800) + + return presigned_url +