| 32 | |
| 33 | |
| 34 | class Trainer: |
| 35 | def __init__(self, config): |
| 36 | self.mode = config.MODE |
| 37 | self.schedule = Schedule(config.SCHEDULE, config.TIMESTEPS) |
| 38 | self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") |
| 39 | in_channels = config.CHANNEL_X + config.CHANNEL_Y |
| 40 | out_channels = config.CHANNEL_Y |
| 41 | self.out_channels = out_channels |
| 42 | self.network = DocDiff( |
| 43 | input_channels=in_channels, |
| 44 | output_channels=out_channels, |
| 45 | n_channels=config.MODEL_CHANNELS, |
| 46 | ch_mults=config.CHANNEL_MULT, |
| 47 | n_blocks=config.NUM_RESBLOCKS |
| 48 | ).to(self.device) |
| 49 | self.diffusion = GaussianDiffusion(self.network.denoiser, config.TIMESTEPS, self.schedule).to(self.device) |
| 50 | self.test_img_save_path = config.TEST_IMG_SAVE_PATH |
| 51 | if not os.path.exists(self.test_img_save_path): |
| 52 | os.makedirs(self.test_img_save_path) |
| 53 | self.pretrained_path_init_predictor = config.PRETRAINED_PATH_INITIAL_PREDICTOR |
| 54 | self.pretrained_path_denoiser = config.PRETRAINED_PATH_DENOISER |
| 55 | self.continue_training = config.CONTINUE_TRAINING |
| 56 | self.continue_training_steps = 0 |
| 57 | self.path_train_gt = config.PATH_GT |
| 58 | self.path_train_img = config.PATH_IMG |
| 59 | self.iteration_max = config.ITERATION_MAX |
| 60 | self.LR = config.LR |
| 61 | self.cross_entropy = nn.BCELoss() |
| 62 | self.num_timesteps = config.TIMESTEPS |
| 63 | self.ema_every = config.EMA_EVERY |
| 64 | self.start_ema = config.START_EMA |
| 65 | self.save_model_every = config.SAVE_MODEL_EVERY |
| 66 | self.EMA_or_not = config.EMA |
| 67 | self.weight_save_path = config.WEIGHT_SAVE_PATH |
| 68 | self.TEST_INITIAL_PREDICTOR_WEIGHT_PATH = config.TEST_INITIAL_PREDICTOR_WEIGHT_PATH |
| 69 | self.TEST_DENOISER_WEIGHT_PATH = config.TEST_DENOISER_WEIGHT_PATH |
| 70 | self.DPM_SOLVER = config.DPM_SOLVER |
| 71 | self.DPM_STEP = config.DPM_STEP |
| 72 | self.test_path_img = config.TEST_PATH_IMG |
| 73 | self.test_path_gt = config.TEST_PATH_GT |
| 74 | self.beta_loss = config.BETA_LOSS |
| 75 | self.pre_ori = config.PRE_ORI |
| 76 | self.high_low_freq = config.HIGH_LOW_FREQ |
| 77 | self.image_size = config.IMAGE_SIZE |
| 78 | self.native_resolution = config.NATIVE_RESOLUTION |
| 79 | if self.mode == 1 and self.continue_training == 'True': |
| 80 | print('Continue Training') |
| 81 | self.network.init_predictor.load_state_dict(torch.load(self.pretrained_path_init_predictor)) |
| 82 | self.network.denoiser.load_state_dict(torch.load(self.pretrained_path_denoiser)) |
| 83 | self.continue_training_steps = config.CONTINUE_TRAINING_STEPS |
| 84 | from data.docdata import DocData |
| 85 | if self.mode == 1: |
| 86 | dataset_train = DocData(self.path_train_img, self.path_train_gt, config.IMAGE_SIZE, self.mode) |
| 87 | self.batch_size = config.BATCH_SIZE |
| 88 | self.dataloader_train = DataLoader(dataset_train, batch_size=self.batch_size, shuffle=True, drop_last=False, |
| 89 | num_workers=config.NUM_WORKERS) |
| 90 | else: |
| 91 | dataset_test = DocData(config.TEST_PATH_IMG, config.TEST_PATH_GT, config.IMAGE_SIZE, self.mode) |