This commit is contained in:
qwaezrx
2023-12-02 03:46:47 +09:00
parent 0391a9ca7d
commit 4f68674ae7
5 changed files with 574 additions and 0 deletions
+161
View File
@@ -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
+5
View File
@@ -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"
}
+128
View File
@@ -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",
}
+3
View File
@@ -0,0 +1,3 @@
boto3
python-dotenv
+277
View File
@@ -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