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
« 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
5# UPDATE:
6# @Time : 2025
7# @Author : Alessandro Soccol
8# @Email : alessandro.soccol@unica.it
10from time import time
12import torch
13from transformers import DataCollatorForLanguageModeling, IntervalStrategy, Trainer, TrainerCallback
15from hopwise.utils import (
16 dict2str,
17 early_stopping,
18 get_gpu_usage,
19 progress_bar,
20 set_color,
21)
24class HFPathTrainer(Trainer):
25 """A HuggingFace Trainer that integrates with hopwise for training and evaluation."""
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 )
39 # Overwrite the callbacks to only use the hopwiseCallback
40 self.callback_handler.callbacks = callbacks
42 def evaluate(self, **kwargs):
43 self.control = self.callback_handler.on_evaluate(self.args, self.state, self.control, metrics=None)
45 return self.control.metrics
48class hopwiseCallback(TrainerCallback):
49 """It handles the training and evaluation communication with the hopwise and HuggingFace trainers."""
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
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
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)
83 def on_epoch_begin(self, args, state, control, **kwargs):
84 self.training_start_time = time()
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 )
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 )
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
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
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"))
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
148 return control
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 )
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")
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}
188 if update_flag:
189 if self.saved:
190 control.should_save = True
191 self.hopwise_trainer._save_checkpoint(epoch_idx, verbose=self.verbose)
193 self.hopwise_trainer.best_valid_result = valid_result
195 if self.callback_fn:
196 self.callback_fn(epoch_idx, valid_score)
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
206 self.valid_step += 1
208 return control