class _Glm5NextMergedColumnParallelLinear(MergedColumnParallelLinear):
"""Merged projection with multiple replicated output shards.
Extends K3's ``_KimiGDNMergedColumnParallelLinear`` to support two
replicated shards (f_a, g_a) instead of one. Pre-multiplies each
replicated entry's output_size by tp_size so the per-rank shard
divides back to the full size, and forces tp_rank=0 during weight
loading for replicated shards.
"""
def __init__(
self,
input_size: int,
output_sizes: list[int],
replicated_shard_ids: tuple[int, ...],
tp_size: int,
**kwargs,
) -> None:
self.replicated_shard_ids = set(replicated_shard_ids)
output_sizes = output_sizes.copy()
for sid in self.replicated_shard_ids:
output_sizes[sid] *= tp_size
super().__init__(input_size, output_sizes, **kwargs)
def weight_loader(
self,
param: nn.Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: tuple[int, ...] | int | None = None,
) -> None:
tp_rank = self.tp_rank
param_tp_rank = getattr(param, "tp_rank", None)
if loaded_shard_id in self.replicated_shard_ids:
self.tp_rank = 0
if param_tp_rank is not None:
param.tp_rank = 0
try:
super().weight_loader(param, loaded_weight, loaded_shard_id)
finally:
self.tp_rank = tp_rank
if param_tp_rank is not None:
param.tp_rank = param_tp_rank
def weight_loader_v2(
self,
param: nn.Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: tuple[int, ...] | int | None = None,
) -> None:
tp_rank = self.tp_rank
param_tp_rank = getattr(param, "tp_rank", None)
if loaded_shard_id in self.replicated_shard_ids:
self.tp_rank = 0
if param_tp_rank is not None:
param.tp_rank = 0
try:
super().weight_loader_v2(param, loaded_weight, loaded_shard_id)
finally:
self.tp_rank = tp_rank
if param_tp_rank is not None:
param.tp_rank = param_tp_rank