File size: 8,867 Bytes
ea00d60
 
 
2df371c
ea00d60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2df371c
ea00d60
2df371c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3c331b7
ea00d60
2df371c
 
 
 
 
 
 
 
 
ea00d60
 
 
 
 
 
 
 
 
3c331b7
 
 
 
 
 
 
 
 
 
 
 
ea00d60
2df371c
 
ea00d60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2df371c
 
ea00d60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2df371c
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
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()
            }