Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
train_sirs_4000.py119 linesDownload Raw Back to DSRNet
1import os
2from os.path import join
3
4import torch.backends.cuda
5import torch.backends.cudnn as cudnn
6
7import data.sirs_dataset as datasets
8import util.util as util
9from data.image_folder import read_fns
10from engine import Engine
11from options.net_options.train_options import TrainOptions
12from tools import mutils
13
14opt = TrainOptions().parse()
15print(opt)
16cudnn.benchmark = False
17torch.backends.cuda.matmul.allow_tf32 = True
18
19opt.display_freq = 10
20
21if opt.debug:
22    opt.display_id = 1
23    opt.display_freq = 1
24    opt.print_freq = 20
25    opt.nEpochs = 40
26    opt.max_dataset_size = 9999
27    opt.no_log = False
28    opt.nThreads = 0
29    opt.decay_iter = 0
30    opt.serial_batches = True
31    opt.no_flip = True
32
33# modify the following code to
34# datadir = os.path.join(os.path.expanduser('~'), 'datasets/reflection-removal')
35datadir = os.path.join(opt.base_dir)
36
37datadir_syn = join(datadir, 'train/VOCdevkit/VOC2012/PNGImages')
38datadir_real = join(datadir, 'train/real')
39datadir_nature = join(datadir, 'train/nature')
40
41train_dataset = datasets.DSRDataset(
42    datadir_syn, read_fns('data/VOC2012_224_train_png.txt'), size=opt.max_dataset_size, enable_transforms=True)
43
44train_dataset_real = datasets.DSRTestDataset(datadir_real, enable_transforms=True, if_align=opt.if_align)
45train_dataset_nature = datasets.DSRTestDataset(datadir_nature, enable_transforms=True, if_align=opt.if_align)
46
47train_dataset_fusion = datasets.FusionDataset([train_dataset,
48                                               train_dataset_real,
49                                               train_dataset_nature], [0.6, 0.2, 0.2],
50                                               size=opt.num_train if opt.num_train > 0 else 4000)
51
52train_dataloader_fusion = datasets.DataLoader(
53    train_dataset_fusion, batch_size=opt.batchSize, shuffle=not opt.serial_batches,
54    num_workers=opt.nThreads, pin_memory=True)
55
56eval_dataset_real = datasets.DSRTestDataset(join(datadir, f'test/real20_{opt.real20_size}'),
57                                            fns=read_fns('data/real_test.txt'), if_align=opt.if_align)
58eval_dataset_solidobject = datasets.DSRTestDataset(join(datadir, 'test/SIR2/SolidObjectDataset'),
59                                                   if_align=opt.if_align)
60eval_dataset_postcard = datasets.DSRTestDataset(join(datadir, 'test/SIR2/PostcardDataset'), if_align=opt.if_align)
61eval_dataset_wild = datasets.DSRTestDataset(join(datadir, 'test/SIR2/WildSceneDataset'), if_align=opt.if_align)
62
63eval_dataloader_real = datasets.DataLoader(
64    eval_dataset_real, batch_size=1, shuffle=False,
65    num_workers=opt.nThreads, pin_memory=True)
66
67eval_dataloader_solidobject = datasets.DataLoader(
68    eval_dataset_solidobject, batch_size=1, shuffle=False,
69    num_workers=opt.nThreads, pin_memory=True)
70eval_dataloader_postcard = datasets.DataLoader(
71    eval_dataset_postcard, batch_size=1, shuffle=False,
72    num_workers=opt.nThreads, pin_memory=True)
73
74eval_dataloader_wild = datasets.DataLoader(
75    eval_dataset_wild, batch_size=1, shuffle=False,
76    num_workers=opt.nThreads, pin_memory=True)
77
78"""Main Loop"""
79engine = Engine(opt)
80result_dir = os.path.join(f'./checkpoints/{opt.name}/results',
81                          mutils.get_formatted_time())
82
83
84def set_learning_rate(lr):
85    for optimizer in engine.model.optimizers:
86        print('[i] set learning rate to {}'.format(lr))
87        util.set_opt_param(optimizer, 'lr', lr)
88
89
90if opt.resume or opt.debug_eval:
91    save_dir = os.path.join(result_dir, '%03d' % engine.epoch)
92    os.makedirs(save_dir, exist_ok=True)
93    engine.save_model()
94
95    engine.eval(eval_dataloader_real, dataset_name='testdata_real20', savedir=save_dir, suffix='real20')
96    engine.eval(eval_dataloader_solidobject, dataset_name='testdata_solidobject', savedir=save_dir,
97                suffix='solidobject')
98    engine.eval(eval_dataloader_postcard, dataset_name='testdata_postcard', savedir=save_dir, suffix='postcard')
99    engine.eval(eval_dataloader_wild, dataset_name='testdata_wild', savedir=save_dir, suffix='wild')
100
101# define training strategy
102engine.model.opt.lambda_gan = 0
103# engine.model.opt.lambda_gan = 0.01
104set_learning_rate(opt.lr)
105
106while engine.epoch < 120:
107    print('random_seed: ', opt.seed)
108    engine.train(train_dataloader_fusion)
109
110    if engine.epoch % 1 == 0:
111        save_dir = os.path.join(result_dir, '%03d' % engine.epoch)
112        os.makedirs(save_dir, exist_ok=True)
113
114        
115        engine.eval(eval_dataloader_real, dataset_name='testdata_real20', savedir=save_dir, suffix='real20')
116        engine.eval(eval_dataloader_solidobject, dataset_name='testdata_solidobject', savedir=save_dir, suffix='solidobject')
117        engine.eval(eval_dataloader_postcard, dataset_name='testdata_postcard', savedir=save_dir, suffix='postcard')
118        engine.eval(eval_dataloader_wild, dataset_name='testdata_wild', savedir=save_dir, suffix='wild')
119