From e91a7176324e3a44fde9276db0af17c93997119b Mon Sep 17 00:00:00 2001 From: Yinghai Lu Date: Fri, 9 Jan 2026 18:32:05 -0800 Subject: [PATCH] [llama] Allow passing tp_rank and tp_size into llama mlp (#16837) --- python/sglang/srt/models/llama.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/python/sglang/srt/models/llama.py b/python/sglang/srt/models/llama.py index ffcf73d46..53761dae5 100644 --- a/python/sglang/srt/models/llama.py +++ b/python/sglang/srt/models/llama.py @@ -71,6 +71,8 @@ class LlamaMLP(nn.Module): quant_config: Optional[QuantizationConfig] = None, prefix: str = "", reduce_results: bool = True, + tp_rank: Optional[int] = None, + tp_size: Optional[int] = None, ) -> None: super().__init__() self.gate_up_proj = MergedColumnParallelLinear( @@ -79,6 +81,8 @@ class LlamaMLP(nn.Module): bias=False, quant_config=quant_config, prefix=add_prefix("gate_up_proj", prefix), + tp_rank=tp_rank, + tp_size=tp_size, ) self.down_proj = RowParallelLinear( intermediate_size, @@ -87,6 +91,8 @@ class LlamaMLP(nn.Module): quant_config=quant_config, prefix=add_prefix("down_proj", prefix), reduce_results=reduce_results, + tp_rank=tp_rank, + tp_size=tp_size, ) if hidden_act != "silu": raise ValueError(