core.ssm.utils#

Module Contents#

Functions#

_split_tensor_factory

Builds a factory that splits a given ShardedTensor into several independent chunks.

API#

core.ssm.utils._split_tensor_factory(
orig_sh_ten: megatron.core.dist_checkpointing.ShardedTensor,
split_sections: list[int],
split_names: list[str],
split_dim: int,
) megatron.core.dist_checkpointing.mapping.ShardedTensorFactory#

Builds a factory that splits a given ShardedTensor into several independent chunks.