Returns a dict of name to shape mapping, only for the flattened fp32 weights saved by the optimizer. the names are exactly as in state_dict. The order is absolutely important, since the saved data is just flattened data with no identifiers and requires reconstruction in the s
(self)
| 5416 | return frozen_param_fragments |
| 5417 | |
| 5418 | def _get_zero_param_shapes(self): |
| 5419 | """Returns a dict of name to shape mapping, only for the flattened fp32 weights saved by the |
| 5420 | optimizer. the names are exactly as in state_dict. The order is absolutely important, since |
| 5421 | the saved data is just flattened data with no identifiers and requires reconstruction in the |
| 5422 | same order it was saved. |
| 5423 | We can't rely on self.module.named_parameters() to get the saved tensors, as some params |
| 5424 | will be missing and others unsaved and then it'd be impossible to reconstruct state_dict |
| 5425 | from the flattened weights. |
| 5426 | optimizer.bit16_groups seems to be the easiest to use as it's in all zeroX versions. |
| 5427 | """ |
| 5428 | param_group_shapes = [] |
| 5429 | cnt = 0 |
| 5430 | numel = 0 |
| 5431 | |
| 5432 | # zero2 started using a round_robin_bit16_groups which is a shuffled version of bit16_groups - |
| 5433 | # if we don't use it, we get parameters ordered incorrectly |
| 5434 | if hasattr(self.optimizer, "round_robin_bit16_groups"): |
| 5435 | bit16_groups = self.optimizer.round_robin_bit16_groups |
| 5436 | elif self.bfloat16_enabled() and hasattr(self.optimizer, "bf16_groups"): |
| 5437 | bit16_groups = self.optimizer.bf16_groups |
| 5438 | else: |
| 5439 | bit16_groups = self.optimizer.bit16_groups if self.zero_optimization_stage( |
| 5440 | ) == 2 else self.optimizer.fp16_groups |
| 5441 | |
| 5442 | for bit16_group in bit16_groups: |
| 5443 | param_shapes = OrderedDict() |
| 5444 | for param in bit16_group: |
| 5445 | cnt += 1 |
| 5446 | numel += param.ds_numel if hasattr(param, "ds_numel") else param.numel() |
| 5447 | shape = param.ds_shape if hasattr(param, "ds_shape") else param.shape |
| 5448 | if param not in self.param_names: |
| 5449 | raise ValueError("failed to find optimizer param in named params") |
| 5450 | name = self.param_names[param] |
| 5451 | param_shapes[name] = shape |
| 5452 | |
| 5453 | # uncomment to debug zero_to_fp32.py problems |
| 5454 | # if self.global_rank == 0: print(f"saving param {name} {shape} (numel={shape.numel()})") |
| 5455 | param_group_shapes.append(param_shapes) |
| 5456 | # if self.global_rank == 0: print(f"Total saved {numel} numels in {cnt} params") |
| 5457 | |
| 5458 | return param_group_shapes |
| 5459 | |
| 5460 | def _get_shared_params(self): |
| 5461 | """ |