[llama] Allow passing tp_rank and tp_size into llama mlp (#16837)

This commit is contained in:
Yinghai Lu
2026-01-10 10:32:05 +08:00
committed by GitHub
parent 1f0ea4f958
commit e91a717632
+6
View File
@@ -71,6 +71,8 @@ class LlamaMLP(nn.Module):
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
reduce_results: bool = True, reduce_results: bool = True,
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
) -> None: ) -> None:
super().__init__() super().__init__()
self.gate_up_proj = MergedColumnParallelLinear( self.gate_up_proj = MergedColumnParallelLinear(
@@ -79,6 +81,8 @@ class LlamaMLP(nn.Module):
bias=False, bias=False,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("gate_up_proj", prefix), prefix=add_prefix("gate_up_proj", prefix),
tp_rank=tp_rank,
tp_size=tp_size,
) )
self.down_proj = RowParallelLinear( self.down_proj = RowParallelLinear(
intermediate_size, intermediate_size,
@@ -87,6 +91,8 @@ class LlamaMLP(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("down_proj", prefix), prefix=add_prefix("down_proj", prefix),
reduce_results=reduce_results, reduce_results=reduce_results,
tp_rank=tp_rank,
tp_size=tp_size,
) )
if hidden_act != "silu": if hidden_act != "silu":
raise ValueError( raise ValueError(