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

Method sync_target_model

trl/trl/trainer/callbacks.py:112–123  ·  view source on GitHub ↗
(model, target_model, alpha)

Source from the content-addressed store, hash-verified

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"]

Callers 1

on_step_endMethod · 0.95

Calls 2

_sync_target_modelMethod · 0.80
get_rankMethod · 0.45

Tested by

no test coverage detected