Get the memory (in bytes) required to store the weights of a transformer layer, given the number of parameters in a transformer layer, the data type used for the weights, the tensor parallelism size, and the DeepSpeed ZeRO stage. WIth ZeRO Stage 3, the weights are sharded acr
(
self,
ds_zero: DSZeRO = DSZeRO.NONE,
return_breakdown: bool = False)
| 349 | self.get_num_params_last_layernorm()) |
| 350 | |
| 351 | def get_weight_memory_per_layer( |
| 352 | self, |
| 353 | ds_zero: DSZeRO = DSZeRO.NONE, |
| 354 | return_breakdown: bool = False) -> Union[float, tuple]: |
| 355 | """Get the memory (in bytes) required to store the weights of a transformer |
| 356 | layer, given the number of parameters in a transformer layer, the data type used |
| 357 | for the weights, the tensor parallelism size, and the DeepSpeed ZeRO stage. WIth |
| 358 | ZeRO Stage 3, the weights are sharded across data parallel groups. |
| 359 | |
| 360 | Args: |
| 361 | ds_zero (DSZeRO, optional): which DeepSpeed ZeRO stage to use. Defaults to DSZeRO.NONE (disabled). |
| 362 | |
| 363 | Returns: |
| 364 | Union[float, tuple]: the memory (in bytes) required to store the weights of a transformer layer, or a tuple of its breakdown |
| 365 | """ |
| 366 | if ds_zero == DSZeRO.STAGE_3: |
| 367 | sharded_dp_size = self.parallelism_config.dp_size |
| 368 | mlp_sharded_dp_size = self.parallelism_config.dp_size / self.parallelism_config.ep_size |
| 369 | else: |
| 370 | sharded_dp_size = 1 |
| 371 | mlp_sharded_dp_size = 1 |
| 372 | |
| 373 | weight_memory_attn_per_layer = self.get_num_params_per_layer_attn( |
| 374 | ) * self.dtype_config.weight_bits / BITS_PER_BYTE / self.parallelism_config.tp_size / sharded_dp_size |
| 375 | |
| 376 | weight_memory_mlp_per_layer = ( |
| 377 | self.get_num_params_per_layer_mlp() / |
| 378 | self.parallelism_config.ep_size + |
| 379 | self.get_num_params_per_layer_router() |
| 380 | ) * self.dtype_config.weight_bits / BITS_PER_BYTE / self.parallelism_config.tp_size / mlp_sharded_dp_size |
| 381 | |
| 382 | weight_memory_layernorm_per_layer = self.get_num_params_per_layer_layernorm( |
| 383 | ) * self.dtype_config.weight_bits / BITS_PER_BYTE / self.parallelism_config.tp_size / sharded_dp_size |
| 384 | |
| 385 | weight_memory_per_layer = weight_memory_attn_per_layer + weight_memory_mlp_per_layer + weight_memory_layernorm_per_layer |
| 386 | |
| 387 | logger.info( |
| 388 | f'weight_memory_attn_per_layer: {_num_to_string(weight_memory_attn_per_layer)}B, weight_memory_mlp_per_layer: {_num_to_string(weight_memory_mlp_per_layer)}B, weight_memory_layernorm_per_layer: {_num_to_string(weight_memory_layernorm_per_layer)}B' |
| 389 | ) |
| 390 | |
| 391 | if return_breakdown: |
| 392 | return weight_memory_per_layer, weight_memory_attn_per_layer, weight_memory_mlp_per_layer, weight_memory_layernorm_per_layer |
| 393 | return weight_memory_per_layer |
| 394 | |
| 395 | def get_weight_memory_last_layernorm(self, ds_zero: DSZeRO = DSZeRO.NONE): |
| 396 | weight_memory_last_layernorm = self.get_num_params_last_layernorm( |
no test coverage detected