Coverage for hopwise/trainer/hf_path_trainer.py: 82%

102 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2025 

2# @Author : Giacomo Medda 

3# @Email : giacomo.medda@unica.it 

4 

5# UPDATE: 

6# @Time : 2025 

7# @Author : Alessandro Soccol 

8# @Email : alessandro.soccol@unica.it 

9 

10from time import time 

11 

12import torch 

13from transformers import DataCollatorForLanguageModeling, IntervalStrategy, Trainer, TrainerCallback 

14 

15from hopwise.utils import ( 

16 dict2str, 

17 early_stopping, 

18 get_gpu_usage, 

19 progress_bar, 

20 set_color, 

21) 

22 

23 

24class HFPathTrainer(Trainer): 

25 """A HuggingFace Trainer that integrates with hopwise for training and evaluation.""" 

26 

27 def __init__(self, model, callbacks, train_data=None, args=None, tokenizer=None): 

28 tokenizer = tokenizer or train_data.dataset.tokenizer 

29 super().__init__( 

30 model=model, 

31 args=args, 

32 callbacks=None, 

33 train_dataset=train_data.dataset, 

34 eval_dataset="none", 

35 processing_class=tokenizer, 

36 data_collator=DataCollatorForLanguageModeling(tokenizer, mlm=False), 

37 ) 

38 

39 # Overwrite the callbacks to only use the hopwiseCallback 

40 self.callback_handler.callbacks = callbacks 

41 

42 def evaluate(self, **kwargs): 

43 self.control = self.callback_handler.on_evaluate(self.args, self.state, self.control, metrics=None) 

44 

45 return self.control.metrics 

46 

47 

48class hopwiseCallback(TrainerCallback): 

49 """It handles the training and evaluation communication with the hopwise and HuggingFace trainers.""" 

50 

51 def __init__( 

52 self, 

53 hopwise_trainer, 

54 train_data=None, 

55 valid_data=None, 

56 verbose=True, 

57 saved=True, 

58 show_progress=False, 

59 callback_fn=None, 

60 model=None, 

61 model_name=None, 

62 ): 

63 self.model = model 

64 self.model_name = model_name 

65 self.hopwise_trainer = hopwise_trainer 

66 self.train_data = train_data 

67 self.valid_data = valid_data 

68 self.verbose = verbose 

69 self.saved = saved 

70 self.show_progress = show_progress 

71 self.callback_fn = callback_fn 

72 

73 def on_train_begin(self, args, state, control, **kwargs): 

74 self.hopwise_trainer.eval_collector.train_data_collect(self.train_data) 

75 if self.hopwise_trainer.config["train_neg_sample_args"].get("dynamic", False): 

76 self.train_data.get_model(self.hopwise_trainer.model) 

77 self.valid_step = 0 

78 

79 def on_train_end(self, args, state, control, **kwargs): 

80 self.hopwise_trainer._add_hparam_to_tensorboard(self.hopwise_trainer.best_valid_score) 

81 return super().on_train_end(args, state, control, **kwargs) 

82 

83 def on_epoch_begin(self, args, state, control, **kwargs): 

84 self.training_start_time = time() 

85 

86 len_hf_dataloader = len(self.train_data.dataset) 

87 steps_in_epoch = len_hf_dataloader // self.hopwise_trainer.config["train_batch_size"] 

88 steps_in_epoch += int(len_hf_dataloader % self.hopwise_trainer.config["train_batch_size"] > 0) 

89 self.progress_bar = ( 

90 progress_bar( 

91 total=steps_in_epoch, 

92 ncols=100, 

93 desc=set_color(f"Train {int(state.epoch):>5}", "magenta", progress=True), 

94 ) 

95 if self.show_progress 

96 else range(steps_in_epoch) 

97 ) 

98 

99 def on_epoch_end(self, args, state, control, **kwargs): 

100 if self.show_progress: 

101 self.progress_bar.close() 

102 training_end_time = time() 

103 # Retrieve training loss and other information 

104 if state.log_history: 

105 epoch_idx = state.epoch 

106 train_loss = state.log_history[-1].get("loss") 

107 self.hopwise_trainer.train_loss_dict[epoch_idx] = train_loss 

108 train_loss_output = self.hopwise_trainer._generate_train_loss_output( 

109 epoch_idx, self.training_start_time, training_end_time, train_loss 

110 ) 

111 if self.verbose: 

112 self.hopwise_trainer.logger.info(train_loss_output) 

113 self.hopwise_trainer._add_train_loss_to_tensorboard(epoch_idx, train_loss) 

114 self.hopwise_trainer.wandblogger.log_metrics( 

115 {"epoch": epoch_idx, "train_loss": train_loss, "train_step": epoch_idx}, 

116 head="train", 

117 ) 

118 

119 if self.hopwise_trainer.eval_step <= 0 or not self.valid_data: 

120 if self.saved: 

121 control.should_save = True 

122 elif epoch_idx % self.hopwise_trainer.eval_step == 0: 

123 control.should_evaluate = True 

124 self.valid_start_time = time() 

125 else: 

126 control.should_evaluate = False 

127 

128 # update attentive-a 

129 if hasattr(self.model, "update_attentive_A"): 

130 with torch.no_grad(): 

131 self.model.update_attentive_A() 

132 return control 

133 

134 def on_step_end(self, args, state, control, **kwargs): 

135 control.should_log = True 

136 if self.show_progress: 

137 self.progress_bar.update(1) 

138 if self.hopwise_trainer.gpu_available and self.show_progress: 

139 gpu_usage = get_gpu_usage(self.hopwise_trainer.device) 

140 self.progress_bar.set_postfix_str(set_color("GPU RAM: " + gpu_usage, "yellow")) 

141 

142 if state.global_step >= state.max_steps: 

143 control.should_training_stop = True 

144 # Save the model at the end if we have a save strategy 

145 if args.save_strategy != IntervalStrategy.NO: 

146 control.should_save = True 

147 

148 return control 

149 

150 def on_evaluate(self, args, state, control, **kwargs): 

151 epoch_idx = state.epoch 

152 valid_score, valid_result = self.hopwise_trainer._valid_epoch( 

153 self.valid_data, show_progress=self.show_progress 

154 ) 

155 

156 ( 

157 self.hopwise_trainer.best_valid_score, 

158 self.hopwise_trainer.cur_step, 

159 stop_flag, 

160 update_flag, 

161 ) = early_stopping( 

162 valid_score, 

163 self.hopwise_trainer.best_valid_score, 

164 self.hopwise_trainer.cur_step, 

165 max_step=self.hopwise_trainer.stopping_step, 

166 bigger=self.hopwise_trainer.valid_metric_bigger, 

167 ) 

168 valid_end_time = time() 

169 valid_score_output = ( 

170 set_color("epoch %d evaluating", "green") 

171 + " [" 

172 + set_color("time", "blue") 

173 + ": %.2fs, " 

174 + set_color("valid_score", "blue") 

175 + ": %f]" 

176 ) % (epoch_idx, valid_end_time - self.valid_start_time, valid_score) 

177 valid_result_output = set_color("valid result", "blue") + ": \n" + dict2str(valid_result) 

178 if self.verbose: 

179 self.hopwise_trainer.logger.info(valid_score_output) 

180 self.hopwise_trainer.logger.info(valid_result_output) 

181 self.hopwise_trainer.tensorboard.add_scalar("Valid_score", valid_score, epoch_idx) 

182 self.hopwise_trainer.wandblogger.log_metrics({**valid_result, "valid_step": self.valid_step}, head="valid") 

183 

184 if not self.hopwise_trainer.valid_metric.startswith("eval_"): 

185 metric_to_check = f"eval_{self.hopwise_trainer.valid_metric}" 

186 control.metrics = {**valid_result, metric_to_check: valid_score} 

187 

188 if update_flag: 

189 if self.saved: 

190 control.should_save = True 

191 self.hopwise_trainer._save_checkpoint(epoch_idx, verbose=self.verbose) 

192 

193 self.hopwise_trainer.best_valid_result = valid_result 

194 

195 if self.callback_fn: 

196 self.callback_fn(epoch_idx, valid_score) 

197 

198 if stop_flag: 

199 stop_output = "Finished training, best eval result in epoch %d" % ( 

200 epoch_idx - self.hopwise_trainer.cur_step * self.hopwise_trainer.eval_step 

201 ) 

202 if self.verbose: 

203 self.hopwise_trainer.logger.info(stop_output) 

204 control.should_training_stop = True 

205 

206 self.valid_step += 1 

207 

208 return control