Tim-canova commited on
Commit
2844be8
·
1 Parent(s): 831471c

Add inference endpoint code for image enhancer

Browse files
Files changed (6) hide show
  1. README.md +41 -0
  2. client_example.py +76 -0
  3. enhancer.py +38 -0
  4. esrgan_model.py +305 -0
  5. inference.py +186 -0
  6. requirements.txt +7 -0
README.md ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Finegrain Image Enhancer - Inference Endpoint
2
+
3
+ This repository contains the Inference Endpoint version of the Finegrain Image Enhancer.
4
+
5
+ ## Usage
6
+
7
+ Send a POST request to the endpoint with:
8
+
9
+ ```json
10
+ {
11
+ "image": "base64_encoded_image_data",
12
+ "prompt": "masterpiece, best quality, highres",
13
+ "upscale_factor": 2.0,
14
+ "seed": 42
15
+ }
16
+ ```
17
+
18
+ Returns:
19
+ ```json
20
+ {
21
+ "enhanced_image": "base64_encoded_enhanced_image",
22
+ "original_size": [width, height],
23
+ "enhanced_size": [width, height]
24
+ }
25
+ ```
26
+
27
+ ## Parameters
28
+
29
+ - `image` (required): Base64 encoded input image
30
+ - `prompt` (optional): Enhancement prompt (default: "masterpiece, best quality, highres")
31
+ - `negative_prompt` (optional): Negative prompt (default: "worst quality, low quality, normal quality")
32
+ - `seed` (optional): Random seed (default: 42)
33
+ - `upscale_factor` (optional): Upscale factor 1-4 (default: 2.0)
34
+ - `controlnet_scale` (optional): ControlNet scale 0-1.5 (default: 0.6)
35
+ - `controlnet_decay` (optional): ControlNet decay 0.5-1 (default: 1.0)
36
+ - `condition_scale` (optional): Condition scale 2-20 (default: 6)
37
+ - `tile_width` (optional): Tile width 64-200 (default: 112)
38
+ - `tile_height` (optional): Tile height 64-200 (default: 144)
39
+ - `denoise_strength` (optional): Denoise strength 0-1 (default: 0.35)
40
+ - `num_inference_steps` (optional): Number of steps 1-30 (default: 18)
41
+ - `solver` (optional): Solver type "DDIM" or "DPMSolver" (default: "DDIM")
client_example.py ADDED
@@ -0,0 +1,76 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import requests
3
+ from PIL import Image
4
+ import io
5
+
6
+ def enhance_image_via_endpoint(
7
+ image_path: str,
8
+ endpoint_url: str,
9
+ hf_token: str,
10
+ prompt: str = "masterpiece, best quality, highres",
11
+ upscale_factor: float = 2.0
12
+ ):
13
+ """
14
+ Call the Hugging Face Inference Endpoint to enhance an image
15
+ """
16
+
17
+ # Load and encode image
18
+ with open(image_path, "rb") as f:
19
+ image_bytes = f.read()
20
+
21
+ image_base64 = base64.b64encode(image_bytes).decode('utf-8')
22
+
23
+ # Prepare request
24
+ headers = {
25
+ "Authorization": f"Bearer {hf_token}",
26
+ "Content-Type": "application/json"
27
+ }
28
+
29
+ payload = {
30
+ "image": image_base64,
31
+ "prompt": prompt,
32
+ "upscale_factor": upscale_factor,
33
+ "seed": 42
34
+ }
35
+
36
+ # Make request
37
+ response = requests.post(endpoint_url, headers=headers, json=payload)
38
+
39
+ if response.status_code == 200:
40
+ result = response.json()
41
+
42
+ if "error" in result:
43
+ raise Exception(f"Enhancement failed: {result['error']}")
44
+
45
+ # Decode enhanced image
46
+ enhanced_base64 = result["enhanced_image"]
47
+ enhanced_bytes = base64.b64decode(enhanced_base64)
48
+ enhanced_image = Image.open(io.BytesIO(enhanced_bytes))
49
+
50
+ print(f"Original size: {result.get('original_size')}")
51
+ print(f"Enhanced size: {result.get('enhanced_size')}")
52
+
53
+ return enhanced_image
54
+ else:
55
+ raise Exception(f"API call failed: {response.status_code} - {response.text}")
56
+
57
+ # Example usage
58
+ if __name__ == "__main__":
59
+ # Replace with your values
60
+ ENDPOINT_URL = "https://YOUR_ENDPOINT_ID.us-east-1.aws.endpoints.huggingface.cloud"
61
+ HF_TOKEN = "hf_YOUR_TOKEN_HERE"
62
+
63
+ try:
64
+ enhanced = enhance_image_via_endpoint(
65
+ image_path="input.jpg",
66
+ endpoint_url=ENDPOINT_URL,
67
+ hf_token=HF_TOKEN,
68
+ prompt="masterpiece, best quality, highres, sharp details",
69
+ upscale_factor=2.0
70
+ )
71
+
72
+ enhanced.save("enhanced_output.png")
73
+ print("Image enhanced and saved!")
74
+
75
+ except Exception as e:
76
+ print(f"Error: {e}")
enhancer.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from dataclasses import dataclass
2
+ from pathlib import Path
3
+ from typing import Any
4
+
5
+ import torch
6
+ from PIL import Image
7
+ from refiners.foundationals.latent_diffusion.stable_diffusion_1.multi_upscaler import (
8
+ MultiUpscaler,
9
+ UpscalerCheckpoints,
10
+ )
11
+
12
+ from esrgan_model import UpscalerESRGAN
13
+
14
+
15
+ @dataclass(kw_only=True)
16
+ class ESRGANUpscalerCheckpoints(UpscalerCheckpoints):
17
+ esrgan: Path
18
+
19
+
20
+ class ESRGANUpscaler(MultiUpscaler):
21
+ def __init__(
22
+ self,
23
+ checkpoints: ESRGANUpscalerCheckpoints,
24
+ device: torch.device,
25
+ dtype: torch.dtype,
26
+ ) -> None:
27
+ super().__init__(checkpoints=checkpoints, device=device, dtype=dtype)
28
+ self.esrgan = UpscalerESRGAN(checkpoints.esrgan, device=self.device, dtype=self.dtype)
29
+
30
+ def to(self, device: torch.device, dtype: torch.dtype):
31
+ self.esrgan.to(device=device, dtype=dtype)
32
+ self.sd = self.sd.to(device=device, dtype=dtype)
33
+ self.device = device
34
+ self.dtype = dtype
35
+
36
+ def pre_upscale(self, image: Image.Image, upscale_factor: float, **_: Any) -> Image.Image:
37
+ image = self.esrgan.upscale_with_tiling(image)
38
+ return super().pre_upscale(image=image, upscale_factor=upscale_factor / 4)
esrgan_model.py ADDED
@@ -0,0 +1,305 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Modified from https://github.com/philz1337x/clarity-upscaler
3
+ which is a copy of https://github.com/AUTOMATIC1111/stable-diffusion-webui
4
+ which is a copy of https://github.com/victorca25/iNNfer
5
+ which is a copy of https://github.com/xinntao/ESRGAN
6
+ """
7
+
8
+ import math
9
+ from pathlib import Path
10
+ from typing import NamedTuple
11
+
12
+ import numpy as np
13
+ import numpy.typing as npt
14
+ import torch
15
+ import torch.nn as nn
16
+ from PIL import Image
17
+
18
+
19
+ def conv_block(in_nc: int, out_nc: int) -> nn.Sequential:
20
+ return nn.Sequential(
21
+ nn.Conv2d(in_nc, out_nc, kernel_size=3, padding=1),
22
+ nn.LeakyReLU(negative_slope=0.2, inplace=True),
23
+ )
24
+
25
+
26
+ class ResidualDenseBlock_5C(nn.Module):
27
+ """
28
+ Residual Dense Block
29
+ The core module of paper: (Residual Dense Network for Image Super-Resolution, CVPR 18)
30
+ Modified options that can be used:
31
+ - "Partial Convolution based Padding" arXiv:1811.11718
32
+ - "Spectral normalization" arXiv:1802.05957
33
+ - "ICASSP 2020 - ESRGAN+ : Further Improving ESRGAN" N. C.
34
+ {Rakotonirina} and A. {Rasoanaivo}
35
+ """
36
+
37
+ def __init__(self, nf: int = 64, gc: int = 32) -> None:
38
+ super().__init__() # type: ignore[reportUnknownMemberType]
39
+
40
+ self.conv1 = conv_block(nf, gc)
41
+ self.conv2 = conv_block(nf + gc, gc)
42
+ self.conv3 = conv_block(nf + 2 * gc, gc)
43
+ self.conv4 = conv_block(nf + 3 * gc, gc)
44
+ # Wrapped in Sequential because of key in state dict.
45
+ self.conv5 = nn.Sequential(nn.Conv2d(nf + 4 * gc, nf, kernel_size=3, padding=1))
46
+
47
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
48
+ x1 = self.conv1(x)
49
+ x2 = self.conv2(torch.cat((x, x1), 1))
50
+ x3 = self.conv3(torch.cat((x, x1, x2), 1))
51
+ x4 = self.conv4(torch.cat((x, x1, x2, x3), 1))
52
+ x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1))
53
+ return x5 * 0.2 + x
54
+
55
+
56
+ class RRDB(nn.Module):
57
+ """
58
+ Residual in Residual Dense Block
59
+ (ESRGAN: Enhanced Super-Resolution Generative Adversarial Networks)
60
+ """
61
+
62
+ def __init__(self, nf: int) -> None:
63
+ super().__init__() # type: ignore[reportUnknownMemberType]
64
+ self.RDB1 = ResidualDenseBlock_5C(nf)
65
+ self.RDB2 = ResidualDenseBlock_5C(nf)
66
+ self.RDB3 = ResidualDenseBlock_5C(nf)
67
+
68
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
69
+ out = self.RDB1(x)
70
+ out = self.RDB2(out)
71
+ out = self.RDB3(out)
72
+ return out * 0.2 + x
73
+
74
+
75
+ class Upsample2x(nn.Module):
76
+ """Upsample 2x."""
77
+
78
+ def __init__(self) -> None:
79
+ super().__init__() # type: ignore[reportUnknownMemberType]
80
+
81
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
82
+ return nn.functional.interpolate(x, scale_factor=2.0) # type: ignore
83
+
84
+
85
+ class ShortcutBlock(nn.Module):
86
+ """Elementwise sum the output of a submodule to its input"""
87
+
88
+ def __init__(self, submodule: nn.Module) -> None:
89
+ super().__init__() # type: ignore[reportUnknownMemberType]
90
+ self.sub = submodule
91
+
92
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
93
+ return x + self.sub(x)
94
+
95
+
96
+ class RRDBNet(nn.Module):
97
+ def __init__(self, in_nc: int, out_nc: int, nf: int, nb: int) -> None:
98
+ super().__init__() # type: ignore[reportUnknownMemberType]
99
+ assert in_nc % 4 != 0 # in_nc is 3
100
+
101
+ self.model = nn.Sequential(
102
+ nn.Conv2d(in_nc, nf, kernel_size=3, padding=1),
103
+ ShortcutBlock(
104
+ nn.Sequential(
105
+ *(RRDB(nf) for _ in range(nb)),
106
+ nn.Conv2d(nf, nf, kernel_size=3, padding=1),
107
+ )
108
+ ),
109
+ Upsample2x(),
110
+ nn.Conv2d(nf, nf, kernel_size=3, padding=1),
111
+ nn.LeakyReLU(negative_slope=0.2, inplace=True),
112
+ Upsample2x(),
113
+ nn.Conv2d(nf, nf, kernel_size=3, padding=1),
114
+ nn.LeakyReLU(negative_slope=0.2, inplace=True),
115
+ nn.Conv2d(nf, nf, kernel_size=3, padding=1),
116
+ nn.LeakyReLU(negative_slope=0.2, inplace=True),
117
+ nn.Conv2d(nf, out_nc, kernel_size=3, padding=1),
118
+ )
119
+
120
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
121
+ return self.model(x)
122
+
123
+
124
+ def infer_params(state_dict: dict[str, torch.Tensor]) -> tuple[int, int, int, int, int]:
125
+ # this code is adapted from https://github.com/victorca25/iNNfer
126
+ scale2x = 0
127
+ scalemin = 6
128
+ n_uplayer = 0
129
+ out_nc = 0
130
+ nb = 0
131
+
132
+ for block in list(state_dict):
133
+ parts = block.split(".")
134
+ n_parts = len(parts)
135
+ if n_parts == 5 and parts[2] == "sub":
136
+ nb = int(parts[3])
137
+ elif n_parts == 3:
138
+ part_num = int(parts[1])
139
+ if part_num > scalemin and parts[0] == "model" and parts[2] == "weight":
140
+ scale2x += 1
141
+ if part_num > n_uplayer:
142
+ n_uplayer = part_num
143
+ out_nc = state_dict[block].shape[0]
144
+ assert "conv1x1" not in block # no ESRGANPlus
145
+
146
+ nf = state_dict["model.0.weight"].shape[0]
147
+ in_nc = state_dict["model.0.weight"].shape[1]
148
+ scale = 2**scale2x
149
+
150
+ assert out_nc > 0
151
+ assert nb > 0
152
+
153
+ return in_nc, out_nc, nf, nb, scale # 3, 3, 64, 23, 4
154
+
155
+
156
+ Tile = tuple[int, int, Image.Image]
157
+ Tiles = list[tuple[int, int, list[Tile]]]
158
+
159
+
160
+ # https://github.com/philz1337x/clarity-upscaler/blob/e0cd797198d1e0e745400c04d8d1b98ae508c73b/modules/images.py#L64
161
+ class Grid(NamedTuple):
162
+ tiles: Tiles
163
+ tile_w: int
164
+ tile_h: int
165
+ image_w: int
166
+ image_h: int
167
+ overlap: int
168
+
169
+
170
+ # adapted from https://github.com/philz1337x/clarity-upscaler/blob/e0cd797198d1e0e745400c04d8d1b98ae508c73b/modules/images.py#L67
171
+ def split_grid(image: Image.Image, tile_w: int = 512, tile_h: int = 512, overlap: int = 64) -> Grid:
172
+ w = image.width
173
+ h = image.height
174
+
175
+ non_overlap_width = tile_w - overlap
176
+ non_overlap_height = tile_h - overlap
177
+
178
+ cols = max(1, math.ceil((w - overlap) / non_overlap_width))
179
+ rows = max(1, math.ceil((h - overlap) / non_overlap_height))
180
+
181
+ dx = (w - tile_w) / (cols - 1) if cols > 1 else 0
182
+ dy = (h - tile_h) / (rows - 1) if rows > 1 else 0
183
+
184
+ grid = Grid([], tile_w, tile_h, w, h, overlap)
185
+ for row in range(rows):
186
+ row_images: list[Tile] = []
187
+ y1 = max(min(int(row * dy), h - tile_h), 0)
188
+ y2 = min(y1 + tile_h, h)
189
+ for col in range(cols):
190
+ x1 = max(min(int(col * dx), w - tile_w), 0)
191
+ x2 = min(x1 + tile_w, w)
192
+ tile = image.crop((x1, y1, x2, y2))
193
+ row_images.append((x1, tile_w, tile))
194
+ grid.tiles.append((y1, tile_h, row_images))
195
+
196
+ return grid
197
+
198
+
199
+ # https://github.com/philz1337x/clarity-upscaler/blob/e0cd797198d1e0e745400c04d8d1b98ae508c73b/modules/images.py#L104
200
+ def combine_grid(grid: Grid):
201
+ def make_mask_image(r: npt.NDArray[np.float32]) -> Image.Image:
202
+ r = r * 255 / grid.overlap
203
+ return Image.fromarray(r.astype(np.uint8), "L")
204
+
205
+ mask_w = make_mask_image(
206
+ np.arange(grid.overlap, dtype=np.float32).reshape((1, grid.overlap)).repeat(grid.tile_h, axis=0)
207
+ )
208
+ mask_h = make_mask_image(
209
+ np.arange(grid.overlap, dtype=np.float32).reshape((grid.overlap, 1)).repeat(grid.image_w, axis=1)
210
+ )
211
+
212
+ combined_image = Image.new("RGB", (grid.image_w, grid.image_h))
213
+ for y, h, row in grid.tiles:
214
+ combined_row = Image.new("RGB", (grid.image_w, h))
215
+ for x, w, tile in row:
216
+ if x == 0:
217
+ combined_row.paste(tile, (0, 0))
218
+ continue
219
+
220
+ combined_row.paste(tile.crop((0, 0, grid.overlap, h)), (x, 0), mask=mask_w)
221
+ combined_row.paste(tile.crop((grid.overlap, 0, w, h)), (x + grid.overlap, 0))
222
+
223
+ if y == 0:
224
+ combined_image.paste(combined_row, (0, 0))
225
+ continue
226
+
227
+ combined_image.paste(
228
+ combined_row.crop((0, 0, combined_row.width, grid.overlap)),
229
+ (0, y),
230
+ mask=mask_h,
231
+ )
232
+ combined_image.paste(
233
+ combined_row.crop((0, grid.overlap, combined_row.width, h)),
234
+ (0, y + grid.overlap),
235
+ )
236
+
237
+ return combined_image
238
+
239
+
240
+ class UpscalerESRGAN:
241
+ def __init__(self, model_path: Path, device: torch.device, dtype: torch.dtype):
242
+ self.model_path = model_path
243
+ self.device = device
244
+ self.model = self.load_model(model_path)
245
+ self.to(device, dtype)
246
+
247
+ def __call__(self, img: Image.Image) -> Image.Image:
248
+ return self.upscale_without_tiling(img)
249
+
250
+ def to(self, device: torch.device, dtype: torch.dtype):
251
+ self.device = device
252
+ self.dtype = dtype
253
+ self.model.to(device=device, dtype=dtype)
254
+
255
+ def load_model(self, path: Path) -> RRDBNet:
256
+ filename = path
257
+ state_dict: dict[str, torch.Tensor] = torch.load(filename, weights_only=True, map_location=self.device) # type: ignore
258
+ in_nc, out_nc, nf, nb, upscale = infer_params(state_dict)
259
+ assert upscale == 4, "Only 4x upscaling is supported"
260
+ model = RRDBNet(in_nc=in_nc, out_nc=out_nc, nf=nf, nb=nb)
261
+ model.load_state_dict(state_dict)
262
+ model.eval()
263
+
264
+ return model
265
+
266
+ def upscale_without_tiling(self, img: Image.Image) -> Image.Image:
267
+ img_np = np.array(img)
268
+ img_np = img_np[:, :, ::-1]
269
+ img_np = np.ascontiguousarray(np.transpose(img_np, (2, 0, 1))) / 255
270
+ img_t = torch.from_numpy(img_np).float() # type: ignore
271
+ img_t = img_t.unsqueeze(0).to(device=self.device, dtype=self.dtype)
272
+ with torch.no_grad():
273
+ output = self.model(img_t)
274
+ output = output.squeeze().float().cpu().clamp_(0, 1).numpy()
275
+ output = 255.0 * np.moveaxis(output, 0, 2)
276
+ output = output.astype(np.uint8)
277
+ output = output[:, :, ::-1]
278
+ return Image.fromarray(output, "RGB")
279
+
280
+ # https://github.com/philz1337x/clarity-upscaler/blob/e0cd797198d1e0e745400c04d8d1b98ae508c73b/modules/esrgan_model.py#L208
281
+ def upscale_with_tiling(self, img: Image.Image) -> Image.Image:
282
+ img = img.convert("RGB")
283
+ grid = split_grid(img)
284
+ newtiles: Tiles = []
285
+ scale_factor: int = 1
286
+
287
+ for y, h, row in grid.tiles:
288
+ newrow: list[Tile] = []
289
+ for tiledata in row:
290
+ x, w, tile = tiledata
291
+ output = self.upscale_without_tiling(tile)
292
+ scale_factor = output.width // tile.width
293
+ newrow.append((x * scale_factor, w * scale_factor, output))
294
+ newtiles.append((y * scale_factor, h * scale_factor, newrow))
295
+
296
+ newgrid = Grid(
297
+ newtiles,
298
+ grid.tile_w * scale_factor,
299
+ grid.tile_h * scale_factor,
300
+ grid.image_w * scale_factor,
301
+ grid.image_h * scale_factor,
302
+ grid.overlap * scale_factor,
303
+ )
304
+ output = combine_grid(newgrid)
305
+ return output
inference.py ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import io
3
+ import torch
4
+ from pathlib import Path
5
+ from PIL import Image
6
+ from huggingface_hub import hf_hub_download
7
+ from refiners.foundationals.latent_diffusion import solvers
8
+
9
+ from enhancer import ESRGANUpscaler, ESRGANUpscalerCheckpoints
10
+
11
+
12
+ class EndpointHandler:
13
+ def __init__(self, path=""):
14
+ """Initialize the handler with model checkpoints"""
15
+
16
+ # Download model checkpoints
17
+ self.checkpoints = ESRGANUpscalerCheckpoints(
18
+ unet=Path(
19
+ hf_hub_download(
20
+ repo_id="refiners/juggernaut.reborn.sd1_5.unet",
21
+ filename="model.safetensors",
22
+ revision="347d14c3c782c4959cc4d1bb1e336d19f7dda4d2",
23
+ )
24
+ ),
25
+ clip_text_encoder=Path(
26
+ hf_hub_download(
27
+ repo_id="refiners/juggernaut.reborn.sd1_5.text_encoder",
28
+ filename="model.safetensors",
29
+ revision="744ad6a5c0437ec02ad826df9f6ede102bb27481",
30
+ )
31
+ ),
32
+ lda=Path(
33
+ hf_hub_download(
34
+ repo_id="refiners/juggernaut.reborn.sd1_5.autoencoder",
35
+ filename="model.safetensors",
36
+ revision="3c1aae3fc3e03e4a2b7e0fa42b62ebb64f1a4c19",
37
+ )
38
+ ),
39
+ controlnet_tile=Path(
40
+ hf_hub_download(
41
+ repo_id="refiners/controlnet.sd1_5.tile",
42
+ filename="model.safetensors",
43
+ revision="48ced6ff8bfa873a8976fa467c3629a240643387",
44
+ )
45
+ ),
46
+ esrgan=Path(
47
+ hf_hub_download(
48
+ repo_id="philz1337x/upscaler",
49
+ filename="4x-UltraSharp.pth",
50
+ revision="011deacac8270114eb7d2eeff4fe6fa9a837be70",
51
+ )
52
+ ),
53
+ negative_embedding=Path(
54
+ hf_hub_download(
55
+ repo_id="philz1337x/embeddings",
56
+ filename="JuggernautNegative-neg.pt",
57
+ revision="203caa7e9cc2bc225031a4021f6ab1ded283454a",
58
+ )
59
+ ),
60
+ negative_embedding_key="string_to_param.*",
61
+ loras={
62
+ "more_details": Path(
63
+ hf_hub_download(
64
+ repo_id="philz1337x/loras",
65
+ filename="more_details.safetensors",
66
+ revision="a3802c0280c0d00c2ab18d37454a8744c44e474e",
67
+ )
68
+ ),
69
+ "sdxl_render": Path(
70
+ hf_hub_download(
71
+ repo_id="philz1337x/loras",
72
+ filename="SDXLrender_v2.0.safetensors",
73
+ revision="a3802c0280c0d00c2ab18d37454a8744c44e474e",
74
+ )
75
+ ),
76
+ },
77
+ )
78
+
79
+ # Initialize device and dtype
80
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
81
+ self.dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float32
82
+
83
+ # Initialize the enhancer
84
+ self.enhancer = ESRGANUpscaler(
85
+ checkpoints=self.checkpoints,
86
+ device=self.device,
87
+ dtype=self.dtype
88
+ )
89
+
90
+ def __call__(self, data):
91
+ """
92
+ Process the input data and return enhanced image
93
+
94
+ Args:
95
+ data (dict): Input data containing:
96
+ - image (str): Base64 encoded input image
97
+ - prompt (str, optional): Enhancement prompt
98
+ - negative_prompt (str, optional): Negative prompt
99
+ - seed (int, optional): Random seed
100
+ - upscale_factor (float, optional): Upscale factor
101
+ - controlnet_scale (float, optional): ControlNet scale
102
+ - controlnet_decay (float, optional): ControlNet decay
103
+ - condition_scale (int, optional): Condition scale
104
+ - tile_width (int, optional): Tile width
105
+ - tile_height (int, optional): Tile height
106
+ - denoise_strength (float, optional): Denoise strength
107
+ - num_inference_steps (int, optional): Number of inference steps
108
+ - solver (str, optional): Solver type
109
+
110
+ Returns:
111
+ dict: Contains enhanced image as base64 string
112
+ """
113
+ try:
114
+ # Extract and decode input image
115
+ image_data = data.get("image")
116
+ if not image_data:
117
+ return {"error": "No image provided"}
118
+
119
+ # Decode base64 image
120
+ if image_data.startswith('data:image'):
121
+ # Remove data:image/...;base64, prefix if present
122
+ image_data = image_data.split(',')[1]
123
+
124
+ image_bytes = base64.b64decode(image_data)
125
+ input_image = Image.open(io.BytesIO(image_bytes)).convert('RGB')
126
+
127
+ # Extract parameters with defaults
128
+ prompt = data.get("prompt", "masterpiece, best quality, highres")
129
+ negative_prompt = data.get("negative_prompt", "worst quality, low quality, normal quality")
130
+ seed = data.get("seed", 42)
131
+ upscale_factor = data.get("upscale_factor", 2)
132
+ controlnet_scale = data.get("controlnet_scale", 0.6)
133
+ controlnet_decay = data.get("controlnet_decay", 1.0)
134
+ condition_scale = data.get("condition_scale", 6)
135
+ tile_width = data.get("tile_width", 112)
136
+ tile_height = data.get("tile_height", 144)
137
+ denoise_strength = data.get("denoise_strength", 0.35)
138
+ num_inference_steps = data.get("num_inference_steps", 18)
139
+ solver_name = data.get("solver", "DDIM")
140
+
141
+ # Get solver type
142
+ solver_type = getattr(solvers, solver_name)
143
+
144
+ # Set up generator
145
+ generator = torch.Generator(device=self.device)
146
+ generator.manual_seed(seed)
147
+
148
+ # Resize input image if too large to avoid VRAM issues
149
+ side_size = min(input_image.size)
150
+ if side_size > 768:
151
+ scale = 768 / side_size
152
+ new_size = (int(input_image.width * scale), int(input_image.height * scale))
153
+ resized_image = input_image.resize(new_size, resample=Image.Resampling.LANCZOS)
154
+ else:
155
+ resized_image = input_image
156
+
157
+ # Enhance the image
158
+ enhanced_image = self.enhancer.upscale(
159
+ image=resized_image,
160
+ prompt=prompt,
161
+ negative_prompt=negative_prompt,
162
+ upscale_factor=upscale_factor,
163
+ controlnet_scale=controlnet_scale,
164
+ controlnet_scale_decay=controlnet_decay,
165
+ condition_scale=condition_scale,
166
+ tile_size=(tile_height, tile_width),
167
+ denoise_strength=denoise_strength,
168
+ num_inference_steps=num_inference_steps,
169
+ loras_scale={"more_details": 0.5, "sdxl_render": 1.0},
170
+ solver_type=solver_type,
171
+ generator=generator,
172
+ )
173
+
174
+ # Convert enhanced image to base64
175
+ buffered = io.BytesIO()
176
+ enhanced_image.save(buffered, format="PNG")
177
+ enhanced_base64 = base64.b64encode(buffered.getvalue()).decode('utf-8')
178
+
179
+ return {
180
+ "enhanced_image": enhanced_base64,
181
+ "original_size": input_image.size,
182
+ "enhanced_size": enhanced_image.size
183
+ }
184
+
185
+ except Exception as e:
186
+ return {"error": f"Enhancement failed: {str(e)}"}
requirements.txt ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ git+https://github.com/finegrain-ai/refiners@cfe8b66ba4f8a906583850ac25e9e89cb83a44b9
2
+ numpy<2.0.0
3
+ pillow>=10.4.0
4
+ pillow-heif>=0.18.0
5
+ torch>=2.0.0
6
+ transformers>=4.21.0
7
+ huggingface-hub>=0.16.0