File size: 7,763 Bytes
2844be8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import base64
import io
import torch
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
        
        Args:
            data (dict): Input data containing:
                - image (str): Base64 encoded input image
                - prompt (str, optional): Enhancement prompt
                - negative_prompt (str, optional): Negative prompt
                - seed (int, optional): Random seed
                - upscale_factor (float, optional): Upscale factor
                - controlnet_scale (float, optional): ControlNet scale
                - controlnet_decay (float, optional): ControlNet decay
                - condition_scale (int, optional): Condition scale
                - tile_width (int, optional): Tile width
                - tile_height (int, optional): Tile height
                - denoise_strength (float, optional): Denoise strength
                - num_inference_steps (int, optional): Number of inference steps
                - solver (str, optional): Solver type
        
        Returns:
            dict: Contains enhanced image as base64 string
        """
        try:
            # Extract and decode input image
            image_data = data.get("image")
            if not image_data:
                return {"error": "No image provided"}
            
            # Decode base64 image
            if image_data.startswith('data:image'):
                # Remove data:image/...;base64, prefix if present
                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 = data.get("prompt", "masterpiece, best quality, highres")
            negative_prompt = data.get("negative_prompt", "worst quality, low quality, normal quality")
            seed = data.get("seed", 42)
            upscale_factor = data.get("upscale_factor", 2)
            controlnet_scale = data.get("controlnet_scale", 0.6)
            controlnet_decay = data.get("controlnet_decay", 1.0)
            condition_scale = data.get("condition_scale", 6)
            tile_width = data.get("tile_width", 112)
            tile_height = data.get("tile_height", 144)
            denoise_strength = data.get("denoise_strength", 0.35)
            num_inference_steps = data.get("num_inference_steps", 18)
            solver_name = data.get("solver", "DDIM")
            
            # 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
            
            # 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:
            return {"error": f"Enhancement failed: {str(e)}"}