Tim-canova
debugs to hadnler.py
2df371c
Raw History Blame Contribute Delete
8.87 kB
import base64
import io
import torch
import json
from pathlib import Path
from PIL import Image
from huggingface_hub import hf_hub_download
from refiners.foundationals.latent_diffusion import solvers
from enhancer import ESRGANUpscaler, ESRGANUpscalerCheckpoints
class EndpointHandler:
def __init__(self, path=""):
"""Initialize the handler with model checkpoints"""
# Download model checkpoints
self.checkpoints = ESRGANUpscalerCheckpoints(
unet=Path(
hf_hub_download(
repo_id="refiners/juggernaut.reborn.sd1_5.unet",
filename="model.safetensors",
revision="347d14c3c782c4959cc4d1bb1e336d19f7dda4d2",
)
),
clip_text_encoder=Path(
hf_hub_download(
repo_id="refiners/juggernaut.reborn.sd1_5.text_encoder",
filename="model.safetensors",
revision="744ad6a5c0437ec02ad826df9f6ede102bb27481",
)
),
lda=Path(
hf_hub_download(
repo_id="refiners/juggernaut.reborn.sd1_5.autoencoder",
filename="model.safetensors",
revision="3c1aae3fc3e03e4a2b7e0fa42b62ebb64f1a4c19",
)
),
controlnet_tile=Path(
hf_hub_download(
repo_id="refiners/controlnet.sd1_5.tile",
filename="model.safetensors",
revision="48ced6ff8bfa873a8976fa467c3629a240643387",
)
),
esrgan=Path(
hf_hub_download(
repo_id="philz1337x/upscaler",
filename="4x-UltraSharp.pth",
revision="011deacac8270114eb7d2eeff4fe6fa9a837be70",
)
),
negative_embedding=Path(
hf_hub_download(
repo_id="philz1337x/embeddings",
filename="JuggernautNegative-neg.pt",
revision="203caa7e9cc2bc225031a4021f6ab1ded283454a",
)
),
negative_embedding_key="string_to_param.*",
loras={
"more_details": Path(
hf_hub_download(
repo_id="philz1337x/loras",
filename="more_details.safetensors",
revision="a3802c0280c0d00c2ab18d37454a8744c44e474e",
)
),
"sdxl_render": Path(
hf_hub_download(
repo_id="philz1337x/loras",
filename="SDXLrender_v2.0.safetensors",
revision="a3802c0280c0d00c2ab18d37454a8744c44e474e",
)
),
},
)
# Initialize device and dtype
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float32
# Initialize the enhancer
self.enhancer = ESRGANUpscaler(
checkpoints=self.checkpoints,
device=self.device,
dtype=self.dtype
)
def __call__(self, data):
"""Process the input data and return enhanced image"""
try:
# DEBUG: Log what we received
print(f"DEBUG: Received data type: {type(data)}")
print(f"DEBUG: Data keys: {list(data.keys()) if isinstance(data, dict) else 'Not a dict'}")
# Handle different input formats
if "inputs" in data:
inputs = data["inputs"]
print("DEBUG: Found 'inputs' key")
else:
inputs = data
print("DEBUG: Using data directly as inputs")
print(f"DEBUG: Inputs type: {type(inputs)}")
print(f"DEBUG: Inputs keys: {list(inputs.keys()) if isinstance(inputs, dict) else 'Not a dict'}")
# Try to find image in different possible locations
image_data = None
if isinstance(inputs, dict):
image_data = inputs.get("image")
if not image_data:
# Try other possible keys
for key in ["img", "input_image", "data"]:
if key in inputs:
image_data = inputs[key]
print(f"DEBUG: Found image data in '{key}' key")
break
print(f"DEBUG: Image data found: {image_data is not None}")
if image_data:
print(f"DEBUG: Image data length: {len(image_data)}")
print(f"DEBUG: Image data starts with: {image_data[:50] if len(image_data) > 50 else image_data}")
if not image_data:
return {
"error": "No image provided",
"debug_info": {
"data_type": str(type(data)),
"data_keys": list(data.keys()) if isinstance(data, dict) else None,
"inputs_type": str(type(inputs)),
"inputs_keys": list(inputs.keys()) if isinstance(inputs, dict) else None
}
}
# Decode base64 image
if image_data.startswith('data:image'):
image_data = image_data.split(',')[1]
image_bytes = base64.b64decode(image_data)
input_image = Image.open(io.BytesIO(image_bytes)).convert('RGB')
# Extract parameters with defaults
prompt = inputs.get("prompt", "masterpiece, best quality, highres")
negative_prompt = inputs.get("negative_prompt", "worst quality, low quality, normal quality")
seed = inputs.get("seed", 42)
upscale_factor = inputs.get("upscale_factor", 2)
controlnet_scale = inputs.get("controlnet_scale", 0.6)
controlnet_decay = inputs.get("controlnet_decay", 1.0)
condition_scale = inputs.get("condition_scale", 6)
tile_width = inputs.get("tile_width", 112)
tile_height = inputs.get("tile_height", 144)
denoise_strength = inputs.get("denoise_strength", 0.35)
num_inference_steps = inputs.get("num_inference_steps", 18)
solver_name = inputs.get("solver", "DDIM")
print(f"DEBUG: Processing image of size: {input_image.size}")
# Get solver type
solver_type = getattr(solvers, solver_name)
# Set up generator
generator = torch.Generator(device=self.device)
generator.manual_seed(seed)
# Resize input image if too large to avoid VRAM issues
side_size = min(input_image.size)
if side_size > 768:
scale = 768 / side_size
new_size = (int(input_image.width * scale), int(input_image.height * scale))
resized_image = input_image.resize(new_size, resample=Image.Resampling.LANCZOS)
else:
resized_image = input_image
print(f"DEBUG: About to enhance image of size: {resized_image.size}")
# Enhance the image
enhanced_image = self.enhancer.upscale(
image=resized_image,
prompt=prompt,
negative_prompt=negative_prompt,
upscale_factor=upscale_factor,
controlnet_scale=controlnet_scale,
controlnet_scale_decay=controlnet_decay,
condition_scale=condition_scale,
tile_size=(tile_height, tile_width),
denoise_strength=denoise_strength,
num_inference_steps=num_inference_steps,
loras_scale={"more_details": 0.5, "sdxl_render": 1.0},
solver_type=solver_type,
generator=generator,
)
# Convert enhanced image to base64
buffered = io.BytesIO()
enhanced_image.save(buffered, format="PNG")
enhanced_base64 = base64.b64encode(buffered.getvalue()).decode('utf-8')
return {
"enhanced_image": enhanced_base64,
"original_size": input_image.size,
"enhanced_size": enhanced_image.size
}
except Exception as e:
import traceback
return {
"error": f"Enhancement failed: {str(e)}",
"traceback": traceback.format_exc()
}