| 32 | |
| 33 | |
| 34 | class LogCallback(TrainerCallback, ExportableState): |
| 35 | def __init__(self, start_time: float = None, elapsed_time: float = None): |
| 36 | |
| 37 | self.start_time = time.time() if start_time is None else start_time |
| 38 | self.elapsed_time = 0 if elapsed_time is None else elapsed_time |
| 39 | self.last_time = self.start_time |
| 40 | |
| 41 | def on_train_begin( |
| 42 | self, |
| 43 | args: TrainingArguments, |
| 44 | state: TrainerState, |
| 45 | control: TrainerControl, |
| 46 | **kwargs |
| 47 | ): |
| 48 | r""" |
| 49 | Event called at the beginning of training. |
| 50 | """ |
| 51 | if state.is_local_process_zero: |
| 52 | if not args.resume_from_checkpoint: |
| 53 | self.start_time = time.time() |
| 54 | self.elapsed_time = 0 |
| 55 | else: |
| 56 | self.start_time = state.stateful_callbacks['LogCallback']['start_time'] |
| 57 | self.elapsed_time = state.stateful_callbacks['LogCallback']['elapsed_time'] |
| 58 | |
| 59 | if args.save_on_each_node: |
| 60 | if not state.is_local_process_zero: |
| 61 | return |
| 62 | else: |
| 63 | if not state.is_world_process_zero: |
| 64 | return |
| 65 | |
| 66 | self.last_time = time.time() |
| 67 | if os.path.exists(os.path.join(args.output_dir, LOG_FILE_NAME)) and args.overwrite_output_dir: |
| 68 | logger.warning("Previous log file in this folder will be deleted.") |
| 69 | os.remove(os.path.join(args.output_dir, LOG_FILE_NAME)) |
| 70 | |
| 71 | def on_log( |
| 72 | self, |
| 73 | args: TrainingArguments, |
| 74 | state: TrainerState, |
| 75 | control: TrainerControl, |
| 76 | logs, |
| 77 | **kwargs |
| 78 | ): |
| 79 | if args.save_on_each_node: |
| 80 | if not state.is_local_process_zero: |
| 81 | return |
| 82 | else: |
| 83 | if not state.is_world_process_zero: |
| 84 | return |
| 85 | |
| 86 | self.elapsed_time += time.time() - self.last_time |
| 87 | self.last_time = time.time() |
| 88 | if 'num_input_tokens_seen' in logs: |
| 89 | logs['num_tokens'] = logs.pop('num_input_tokens_seen') |
| 90 | state.log_history[-1].pop('num_input_tokens_seen') |
| 91 | throughput = logs['num_tokens'] / args.world_size / self.elapsed_time |