Co-authored-by: cao1zhg <114661107+cao1zhg@users.noreply.github.com> Co-authored-by: ispobock <ispobaoke@gmail.com> Co-authored-by: Binyao Jiang <byjiang1996@gmail.com> Co-authored-by: hebiao064 <hebiaobuaa@gmail.com> Co-authored-by: Lifu Huang <lifu.hlf@gmail.com> Co-authored-by: qingquansong <ustcsqq@gmail.com> Co-authored-by: Yaoyao Ding <dingyaoyao.cs@gmail.com> Co-authored-by: Ke Bao <ISPObaoke@163.com> Co-authored-by: Minglei Zhu <mingleizhu1122@gmail.com>
65 lines
2.4 KiB
Python
65 lines
2.4 KiB
Python
from typing import Callable, List, Tuple
|
|
|
|
import torch
|
|
|
|
LoaderFunction = Callable[[torch.Tensor, torch.Tensor], None]
|
|
|
|
|
|
def mamba_v2_sharded_weight_loader(
|
|
shard_spec: List[Tuple[int, int, float]],
|
|
tp_size: int,
|
|
tp_rank: int,
|
|
) -> LoaderFunction:
|
|
"""Create a weight loader for mamba v2. This ensures that the projections
|
|
are correctly sharded so that they can be split into x, B, C. It also
|
|
ensures the the all the groups corresponding to a head shard is placed
|
|
together with it.
|
|
"""
|
|
|
|
def loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
|
|
|
|
# - track boundary of (sharded) param, and loaded_weight, respectively
|
|
boundary, loaded_boundary = 0, 0
|
|
|
|
# - iterate over the shard specs
|
|
for full_dim, extra, duplicate_groups in shard_spec:
|
|
# - full dim is the model dim (before TP).
|
|
# - extra > 0, means there is expected overall increase
|
|
# of dimensions. This is so because of replication.
|
|
# - ratio is used map the tp_rank to the actual shard
|
|
# rank. This is useful when there is replication of
|
|
# groups to accompany head shards.
|
|
|
|
# - size of the loaded shard
|
|
shard_size = full_dim // tp_size
|
|
|
|
# - compute the rank into the loaded shard.
|
|
# - if there is replication, different TP shards will
|
|
# take from the same rank.
|
|
# NOTE: currently we only support duplication
|
|
# in the case where num_groups == 1
|
|
rank = 0 if duplicate_groups else tp_rank
|
|
|
|
# - leftmost boundary index into loaded weight.
|
|
loaded_skip = rank * shard_size
|
|
loaded_start_idx = loaded_boundary + loaded_skip
|
|
|
|
# - take these many dims from the loaded weight.
|
|
take = min(shard_size, full_dim - extra - loaded_skip)
|
|
|
|
# - always shard on dim 0
|
|
# - the ignore is for a mundane mypy error as it does not
|
|
# seem to handle slices well.
|
|
# https://github.com/python/mypy/issues/2410
|
|
param.data[
|
|
boundary : (boundary + take), ... # type: ignore[misc]
|
|
] = loaded_weight[
|
|
loaded_start_idx : (loaded_start_idx + take) # type: ignore[misc]
|
|
] # type: ignore[misc]
|
|
|
|
# move indexing boundaries
|
|
boundary += shard_size
|
|
loaded_boundary += full_dim - extra
|
|
|
|
return loader
|