diff --git a/inference.py b/inference.py new file mode 100644 index 0000000..fed6ab1 --- /dev/null +++ b/inference.py @@ -0,0 +1,140 @@ +import argparse +import chainer +import numpy as np +import os +from tqdm import tqdm +from brisque import BRISQUE +from mini_batch_loader import MiniBatchLoader +import State +from MyFCN import * +from pixelwise_a3c import * +import torch +import torch.optim as optim +from utils.traj_dataset import TrajectoryDataset +from utils.trajectory import list_of_tuple_to_traj +import cv2 +import json +from utils.compute_Rbrisque import brisque_reward +from utils.augmentation import aug_mask + +def overlapped_process(raw_x, mask, agent, current_state, patch_size, stride): + _, _, h, w = raw_x.shape + output = np.zeros_like(raw_x) + counter = np.zeros_like(raw_x) + for x in range(0, h, stride): + for y in range(0, w, stride): + x_end = min(x + patch_size, h) + y_end = min(y + patch_size, w) + patch_image = raw_x[:, :, x:x_end, y:y_end] + patch_mask = mask[:, :, x:x_end, y:y_end] + patch_output = inference_patch(agent, current_state, patch_image, patch_mask) + output[:, :, x:x_end, y:y_end] += patch_output + counter[:, :, x:x_end, y:y_end] += 1 + output = np.divide(output, counter) + return output + +def inference_patch(agent, current_state, patch_image, mask): + # only return the final result + current_state.reset(patch_image) + mask_squeeze = np.squeeze(mask, axis=1) + for t in range(0, args.episode_len): + action, inner_state = agent.act(current_state.tensor) + action = np.where(mask_squeeze==0, 1, action) # if mask equal to 0, the pixel isn't a rain, so the act should be id==1:"do nothing" + current_state.step(action, inner_state) + agent.stop_episode() + return current_state.image + +def inference(agent, raw_x, mask, name): + os.makedirs(os.path.join(args.save_dir_path, 'derained_result'), exist_ok=True) + current_state = State.State(args.move_range) + B, C, H, W = raw_x.shape + if H*W > 535000: # for high resolurion images, we use overlapped inference due to GPU limitations + output = overlapped_process(raw_x, mask, agent, current_state, patch_size=128, stride=64) + p = np.maximum(0,output) + p = np.minimum(1,p) + p = (p*255).astype(np.uint8) + p = np.transpose(p[0], [1,2,0]) + else: + current_state.reset(raw_x) + mask_squeeze = np.squeeze(mask, axis=1) + for t in range(0, args.episode_len): + action, inner_state = agent.act(current_state.tensor) + action = np.where(mask_squeeze==0, 1, action) # if mask equal to 0, the pixel isn't a rain, so the act should be id==1:"do nothing" + current_state.step(action, inner_state) + agent.stop_episode() + + p = np.maximum(0,current_state.image) + p = np.minimum(1,p) + p = (p*255).astype(np.uint8) + p = np.transpose(p[0], [1,2,0]) + + cv2.imwrite(os.path.join(args.save_dir_path, 'derained_result', name), p) + +def compute_diff(image1, image2): + return np.mean(np.abs(image1 - image2)) + +def main(args): + #_/_/_/ load dataset _/_/_/ + mini_batch_loader = MiniBatchLoader( + args.data_path, + args.image_dir_path) + + brisque_metrics = BRISQUE(url=False) + + chainer.cuda.get_device_from_id(args.gpu_id).use() + + current_state = State.State(args.move_range) + + train_data_size = MiniBatchLoader.count_paths(args.data_path) + + # criterion for training Rnet + CE = torch.nn.CrossEntropyLoss() + + for data_idx in range(0, train_data_size): + # train + raw_x, pseudo_ys, mask, name = mini_batch_loader.load_training_data(index=data_idx) + model = MyFcn(args.n_actions) + optimizer = chainer.optimizers.Adam(alpha=args.lr) + optimizer.setup(model) + agent = PixelWiseA3C_InnerState(model, optimizer, 5, args.gamma) + agent.act_deterministically = True + agent.model.to_gpu() + agent.load(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) + print("===== Process for {} =====".format(name)) + # init Reward function + r_net_brisque = Reward_Predictor(image_size=(args.pretrained_img_size, args.pretrained_img_size)).cuda() + r_net_brisque.load_state_dict(torch.load(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'Rnet', 'rnet_brisque.pt'))) + + inference(agent, raw_x, mask, name) + + agent.save(os.path.join(args.save_dir_path, 'model_weight', name)) + + # save Rnet model + os.makedirs(os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet'), exist_ok=True) + torch.save(r_net_brisque.state_dict(), os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet', 'rnet_brisque.pt')) + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='parameters for training') + # seed + parser.add_argument('--random_seed', type=int, default=1) + # Directories + parser.add_argument('--image_dir_path', type=str, default='dataset/') + parser.add_argument('--data_path', type=str, default='dataset/Rain12/testing.txt') + parser.add_argument('--gt_dir_path', type=str, default='dataset/Rain12/test/gt/') + parser.add_argument('--save_dir_path', type=str, default='./Results/Rain12/test/SRL-Derain/') + parser.add_argument('--checkpoint_dir_path', type=str, default='./Checkpoints/Rain12/SRL-Derain/') + # config + parser.add_argument('--gpu_id', type=int, default=0) + parser.add_argument('--move_range', type=int, default=3) + parser.add_argument('--episode_len', type=int, default=15) + parser.add_argument('--max_episode', type=int, default=150) + parser.add_argument('--gamma', type=float, default=0.99) + parser.add_argument('--n_actions', type=int, default=9) + parser.add_argument('--lr', type=float, default=1e-3) + parser.add_argument('--ld', type=float, default=0.05, help='lambda, the weight for reward') + parser.add_argument('--N_pre', type=int, default=6000) + parser.add_argument('--pretrained_img_size', type=int, default=128) + parser.add_argument('--pretrained_batch_size', type=int, default=64) + args = parser.parse_args() + + main(args) \ No newline at end of file diff --git a/inference_plus.py b/inference_plus.py new file mode 100644 index 0000000..97ac9be --- /dev/null +++ b/inference_plus.py @@ -0,0 +1,132 @@ +import argparse +import json +import chainer +import numpy as np +import os +from tqdm import tqdm +from brisque import BRISQUE +from mini_batch_loader import MiniBatchLoader +import State +from MyFCN import * +from pixelwise_a3c import * +import torch +from utils.trajectory import list_of_tuple_to_traj +import cv2 +from utils.compute_Rbrisque import brisque_reward + +def overlapped_process(raw_x, mask, agent, current_state, patch_size, stride): + _, _, h, w = raw_x.shape + output = np.zeros_like(raw_x) + counter = np.zeros_like(raw_x) + for x in range(0, h, stride): + for y in range(0, w, stride): + x_end = min(x + patch_size, h) + y_end = min(y + patch_size, w) + patch_image = raw_x[:, :, x:x_end, y:y_end] + patch_mask = mask[:, :, x:x_end, y:y_end] + patch_output = inference_patch(agent, current_state, patch_image, patch_mask) + output[:, :, x:x_end, y:y_end] += patch_output + counter[:, :, x:x_end, y:y_end] += 1 + output = np.divide(output, counter) + return output + +def inference_patch(agent, current_state, patch_image, mask): + # only return the final result + current_state.reset(patch_image) + mask_squeeze = np.squeeze(mask, axis=1) + for t in range(0, args.episode_len): + action, inner_state = agent.act(current_state.tensor) + action = np.where(mask_squeeze==0, 1, action) # if mask equal to 0, the pixel isn't a rain, so the act should be id==1:"do nothing" + current_state.step(action, inner_state) + agent.stop_episode() + return current_state.image + +def inference(agent, raw_x, mask, name): + os.makedirs(os.path.join(args.save_dir_path, 'derained_result'), exist_ok=True) + current_state = State.State(args.move_range) + B, C, H, W = raw_x.shape + if H*W > 535000: # for high resolurion images, we use overlapped inference due to GPU limitations + output = overlapped_process(raw_x, mask, agent, current_state, patch_size=128, stride=64) + p = np.maximum(0,output) + p = np.minimum(1,p) + p = (p*255).astype(np.uint8) + p = np.transpose(p[0], [1,2,0]) + else: + current_state.reset(raw_x) + mask_squeeze = np.squeeze(mask, axis=1) + for t in range(0, args.episode_len): + action, inner_state = agent.act(current_state.tensor) + action = np.where(mask_squeeze==0, 1, action) # if mask equal to 0, the pixel isn't a rain, so the act should be id==1:"do nothing" + current_state.step(action, inner_state) + agent.stop_episode() + + p = np.maximum(0,current_state.image) + p = np.minimum(1,p) + p = (p*255).astype(np.uint8) + p = np.transpose(p[0], [1,2,0]) + + cv2.imwrite(os.path.join(args.save_dir_path, 'derained_result', name), p) + +def compute_diff(image1, image2): + return np.mean(np.abs(image1 - image2)) + +def main(args): + #_/_/_/ load dataset _/_/_/ + mini_batch_loader = MiniBatchLoader( + args.data_path, + args.image_dir_path) + + brisque_metrics = BRISQUE(url=False) + + chainer.cuda.get_device_from_id(args.gpu_id).use() + + current_state = State.State(args.move_range) + + train_data_size = MiniBatchLoader.count_paths(args.data_path) + + for i in range(0, train_data_size): + # train + raw_x, pseudo_ys, mask, name = mini_batch_loader.load_training_data(index=i) + model = MyFcn(args.n_actions) + optimizer = chainer.optimizers.Adam(alpha=args.lr) + optimizer.setup(model) + agent = PixelWiseA3C_InnerState(model, optimizer, 5, args.gamma) + agent.act_deterministically = True + agent.model.to_gpu() + agent.load(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) + print("===== Process for {} =====".format(name)) + r_net_brisque = Reward_Predictor(image_size=(args.pretrained_img_size, args.pretrained_img_size)).cuda() + rnet_brisque_state_dict = torch.load(os.path.join(args.rnet_weight_dir, 'rnet_brisque.pt')) + r_net_brisque.load_state_dict(rnet_brisque_state_dict) + + inference(agent, raw_x, mask, name) + + agent.save(os.path.join(args.save_dir_path, 'model_weight', name)) + # save Rnet model + os.makedirs(os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet'), exist_ok=True) + torch.save(r_net_brisque.state_dict(), os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet', 'rnet_brisque.pt')) + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='parameters for training') + # seed + parser.add_argument('--random_seed', type=int, default=1) + # Directories + parser.add_argument('--image_dir_path', type=str, default='dataset/') + parser.add_argument('--data_path', type=str, default='dataset/Rain12/testing.txt') + parser.add_argument('--gt_dir_path', type=str, default='dataset/Rain12/test/gt/') + parser.add_argument('--save_dir_path', type=str, default='./Results/Rain12/test/SRL-Derain+/') + parser.add_argument('--rnet_weight_dir', type=str, default='./Checkpoints/Rain12/Rnet+/model_weight_best/') + parser.add_argument('--checkpoint_dir_path', type=str, default='./Checkpoints/Rain12/SRL-Derain+/') + # config + parser.add_argument('--gpu_id', type=int, default=0) + parser.add_argument('--move_range', type=int, default=3) + parser.add_argument('--episode_len', type=int, default=15) + parser.add_argument('--max_episode', type=int, default=150) + parser.add_argument('--gamma', type=float, default=0.99) + parser.add_argument('--n_actions', type=int, default=9) + parser.add_argument('--lr', type=float, default=1e-3) + parser.add_argument('--ld', type=float, default=0.05, help='lambda, the weight for reward') + parser.add_argument('--pretrained_img_size', type=int, default=128) + args = parser.parse_args() + + main(args) \ No newline at end of file diff --git a/main_srl_derain.py b/main_srl_derain.py index 218ca30..b675f78 100755 --- a/main_srl_derain.py +++ b/main_srl_derain.py @@ -13,8 +13,15 @@ from utils.traj_dataset import TrajectoryDataset from utils.trajectory import list_of_tuple_to_traj import cv2 +import json +import lpips +import skimage.io from utils.compute_Rbrisque import brisque_reward from utils.augmentation import aug_mask +from skimage.metrics import peak_signal_noise_ratio as psnr +from skimage.metrics import structural_similarity as ssim +from PIL import Image +from torchvision import transforms def overlapped_process(raw_x, mask, agent, current_state, patch_size, stride): _, _, h, w = raw_x.shape @@ -72,6 +79,50 @@ def inference(agent, raw_x, mask, name): def compute_diff(image1, image2): return np.mean(np.abs(image1 - image2)) +def evaluate_syn(agent, name, metrics): + gt_path = os.path.join(args.gt_dir_path, name) + result_path = os.path.join(args.save_dir_path, 'derained_result', name) + gt_img = cv2.imread(gt_path) + result_img = cv2.imread(result_path) + psnr_val = psnr(gt_img, result_img) + ssim_val = ssim(gt_img, result_img, channel_axis=2) + # transform2 = transforms.Compose([transforms.ToTensor()]) + # gt_tensor = transform2(Image.open(gt_path).convert('RGB')) + # result_tensor = transform2(Image.open(result_path).convert('RGB')) + # loss_fn_alex = lpips.LPIPS(net='alex') + # lpips_val = loss_fn_alex(gt_tensor, result_tensor).item() + + total_old = metrics['psnr'] + metrics['ssim'] + total_new = psnr_val + ssim_val + + if total_new > total_old: + agent.save(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) + latest_metrics = {'psnr': psnr_val, 'ssim': ssim_val} + else: + latest_metrics = metrics + + print(f"Image {name} metrics: PSNR {psnr_val}, SSIM {ssim_val}") + + return latest_metrics + +def evaluate_real(agent, name, metrics): + result_path = os.path.join(args.save_dir_path, 'derained_result', name) + result_img = skimage.io.imread(result_path) + brisque_val = BRISQUE(url=False).score(result_img) + + total_old = metrics['brisque'] + total_new = brisque_val + + if total_new < total_old: + agent.save(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) + latest_metrics = {'brisque': brisque_val} + else: + latest_metrics = metrics + + print(f"Image {name} metrics: BRISQUE {brisque_val}") + + return latest_metrics + def main(args): #_/_/_/ load dataset _/_/_/ mini_batch_loader = MiniBatchLoader( @@ -98,11 +149,14 @@ def main(args): agent = PixelWiseA3C_InnerState(model, optimizer, 5, args.gamma) agent.act_deterministically = True agent.model.to_gpu() + agent.load(os.path.join(args.checkpoint_dir_path, 'model_weight', name)) + # agent.load(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) print("===== Process for {} =====".format(name)) # for saving trajectories trajectory_dataset = TrajectoryDataset() # init Reward function r_net_brisque = Reward_Predictor(image_size=(args.pretrained_img_size, args.pretrained_img_size)).cuda() + # r_net_brisque.load_state_dict(torch.load(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'Rnet', 'rnet_brisque.pt'))) optimizer_rnet_brisque = optim.Adam(r_net_brisque.parameters(), lr=1e-5) print("----aug_rainy_imgs----") aug_rain_image_list = [] @@ -170,7 +224,13 @@ def main(args): rnet_brisque_loss.backward() optimizer_rnet_brisque.step() print("----start training agent----") - os.makedirs(os.path.join(args.save_dir_path, 'model_weight', name), exist_ok=True) + os.makedirs(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name), exist_ok=True) + metrics_path = os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'metrics.json') + if os.path.exists(metrics_path): + with open(metrics_path, 'r') as f: + metrics = json.load(f) + else: + metrics = {'psnr': 0, 'ssim': 0, 'lpips': 0, 'brisque': 1000} for episode in tqdm(range(1, args.max_episode+1)): # random crop args.pretrained_img_size x args.pretrained_img_size img_size = args.pretrained_img_size @@ -216,13 +276,18 @@ def main(args): sum_reward += np.mean(reward)*np.power(args.gamma,t) agent.stop_episode_and_train(current_state.tensor, reward, True) optimizer.alpha = args.lr*((1-episode/args.max_episode)**0.9) - agent.save(os.path.join(args.save_dir_path, 'model_weight', name)) - inference(agent, raw_x, mask, name) + if (episode + 1) % 10 == 0: + inference(agent, raw_x, mask, name) + metrics = evaluate_syn(agent, name, metrics) + # metrics = evaluate_real(agent, name, metrics) + with open(metrics_path, 'w') as f: + json.dump(metrics, f) + # save Rnet model - os.makedirs(os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet'), exist_ok=True) - torch.save(r_net_brisque.state_dict(), os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet', 'rnet_brisque.pt')) + os.makedirs(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'Rnet'), exist_ok=True) + torch.save(r_net_brisque.state_dict(), os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'Rnet', 'rnet_brisque.pt')) if __name__ == '__main__': parser = argparse.ArgumentParser(description='parameters for training') @@ -230,8 +295,10 @@ def main(args): parser.add_argument('--random_seed', type=int, default=1) # Directories parser.add_argument('--image_dir_path', type=str, default='dataset/') - parser.add_argument('--data_path', type=str, default='dataset/Rain100L/testing.txt') - parser.add_argument('--save_dir_path', type=str, default='./Results/Rain100L/test/SRL-Derain/') + parser.add_argument('--data_path', type=str, default='dataset/Rain12/testing.txt') + parser.add_argument('--gt_dir_path', type=str, default='dataset/Rain12/test/gt/') + parser.add_argument('--save_dir_path', type=str, default='./Results/Rain12/test/SRL-Derain/') + parser.add_argument('--checkpoint_dir_path', type=str, default='./Checkpoints/Rain12/SRL-Derain/') # config parser.add_argument('--gpu_id', type=int, default=0) parser.add_argument('--move_range', type=int, default=3) diff --git a/main_srl_derain_plus.py b/main_srl_derain_plus.py index ddb051c..b3a1446 100755 --- a/main_srl_derain_plus.py +++ b/main_srl_derain_plus.py @@ -1,4 +1,5 @@ import argparse +import json import chainer import numpy as np import os @@ -11,7 +12,10 @@ import torch from utils.trajectory import list_of_tuple_to_traj import cv2 +import skimage.io from utils.compute_Rbrisque import brisque_reward +from skimage.metrics import peak_signal_noise_ratio as psnr +from skimage.metrics import structural_similarity as ssim def overlapped_process(raw_x, mask, agent, current_state, patch_size, stride): _, _, h, w = raw_x.shape @@ -69,6 +73,50 @@ def inference(agent, raw_x, mask, name): def compute_diff(image1, image2): return np.mean(np.abs(image1 - image2)) +def evaluate_syn(agent, name, metrics): + gt_path = os.path.join(args.gt_dir_path, name) + result_path = os.path.join(args.save_dir_path, 'derained_result', name) + gt_img = cv2.imread(gt_path) + result_img = cv2.imread(result_path) + psnr_val = psnr(gt_img, result_img) + ssim_val = ssim(gt_img, result_img, channel_axis=2) + # transform2 = transforms.Compose([transforms.ToTensor()]) + # gt_tensor = transform2(Image.open(gt_path).convert('RGB')) + # result_tensor = transform2(Image.open(result_path).convert('RGB')) + # loss_fn_alex = lpips.LPIPS(net='alex') + # lpips_val = loss_fn_alex(gt_tensor, result_tensor).item() + + total_old = metrics['psnr'] + metrics['ssim'] #+ metrics['lpips'] + total_new = psnr_val + ssim_val #+ lpips_val + + if total_new > total_old: + agent.save(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) + latest_metrics = {'psnr': psnr_val, 'ssim': ssim_val} + else: + latest_metrics = metrics + + print(f"Image {name} metrics: PSNR {psnr_val}, SSIM {ssim_val}") + + return latest_metrics + +def evaluate_real(agent, name, metrics): + result_path = os.path.join(args.save_dir_path, 'derained_result', name) + result_img = skimage.io.imread(result_path) + brisque_val = BRISQUE(url=False).score(result_img) + + total_old = metrics['brisque'] + total_new = brisque_val + + if total_new < total_old: + agent.save(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) + latest_metrics = {'brisque': brisque_val} + else: + latest_metrics = metrics + + print(f"Image {name} metrics: BRISQUE {brisque_val}") + + return latest_metrics + def main(args): #_/_/_/ load dataset _/_/_/ mini_batch_loader = MiniBatchLoader( @@ -92,13 +140,21 @@ def main(args): agent = PixelWiseA3C_InnerState(model, optimizer, 5, args.gamma) agent.act_deterministically = True agent.model.to_gpu() + agent.load(os.path.join(args.checkpoint_dir_path, 'model_weight', name)) + # agent.load(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name)) print("===== Process for {} =====".format(name)) r_net_brisque = Reward_Predictor(image_size=(args.pretrained_img_size, args.pretrained_img_size)).cuda() rnet_brisque_state_dict = torch.load(os.path.join(args.rnet_weight_dir, 'rnet_brisque.pt')) r_net_brisque.load_state_dict(rnet_brisque_state_dict) print("----start training agent----") - os.makedirs(os.path.join(args.save_dir_path, 'model_weight', name), exist_ok=True) + os.makedirs(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name), exist_ok=True) + metrics_path = os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'metrics.json') + if os.path.exists(metrics_path): + with open(metrics_path, 'r') as f: + metrics = json.load(f) + else: + metrics = {'psnr': 0, 'ssim': 0, 'lpips': 0, 'brisque': 1000} for episode in tqdm(range(1, args.max_episode+1)): # random crop args.pretrained_img_size x args.pretrained_img_size img_size = args.pretrained_img_size @@ -141,13 +197,18 @@ def main(args): sum_reward += np.mean(reward)*np.power(args.gamma,t) agent.stop_episode_and_train(current_state.tensor, reward, True) optimizer.alpha = args.lr*((1-episode/args.max_episode)**0.9) - agent.save(os.path.join(args.save_dir_path, 'model_weight', name)) - - inference(agent, raw_x, mask, name) + + if (episode + 1) % 10 == 0: + inference(agent, raw_x, mask, name) + metrics = evaluate_syn(agent, name, metrics) + # metrics = evaluate_real(agent, name, metrics) + + with open(metrics_path, 'w') as f: + json.dump(metrics, f) # save Rnet model - os.makedirs(os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet'), exist_ok=True) - torch.save(r_net_brisque.state_dict(), os.path.join(args.save_dir_path, 'model_weight', name, 'Rnet', 'rnet_brisque.pt')) + os.makedirs(os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'Rnet'), exist_ok=True) + torch.save(r_net_brisque.state_dict(), os.path.join(args.checkpoint_dir_path, 'model_weight_best', name, 'Rnet', 'rnet_brisque.pt')) if __name__ == '__main__': parser = argparse.ArgumentParser(description='parameters for training') @@ -155,9 +216,11 @@ def main(args): parser.add_argument('--random_seed', type=int, default=1) # Directories parser.add_argument('--image_dir_path', type=str, default='dataset/') - parser.add_argument('--data_path', type=str, default='dataset/Rain100L/testing.txt') - parser.add_argument('--save_dir_path', type=str, default='./Results/Rain100L/test/SRL-Derain+/') - parser.add_argument('--rnet_weight_dir', type=str, default='./Results/Rain100L/test/Rnet+/model_weight/') + parser.add_argument('--data_path', type=str, default='dataset/Rain12/testing.txt') + parser.add_argument('--gt_dir_path', type=str, default='dataset/Rain12/test/gt/') + parser.add_argument('--save_dir_path', type=str, default='./Results/Rain12/test/SRL-Derain+/') + parser.add_argument('--rnet_weight_dir', type=str, default='./Checkpoints/Rain12/Rnet+/model_weight_best/') + parser.add_argument('--checkpoint_dir_path', type=str, default='./Checkpoints/Rain12/SRL-Derain+/') # config parser.add_argument('--gpu_id', type=int, default=0) parser.add_argument('--move_range', type=int, default=3) diff --git a/mini_batch_loader.py b/mini_batch_loader.py index 7175258..d950786 100755 --- a/mini_batch_loader.py +++ b/mini_batch_loader.py @@ -36,8 +36,8 @@ def load_training_data(self, index): def load_data(self, path_infos, index): in_channels = 3 path = path_infos[index] - mask_path = path.replace('input', 'RDP_w_sampling') - pseudo_gt_dir = path.replace('input', 'Pseudo_Reference_RDP_w_sampling')[:-4] + mask_path = path.replace('input', 'RDP') + pseudo_gt_dir = path.replace('input', 'Pseudo_Reference_RDP')[:-4] img = cv2.imread(path) if '.jpg' in path: @@ -78,8 +78,8 @@ def load_batch_data(self, path_infos, indices, img_size=None, augment=False): for i, index in enumerate(indices): path = path_infos[index] - mask_path = path.replace('input', 'RDP_w_sampling') - pseudo_gt_dir = path.replace('input', 'Pseudo_Reference_RDP_w_sampling')[:-4] + mask_path = path.replace('input', 'RDP') + pseudo_gt_dir = path.replace('input', 'Pseudo_Reference_RDP')[:-4] img = cv2.imread(path) if '.jpg' in path: @@ -128,7 +128,7 @@ def load_batch_data(self, path_infos, indices, img_size=None, augment=False): elif mini_batch_size == 1: for i, index in enumerate(indices): path = path_infos[index] - mask_path = path.replace('input', 'RDP_w_sampling') + mask_path = path.replace('input', 'RDP') img = cv2.imread(path) if '.jpg' in path: diff --git a/pretrain_rnet_plus.py b/pretrain_rnet_plus.py index e456ce2..502a957 100755 --- a/pretrain_rnet_plus.py +++ b/pretrain_rnet_plus.py @@ -106,15 +106,16 @@ def main(args): optimizer_rnet_brisque.step() # save Rnet model - os.makedirs(os.path.join(args.save_dir_path, 'model_weight'), exist_ok=True) - torch.save(r_net_brisque.state_dict(), os.path.join(args.save_dir_path, 'model_weight', 'rnet_brisque.pt')) + os.makedirs(os.path.join(args.checkpoint_dir_path, 'model_weight_best'), exist_ok=True) + torch.save(r_net_brisque.state_dict(), os.path.join(args.checkpoint_dir_path, 'model_weight_best', 'rnet_brisque.pt')) if __name__ == '__main__': parser = argparse.ArgumentParser(description='parameters for training') # Directories parser.add_argument('--image_dir_path', type=str, default='dataset/') - parser.add_argument('--data_path', type=str, default='dataset/Rain100L/testing.txt') - parser.add_argument('--save_dir_path', type=str, default='./Results/Rain100L/test/Rnet+/') + parser.add_argument('--data_path', type=str, default='dataset/Rain12/testing.txt') + parser.add_argument('--save_dir_path', type=str, default='./Results/Rain12/test/Rnet+/') + parser.add_argument('--checkpoint_dir_path', type=str, default='./Checkpoints/Rain12/Rnet+/') # config parser.add_argument('--batch_size', type=int, default=64) parser.add_argument('--N_pre', type=int, default=6000) diff --git a/stochastic_filling.py b/stochastic_filling.py index b0e5d4f..80f6de2 100755 --- a/stochastic_filling.py +++ b/stochastic_filling.py @@ -5,7 +5,7 @@ dataset_path = './dataset/Rain100L/test/' save_path = './dataset/Rain100L/test/' -target_path = 'Pseudo_Reference_RDP_w_sampling/' +target_path = 'Pseudo_Reference_RDP/' def make_folder(path): try: @@ -42,7 +42,7 @@ def compute_similarity(rainy_image, rdp_image, j, i): ### MAIN PROCESS GOES HERE ### rainy_path = os.path.join(dataset_path, "input") -rdp_path = os.path.join(dataset_path, "RDP_w_sampling") +rdp_path = os.path.join(dataset_path, "RDP") rainy_folder = os.listdir(rainy_path) print(rainy_folder) print(len(rainy_folder))