Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
134 changes: 66 additions & 68 deletions workers/DreamBooth-v1/docker_example/rp_custom_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,15 +3,37 @@
'''

import os
import re
import wget
import subprocess
from subprocess import call, check_output


# Only allow "org/repo"-style Hugging Face identifiers: no shell metacharacters,
# no path traversal, no flag-injection via a leading "-".
HF_REPO_ID_RE = re.compile(r"^[A-Za-z0-9_.\-]+/[A-Za-z0-9_.\-]+$")


def _run(command, error_prefix):
'''
Run a command as an argv list (never shell=True) so that untrusted job
input interpolated into any argument can't be interpreted as additional
shell commands.
'''
result = subprocess.run(command, shell=False, stderr=subprocess.PIPE, check=False)
if result.returncode != 0:
raise RuntimeError(
f"{error_prefix}: {' '.join(command)}\nError message: {result.stderr.decode('utf-8')}")
return result


def downloadmodel_hf(Path_to_HuggingFace, huggingface_token=None):
'''
Download model from HuggingFace.
'''
if not HF_REPO_ID_RE.match(Path_to_HuggingFace or ""):
raise ValueError(
f"Invalid Hugging Face repo id, expected 'org/repo': {Path_to_HuggingFace!r}")

if huggingface_token:
auth = f'https://USER:{huggingface_token}@'
else:
Expand All @@ -23,26 +45,30 @@ def downloadmodel_hf(Path_to_HuggingFace, huggingface_token=None):
print(f"Current working directory: {os.getcwd()}")

os.chdir(custom_path)
commands = [
"git init",
"git lfs install --system --skip-repo",
f'git remote add -f origin {auth}huggingface.co/{Path_to_HuggingFace}',
"git config core.sparsecheckout true",
'echo -e "\nscheduler\ntext_encoder\ntokenizer\nunet\nvae\nmodel_index.json\n!*.safetensors" > .git/info/sparse-checkout',
"git pull origin main"
]

for command in commands:
result = subprocess.run(command, shell=True, stderr=subprocess.PIPE, check=False)
if result.returncode != 0:
raise RuntimeError(
f"Error executing command: {command}\nError message: {result.stderr.decode('utf-8')}")

_run(["git", "init"], "Error executing command")
_run(["git", "lfs", "install", "--system", "--skip-repo"], "Error executing command")
_run(
["git", "remote", "add", "-f", "origin", f'{auth}huggingface.co/{Path_to_HuggingFace}'],
"Error executing command"
)
_run(["git", "config", "core.sparsecheckout", "true"], "Error executing command")

# Write the sparse-checkout config directly instead of shelling out to
# `echo ... > file`, which requires shell interpretation of redirection.
os.makedirs(".git/info", exist_ok=True)
with open(".git/info/sparse-checkout", "w", encoding="utf-8") as sparse_checkout_file:
sparse_checkout_file.write(
"\nscheduler\ntext_encoder\ntokenizer\nunet\nvae\nmodel_index.json\n!*.safetensors\n"
)

_run(["git", "pull", "origin", "main"], "Error executing command")

print("Successfully downloaded model from HuggingFace.")

if os.path.exists('unet/diffusion_pytorch_model.bin'):
call("rm -r .git", shell=True)
call("rm model_index.json", shell=True)
_run(["rm", "-r", ".git"], "Error executing command")
_run(["rm", "model_index.json"], "Error executing command")
wget.download(
'https://raw.githubusercontent.com/TheLastBen/fast-stable-diffusion/main/Dreambooth/model_index.json')
os.chdir('/src')
Expand All @@ -57,53 +83,27 @@ def downloadmodel_lnk(ckpt_link):
'''
Download a model from a ckpt link.
'''
result = subprocess.run(
f"gdown --fuzzy -O model.ckpt {ckpt_link}",
shell=True, stderr=subprocess.PIPE, check=False
if not re.match(r"^https?://", ckpt_link or ""):
raise ValueError(f"Invalid checkpoint link, must be an http(s) URL: {ckpt_link!r}")

_run(
["gdown", "--fuzzy", "-O", "model.ckpt", ckpt_link],
f"Error downloading model from link: {ckpt_link}"
)
if result.returncode != 0:
raise RuntimeError(
f"Error downloading model from link: {ckpt_link}\nError message: {result.stderr.decode('utf-8')}")

if os.path.exists('model.ckpt') and os.path.getsize("model.ckpt") > 1810671599:
# wget.download('https://github.com/TheLastBen/fast-stable-diffusion/raw/main/Dreambooth/det.py')
# custom_model_version = check_output(
# 'python det.py --MODEL_PATH /src/model.ckpt', shell=True).decode('utf-8').replace('\n', '')

# if custom_model_version == 'v1.5':
wget.download(
'https://github.com/CompVis/stable-diffusion/raw/main/configs/stable-diffusion/v1-inference.yaml', 'config.yaml')
subprocess.run(
'python /src/diffusers/scripts/convert_original_stable_diffusion_to_diffusers.py --checkpoint_path /src/model.ckpt --dump_path /src/stable-diffusion-custom --original_config_file config.yaml',
shell=True, check=True)

# refmdlz_file = 'refmdlz'
# wget.download(
# f'https://github.com/TheLastBen/fast-stable-diffusion/raw/main/Dreambooth/{refmdlz_file}')

# if not os.path.exists(refmdlz_file):
# raise RuntimeError(f"Error downloading {refmdlz_file}")

# subprocess.run(f'unzip -o -q {refmdlz_file}', shell=True, check=True)
# subprocess.run(f'rm -f {refmdlz_file}', shell=True, check=True)

# wget.download(
# 'https://raw.githubusercontent.com/TheLastBen/fast-stable-diffusion/main/Dreambooth/convertodiffv1.py')

# # result = subprocess.run(
# # 'python convertodiffv1.py model.ckpt /src/stable-diffusion-custom --v1',
# # shell=True, stderr=subprocess.PIPE, check=False
# # )
# result = subprocess.run(
# '/src/diffusers/scripts/convert_original_stable_diffusion_to_diffusers.py - -checkpoint_path /src/model.ckpt - -dump_path /src/stable-diffusion-custom - -original_config_file config.yaml ',
# shell=True, stderr=subprocess.PIPE, check=False)

# if result.returncode != 0:
# raise RuntimeError(
# f"Error executing convert_original_stable_diffusion_to_diffusers.py\nError message: {result.stderr.decode('utf-8')}")

# subprocess.run('rm convertodiffv1.py', shell=True, check=True)
# subprocess.run('rm -r refmdl', shell=True, check=True)
'https://github.com/CompVis/stable-diffusion/raw/main/configs/stable-diffusion/v1-inference.yaml',
'config.yaml')
_run(
[
"python", "/src/diffusers/scripts/convert_original_stable_diffusion_to_diffusers.py",
"--checkpoint_path", "/src/model.ckpt",
"--dump_path", "/src/stable-diffusion-custom",
"--original_config_file", "config.yaml"
],
"Error converting checkpoint to diffusers format"
)


def selected_model(path_to_huggingface=None, ckpt_link=None, huggingface_token=None):
Expand All @@ -121,14 +121,12 @@ def selected_model(path_to_huggingface=None, ckpt_link=None, huggingface_token=N
downloadmodel_lnk(ckpt_link)
model_name = "/src/stable-diffusion-custom"

# Modify the config.json file
result = subprocess.run(
f"sed -i 's@\"sample_size\": 256,@\"sample_size\": 512,@g' {model_name}/vae/config.json",
shell=True, stderr=subprocess.PIPE, check=False
# Modify the config.json file. model_name is one of the two hardcoded
# paths above (never derived from job input), but this is converted to
# argv form too so the file has no remaining shell=True usage.
_run(
["sed", "-i", 's@"sample_size": 256,@"sample_size": 512,@g', f"{model_name}/vae/config.json"],
"Error modifying config.json"
)

if result.returncode != 0:
raise RuntimeError(
f"Error modifying config.json\nError message: {result.stderr.decode('utf-8')}")

return model_name
30 changes: 29 additions & 1 deletion workers/Real-ESRGAN/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from runpod.serverless.utils.rp_validator import validate

import os
import stat
import cv2
import zipfile
from basicsr.archs.rrdbnet_arch import RRDBNet
Expand Down Expand Up @@ -64,6 +65,33 @@ def is_image_file(filename):
return any(filename.endswith(extension) for extension in [".png", ".jpg", ".jpeg", ".bmp", ".tiff"])


def safe_extract(zip_ref, dest_dir):
'''
Safely extract a zip archive to dest_dir.

Rejects any member whose resolved path would escape dest_dir (zip-slip /
path traversal via "../" or absolute paths in the archive entry name),
and rejects symlink entries outright, since a symlink can point outside
dest_dir and defeat a write-time path check. `data_url` is untrusted,
user-supplied job input, so the archive it points to must be treated as
hostile.
'''
dest_dir = os.path.realpath(dest_dir)

for member in zip_ref.infolist():
# Reject symlink entries - they can be extracted pointing outside
# dest_dir regardless of what their declared name looks like.
mode = (member.external_attr >> 16) & 0xFFFF
if stat.S_ISLNK(mode):
raise ValueError(f"Unsafe zip entry (symlink) rejected: {member.filename}")

member_path = os.path.realpath(os.path.join(dest_dir, member.filename))
if member_path != dest_dir and not member_path.startswith(dest_dir + os.sep):
raise ValueError(f"Unsafe zip entry (path traversal) rejected: {member.filename}")

zip_ref.extractall(dest_dir)


def process_image(upsampler, validated_input, image_path, job_id):
img = cv2.imread(image_path, cv2.IMREAD_UNCHANGED)
img_mode = 'RGBA' if len(img.shape) == 3 and img.shape[2] == 4 else None
Expand Down Expand Up @@ -137,7 +165,7 @@ def handler(job):
result_paths = []
if zipfile.is_zipfile(data_path):
with zipfile.ZipFile(data_path, 'r') as zip_ref:
zip_ref.extractall(temp_dir)
safe_extract(zip_ref, temp_dir)

for file in os.listdir(temp_dir):
if not file.startswith("__MACOSX") and is_image_file(file):
Expand Down
Loading