MCPcopy Create free account
hub / github.com/Royalvice/DocDiff / Trainer

Class Trainer

src/trainer.py:34–240  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

32
33
34class 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)

Callers 2

trainFunction · 0.90
testFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected