Callback to synchronize the model with a reference model.
| 91 | |
| 92 | |
| 93 | class SyncRefModelCallback(TrainerCallback): |
| 94 | """ |
| 95 | Callback to synchronize the model with a reference model. |
| 96 | """ |
| 97 | |
| 98 | def __init__( |
| 99 | self, |
| 100 | ref_model: Union[PreTrainedModel, torch.nn.Module], |
| 101 | accelerator: Optional[Accelerator], |
| 102 | ): |
| 103 | self.accelerator = accelerator |
| 104 | self.ref_model = ref_model |
| 105 | |
| 106 | @staticmethod |
| 107 | def _sync_target_model(model, target_model, alpha): |
| 108 | for target_param, copy_param in zip(target_model.parameters(), model.parameters()): |
| 109 | target_param.data.mul_(1.0 - alpha).add_(copy_param.data, alpha=alpha) |
| 110 | |
| 111 | @staticmethod |
| 112 | def sync_target_model(model, target_model, alpha): |
| 113 | deepspeed_plugin = AcceleratorState().deepspeed_plugin |
| 114 | if deepspeed_plugin is not None and deepspeed_plugin.zero_stage == 3: |
| 115 | import deepspeed |
| 116 | |
| 117 | with deepspeed.zero.GatheredParameters( |
| 118 | list(model.parameters()) + list(target_model.parameters()), modifier_rank=0 |
| 119 | ): |
| 120 | if deepspeed.comm.get_rank() == 0: |
| 121 | SyncRefModelCallback._sync_target_model(model, target_model, alpha) |
| 122 | else: |
| 123 | SyncRefModelCallback._sync_target_model(model, target_model, alpha) |
| 124 | |
| 125 | def on_step_end(self, args, state, control, **kwargs): |
| 126 | model: PreTrainedModel = kwargs["model"] |
| 127 | |
| 128 | if self.ref_model is not None and state.global_step % args.ref_model_sync_steps == 0: |
| 129 | if self.accelerator: |
| 130 | model = self.accelerator.unwrap_model(model) |
| 131 | self.sync_target_model(model, self.ref_model, args.ref_model_mixup_alpha) |
| 132 | |
| 133 | |
| 134 | class RichProgressCallback(TrainerCallback): |