[feat] update bucketed weights from distributed (#13824)
Co-authored-by: Stefan He <hebiaobuaa@gmail.com>
This commit is contained in:
@@ -1153,7 +1153,14 @@ class ModelRunner:
|
||||
logger.error(message)
|
||||
return False, message
|
||||
|
||||
def update_weights_from_distributed(self, names, dtypes, shapes, group_name):
|
||||
def update_weights_from_distributed(
|
||||
self,
|
||||
names,
|
||||
dtypes,
|
||||
shapes,
|
||||
group_name,
|
||||
load_format: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Update specific parameter in the model weights online
|
||||
through `_model_update_group` process group.
|
||||
@@ -1169,6 +1176,10 @@ class ModelRunner:
|
||||
"Please call `init_weights_update_group` first."
|
||||
)
|
||||
|
||||
if load_format == "flattened_bucket":
|
||||
return self._update_bucketed_weights_from_distributed(
|
||||
names, dtypes, shapes, group_name
|
||||
)
|
||||
try:
|
||||
weights = []
|
||||
handles = []
|
||||
@@ -1201,6 +1212,37 @@ class ModelRunner:
|
||||
logger.error(error_msg)
|
||||
return False, error_msg
|
||||
|
||||
def _update_bucketed_weights_from_distributed(
|
||||
self, names, dtypes, shapes, group_name
|
||||
):
|
||||
try:
|
||||
named_tensors = []
|
||||
for name, dtype, shape in zip(names, dtypes, shapes):
|
||||
target_dtype = (
|
||||
dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
|
||||
)
|
||||
named_tensors.append(
|
||||
(name, torch.empty(shape, dtype=target_dtype, device=self.device))
|
||||
)
|
||||
bucket = FlattenedTensorBucket(named_tensors=named_tensors)
|
||||
flattened_tensor = bucket.get_flattened_tensor()
|
||||
torch.distributed.broadcast(
|
||||
flattened_tensor,
|
||||
src=0,
|
||||
group=self._model_update_group[group_name],
|
||||
)
|
||||
reconstructed_tensors = bucket.reconstruct_tensors()
|
||||
self.model.load_weights(reconstructed_tensors)
|
||||
return True, f"Succeeded to update parameter online."
|
||||
except Exception as e:
|
||||
error_msg = (
|
||||
f"Failed to update parameter online: {e}. "
|
||||
f"The full weights of the ModelRunner are partially updated. "
|
||||
f"Please discard the whole weights."
|
||||
)
|
||||
logger.error(error_msg)
|
||||
return False, error_msg
|
||||
|
||||
def update_weights_from_tensor(
|
||||
self,
|
||||
named_tensors: List[Tuple[str, Union[torch.Tensor, "LocalSerializedTensor"]]],
|
||||
|
||||
Reference in New Issue
Block a user