@@ -51,6 +51,8 @@ def split(self, index_cross):
5151
5252 self .logger = utils .Logger (self .path , utils .path .get_filename (self .model_cfg ._path ))
5353 self .dataset .set_logger (self .logger )
54+ self .wandb = utils .Wandb (self .model_cfg , self .dataset_cfg , self .run_cfg )
55+ self .dataset .set_wandb (self .wandb )
5456 self .summary = utils .Summary (self .path , dataset = self .dataset )
5557 self .dataset .set_summary (self .summary )
5658
@@ -97,7 +99,8 @@ def split(self, index_cross):
9799 )
98100
99101 self .model = models .functional .common .find (self .model_cfg .name )(
100- self .model_cfg , self .dataset .cfg , self .run_cfg , logger = self .logger , summary = self .summary , main_msg = self .msg )
102+ self .model_cfg , self .dataset .cfg , self .run_cfg ,
103+ logger = self .logger , wandb = self .wandb , summary = self .summary , main_msg = self .msg )
101104 self .start_epoch = self .model .load (self .args .test_epoch )
102105
103106 if not self .run_cfg .distributed or (self .run_cfg .distributed and self .run_cfg .local_rank == 0 ):
@@ -121,7 +124,7 @@ def train(self, epoch):
121124 if self .run_cfg .distributed :
122125 self .train_loader .sampler .set_epoch (epoch )
123126 self .train_loader = self .model .train_loader_hook (self .train_loader )
124- batch_per_epoch , count_data = len (self .train_loader ), len (self .train_loader .dataset )
127+ batch_per_epoch , count_data = len (self .train_loader ), len (self .train_loader .sampler )
125128 log_step = 1 #max(int(np.power(10, np.floor(np.log10(batch_per_epoch / 10)))), 1) if batch_per_epoch > 0 else 1
126129 epoch_info = {'epoch' : epoch , 'batch_per_epoch' : batch_per_epoch , 'count_data' : count_data }
127130 epoch_info ['local_rank' ] = self .run_cfg .local_rank if self .run_cfg .distributed else 1
@@ -165,8 +168,10 @@ def train(self, epoch):
165168 if self .run_cfg .distributed :
166169 with utils .ddp .sequence ():
167170 self .logger .info_scalars ('Train Epoch: {} rank {}\t ' , (epoch , self .run_cfg .local_rank ), loss_all )
171+ self .wandb .info (loss_all )
168172 else :
169173 self .logger .info_scalars ('Train Epoch: {}\t ' , (epoch ,), loss_all )
174+ self .wandb .info (loss_all )
170175 if epoch % self .run_cfg .save_step == 0 :
171176 self .model .save (epoch )
172177
@@ -182,7 +187,7 @@ def test(self, epoch, data_loader=None, log_text='Test'):
182187 if self .run_cfg .distributed :
183188 data_loader .sampler .set_epoch (epoch )
184189 data_loader = self .model .test_loader_hook (data_loader )
185- batch_per_epoch , count_data = len (data_loader ), len (data_loader .dataset )
190+ batch_per_epoch , count_data = len (data_loader ), len (data_loader .sampler )
186191 log_step = max (int (np .power (10 , np .floor (np .log10 (batch_per_epoch / 10 )))), 1 )
187192 epoch_info = {'epoch' : epoch , 'batch_per_epoch' : batch_per_epoch , 'count_data' : count_data , 'log_text' : log_text }
188193 epoch_info ['local_rank' ] = self .run_cfg .local_rank if self .run_cfg .distributed else 1
@@ -269,6 +274,7 @@ def test(self, epoch, data_loader=None, log_text='Test'):
269274 msgs_dict [name ] = 100. * value / count
270275 accuracy .append (msgs_dict [name ])
271276 self .logger .info (log_msg .format (log_text , epoch , * accuracy ))
277+ self .wandb .info ({f'{ log_text } _accuracy' : accuracy })
272278 self .summary .add_scalars ('Accuracy' , msgs_dict , epoch )
273279
274280 # TODO do not support chain norm and renorm
@@ -277,7 +283,7 @@ def test(self, epoch, data_loader=None, log_text='Test'):
277283 # dataset.append(dataset[-1].super_dataset)
278284 # dataset_cfg.append(dataset[-1].cfg)
279285 for name , value in predict .items ():
280- predict [name ] = np . array ( value .cpu ())
286+ predict [name ] = value .cpu (). numpy ( )
281287 for d_cfg , d in zip (dataset_cfg , dataset ):
282288 if d_cfg .norm and d .need_norm (value .shape ):
283289 data_type , data_cfg = self ._get_type (d_cfg , name , test = False )
@@ -339,17 +345,22 @@ def run():
339345 main .split (index_cross )
340346 main .model .process_pre_hook ()
341347 if args .test_epoch is None :
348+ main .model .main_msg .update (dict (only_test = False ))
342349 if main .start_epoch == 0 :
343- main .val_test (main .start_epoch ) # TODO set flag only_test=False
350+ main .model .process_test_msg_pre_hook (main .model .main_msg )
351+ while main .model .main_msg ['test_flag' ]:
352+ main .val_test (main .start_epoch )
353+ main .model .process_test_msg_hook (main .model .main_msg )
344354 for epoch in range (main .start_epoch + 1 , main .run_cfg .epochs + 1 ):
345355 main .train (epoch )
346356 if epoch % main .run_cfg .save_step == 0 :
347- main .model .main_msg . update ( dict ( test_idx = 1 , test_flag = True , only_test = False ) )
357+ main .model .process_test_msg_pre_hook ( main . model . main_msg )
348358 while main .model .main_msg ['test_flag' ]:
349359 main .val_test (epoch )
350360 main .model .process_test_msg_hook (main .model .main_msg )
351361 else :
352- main .model .main_msg .update (dict (test_idx = 1 , test_flag = True , only_test = True ))
362+ main .model .main_msg .update (dict (only_test = True ))
363+ main .model .process_test_msg_pre_hook (main .model .main_msg )
353364 while main .model .main_msg ['test_flag' ]:
354365 main .val_test (main .start_epoch )
355366 main .model .process_test_msg_hook (main .model .main_msg )
0 commit comments