jiachen commited on
Commit
b6ee05f
·
1 Parent(s): d4ac85e
add_noise.py CHANGED
@@ -22,6 +22,6 @@ def add_noise(image_path, output_path, sigma=50):
22
  noisy_image.save(output_path)
23
 
24
  # Example usage
25
- input_image_path = '/home/jiachen/MyGradio/test_images/myclean.jpg'
26
- output_image_path = '/home/jiachen/MyGradio/test_images/noisy_myimage.jpg'
27
  add_noise(input_image_path, output_image_path, sigma=15)
 
22
  noisy_image.save(output_path)
23
 
24
  # Example usage
25
+ input_image_path = '/home/jiachen/MyGradio/test_images/0059.png'
26
+ output_image_path = '/home/jiachen/MyGradio/test_images/noisy_0059.png'
27
  add_noise(input_image_path, output_image_path, sigma=15)
app.py CHANGED
@@ -1,16 +1,17 @@
 
1
  import numpy as np
2
  import gradio as gr
3
- import numpy as np
4
  import torch
5
- import spaces
6
-
7
 
8
  from PIL import Image
9
  from torchvision.transforms import ToTensor
10
-
11
- from net.prompt_xrestormer import PromptXRestormer
12
  import lightning.pytorch as pl
13
 
 
 
 
 
14
 
15
 
16
  # crop an image to the multiple of base
@@ -21,27 +22,33 @@ def crop_img(image, base=64):
21
  crop_w = w % base
22
  return image[crop_h // 2:h - crop_h + crop_h // 2, crop_w // 2:w - crop_w + crop_w // 2, :]
23
 
24
- class PromptXRestormerIRModel(pl.LightningModule):
25
  def __init__(self):
26
  super().__init__()
27
- self.net = PromptXRestormer(
28
  inp_channels=3,
29
  out_channels=3,
30
  dim = 48,
31
  num_blocks = [2,4,4,4],
32
  num_refinement_blocks = 4,
33
- channel_heads= [1,1,1,1],
34
- spatial_heads= [1,2,4,8],
35
- overlap_ratio= [0.5, 0.5, 0.5, 0.5],
36
- ffn_expansion_factor = 2.66,
 
 
37
  bias = False,
 
38
  LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
39
  dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
40
- scale = 1,prompt = True
 
 
41
  )
 
42
 
43
- def forward(self,x):
44
- return self.net(x)
45
 
46
  def np_to_pil(img_np):
47
  """
@@ -78,15 +85,12 @@ def restore_image(input_img):
78
  np.random.seed(0)
79
  torch.manual_seed(0)
80
 
81
- #ckpt_path = "/home/jiachen/MyGradio/ckpt/promptxrestormer_epoch=64-step=578630.ckpt"
82
- ckpt_path = "ckpt/promptxrestormer_epoch=64-step=578630.ckpt"
83
  print("CKPT name : {}".format(ckpt_path))
84
 
85
- #net = PromptXRestormerIRModel().load_from_checkpoint(ckpt_path).cuda()
86
- net = PromptXRestormerIRModel.load_from_checkpoint(ckpt_path).cuda()
87
  net.eval()
88
 
89
- #degraded_path = "/home/jiachen/MyGradio/test_images/rain-070.png"
90
 
91
  degraded_img = crop_img(np.array(input_img.convert('RGB')), base=16)
92
  toTensor = ToTensor()
@@ -104,26 +108,28 @@ def restore_image(input_img):
104
  degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [2])], 2)[:,:,:H_old+h_pad,:]
105
  degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [3])], 3)[:,:,:,:W_old+w_pad]
106
 
107
- restored = net(degraded_img)
108
- restored = restored[:,:,:H_old:,:W_old]
 
 
 
 
 
 
 
 
109
 
 
 
110
  restored_image = torch_to_np(restored)
111
- # change shape from [C, H, W] to [H, W, C]
112
  restored_image = restored_image.transpose(1, 2, 0)
113
  restored_image = np.clip(restored_image * 255, 0, 255).astype(np.uint8)
114
 
115
- # restored_image = Image.fromarray(restored_image)
116
- # print("restored shape : {}".format(restored_image.size))
117
-
118
- return restored_image
119
-
120
- # degraded_path = "/home/jiachen/MyGradio/test_images/rain-070.png"
121
- # input_img = np.array(Image.open(degraded_path).convert('RGB'))
122
- # print(input_img)
123
- # restored_image = restore_image(input_img)
124
- # print(restored_image)
125
-
126
-
127
 
128
 
129
  title = "Content & Task Awareness All-In-One Image Restoration✏️🖼️ 🤗"
@@ -147,17 +153,16 @@ examples = [['test_images/noisy_0000.png'],
147
  ['test_images/noisy_0002.png'],
148
  ['test_images/noisy_0003.png'],
149
  ['test_images/noisy_0004.png'],
150
- ['test_images/rain-01.png'],
151
- ['test_images/rain-02.png'],
152
- ['test_images/rain-03.png'],
153
- ['test_images/rain-04.png'],
154
- ['test_images/rain-05.png'],
155
- ['test_images/rain-06.png'],
156
- ['test_images/hazy-00.jpg'],
157
  ['test_images/hazy-01.jpg'],
158
  ['test_images/hazy-02.jpg'],
159
  ['test_images/hazy-03.jpg'],
160
- ['test_images/hazy-04.jpg'],
 
161
  ]
162
  css = """
163
  .image-frame img, .image-container img {
@@ -170,7 +175,7 @@ css = """
170
  demo = gr.Interface(
171
  fn=restore_image,
172
  inputs=[gr.Image(type="pil", label="Input")],
173
- outputs=[gr.Image(type="pil", label="Ouput")],
174
  title=title,
175
  description=description,
176
  article=article,
@@ -179,5 +184,5 @@ demo = gr.Interface(
179
  )
180
 
181
 
182
- # if __name__ == "__main__":
183
- demo.launch(debug=True, show_error=True)
 
1
+
2
  import numpy as np
3
  import gradio as gr
 
4
  import torch
5
+ import torch.nn as nn
 
6
 
7
  from PIL import Image
8
  from torchvision.transforms import ToTensor
 
 
9
  import lightning.pytorch as pl
10
 
11
+ from net.cata_prompt_xrestormer import CATAPromptXRestormerOnlyAttn
12
+ from einops import rearrange
13
+ import spaces
14
+
15
 
16
 
17
  # crop an image to the multiple of base
 
22
  crop_w = w % base
23
  return image[crop_h // 2:h - crop_h + crop_h // 2, crop_w // 2:w - crop_w + crop_w // 2, :]
24
 
25
+ class CATAPromptXRestormerIRModel(pl.LightningModule):
26
  def __init__(self):
27
  super().__init__()
28
+ self.net = CATAPromptXRestormerOnlyAttn(
29
  inp_channels=3,
30
  out_channels=3,
31
  dim = 48,
32
  num_blocks = [2,4,4,4],
33
  num_refinement_blocks = 4,
34
+ channel_heads = [1,1,1,1],
35
+ spatial_heads = [1,2,4,8],
36
+ overlap_ratio = 0.5,
37
+ dim_head = 16,
38
+ ratio = 0.5,
39
+ window_size = 8,
40
  bias = False,
41
+ ffn_expansion_factor = 2.66,
42
  LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
43
  dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
44
+ scale = 1,
45
+ prompt = True,
46
+ hard_ratio = 0.5
47
  )
48
+ self.loss_fn = nn.L1Loss()
49
 
50
+ def forward(self,x, training=False):
51
+ return self.net(x, training)
52
 
53
  def np_to_pil(img_np):
54
  """
 
85
  np.random.seed(0)
86
  torch.manual_seed(0)
87
 
88
+ ckpt_path = "ckpt/cata_promptxrestormeronlyattn_epoch=30-step=275962.ckpt"
 
89
  print("CKPT name : {}".format(ckpt_path))
90
 
91
+ net = CATAPromptXRestormerIRModel.load_from_checkpoint(ckpt_path).cuda()
 
92
  net.eval()
93
 
 
94
 
95
  degraded_img = crop_img(np.array(input_img.convert('RGB')), base=16)
96
  toTensor = ToTensor()
 
108
  degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [2])], 2)[:,:,:H_old+h_pad,:]
109
  degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [3])], 3)[:,:,:,:W_old+w_pad]
110
 
111
+ restored, spatial_mask, channel_mask = net(degraded_img, training=False)
112
+
113
+ # process spatial mask
114
+ encoder_level1_mask = spatial_mask['encoder_level1'][0][0][0]
115
+ window_size = 8
116
+ _, c, h, w = restored.shape
117
+ restored_windows = rearrange(restored, 'b c (h w1) (w w2) -> b c (h w) w1 w2', w1=window_size, w2=window_size)
118
+ for idx in encoder_level1_mask:
119
+ restored_windows[:, :, idx, :, :] = 1 # Mask out the window by setting it to one
120
+ restored_masked = rearrange(restored_windows, 'b c (h w) w1 w2 -> b c (h w1) (w w2)', h=h // window_size, w=w // window_size)
121
 
122
+
123
+ restored = restored[:,:,:H_old:,:W_old]
124
  restored_image = torch_to_np(restored)
 
125
  restored_image = restored_image.transpose(1, 2, 0)
126
  restored_image = np.clip(restored_image * 255, 0, 255).astype(np.uint8)
127
 
128
+ restored_masked = restored_masked[:,:,:H_old:,:W_old]
129
+ restored_masked_image = torch_to_np(restored_masked)
130
+ restored_masked_image = restored_masked_image.transpose(1, 2, 0)
131
+ restored_masked_image = np.clip(restored_masked_image * 255, 0, 255).astype(np.uint8)
132
+ return restored_image, restored_masked_image
 
 
 
 
 
 
 
133
 
134
 
135
  title = "Content & Task Awareness All-In-One Image Restoration✏️🖼️ 🤗"
 
153
  ['test_images/noisy_0002.png'],
154
  ['test_images/noisy_0003.png'],
155
  ['test_images/noisy_0004.png'],
156
+ ['test_images/rain-001.png'],
157
+ ['test_images/rain-002.png'],
158
+ ['test_images/rain-003.png'],
159
+ ['test_images/rain-004.png'],
160
+ ['test_images/rain-005.png'],
 
 
161
  ['test_images/hazy-01.jpg'],
162
  ['test_images/hazy-02.jpg'],
163
  ['test_images/hazy-03.jpg'],
164
+ ['test_images/hazy-04.jpg'],
165
+ ['test_images/hazy-05.jpg'],
166
  ]
167
  css = """
168
  .image-frame img, .image-container img {
 
175
  demo = gr.Interface(
176
  fn=restore_image,
177
  inputs=[gr.Image(type="pil", label="Input")],
178
+ outputs=[gr.Image(type="pil", label="Ouput"), gr.Image(type="pil", label="Output-Mask")],
179
  title=title,
180
  description=description,
181
  article=article,
 
184
  )
185
 
186
 
187
+ if __name__ == "__main__":
188
+ demo.launch(debug=True, show_error=True)
app_local.py ADDED
@@ -0,0 +1,182 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import gradio as gr
3
+ import numpy as np
4
+ import torch
5
+
6
+
7
+ from PIL import Image
8
+ from torchvision.transforms import ToTensor
9
+
10
+ from net.prompt_xrestormer import PromptXRestormer
11
+ import lightning.pytorch as pl
12
+
13
+
14
+
15
+ # crop an image to the multiple of base
16
+ def crop_img(image, base=64):
17
+ h = image.shape[0]
18
+ w = image.shape[1]
19
+ crop_h = h % base
20
+ crop_w = w % base
21
+ return image[crop_h // 2:h - crop_h + crop_h // 2, crop_w // 2:w - crop_w + crop_w // 2, :]
22
+
23
+ class PromptXRestormerIRModel(pl.LightningModule):
24
+ def __init__(self):
25
+ super().__init__()
26
+ self.net = PromptXRestormer(
27
+ inp_channels=3,
28
+ out_channels=3,
29
+ dim = 48,
30
+ num_blocks = [2,4,4,4],
31
+ num_refinement_blocks = 4,
32
+ channel_heads= [1,1,1,1],
33
+ spatial_heads= [1,2,4,8],
34
+ overlap_ratio= [0.5, 0.5, 0.5, 0.5],
35
+ ffn_expansion_factor = 2.66,
36
+ bias = False,
37
+ LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
38
+ dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
39
+ scale = 1,prompt = True
40
+ )
41
+
42
+ def forward(self,x):
43
+ return self.net(x)
44
+
45
+ def np_to_pil(img_np):
46
+ """
47
+ Converts image in np.array format to PIL image.
48
+
49
+ From C x W x H [0..1] to W x H x C [0...255]
50
+ :param img_np:
51
+ :return:
52
+ """
53
+ ar = np.clip(img_np * 255, 0, 255).astype(np.uint8)
54
+
55
+ if img_np.shape[0] == 1:
56
+ ar = ar[0]
57
+ else:
58
+ assert img_np.shape[0] == 3, img_np.shape
59
+ ar = ar.transpose(1, 2, 0)
60
+
61
+ return Image.fromarray(ar)
62
+
63
+ def torch_to_np(img_var):
64
+ """
65
+ Converts an image in torch.Tensor format to np.array.
66
+
67
+ From 1 x C x W x H [0..1] to C x W x H [0..1]
68
+ :param img_var:
69
+ :return:
70
+ """
71
+ return img_var.detach().cpu().numpy()[0]
72
+
73
+
74
+
75
+ #@spaces.GPU(duration=200)
76
+ def restore_image(input_img):
77
+ np.random.seed(0)
78
+ torch.manual_seed(0)
79
+
80
+ #ckpt_path = "/home/jiachen/MyGradio/ckpt/promptxrestormer_epoch=64-step=578630.ckpt"
81
+ ckpt_path = "ckpt/promptxrestormer_epoch=64-step=578630.ckpt"
82
+ print("CKPT name : {}".format(ckpt_path))
83
+
84
+ #net = PromptXRestormerIRModel().load_from_checkpoint(ckpt_path).cuda()
85
+ net = PromptXRestormerIRModel.load_from_checkpoint(ckpt_path).cuda()
86
+ net.eval()
87
+
88
+ #degraded_path = "/home/jiachen/MyGradio/test_images/rain-070.png"
89
+
90
+ degraded_img = crop_img(np.array(input_img.convert('RGB')), base=16)
91
+ toTensor = ToTensor()
92
+ degraded_img = toTensor(degraded_img)
93
+
94
+
95
+ with torch.no_grad():
96
+ degraded_img = degraded_img.unsqueeze(0).cuda()
97
+
98
+ _, _, H_old, W_old = degraded_img.shape
99
+
100
+
101
+ h_pad = (H_old // 64 + 1) * 64 - H_old
102
+ w_pad = (W_old // 64 + 1) * 64 - W_old
103
+ degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [2])], 2)[:,:,:H_old+h_pad,:]
104
+ degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [3])], 3)[:,:,:,:W_old+w_pad]
105
+
106
+ restored = net(degraded_img)
107
+ restored = restored[:,:,:H_old:,:W_old]
108
+
109
+ restored_image = torch_to_np(restored)
110
+ # change shape from [C, H, W] to [H, W, C]
111
+ restored_image = restored_image.transpose(1, 2, 0)
112
+ restored_image = np.clip(restored_image * 255, 0, 255).astype(np.uint8)
113
+
114
+ # restored_image = Image.fromarray(restored_image)
115
+ # print("restored shape : {}".format(restored_image.size))
116
+
117
+ return restored_image
118
+
119
+ # degraded_path = "/home/jiachen/MyGradio/test_images/rain-070.png"
120
+ # input_img = np.array(Image.open(degraded_path).convert('RGB'))
121
+ # print(input_img)
122
+ # restored_image = restore_image(input_img)
123
+ # print(restored_image)
124
+
125
+
126
+
127
+
128
+ title = "Content & Task Awareness All-In-One Image Restoration✏️🖼️ 🤗"
129
+ description = ''' ## [Content & Task Awareness All-In-One Image Restoration]
130
+
131
+ The Ohio State Unviersity | Microsoft Research
132
+
133
+ ### TL;DR: quickstart
134
+ ***One single model can perform several restoration tasks including image denoising, deraining and dehazing 🚀 . Our content & task awareness model would have better efficiency***
135
+ The (single) neural model performs all-in-one image restoration.
136
+ **🚀 You can start with the [demo tutorial.]** Check [our github] for more information.
137
+ <br>
138
+ '''
139
+
140
+
141
+ article = "<p style='text-align: center'><a href='https://github.com/mv-lab/InstructIR' target='_blank'>Content & Task Awareness All-In-One Image Restoration</a></p>"
142
+
143
+ #### Image,Prompts examples
144
+ examples = [['test_images/noisy_0000.png'],
145
+ ['test_images/noisy_0001.png'],
146
+ ['test_images/noisy_0002.png'],
147
+ ['test_images/noisy_0003.png'],
148
+ ['test_images/noisy_0004.png'],
149
+ ['test_images/rain-01.png'],
150
+ ['test_images/rain-02.png'],
151
+ ['test_images/rain-03.png'],
152
+ ['test_images/rain-04.png'],
153
+ ['test_images/rain-05.png'],
154
+ ['test_images/rain-06.png'],
155
+ ['test_images/hazy-00.jpg'],
156
+ ['test_images/hazy-01.jpg'],
157
+ ['test_images/hazy-02.jpg'],
158
+ ['test_images/hazy-03.jpg'],
159
+ ['test_images/hazy-04.jpg'],
160
+ ]
161
+ css = """
162
+ .image-frame img, .image-container img {
163
+ width: auto;
164
+ height: auto;
165
+ max-width: none;
166
+ }
167
+ """
168
+
169
+ demo = gr.Interface(
170
+ fn=restore_image,
171
+ inputs=[gr.Image(type="pil", label="Input")],
172
+ outputs=[gr.Image(type="pil", label="Ouput")],
173
+ title=title,
174
+ description=description,
175
+ article=article,
176
+ examples=examples,
177
+ css=css,
178
+ )
179
+
180
+
181
+ if __name__ == "__main__":
182
+ demo.launch(server_port=8085)
app_xrestormer.py ADDED
@@ -0,0 +1,183 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import gradio as gr
3
+ import numpy as np
4
+ import torch
5
+ import spaces
6
+
7
+
8
+ from PIL import Image
9
+ from torchvision.transforms import ToTensor
10
+
11
+ from net.prompt_xrestormer import PromptXRestormer
12
+ import lightning.pytorch as pl
13
+
14
+
15
+
16
+ # crop an image to the multiple of base
17
+ def crop_img(image, base=64):
18
+ h = image.shape[0]
19
+ w = image.shape[1]
20
+ crop_h = h % base
21
+ crop_w = w % base
22
+ return image[crop_h // 2:h - crop_h + crop_h // 2, crop_w // 2:w - crop_w + crop_w // 2, :]
23
+
24
+ class PromptXRestormerIRModel(pl.LightningModule):
25
+ def __init__(self):
26
+ super().__init__()
27
+ self.net = PromptXRestormer(
28
+ inp_channels=3,
29
+ out_channels=3,
30
+ dim = 48,
31
+ num_blocks = [2,4,4,4],
32
+ num_refinement_blocks = 4,
33
+ channel_heads= [1,1,1,1],
34
+ spatial_heads= [1,2,4,8],
35
+ overlap_ratio= [0.5, 0.5, 0.5, 0.5],
36
+ ffn_expansion_factor = 2.66,
37
+ bias = False,
38
+ LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
39
+ dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
40
+ scale = 1,prompt = True
41
+ )
42
+
43
+ def forward(self,x):
44
+ return self.net(x)
45
+
46
+ def np_to_pil(img_np):
47
+ """
48
+ Converts image in np.array format to PIL image.
49
+
50
+ From C x W x H [0..1] to W x H x C [0...255]
51
+ :param img_np:
52
+ :return:
53
+ """
54
+ ar = np.clip(img_np * 255, 0, 255).astype(np.uint8)
55
+
56
+ if img_np.shape[0] == 1:
57
+ ar = ar[0]
58
+ else:
59
+ assert img_np.shape[0] == 3, img_np.shape
60
+ ar = ar.transpose(1, 2, 0)
61
+
62
+ return Image.fromarray(ar)
63
+
64
+ def torch_to_np(img_var):
65
+ """
66
+ Converts an image in torch.Tensor format to np.array.
67
+
68
+ From 1 x C x W x H [0..1] to C x W x H [0..1]
69
+ :param img_var:
70
+ :return:
71
+ """
72
+ return img_var.detach().cpu().numpy()[0]
73
+
74
+
75
+
76
+ @spaces.GPU(duration=200)
77
+ def restore_image(input_img):
78
+ np.random.seed(0)
79
+ torch.manual_seed(0)
80
+
81
+ #ckpt_path = "/home/jiachen/MyGradio/ckpt/promptxrestormer_epoch=64-step=578630.ckpt"
82
+ ckpt_path = "ckpt/promptxrestormer_epoch=64-step=578630.ckpt"
83
+ print("CKPT name : {}".format(ckpt_path))
84
+
85
+ #net = PromptXRestormerIRModel().load_from_checkpoint(ckpt_path).cuda()
86
+ net = PromptXRestormerIRModel.load_from_checkpoint(ckpt_path).cuda()
87
+ net.eval()
88
+
89
+ #degraded_path = "/home/jiachen/MyGradio/test_images/rain-070.png"
90
+
91
+ degraded_img = crop_img(np.array(input_img.convert('RGB')), base=16)
92
+ toTensor = ToTensor()
93
+ degraded_img = toTensor(degraded_img)
94
+
95
+
96
+ with torch.no_grad():
97
+ degraded_img = degraded_img.unsqueeze(0).cuda()
98
+
99
+ _, _, H_old, W_old = degraded_img.shape
100
+
101
+
102
+ h_pad = (H_old // 64 + 1) * 64 - H_old
103
+ w_pad = (W_old // 64 + 1) * 64 - W_old
104
+ degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [2])], 2)[:,:,:H_old+h_pad,:]
105
+ degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [3])], 3)[:,:,:,:W_old+w_pad]
106
+
107
+ restored = net(degraded_img)
108
+ restored = restored[:,:,:H_old:,:W_old]
109
+
110
+ restored_image = torch_to_np(restored)
111
+ # change shape from [C, H, W] to [H, W, C]
112
+ restored_image = restored_image.transpose(1, 2, 0)
113
+ restored_image = np.clip(restored_image * 255, 0, 255).astype(np.uint8)
114
+
115
+ # restored_image = Image.fromarray(restored_image)
116
+ # print("restored shape : {}".format(restored_image.size))
117
+
118
+ return restored_image
119
+
120
+ # degraded_path = "/home/jiachen/MyGradio/test_images/rain-070.png"
121
+ # input_img = np.array(Image.open(degraded_path).convert('RGB'))
122
+ # print(input_img)
123
+ # restored_image = restore_image(input_img)
124
+ # print(restored_image)
125
+
126
+
127
+
128
+
129
+ title = "Content & Task Awareness All-In-One Image Restoration✏️🖼️ 🤗"
130
+ description = ''' ## [Content & Task Awareness All-In-One Image Restoration]
131
+
132
+ The Ohio State Unviersity | Microsoft Research
133
+
134
+ ### TL;DR: quickstart
135
+ ***One single model can perform several restoration tasks including image denoising, deraining and dehazing 🚀 . Our content & task awareness model would have better efficiency***
136
+ The (single) neural model performs all-in-one image restoration.
137
+ **🚀 You can start with the [demo tutorial.]** Check [our github] for more information.
138
+ <br>
139
+ '''
140
+
141
+
142
+ article = "<p style='text-align: center'><a href='https://github.com/mv-lab/InstructIR' target='_blank'>Content & Task Awareness All-In-One Image Restoration</a></p>"
143
+
144
+ #### Image,Prompts examples
145
+ examples = [['test_images/noisy_0000.png'],
146
+ ['test_images/noisy_0001.png'],
147
+ ['test_images/noisy_0002.png'],
148
+ ['test_images/noisy_0003.png'],
149
+ ['test_images/noisy_0004.png'],
150
+ ['test_images/rain-01.png'],
151
+ ['test_images/rain-02.png'],
152
+ ['test_images/rain-03.png'],
153
+ ['test_images/rain-04.png'],
154
+ ['test_images/rain-05.png'],
155
+ ['test_images/rain-06.png'],
156
+ ['test_images/hazy-00.jpg'],
157
+ ['test_images/hazy-01.jpg'],
158
+ ['test_images/hazy-02.jpg'],
159
+ ['test_images/hazy-03.jpg'],
160
+ ['test_images/hazy-04.jpg'],
161
+ ]
162
+ css = """
163
+ .image-frame img, .image-container img {
164
+ width: auto;
165
+ height: auto;
166
+ max-width: none;
167
+ }
168
+ """
169
+
170
+ demo = gr.Interface(
171
+ fn=restore_image,
172
+ inputs=[gr.Image(type="pil", label="Input")],
173
+ outputs=[gr.Image(type="pil", label="Ouput")],
174
+ title=title,
175
+ description=description,
176
+ article=article,
177
+ examples=examples,
178
+ css=css,
179
+ )
180
+
181
+
182
+ # if __name__ == "__main__":
183
+ demo.launch(debug=True, show_error=True)
ckpt/cata_promptxrestormeronlyattn_epoch=30-step=275962.ckpt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b9657e922beea4c56f90cbcec6708a9cebdf79f48b827a422bffdc0e9c2cc3b0
3
+ size 436105069
net/__pycache__/cata_prompt_xrestormer.cpython-38.pyc ADDED
Binary file (29.7 kB). View file
 
net/cata_prompt_xrestormer.py ADDED
@@ -0,0 +1,1009 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ from torch import einsum
4
+ import torch.nn.functional as F
5
+ from pdb import set_trace as stx
6
+ import numbers
7
+ from einops import rearrange
8
+ import math
9
+
10
+
11
+ def to(x):
12
+ return {'device': x.device, 'dtype': x.dtype}
13
+
14
+ def pair(x):
15
+ return (x, x) if not isinstance(x, tuple) else x
16
+
17
+ def expand_dim(t, dim, k):
18
+ t = t.unsqueeze(dim = dim)
19
+ expand_shape = [-1] * len(t.shape)
20
+ expand_shape[dim] = k
21
+ return t.expand(*expand_shape)
22
+
23
+ def rel_to_abs(x):
24
+ b, l, m = x.shape
25
+ r = (m + 1) // 2
26
+
27
+ col_pad = torch.zeros((b, l, 1), **to(x))
28
+ x = torch.cat((x, col_pad), dim = 2)
29
+ flat_x = rearrange(x, 'b l c -> b (l c)')
30
+ flat_pad = torch.zeros((b, m - l), **to(x))
31
+ flat_x_padded = torch.cat((flat_x, flat_pad), dim = 1)
32
+ final_x = flat_x_padded.reshape(b, l + 1, m)
33
+ final_x = final_x[:, :l, -r:]
34
+ return final_x
35
+
36
+ def relative_logits_1d(q, rel_k):
37
+ b, h, w, _ = q.shape
38
+ r = (rel_k.shape[0] + 1) // 2
39
+
40
+ logits = einsum('b x y d, r d -> b x y r', q, rel_k)
41
+ logits = rearrange(logits, 'b x y r -> (b x) y r')
42
+ logits = rel_to_abs(logits)
43
+
44
+ logits = logits.reshape(b, h, w, r)
45
+ logits = expand_dim(logits, dim = 2, k = r)
46
+ return logits
47
+
48
+ class RelPosEmb(nn.Module):
49
+ def __init__(
50
+ self,
51
+ block_size,
52
+ rel_size,
53
+ dim_head
54
+ ):
55
+ super().__init__()
56
+ height = width = rel_size
57
+ scale = dim_head ** -0.5
58
+
59
+ self.block_size = block_size
60
+ self.rel_height = nn.Parameter(torch.randn(height * 2 - 1, dim_head) * scale)
61
+ self.rel_width = nn.Parameter(torch.randn(width * 2 - 1, dim_head) * scale)
62
+
63
+ def forward(self, q):
64
+ block = self.block_size
65
+
66
+ q = rearrange(q, 'b (x y) c -> b x y c', x = block)
67
+ rel_logits_w = relative_logits_1d(q, self.rel_width)
68
+ rel_logits_w = rearrange(rel_logits_w, 'b x i y j-> b (x y) (i j)')
69
+
70
+ q = rearrange(q, 'b x y d -> b y x d')
71
+ rel_logits_h = relative_logits_1d(q, self.rel_height)
72
+ rel_logits_h = rearrange(rel_logits_h, 'b x i y j -> b (y x) (j i)')
73
+ return rel_logits_w + rel_logits_h
74
+
75
+ ##########################################################################
76
+ ## Layer Norm
77
+
78
+ def to_3d(x):
79
+ return rearrange(x, 'b c h w -> b (h w) c')
80
+
81
+ def to_4d(x,h,w):
82
+ return rearrange(x, 'b (h w) c -> b c h w',h=h,w=w)
83
+
84
+ class BiasFree_LayerNorm(nn.Module):
85
+ def __init__(self, normalized_shape):
86
+ super(BiasFree_LayerNorm, self).__init__()
87
+ if isinstance(normalized_shape, numbers.Integral):
88
+ normalized_shape = (normalized_shape,)
89
+ normalized_shape = torch.Size(normalized_shape)
90
+
91
+ assert len(normalized_shape) == 1
92
+
93
+ self.weight = nn.Parameter(torch.ones(normalized_shape))
94
+ self.normalized_shape = normalized_shape
95
+
96
+ def forward(self, x):
97
+ sigma = x.var(-1, keepdim=True, unbiased=False)
98
+ return x / torch.sqrt(sigma+1e-5) * self.weight
99
+
100
+ class WithBias_LayerNorm(nn.Module):
101
+ def __init__(self, normalized_shape):
102
+ super(WithBias_LayerNorm, self).__init__()
103
+ if isinstance(normalized_shape, numbers.Integral):
104
+ normalized_shape = (normalized_shape,)
105
+ normalized_shape = torch.Size(normalized_shape)
106
+
107
+ assert len(normalized_shape) == 1
108
+
109
+ self.weight = nn.Parameter(torch.ones(normalized_shape))
110
+ self.bias = nn.Parameter(torch.zeros(normalized_shape))
111
+ self.normalized_shape = normalized_shape
112
+
113
+ def forward(self, x):
114
+ mu = x.mean(-1, keepdim=True)
115
+ sigma = x.var(-1, keepdim=True, unbiased=False)
116
+ return (x - mu) / torch.sqrt(sigma+1e-5) * self.weight + self.bias
117
+
118
+ class RestormerLayerNorm(nn.Module):
119
+ def __init__(self, dim, LayerNorm_type):
120
+ super(RestormerLayerNorm, self).__init__()
121
+ if LayerNorm_type =='BiasFree':
122
+ self.body = BiasFree_LayerNorm(dim)
123
+ else:
124
+ self.body = WithBias_LayerNorm(dim)
125
+
126
+ def forward(self, x):
127
+ h, w = x.shape[-2:]
128
+ return to_4d(self.body(to_3d(x)), h, w)
129
+
130
+ ##########################################################################
131
+ ## Gated-Dconv Feed-Forward Network (GDFN)
132
+ class HardFeedForward(nn.Module):
133
+ def __init__(self, dim, ffn_expansion_factor, bias):
134
+ super(HardFeedForward, self).__init__()
135
+
136
+ hidden_features = int(dim*ffn_expansion_factor)
137
+
138
+ self.project_in = nn.Conv2d(dim, hidden_features*2, kernel_size=1, bias=bias)
139
+
140
+ self.dwconv = nn.Conv2d(hidden_features*2, hidden_features*2, kernel_size=3, stride=1, padding=1, groups=hidden_features*2, bias=bias)
141
+
142
+ self.project_out = nn.Conv2d(hidden_features, dim, kernel_size=1, bias=bias)
143
+
144
+ def forward(self, x):
145
+ x = self.project_in(x)
146
+ x1, x2 = self.dwconv(x).chunk(2, dim=1)
147
+ x = F.gelu(x1) * x2
148
+ x = self.project_out(x)
149
+ return x
150
+
151
+ def round_to_nearest_power_of_2(x):
152
+ if x & (x - 1) == 0: # Step 1: Check if x is already a power of 2
153
+ return x
154
+ msb_pos = x.bit_length() - 1 # Step 2: Find MSB position
155
+ lower_bound = 1 << msb_pos # Step 3: Calculate lower bound
156
+ upper_bound = 1 << (msb_pos + 1) # Step 4: Calculate upper bound
157
+ midpoint = (upper_bound + lower_bound) // 2 # Calculate midpoint
158
+ if x < midpoint: # Step 5 & 6: Compare and decide to round down or up
159
+ return lower_bound
160
+ else:
161
+ return upper_bound
162
+
163
+ ##########################################################################
164
+ ## Gated-Dconv Feed-Forward Network (GDFN)
165
+ class EasyFeedForward(nn.Module):
166
+ def __init__(self, dim, ffn_expansion_factor, bias):
167
+ super(EasyFeedForward, self).__init__()
168
+
169
+ ffn_channel = int(ffn_expansion_factor * dim)
170
+ ffn_channel = round_to_nearest_power_of_2(ffn_channel)
171
+ #print("FFN Channel: ", ffn_channel)
172
+ self.conv1 = nn.Conv2d(in_channels=dim, out_channels=ffn_channel, kernel_size=1, padding=0, stride=1, groups=1, bias=True)
173
+ self.conv2 = nn.Conv2d(in_channels=ffn_channel // 2, out_channels=dim, kernel_size=1, padding=0, stride=1, groups=1, bias=True)
174
+
175
+ self.sg = SimpleGate()
176
+ self.project_out = nn.Conv2d(dim, dim, kernel_size=1, bias=bias)
177
+
178
+ def forward(self, x):
179
+
180
+ x = self.conv1(x)
181
+ x = self.sg(x)
182
+ x = self.conv2(x)
183
+ x = self.project_out(x)
184
+ return x
185
+
186
+
187
+ class SimpleGate(nn.Module):
188
+ def forward(self, x):
189
+ x1, x2 = x.chunk(2, dim=1)
190
+ return x1 * x2
191
+
192
+ def batch_index_select(x, idx):
193
+ if len(x.size()) == 3:
194
+ B, N, C = x.size()
195
+ N_new = idx.size(1)
196
+ offset = torch.arange(B, dtype=torch.long, device=x.device).view(B, 1) * N
197
+ idx = idx + offset
198
+ out = x.reshape(B*N, C)[idx.reshape(-1)].reshape(B, N_new, C)
199
+ return out
200
+ elif len(x.size()) == 2:
201
+ B, N = x.size()
202
+ N_new = idx.size(1)
203
+ offset = torch.arange(B, dtype=torch.long, device=x.device).view(B, 1) * N
204
+ idx = idx + offset
205
+ out = x.reshape(B*N)[idx.reshape(-1)].reshape(B, N_new)
206
+ return out
207
+ else:
208
+ raise NotImplementedError
209
+
210
+ def batch_index_fill(x, x1, x2, idx1, idx2):
211
+ B, N, C = x.size()
212
+ B, N1, C = x1.size()
213
+ B, N2, C = x2.size()
214
+
215
+ offset = torch.arange(B, dtype=torch.long, device=x.device).view(B, 1)
216
+ idx1 = idx1 + offset * N
217
+ idx2 = idx2 + offset * N
218
+
219
+ x = x.reshape(B*N, C)
220
+
221
+ x[idx1.reshape(-1)] = x1.reshape(B*N1, C)
222
+ x[idx2.reshape(-1)] = x2.reshape(B*N2, C)
223
+
224
+ x = x.reshape(B, N, C)
225
+ return x
226
+
227
+
228
+
229
+ class LayerNorm(nn.Module):
230
+ r""" LayerNorm that supports two data formats: channels_last (default) or channels_first.
231
+ The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
232
+ shape (batch_size, height, width, channels) while channels_first corresponds to inputs
233
+ with shape (batch_size, channels, height, width).
234
+ """
235
+ def __init__(self, normalized_shape, eps=1e-6, data_format="channels_first"):
236
+ super().__init__()
237
+ self.weight = nn.Parameter(torch.ones(normalized_shape))
238
+ self.bias = nn.Parameter(torch.zeros(normalized_shape))
239
+ self.eps = eps
240
+ self.data_format = data_format
241
+ if self.data_format not in ["channels_last", "channels_first"]:
242
+ raise NotImplementedError
243
+ self.normalized_shape = (normalized_shape, )
244
+
245
+ def forward(self, x):
246
+ if self.data_format == "channels_last":
247
+ return F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
248
+ elif self.data_format == "channels_first":
249
+ u = x.mean(1, keepdim=True)
250
+ s = (x - u).pow(2).mean(1, keepdim=True)
251
+ x = (x - u) / torch.sqrt(s + self.eps)
252
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
253
+ return x
254
+
255
+ class PredictorLG(nn.Module):
256
+ """ Importance Score Predictor
257
+ """
258
+ def __init__(self, dim, window_size=8, k=4,ratio=0.5):
259
+ super().__init__()
260
+
261
+ self.ratio = ratio
262
+ self.window_size = window_size
263
+ cdim = dim + k
264
+ embed_dim = window_size**2
265
+
266
+ self.in_conv = nn.Sequential(
267
+ nn.Conv2d(cdim, cdim//4, 1),
268
+ LayerNorm(cdim//4),
269
+ nn.LeakyReLU(negative_slope=0.1, inplace=True),
270
+ )
271
+
272
+ self.out_mask = nn.Sequential(
273
+ nn.Linear(embed_dim, window_size),
274
+ nn.LeakyReLU(negative_slope=0.1, inplace=True),
275
+ nn.Linear(window_size, 2),
276
+ nn.Softmax(dim=-1)
277
+ )
278
+
279
+ self.out_SA = nn.Sequential(
280
+ nn.Conv2d(cdim//4, 1, 3, 1, 1),
281
+ nn.Sigmoid(),
282
+ )
283
+
284
+
285
+ def forward(self, input_x, mask=None, ratio=0.5, training = False):
286
+
287
+ x = self.in_conv(input_x)
288
+
289
+ sa = self.out_SA(x)
290
+
291
+ x = torch.mean(x, keepdim=True, dim=1)
292
+
293
+ x = rearrange(x,'b c (h dh) (w dw) -> b (h w) (dh dw c)', dh=self.window_size, dw=self.window_size)
294
+ B, N, C = x.size()
295
+
296
+ pred_score = self.out_mask(x)
297
+ mask = F.gumbel_softmax(pred_score, hard=True, dim=2)[:, :, 0:1]
298
+
299
+ if training:
300
+ return mask, sa
301
+ else:
302
+ score = pred_score[:, : , 0]
303
+ B, N = score.shape
304
+ r = torch.mean(mask,dim=(0,1))*1.0
305
+ if self.ratio == 1:
306
+ num_keep_node = N #int(N * r) #int(N * r)
307
+ else:
308
+ num_keep_node = min(int(N * r * 2 * self.ratio), N)
309
+ idx = torch.argsort(score, dim=1, descending=True)
310
+ idx1 = idx[:, :num_keep_node]
311
+ idx2 = idx[:, num_keep_node:]
312
+ return [idx1, idx2], sa
313
+
314
+
315
+ class BranchSelector(nn.Module):
316
+ def __init__(self, dim, hard_ratio = 0.5):
317
+ super(BranchSelector, self).__init__()
318
+ self.dim = dim
319
+ self.hard_ratio = hard_ratio
320
+
321
+ self.in_conv = nn.Sequential(
322
+ nn.Conv2d(dim, dim//4, 1),
323
+ LayerNorm(dim//4),
324
+ nn.LeakyReLU(negative_slope=0.1, inplace=True),
325
+ )
326
+
327
+
328
+ self.se = nn.Sequential(
329
+ nn.AdaptiveAvgPool2d(1),
330
+ nn.Conv2d(dim//4, dim//4, 1, bias=False),
331
+ nn.LeakyReLU(0.1, True),
332
+ nn.Conv2d(dim//4, dim//4, 1, bias=False),
333
+ )
334
+
335
+ self.classifier = nn.Sequential(
336
+ nn.Linear(dim//4, 1),
337
+ nn.Sigmoid()
338
+ )
339
+
340
+
341
+ def forward(self, x, training = False):
342
+ N, C, H, W = x.shape
343
+ x = self.in_conv(x)
344
+ x = self.se(x)
345
+ x = x.mean([2, 3])
346
+ label = self.classifier(x) #[B, 1]
347
+ label = F.gumbel_softmax(label, hard=True, dim=0).squeeze(1)
348
+ if training:
349
+ return label
350
+ else:
351
+ num_keep_node = min(int(N * self.hard_ratio), N)
352
+ idx = torch.argsort(label, descending=True)
353
+ idx1 = idx[:num_keep_node]
354
+ idx2 = idx[num_keep_node:]
355
+ return [idx1, idx2]
356
+
357
+ class CAMixer(nn.Module):
358
+ def __init__(self, dim, window_size=8, bias=True, is_deformable=True, num_heads = 4, dim_head = 16,overlap_ratio = 0.5, ratio=0.5):
359
+ super().__init__()
360
+
361
+ self.dim = dim
362
+ self.window_size = window_size
363
+ self.is_deformable = is_deformable
364
+ self.ratio = ratio
365
+
366
+ self.num_heads = num_heads
367
+ self.overlap_win_size = int(window_size * overlap_ratio) + window_size
368
+ self.dim_head = dim_head
369
+ self.inner_dim = self.dim_head * self.num_heads
370
+ self.scale = self.dim_head**-0.5
371
+
372
+
373
+ self.inner_dim = self.dim_head * self.num_heads
374
+
375
+ k = 3
376
+ d = 2
377
+
378
+ #self.proj_qkv = nn.Conv2d(self.dim, self.inner_dim*3, kernel_size=1, bias=bias)
379
+ self.proj_v = nn.Conv2d(self.dim, self.inner_dim, kernel_size=1, bias=bias)
380
+ self.proj_q = nn.Conv2d(self.dim, self.inner_dim, kernel_size=1, bias=bias)
381
+ self.proj_k = nn.Conv2d(self.dim, self.inner_dim, kernel_size=1, bias=bias)
382
+
383
+ self.unfold = nn.Unfold(kernel_size=(self.overlap_win_size, self.overlap_win_size), stride=window_size, padding=(self.overlap_win_size-window_size)//2)
384
+ self.project_out = nn.Conv2d(self.inner_dim, dim, kernel_size=1, bias=bias)
385
+ self.rel_pos_emb = RelPosEmb(
386
+ block_size = window_size,
387
+ rel_size = window_size + (self.overlap_win_size - window_size),
388
+ dim_head = self.dim_head
389
+ )
390
+
391
+ # Predictor
392
+ self.route = PredictorLG(dim = self.inner_dim,window_size = window_size,ratio=ratio)
393
+
394
+ def forward(self,x,condition_global=None, mask=None, training = False):
395
+ N,C,H,W = x.shape
396
+
397
+ qs = self.proj_q(x)
398
+ ks = self.proj_k(x)
399
+ vs = self.proj_v(x)
400
+
401
+ if self.is_deformable:
402
+ condition_wind = torch.stack(torch.meshgrid(torch.linspace(-1,1,self.window_size),torch.linspace(-1,1,self.window_size)))\
403
+ .type_as(x).unsqueeze(0).repeat(N, 1, H//self.window_size, W//self.window_size)
404
+ if condition_global is None:
405
+ _condition = torch.cat([vs, condition_wind], dim=1)
406
+ else:
407
+ _condition = torch.cat([vs, condition_global, condition_wind], dim=1)
408
+
409
+ mask, sa = self.route(_condition,ratio=self.ratio, training=training)
410
+ # easy attn
411
+ v_out_easy = vs*sa
412
+ # #print("mask", mask)
413
+ if training:
414
+ # spatial attention
415
+ qs = rearrange(qs, 'b c (h p1) (w p2) -> (b h w) (p1 p2) c', p1 = self.window_size, p2 = self.window_size)
416
+ ks, vs = map(lambda t: self.unfold(t), (ks, vs))
417
+ ks, vs = map(lambda t: rearrange(t, 'b (c j) i -> (b i) j c', c = self.inner_dim), (ks, vs))
418
+
419
+ # print(f'qs.shape:{qs.shape}, ks.shape:{ks.shape}, vs.shape:{vs.shape}')
420
+ #split heads
421
+ qs, ks, vs = map(lambda t: rearrange(t, 'b n (head c) -> (b head) n c', head = self.num_heads), (qs, ks, vs))
422
+
423
+ # attention
424
+ qs = qs * self.scale
425
+ spatial_attn = (qs @ ks.transpose(-2, -1))
426
+ spatial_attn += self.rel_pos_emb(qs)
427
+ spatial_attn = spatial_attn.softmax(dim=-1)
428
+
429
+ v_out_hard = (spatial_attn @ vs)
430
+ v_out_hard = rearrange(v_out_hard, '(b h w head) (p1 p2) c -> b (head c) (h p1) (w p2)', head = self.num_heads, h = H // self.window_size, w = W // self.window_size, p1 = self.window_size, p2 = self.window_size)
431
+
432
+ v_out_easy = rearrange(v_out_easy,'b c (h dh) (w dw) -> b (h w) (dh dw c)', dh=self.window_size, dw=self.window_size)
433
+ v_out_hard = rearrange(v_out_hard,'b c (h dh) (w dw) -> b (h w) (dh dw c)', dh=self.window_size, dw=self.window_size)
434
+
435
+
436
+ out = v_out_hard*mask + v_out_easy*(1-mask)
437
+
438
+ out = rearrange(out, 'b (h w) (dh dw c) -> b c (h dh) (w dw)', dh=self.window_size, dw=self.window_size,h = H // self.window_size, w = W // self.window_size)
439
+ out = self.project_out(out)
440
+
441
+ return out, torch.mean(mask,dim=1)
442
+
443
+ else:
444
+
445
+ qs = rearrange(qs, 'b c (h p1) (w p2) -> b (h w) (p1 p2 c)', p1 = self.window_size, p2 = self.window_size)
446
+ ks, vs = map(lambda t: self.unfold(t), (ks, vs))
447
+ ks, vs = map(lambda t: rearrange(t, 'b (c j) i -> b i (j c)', c = self.inner_dim), (ks, vs))
448
+
449
+ idx1, idx2 = mask
450
+ qs = batch_index_select(qs, idx1)
451
+ ks = batch_index_select(ks, idx1)
452
+ vs = batch_index_select(vs, idx1)
453
+
454
+ qs, ks, vs = map(lambda t: rearrange(t, 'b n (j c) -> (b n) j c', c = self.inner_dim), (qs, ks, vs))
455
+
456
+ qs, ks, vs = map(lambda t: rearrange(t, 'b n (head c) -> (b head) n c', head = self.num_heads), (qs, ks, vs))
457
+
458
+ # attention
459
+ #print(f'qs.shape:{qs.shape}, ks.shape:{ks.shape}, vs.shape:{vs.shape}')
460
+ qs = qs * self.scale
461
+ spatial_attn = (qs @ ks.transpose(-2, -1))
462
+ spatial_attn += self.rel_pos_emb(qs)
463
+ spatial_attn = spatial_attn.softmax(dim=-1)
464
+
465
+ v_out_hard = (spatial_attn @ vs)
466
+ v1 = rearrange(v_out_hard, '(b j head) (p1 p2) c -> b j (p1 p2 head c)',b = N, head = self.num_heads, p1 = self.window_size, p2 = self.window_size)
467
+ v2 = rearrange(v_out_easy,'b c (h dh) (w dw) -> b (h w) (dh dw c)', dh=self.window_size, dw=self.window_size)
468
+ v2 = batch_index_select(v2, idx2)
469
+ v_out = torch.cat([v1, v2], dim=1)
470
+ out = batch_index_fill(v_out.clone(), v1.clone(), v2.clone(), idx1, idx2)
471
+
472
+ out = rearrange(out, 'b (h w) (dh dw c) -> b c (h dh) (w dw)', dh=self.window_size, dw=self.window_size,h = H // self.window_size, w = W // self.window_size)
473
+ out = self.project_out(out)
474
+
475
+
476
+ return out, mask
477
+
478
+
479
+ ##########################################################################
480
+ ## Multi-DConv Head Transposed Self-Attention (MDTA)
481
+ class HardChannelAttention(nn.Module):
482
+ def __init__(self, dim, num_heads, bias):
483
+ super(HardChannelAttention, self).__init__()
484
+ self.num_heads = num_heads
485
+ self.temperature = nn.Parameter(torch.ones(num_heads, 1, 1))
486
+
487
+ self.qkv = nn.Conv2d(dim, dim*3, kernel_size=1, bias=bias)
488
+ self.qkv_dwconv = nn.Conv2d(dim*3, dim*3, kernel_size=3, stride=1, padding=1, groups=dim*3, bias=bias)
489
+ self.project_out = nn.Conv2d(dim, dim, kernel_size=1, bias=bias)
490
+
491
+ def forward(self, x):
492
+ b,c,h,w = x.shape
493
+
494
+ qkv = self.qkv_dwconv(self.qkv(x))
495
+ q,k,v = qkv.chunk(3, dim=1)
496
+
497
+ q = rearrange(q, 'b (head c) h w -> b head c (h w)', head=self.num_heads)
498
+ k = rearrange(k, 'b (head c) h w -> b head c (h w)', head=self.num_heads)
499
+ v = rearrange(v, 'b (head c) h w -> b head c (h w)', head=self.num_heads)
500
+
501
+ q = torch.nn.functional.normalize(q, dim=-1)
502
+ k = torch.nn.functional.normalize(k, dim=-1)
503
+
504
+ attn = (q @ k.transpose(-2, -1)) * self.temperature
505
+ attn = attn.softmax(dim=-1)
506
+
507
+ out = (attn @ v)
508
+
509
+ out = rearrange(out, 'b head c (h w) -> b (head c) h w', head=self.num_heads, h=h, w=w)
510
+
511
+ out = self.project_out(out)
512
+ return out
513
+
514
+ def image_idx_fill(x1, x2, idx1, idx2):
515
+ B1 = x1.shape[0]
516
+ B2 = x2.shape[0]
517
+ B = B1 + B2
518
+ x_combined = torch.zeros(B, *x1.shape[1:], device=x1.device, dtype=x1.dtype)
519
+ x_combined[idx1] = x1
520
+ x_combined[idx2] = x2
521
+ return x_combined
522
+
523
+ ##########################################################################
524
+ ## Multi-DConv Head Transposed Self-Attention (MDTA)
525
+ class EasyChannelAttention(nn.Module):
526
+ def __init__(self, dim, num_channel_heads, bias):
527
+ super(EasyChannelAttention, self).__init__()
528
+ dw_channel = dim
529
+ self.conv1 = nn.Conv2d(in_channels=dim, out_channels=dw_channel, kernel_size=1, padding=0, stride=1, groups=1, bias=True)
530
+ self.conv2 = nn.Conv2d(in_channels=dw_channel, out_channels=dw_channel, kernel_size=3, padding=1, stride=1, groups=dw_channel,
531
+ bias=True)
532
+ self.conv3 = nn.Conv2d(in_channels=dw_channel // 2, out_channels=dim, kernel_size=1, padding=0, stride=1, groups=1, bias=True)
533
+
534
+ # Simplified Channel Attention
535
+ self.sca = nn.Sequential(
536
+ nn.AdaptiveAvgPool2d(1),
537
+ nn.Conv2d(in_channels=dw_channel // 2, out_channels=dw_channel // 2, kernel_size=1, padding=0, stride=1,
538
+ groups=1, bias=True),
539
+ )
540
+
541
+ # SimpleGate
542
+ self.sg = SimpleGate()
543
+ self.project_out = nn.Conv2d(dim, dim, kernel_size=1, bias=bias)
544
+
545
+ def forward(self, x):
546
+ x = self.conv1(x)
547
+ x = self.conv2(x)
548
+ x = self.sg(x)
549
+ x = x * self.sca(x)
550
+ x = self.conv3(x)
551
+
552
+ out = self.project_out(x)
553
+ return out
554
+
555
+
556
+ ##########################################################################
557
+ class CATransformerBlock(nn.Module):
558
+ def __init__(self, dim, window_size, ratio, num_channel_heads, ffn_expansion_factor, bias, LayerNorm_type, num_heads = 4, dim_head = 16, overlap_ratio = 0.5, hard_ratio = 0.5):
559
+ super(CATransformerBlock, self).__init__()
560
+
561
+ self.spatial_attn = CAMixer(dim,window_size=window_size,ratio=ratio,num_heads = num_heads, dim_head = dim_head,overlap_ratio = overlap_ratio)
562
+ self.hard_channel_attn = HardChannelAttention(dim, num_channel_heads, bias)
563
+ self.easy_channel_attn = EasyChannelAttention(dim, num_channel_heads, bias)
564
+
565
+ self.norm1 = RestormerLayerNorm(dim, LayerNorm_type)
566
+ self.norm2 = RestormerLayerNorm(dim, LayerNorm_type)
567
+ self.norm3 = RestormerLayerNorm(dim, LayerNorm_type)
568
+ self.norm4 = RestormerLayerNorm(dim, LayerNorm_type)
569
+
570
+ self.channel_ffn = HardFeedForward(dim, ffn_expansion_factor, bias)
571
+ self.spatial_ffn = HardFeedForward(dim, ffn_expansion_factor, bias)
572
+
573
+ self.branch_selector = BranchSelector(dim, hard_ratio = hard_ratio)
574
+
575
+
576
+
577
+ def forward(self, x, global_condition =None, training = False):
578
+ label = self.branch_selector(x, training=training)
579
+ if training:
580
+ x_hard = x + self.hard_channel_attn(self.norm1(x))
581
+ x_easy = x + self.easy_channel_attn(self.norm1(x))
582
+ label = label.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
583
+ x = x_hard * label + x_easy * (1-label)
584
+
585
+ x = x + self.channel_ffn(self.norm2(x))
586
+
587
+ y, decision = self.spatial_attn(self.norm3(x), global_condition, training=training)
588
+ x = x + y
589
+
590
+ x = x + self.spatial_ffn(self.norm4(x))
591
+ return x, decision, torch.mean(label)
592
+ else:
593
+
594
+ idx1, idx2 = label
595
+ x_hard = torch.index_select(x, 0, idx1)
596
+ x_easy = torch.index_select(x, 0, idx2)
597
+
598
+ x_hard = x_hard + self.hard_channel_attn(self.norm1(x_hard))
599
+ x_easy = x_easy + self.easy_channel_attn(self.norm1(x_easy))
600
+
601
+ x = image_idx_fill(x_hard, x_easy, idx1, idx2)
602
+
603
+ x = x + self.channel_ffn(self.norm2(x))
604
+
605
+ y, spatial_mask = self.spatial_attn(self.norm3(x), global_condition, training=training)
606
+ x = x + y
607
+
608
+ x = x + self.spatial_ffn(self.norm4(x))
609
+ return x, spatial_mask, label
610
+
611
+
612
+ ##########################################################################
613
+ class ChannelTransformerBlock(nn.Module):
614
+ def __init__(self, dim, num_channel_heads, ffn_expansion_factor, bias, LayerNorm_type):
615
+ super(ChannelTransformerBlock, self).__init__()
616
+
617
+ self.channel_attn = EasyChannelAttention(dim, num_channel_heads, bias)
618
+ self.norm1 = RestormerLayerNorm(dim, LayerNorm_type)
619
+ self.norm2 = RestormerLayerNorm(dim, LayerNorm_type)
620
+
621
+ self.channel_ffn = EasyFeedForward(dim, ffn_expansion_factor, bias)
622
+
623
+ def forward(self, x):
624
+ x = x + self.channel_attn(self.norm1(x))
625
+ x = x + self.channel_ffn(self.norm2(x))
626
+ return x
627
+
628
+
629
+
630
+ ##########################################################################
631
+ ## Overlapped image patch embedding with 3x3 Conv
632
+ class OverlapPatchEmbed(nn.Module):
633
+ def __init__(self, in_c=3, embed_dim=48, bias=False):
634
+ super(OverlapPatchEmbed, self).__init__()
635
+
636
+ self.proj = nn.Conv2d(in_c, embed_dim, kernel_size=3, stride=1, padding=1, bias=bias)
637
+
638
+ def forward(self, x):
639
+ x = self.proj(x)
640
+
641
+ return x
642
+
643
+ ##########################################################################
644
+ ## Resizing modules
645
+ class Downsample(nn.Module):
646
+ def __init__(self, n_feat):
647
+ super(Downsample, self).__init__()
648
+
649
+ self.body = nn.Sequential(nn.Conv2d(n_feat, n_feat//2, kernel_size=3, stride=1, padding=1, bias=False),
650
+ nn.PixelUnshuffle(2))
651
+
652
+ def forward(self, x):
653
+ return self.body(x)
654
+
655
+
656
+ class Upsample(nn.Module):
657
+ def __init__(self, n_feat):
658
+ super(Upsample, self).__init__()
659
+
660
+ self.body = nn.Sequential(nn.Conv2d(n_feat, n_feat*2, kernel_size=3, stride=1, padding=1, bias=False),
661
+ nn.PixelShuffle(2))
662
+
663
+ def forward(self, x):
664
+ return self.body(x)
665
+
666
+
667
+ class SR_Upsample(nn.Sequential):
668
+ """SR_Upsample module.
669
+ Args:
670
+ scale (int): Scale factor. Supported scales: 2^n and 3.
671
+ num_feat (int): Channel number of features.
672
+ """
673
+
674
+ def __init__(self, scale, num_feat):
675
+ m = []
676
+
677
+ if (scale & (scale - 1)) == 0: # scale = 2^n
678
+ for _ in range(int(math.log(scale, 2))):
679
+ m.append(nn.Conv2d(num_feat, 4 * num_feat, kernel_size = 3, stride = 1, padding = 1))
680
+ m.append(nn.PixelShuffle(2))
681
+ elif scale == 3:
682
+ m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1))
683
+ m.append(nn.PixelShuffle(3))
684
+ else:
685
+ raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.')
686
+ super(SR_Upsample, self).__init__(*m)
687
+
688
+ ##---------- Prompt Gen Module -----------------------
689
+ class PromptGenBlock(nn.Module):
690
+ def __init__(self,prompt_dim=128,prompt_len=5,prompt_size = 96,lin_dim = 192):
691
+ super(PromptGenBlock,self).__init__()
692
+ self.prompt_param = nn.Parameter(torch.rand(1,prompt_len,prompt_dim,prompt_size,prompt_size))
693
+ self.linear_layer = nn.Linear(lin_dim,prompt_len)
694
+ self.conv3x3 = nn.Conv2d(prompt_dim,prompt_dim,kernel_size=3,stride=1,padding=1,bias=False)
695
+
696
+
697
+ def forward(self,x):
698
+ B,C,H,W = x.shape
699
+ emb = x.mean(dim=(-2,-1))
700
+ prompt_weights = F.softmax(self.linear_layer(emb),dim=1)
701
+ prompt = prompt_weights.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) * self.prompt_param.unsqueeze(0).repeat(B,1,1,1,1,1).squeeze(1)
702
+ prompt = torch.sum(prompt,dim=1)
703
+ prompt = F.interpolate(prompt,(H,W),mode="bilinear")
704
+ prompt = self.conv3x3(prompt)
705
+
706
+ return prompt
707
+
708
+
709
+
710
+ class XRestormerLayer(nn.Module):
711
+ def __init__(self, dim, depth, window_size, ratio, num_channel_heads, ffn_expansion_factor, bias, LayerNorm_type, num_heads, dim_head, overlap_ratio, hard_ratio):
712
+ super(XRestormerLayer, self).__init__()
713
+ self.layer = nn.Sequential(*[CATransformerBlock(dim=dim, window_size = window_size, ratio = ratio, num_channel_heads=num_channel_heads, ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, num_heads = num_heads, dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio=hard_ratio) for i in range(depth)])
714
+
715
+ def forward(self, x, global_condition=None, training = False):
716
+ if training:
717
+ decision_avg = 0
718
+ hard_ratio_avg = 0
719
+ for layer in self.layer:
720
+ x, decision, hard_ratio = layer(x, global_condition, training = training)
721
+ decision_avg += decision
722
+ hard_ratio_avg += hard_ratio
723
+ decision_avg /= len(self.layer)
724
+ hard_ratio_avg /= len(self.layer)
725
+ return x, decision_avg, hard_ratio_avg
726
+ else:
727
+ spatial_mask_list = []
728
+ channel_mask_list = []
729
+ for layer in self.layer:
730
+ x, spatial_mask, channel_mask = layer(x, global_condition, training = training)
731
+ spatial_mask_list.append(spatial_mask)
732
+ channel_mask_list.append(channel_mask)
733
+ return x, spatial_mask_list, channel_mask_list
734
+
735
+
736
+
737
+
738
+ ##########################################################################
739
+
740
+
741
+ class CATAPromptXRestormerOnlyAttn(nn.Module):
742
+ def __init__(self,
743
+ inp_channels=3,
744
+ out_channels=3,
745
+ dim = 48,
746
+ num_blocks = [4,6,6,8],
747
+ num_refinement_blocks = 4,
748
+ channel_heads = [1,2,4,8],
749
+ spatial_heads = [1,2,4,8],
750
+ overlap_ratio = 0.5,
751
+ dim_head = 16,
752
+ ratio = 0.5,
753
+ window_size = 8,
754
+ bias = False,
755
+ ffn_expansion_factor = 2.66,
756
+ LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
757
+ dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
758
+ scale = 1,
759
+ prompt = True,
760
+ hard_ratio = 0.5
761
+ ):
762
+
763
+ super(CATAPromptXRestormerOnlyAttn, self).__init__()
764
+ print("Initializing XRestormer")
765
+ self.scale = scale
766
+ self.ratio = ratio
767
+ self.hard_ratio = hard_ratio
768
+
769
+ self.patch_embed = OverlapPatchEmbed(inp_channels, dim)
770
+ self.encoder_level1 = XRestormerLayer(dim=dim, window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[0], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth=num_blocks[0], num_heads = spatial_heads[0], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
771
+
772
+ self.down1_2 = Downsample(dim) ## From Level 1 to Level 2
773
+ self.encoder_level2= XRestormerLayer(dim=int(dim*2**1), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[1], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth=num_blocks[1], num_heads = spatial_heads[1], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
774
+
775
+ self.down2_3 = Downsample(int(dim*2**1)) ## From Level 2 to Level 3
776
+ self.encoder_level3 = XRestormerLayer(dim=int(dim*2**2), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[2], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth=num_blocks[2], num_heads = spatial_heads[2], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
777
+
778
+ self.down3_4 = Downsample(int(dim*2**2)) ## From Level 3 to Level 4
779
+ self.latent = XRestormerLayer(dim=int(dim*2**3), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[3], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth=num_blocks[3], num_heads = spatial_heads[3], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
780
+
781
+ #self.latent = nn.Sequential(*[TransformerBlock(dim=int(dim*2**3), window_size = window_size, overlap_ratio=0.5, num_channel_heads=channel_heads[3], num_spatial_heads=8, spatial_dim_head = 16, ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type) for i in range(num_blocks[3])])
782
+
783
+
784
+ self.up4_3 = Upsample(int(dim*2**2)) ## From Level 4 to Level 3
785
+ self.reduce_chan_level3 = nn.Conv2d(int(dim*2**1) + 192, int(dim*2**2), kernel_size=1, bias=bias)
786
+ self.decoder_level3 = XRestormerLayer(dim=int(dim*2**2), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[2], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth = num_blocks[2], num_heads = spatial_heads[2], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
787
+
788
+
789
+ self.up3_2 = Upsample(int(dim*2**2)) ## From Level 3 to Level 2
790
+ self.reduce_chan_level2 = nn.Conv2d(int(dim*2**2), int(dim*2**1), kernel_size=1, bias=bias)
791
+ self.decoder_level2 = XRestormerLayer(dim=int(dim*2**1), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[1], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth = num_blocks[1], num_heads = spatial_heads[1], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
792
+
793
+ self.up2_1 = Upsample(int(dim*2**1)) ## From Level 2 to Level 1 (NO 1x1 conv to reduce channels)
794
+
795
+ self.decoder_level1 = XRestormerLayer(dim=int(dim*2**1), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[0], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth = num_blocks[0], num_heads = spatial_heads[0], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
796
+
797
+ self.refinement = XRestormerLayer(dim=int(dim*2**1), window_size = window_size, ratio = ratio, num_channel_heads=channel_heads[0], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type, depth=num_refinement_blocks, num_heads = spatial_heads[0], dim_head = dim_head, overlap_ratio = overlap_ratio, hard_ratio = hard_ratio)
798
+
799
+ self.output = nn.Conv2d(int(dim*2**1), out_channels, kernel_size=3, stride=1, padding=1, bias=bias)
800
+
801
+ self.prompt = prompt
802
+ if prompt:
803
+ self.prompt1 = PromptGenBlock(prompt_dim=64,prompt_len=5,prompt_size = 64,lin_dim = 96)
804
+ self.prompt2 = PromptGenBlock(prompt_dim=128,prompt_len=5,prompt_size = 32,lin_dim = 192)
805
+ self.prompt3 = PromptGenBlock(prompt_dim=320,prompt_len=5,prompt_size = 16,lin_dim = 384)
806
+
807
+ self.noise_level1 = ChannelTransformerBlock(dim=int(dim*2**1)+64, num_channel_heads = 1, ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type)
808
+ self.reduce_noise_level1 = nn.Conv2d(int(dim*2**1)+64,int(dim*2**1),kernel_size=1,bias=bias)
809
+
810
+ self.noise_level2 = ChannelTransformerBlock(dim=int(dim*2**1) + 224, num_channel_heads = 1, ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type)
811
+ self.reduce_noise_level2 = nn.Conv2d(int(dim*2**1)+224,int(dim*2**2),kernel_size=1,bias=bias)
812
+
813
+ self.noise_level3 = ChannelTransformerBlock(dim=int(dim*2**2) + 512, num_channel_heads = 1, ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type)
814
+ self.reduce_noise_level3 = nn.Conv2d(int(dim*2**2)+512,int(dim*2**2),kernel_size=1,bias=bias)
815
+
816
+ self.global_predictor = nn.Sequential(nn.Conv2d(dim, 8, 1, 1, 0, bias=True),
817
+ nn.LeakyReLU(negative_slope=0.1, inplace=True),
818
+ nn.Conv2d(8, 2, 3, 1, 1, bias=True),
819
+ nn.LeakyReLU(negative_slope=0.1, inplace=True))
820
+
821
+
822
+
823
+ def forward(self, inp_img, training = False):
824
+ all_spatial_mask = {}
825
+ all_channel_mask = {}
826
+ if self.scale > 1:
827
+ inp_img = F.interpolate(inp_img, scale_factor=self.scale, mode='bilinear', align_corners=False)
828
+ B, C, H, W = inp_img.shape
829
+ inp_enc_level1 = self.patch_embed(inp_img)
830
+ condition_global = self.global_predictor(inp_enc_level1)
831
+ condition_global_level2 = F.interpolate(condition_global, size=(H //2, W //2), mode='bilinear', align_corners=False)
832
+ condition_global_level3 = F.interpolate(condition_global, size=(H //4, W //4), mode='bilinear', align_corners=False)
833
+ condition_global_level4 = F.interpolate(condition_global, size=(H //8, W //8), mode='bilinear', align_corners=False)
834
+ if training:
835
+ # Encoder1
836
+
837
+ decision_avg = 0
838
+ hard_ratio_avg = 0
839
+ out_enc_level1, decision, hard_ratio = self.encoder_level1(inp_enc_level1, condition_global, training = training)
840
+ inp_enc_level2 = self.down1_2(out_enc_level1)
841
+ decision_avg += decision
842
+ hard_ratio_avg += hard_ratio
843
+
844
+
845
+ # Encoder2
846
+ out_enc_level2, decision, hard_ratio = self.encoder_level2(inp_enc_level2, condition_global_level2, training = training)
847
+ inp_enc_level3 = self.down2_3(out_enc_level2)
848
+ decision_avg += decision
849
+ hard_ratio_avg += hard_ratio
850
+
851
+ # Encoder3
852
+ out_enc_level3, decision, hard_ratio = self.encoder_level3(inp_enc_level3, condition_global_level3, training = training)
853
+ inp_enc_level4 = self.down3_4(out_enc_level3)
854
+ decision_avg += decision
855
+ hard_ratio_avg += hard_ratio
856
+
857
+ # Bottleneck
858
+ latent, decision, hard_ratio = self.latent(inp_enc_level4, condition_global_level4, training = training)
859
+ decision_avg += decision
860
+ hard_ratio_avg += hard_ratio
861
+
862
+ if self.prompt:
863
+ dec3_param = self.prompt3(latent)
864
+ latent = torch.cat([latent, dec3_param], 1)
865
+ latent = self.noise_level3(latent)
866
+ latent = self.reduce_noise_level3(latent)
867
+
868
+
869
+ inp_dec_level3 = self.up4_3(latent)
870
+ inp_dec_level3 = torch.cat([inp_dec_level3, out_enc_level3], 1)
871
+ inp_dec_level3 = self.reduce_chan_level3(inp_dec_level3)
872
+ out_dec_level3, decision, hard_ratio = self.decoder_level3(inp_dec_level3, condition_global_level3, training = training)
873
+ decision_avg += decision
874
+ hard_ratio_avg += hard_ratio
875
+
876
+ if self.prompt:
877
+ dec2_param = self.prompt2(out_dec_level3)
878
+ out_dec_level3 = torch.cat([out_dec_level3, dec2_param], 1)
879
+ out_dec_level3 = self.noise_level2(out_dec_level3)
880
+ out_dec_level3 = self.reduce_noise_level2(out_dec_level3)
881
+
882
+
883
+ inp_dec_level2 = self.up3_2(out_dec_level3)
884
+ inp_dec_level2 = torch.cat([inp_dec_level2, out_enc_level2], 1)
885
+ inp_dec_level2 = self.reduce_chan_level2(inp_dec_level2)
886
+ out_dec_level2, decision, hard_ratio = self.decoder_level2(inp_dec_level2, condition_global_level2, training = training)
887
+ decision_avg += decision
888
+ hard_ratio_avg += hard_ratio
889
+
890
+ if self.prompt:
891
+ dec1_param = self.prompt1(out_dec_level2)
892
+ out_dec_level2 = torch.cat([out_dec_level2, dec1_param], 1)
893
+ out_dec_level2 = self.noise_level1(out_dec_level2)
894
+ out_dec_level2 = self.reduce_noise_level1(out_dec_level2)
895
+
896
+
897
+ inp_dec_level1 = self.up2_1(out_dec_level2)
898
+ inp_dec_level1 = torch.cat([inp_dec_level1, out_enc_level1], 1)
899
+ out_dec_level1, decision, hard_ratio = self.decoder_level1(inp_dec_level1, condition_global, training = training)
900
+ decision_avg += decision
901
+ hard_ratio_avg += hard_ratio
902
+
903
+ out_dec_level1, decision, hard_ratio = self.refinement(out_dec_level1, condition_global, training = training)
904
+ decision_avg += decision
905
+ hard_ratio_avg += hard_ratio
906
+
907
+ out_dec_level1 = self.output(out_dec_level1) + inp_img
908
+
909
+ decision_avg /= 8
910
+ hard_ratio_avg /= 8
911
+
912
+ ratio_loss = 2*self.ratio*(torch.mean(decision_avg)-0.5)**2
913
+ hard_ratio_loss = 2*self.hard_ratio*(torch.mean(hard_ratio_avg)-0.5)**2
914
+
915
+ return out_dec_level1, ratio_loss, hard_ratio_loss
916
+
917
+ else:
918
+ out_enc_level1, spatial_mask1, channel_mask1 = self.encoder_level1(inp_enc_level1, condition_global, training = training)
919
+ inp_enc_level2 = self.down1_2(out_enc_level1)
920
+ out_enc_level2, spatial_mask2, channel_mask2 = self.encoder_level2(inp_enc_level2, condition_global_level2, training = training)
921
+ inp_enc_level3 = self.down2_3(out_enc_level2)
922
+ out_enc_level3, spatial_mask3, channel_mask3 = self.encoder_level3(inp_enc_level3, condition_global_level3, training = training)
923
+ inp_enc_level4 = self.down3_4(out_enc_level3)
924
+ latent, spatial_mask4, channel_mask4 = self.latent(inp_enc_level4, condition_global_level4, training = training)
925
+
926
+ if self.prompt:
927
+ dec3_param = self.prompt3(latent)
928
+ latent = torch.cat([latent, dec3_param], 1)
929
+ latent = self.noise_level3(latent)
930
+ latent = self.reduce_noise_level3(latent)
931
+
932
+ inp_dec_level3 = self.up4_3(latent)
933
+ inp_dec_level3 = torch.cat([inp_dec_level3, out_enc_level3], 1)
934
+ inp_dec_level3 = self.reduce_chan_level3(inp_dec_level3)
935
+ out_dec_level3, spatial_mask5, channel_mask5 = self.decoder_level3(inp_dec_level3, condition_global_level3, training = training)
936
+
937
+ if self.prompt:
938
+ dec2_param = self.prompt2(out_dec_level3)
939
+ out_dec_level3 = torch.cat([out_dec_level3, dec2_param], 1)
940
+ out_dec_level3 = self.noise_level2(out_dec_level3)
941
+ out_dec_level3 = self.reduce_noise_level2(out_dec_level3)
942
+
943
+ inp_dec_level2 = self.up3_2(out_dec_level3)
944
+ inp_dec_level2 = torch.cat([inp_dec_level2, out_enc_level2], 1)
945
+ inp_dec_level2 = self.reduce_chan_level2(inp_dec_level2)
946
+ out_dec_level2, spatial_mask6, channel_mask6 = self.decoder_level2(inp_dec_level2, condition_global_level2, training = training)
947
+
948
+ if self.prompt:
949
+ dec1_param = self.prompt1(out_dec_level2)
950
+ out_dec_level2 = torch.cat([out_dec_level2, dec1_param], 1)
951
+ out_dec_level2 = self.noise_level1(out_dec_level2)
952
+ out_dec_level2 = self.reduce_noise_level1(out_dec_level2)
953
+
954
+ inp_dec_level1 = self.up2_1(out_dec_level2)
955
+ inp_dec_level1 = torch.cat([inp_dec_level1, out_enc_level1], 1)
956
+ out_dec_level1, spatial_mask7, channel_mask7 = self.decoder_level1(inp_dec_level1, condition_global, training = training)
957
+
958
+ out_dec_level1, spatial_mask8, channel_mask8 = self.refinement(out_dec_level1, condition_global, training = training)
959
+ out_dec_level1 = self.output(out_dec_level1) + inp_img
960
+
961
+ all_spatial_mask["encoder_level1"] = spatial_mask1
962
+ all_spatial_mask["encoder_level2"] = spatial_mask2
963
+ all_spatial_mask["encoder_level3"] = spatial_mask3
964
+ all_spatial_mask["latent"] = spatial_mask4
965
+ all_spatial_mask["decoder_level3"] = spatial_mask5
966
+ all_spatial_mask["decoder_level2"] = spatial_mask6
967
+ all_spatial_mask["decoder_level1"] = spatial_mask7
968
+ all_spatial_mask["refinement"] = spatial_mask8
969
+
970
+ all_channel_mask["encoder_level1"] = channel_mask1
971
+ all_channel_mask["encoder_level2"] = channel_mask2
972
+ all_channel_mask["encoder_level3"] = channel_mask3
973
+ all_channel_mask["latent"] = channel_mask4
974
+ all_channel_mask["decoder_level3"] = channel_mask5
975
+ all_channel_mask["decoder_level2"] = channel_mask6
976
+ all_channel_mask["decoder_level1"] = channel_mask7
977
+ all_channel_mask["refinement"] = channel_mask8
978
+
979
+
980
+ return out_dec_level1, all_spatial_mask, all_channel_mask
981
+
982
+ if __name__ == "__main__":
983
+ training = False
984
+ model = CATAPromptXRestormerOnlyAttn(
985
+ inp_channels=3,
986
+ out_channels=3,
987
+ dim = 48,
988
+ num_blocks = [2,4,4,4],
989
+ num_refinement_blocks = 4,
990
+ channel_heads = [1,1,1,1],
991
+ spatial_heads = [1,2,4,8],
992
+ overlap_ratio = 0.5,
993
+ dim_head = 16,
994
+ ratio = 0.5,
995
+ window_size = 8,
996
+ bias = False,
997
+ ffn_expansion_factor = 2.66,
998
+ LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
999
+ dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
1000
+ scale = 1,
1001
+ prompt = True,
1002
+ hard_ratio = 0.5
1003
+ )
1004
+
1005
+ # torchstat
1006
+ x = torch.randn(8, 3, 64, 64)
1007
+
1008
+ y, all_spatial_mask, all_channel_mask = model(x, training=training)
1009
+ print("output shape", y.shape)
output.png CHANGED
output_masked.png ADDED
test_cata.py ADDED
@@ -0,0 +1,155 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import subprocess
3
+ from tqdm import tqdm
4
+ import numpy as np
5
+
6
+ import torch
7
+ from torch.utils.data import DataLoader
8
+ import os
9
+ import torch.nn as nn
10
+
11
+ # from utils.dataset_utils import DenoiseTestDataset, DerainDehazeDataset
12
+ # from utils.val_utils import AverageMeter, compute_psnr_ssim
13
+ # from utils.image_io import save_image_tensor
14
+
15
+ from PIL import Image
16
+ from torchvision.transforms import ToTensor
17
+
18
+
19
+ import lightning.pytorch as pl
20
+ import torch.nn.functional as F
21
+
22
+ from net.cata_prompt_xrestormer import CATAPromptXRestormerOnlyAttn
23
+ from einops import rearrange
24
+
25
+ # crop an image to the multiple of base
26
+ def crop_img(image, base=64):
27
+ h = image.shape[0]
28
+ w = image.shape[1]
29
+ crop_h = h % base
30
+ crop_w = w % base
31
+ return image[crop_h // 2:h - crop_h + crop_h // 2, crop_w // 2:w - crop_w + crop_w // 2, :]
32
+
33
+ class CATAPromptXRestormerIRModel(pl.LightningModule):
34
+ def __init__(self):
35
+ super().__init__()
36
+ self.net = CATAPromptXRestormerOnlyAttn(
37
+ inp_channels=3,
38
+ out_channels=3,
39
+ dim = 48,
40
+ num_blocks = [2,4,4,4],
41
+ num_refinement_blocks = 4,
42
+ channel_heads = [1,1,1,1],
43
+ spatial_heads = [1,2,4,8],
44
+ overlap_ratio = 0.5,
45
+ dim_head = 16,
46
+ ratio = 0.5,
47
+ window_size = 8,
48
+ bias = False,
49
+ ffn_expansion_factor = 2.66,
50
+ LayerNorm_type = 'WithBias', ## Other option 'BiasFree'
51
+ dual_pixel_task = False, ## True for dual-pixel defocus deblurring only. Also set inp_channels=6
52
+ scale = 1,
53
+ prompt = True,
54
+ hard_ratio = 0.5
55
+ )
56
+ self.loss_fn = nn.L1Loss()
57
+
58
+ def forward(self,x, training=False):
59
+ return self.net(x, training)
60
+
61
+ def np_to_pil(img_np):
62
+ """
63
+ Converts image in np.array format to PIL image.
64
+
65
+ From C x W x H [0..1] to W x H x C [0...255]
66
+ :param img_np:
67
+ :return:
68
+ """
69
+ ar = np.clip(img_np * 255, 0, 255).astype(np.uint8)
70
+
71
+ if img_np.shape[0] == 1:
72
+ ar = ar[0]
73
+ else:
74
+ assert img_np.shape[0] == 3, img_np.shape
75
+ ar = ar.transpose(1, 2, 0)
76
+
77
+ return Image.fromarray(ar)
78
+
79
+ def torch_to_np(img_var):
80
+ """
81
+ Converts an image in torch.Tensor format to np.array.
82
+
83
+ From 1 x C x W x H [0..1] to C x W x H [0..1]
84
+ :param img_var:
85
+ :return:
86
+ """
87
+ return img_var.detach().cpu().numpy()[0]
88
+
89
+ def save_image_tensor(image_tensor, output_path="output/"):
90
+ image_np = torch_to_np(image_tensor)
91
+ # print(image_np.shape)
92
+ p = np_to_pil(image_np)
93
+ p.save(output_path)
94
+
95
+
96
+
97
+ if __name__ == '__main__':
98
+
99
+ np.random.seed(0)
100
+ torch.manual_seed(0)
101
+ torch.cuda.set_device(0)
102
+
103
+ ckpt_path = "ckpt/cata_promptxrestormeronlyattn_epoch=30-step=275962.ckpt"
104
+ print("CKPT name : {}".format(ckpt_path))
105
+
106
+ net = CATAPromptXRestormerIRModel.load_from_checkpoint(ckpt_path).cuda()
107
+ net.eval()
108
+
109
+ degraded_path = "/home/jiachen/MyGradio/test_images/rain-01.png"
110
+
111
+ degraded_img = crop_img(np.array(Image.open(degraded_path).convert('RGB')), base=16)
112
+ toTensor = ToTensor()
113
+ degraded_img = toTensor(degraded_img)
114
+ print(degraded_img.shape)
115
+
116
+ with torch.no_grad():
117
+ degraded_img = degraded_img.unsqueeze(0).cuda()
118
+
119
+ _, _, H_old, W_old = degraded_img.shape
120
+
121
+
122
+ h_pad = (H_old // 64 + 1) * 64 - H_old
123
+ w_pad = (W_old // 64 + 1) * 64 - W_old
124
+ degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [2])], 2)[:,:,:H_old+h_pad,:]
125
+ degraded_img = torch.cat([degraded_img, torch.flip(degraded_img, [3])], 3)[:,:,:,:W_old+w_pad]
126
+
127
+ print("inputImage size", degraded_img.shape)
128
+ restored, spatial_mask, channel_mask = net(degraded_img, training=False)
129
+
130
+
131
+ encoder_level1_mask = spatial_mask['encoder_level1'][0][0][0]
132
+ window_size = 8
133
+ _, c, h, w = restored.shape
134
+
135
+ # Split the restored image into 8x8 windows
136
+ restored_windows = rearrange(restored, 'b c (h w1) (w w2) -> b c (h w) w1 w2', w1=window_size, w2=window_size)
137
+
138
+ # Mask out the windows according to the indices in encoder_level1_mask
139
+ for idx in encoder_level1_mask:
140
+ restored_windows[:, :, idx, :, :] = 1 # Mask out the window by setting it to one
141
+
142
+ # Reconstruct the image from the masked windows
143
+ restored_masked = rearrange(restored_windows, 'b c (h w) w1 w2 -> b c (h w1) (w w2)', h=h // window_size, w=w // window_size)
144
+
145
+ restored = restored[:,:,:H_old:,:W_old]
146
+ restored_masked = restored_masked[:,:,:H_old:,:W_old]
147
+ save_image_tensor(restored, "output.png")
148
+ save_image_tensor(restored_masked, "output_masked.png")
149
+
150
+
151
+
152
+
153
+
154
+
155
+
test_images/hazy-00.jpg DELETED
Binary file (164 kB)
 
test_images/hazy-01.jpg CHANGED
test_images/hazy-02.jpg CHANGED
test_images/hazy-05.jpg ADDED
test_images/noisy_0000.png CHANGED
test_images/noisy_0001.png CHANGED
test_images/noisy_0002.png CHANGED
test_images/noisy_0003.png CHANGED
test_images/noisy_0004.png CHANGED
test_images/{rain-03.png → rain-001.png} RENAMED
File without changes
test_images/rain-002.png ADDED
test_images/rain-003.png ADDED
test_images/rain-004.png ADDED
test_images/rain-005.png ADDED
test_images/rain-01.png DELETED
Binary file (293 kB)
 
test_images/rain-02.png DELETED
Binary file (254 kB)
 
test_images/rain-04.png DELETED
Binary file (232 kB)
 
test_images/rain-05.png DELETED
Binary file (330 kB)
 
test_images/rain-06.png DELETED
Binary file (254 kB)