hugging-apps/echo-memory
0
1import torch2from einops import repeat3from PIL import Image4import numpy as np5 6 7class ResidualDenseBlock(torch.nn.Module):8 9 def __init__(self, num_feat=64, num_grow_ch=32):10 super(ResidualDenseBlock, self).__init__()11 self.conv1 = torch.nn.Conv2d(num_feat, num_grow_ch, 3, 1, 1)12 self.conv2 = torch.nn.Conv2d(num_feat + num_grow_ch, num_grow_ch, 3, 1, 1)13 self.conv3 = torch.nn.Conv2d(num_feat + 2 * num_grow_ch, num_grow_ch, 3, 1, 1)14 self.conv4 = torch.nn.Conv2d(num_feat + 3 * num_grow_ch, num_grow_ch, 3, 1, 1)15 self.conv5 = torch.nn.Conv2d(num_feat + 4 * num_grow_ch, num_feat, 3, 1, 1)16 self.lrelu = torch.nn.LeakyReLU(negative_slope=0.2, inplace=True)17 18 def forward(self, x):19 x1 = self.lrelu(self.conv1(x))20 x2 = self.lrelu(self.conv2(torch.cat((x, x1), 1)))21 x3 = self.lrelu(self.conv3(torch.cat((x, x1, x2), 1)))22 x4 = self.lrelu(self.conv4(torch.cat((x, x1, x2, x3), 1)))23 x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1))24 return x5 * 0.2 + x25 26 27class RRDB(torch.nn.Module):28 29 def __init__(self, num_feat, num_grow_ch=32):30 super(RRDB, self).__init__()31 self.rdb1 = ResidualDenseBlock(num_feat, num_grow_ch)32 self.rdb2 = ResidualDenseBlock(num_feat, num_grow_ch)33 self.rdb3 = ResidualDenseBlock(num_feat, num_grow_ch)34 35 def forward(self, x):36 out = self.rdb1(x)37 out = self.rdb2(out)38 out = self.rdb3(out)39 return out * 0.2 + x40 41 42class RRDBNet(torch.nn.Module):43 44 def __init__(self, num_in_ch=3, num_out_ch=3, num_feat=64, num_block=23, num_grow_ch=32, **kwargs):45 super(RRDBNet, self).__init__()46 self.conv_first = torch.nn.Conv2d(num_in_ch, num_feat, 3, 1, 1)47 self.body = torch.torch.nn.Sequential(*[RRDB(num_feat=num_feat, num_grow_ch=num_grow_ch) for _ in range(num_block)])48 self.conv_body = torch.nn.Conv2d(num_feat, num_feat, 3, 1, 1)49 # upsample50 self.conv_up1 = torch.nn.Conv2d(num_feat, num_feat, 3, 1, 1)51 self.conv_up2 = torch.nn.Conv2d(num_feat, num_feat, 3, 1, 1)52 self.conv_hr = torch.nn.Conv2d(num_feat, num_feat, 3, 1, 1)53 self.conv_last = torch.nn.Conv2d(num_feat, num_out_ch, 3, 1, 1)54 self.lrelu = torch.nn.LeakyReLU(negative_slope=0.2, inplace=True)55 56 def forward(self, x):57 feat = x58 feat = self.conv_first(feat)59 body_feat = self.conv_body(self.body(feat))60 feat = feat + body_feat61 # upsample62 feat = repeat(feat, "B C H W -> B C (H 2) (W 2)")63 feat = self.lrelu(self.conv_up1(feat))64 feat = repeat(feat, "B C H W -> B C (H 2) (W 2)")65 feat = self.lrelu(self.conv_up2(feat))66 out = self.conv_last(self.lrelu(self.conv_hr(feat)))67 return out68 69 @staticmethod70 def state_dict_converter():71 return RRDBNetStateDictConverter()72 73 74class RRDBNetStateDictConverter:75 def __init__(self):76 pass77 78 def from_diffusers(self, state_dict):79 return state_dict, {"upcast_to_float32": True}80 81 def from_civitai(self, state_dict):82 return state_dict, {"upcast_to_float32": True}83 84 85class ESRGAN(torch.nn.Module):86 def __init__(self, model):87 super().__init__()88 self.model = model89 90 @staticmethod91 def from_model_manager(model_manager):92 return ESRGAN(model_manager.fetch_model("esrgan"))93 94 def process_image(self, image):95 image = torch.Tensor(np.array(image, dtype=np.float32) / 255).permute(2, 0, 1)96 return image97 98 def process_images(self, images):99 images = [self.process_image(image) for image in images]100 images = torch.stack(images)101 return images102 103 def decode_images(self, images):104 images = (images.permute(0, 2, 3, 1) * 255).clip(0, 255).numpy().astype(np.uint8)105 images = [Image.fromarray(image) for image in images]106 return images107 108 @torch.no_grad()109 def upscale(self, images, batch_size=4, progress_bar=lambda x:x):110 if not isinstance(images, list):111 images = [images]112 is_single_image = True113 else:114 is_single_image = False115 116 # Preprocess117 input_tensor = self.process_images(images)118 119 # Interpolate120 output_tensor = []121 for batch_id in progress_bar(range(0, input_tensor.shape[0], batch_size)):122 batch_id_ = min(batch_id + batch_size, input_tensor.shape[0])123 batch_input_tensor = input_tensor[batch_id: batch_id_]124 batch_input_tensor = batch_input_tensor.to(125 device=self.model.conv_first.weight.device,126 dtype=self.model.conv_first.weight.dtype)127 batch_output_tensor = self.model(batch_input_tensor)128 output_tensor.append(batch_output_tensor.cpu())129 130 # Output131 output_tensor = torch.concat(output_tensor, dim=0)132 133 # To images134 output_images = self.decode_images(output_tensor)135 if is_single_image:136 output_images = output_images[0]137 return output_images138 