MCPcopy Create free account
hub / github.com/cli99/llm-analysis / get_weight_memory_per_layer

Method get_weight_memory_per_layer

llm_analysis/analysis.py:351–393  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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(

Callers 2

inferenceMethod · 0.95
trainingMethod · 0.95

Tested by

no test coverage detected