MCPcopy Create free account
hub / github.com/AMAP-ML/EMF / SyncRefModelCallback

Class SyncRefModelCallback

trl/trl/trainer/callbacks.py:93–131  ·  view source on GitHub ↗

Callback to synchronize the model with a reference model.

Source from the content-addressed store, hash-verified

91
92
93class 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
134class RichProgressCallback(TrainerCallback):

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected