Skip to content
Merged
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
bdca8ee
[Model] Support EAGLE3 for Sarvam
mohit-sarvam Aug 20, 2026
1c5b972
[Model] Add EAGLE3 unit tests for Sarvam
mohit-sarvam Aug 20, 2026
c46e3bb
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Aug 20, 2026
5dcea91
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Aug 24, 2026
c3267ea
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Aug 25, 2026
6758be7
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Aug 25, 2026
db9772c
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Aug 25, 2026
d4d75bb
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Aug 27, 2026
b6f5efb
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 1, 2026
11bf00b
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 1, 2026
b142fec
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 2, 2026
51f652b
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 3, 2026
1c6a4ee
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 3, 2026
85fe491
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 4, 2026
761f0ef
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 4, 2026
6c64b61
Document Sarvam EAGLE3 pipeline scope and test ordinary PP
mohit-sarvam Sep 5, 2026
c112d91
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 5, 2026
176b762
Document Sarvam EAGLE3 pipeline requirement in README
mohit-sarvam Sep 5, 2026
6dfe671
Merge remote-tracking branch 'fork/mohit/sarvam-eagle3' into mohit/sa…
mohit-sarvam Sep 5, 2026
4e1c874
Revert Sarvam-specific README note
mohit-sarvam Sep 5, 2026
3c9d8b7
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 7, 2026
e336e1e
Remove Sarvam EAGLE3 unit test file
mohit-sarvam Sep 7, 2026
5d86fd4
Merge branch 'main' into mohit/sarvam-eagle3
mohit-sarvam Sep 7, 2026
59465fe
Merge branch 'main' into mohit/sarvam-eagle3
vadiklyutiy Sep 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 40 additions & 6 deletions vllm/model_executor/models/sarvam.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,13 @@
from vllm.sequence import IntermediateTensors

from .bailing_moe import BailingMoeForCausalLM
from .interfaces import MixtureOfExperts, SupportsLoRA, SupportsPP
from .interfaces import (
EagleModelMixin,
MixtureOfExperts,
SupportsEagle3,
SupportsLoRA,
SupportsPP,
)
from .utils import (
AutoWeightsLoader,
PPMissingLayer,
Expand Down Expand Up @@ -442,7 +448,9 @@ def forward(
return hidden_states, residual


class SarvamMLAModel(nn.Module):
class SarvamMLAModel(nn.Module, EagleModelMixin):
"""Sarvam MLA backbone with stage-local EAGLE3 auxiliary capture."""

hf_to_vllm_mapper = WeightsMapper(
orig_to_new_stacked={
# .experts.gate_up_proj must be handled by MoERunner.load_weights for EP
Expand Down Expand Up @@ -508,7 +516,17 @@ def forward(
positions: torch.Tensor,
intermediate_tensors: IntermediateTensors | None,
inputs_embeds: torch.Tensor | None = None,
) -> torch.Tensor | IntermediateTensors:
) -> torch.Tensor | IntermediateTensors | tuple[torch.Tensor, list[torch.Tensor]]:
"""Run this stage and optionally return its auxiliary hidden states.

Auxiliary captures are local to this stage, matching Qwen3 MoE.
Pipeline transport carries only hidden states and residual; EAGLE3
capture across pipeline stages is not supported by this model.

Returns:
Intermediate tensors on non-final stages. On the final stage,
normalized hidden states, paired with auxiliary states if captured.
"""
if get_pp_group().is_first_rank:
if inputs_embeds is not None:
hidden_states = inputs_embeds
Expand All @@ -521,12 +539,22 @@ def forward(
hidden_states = intermediate_tensors["hidden_states"]
residual = intermediate_tensors["residual"]

for layer in islice(self.layers, self.start_layer, self.end_layer):
aux_hidden_states = self._maybe_add_hidden_state(
[], self.start_layer, hidden_states, residual
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
for layer_idx, layer in enumerate(
islice(self.layers, self.start_layer, self.end_layer),
start=self.start_layer,
):
hidden_states, residual = layer(
hidden_states,
positions,
residual,
)
self._maybe_add_hidden_state(
aux_hidden_states, layer_idx + 1, hidden_states, residual
)

if not get_pp_group().is_last_rank:
return IntermediateTensors(
{"hidden_states": hidden_states, "residual": residual}
Expand All @@ -535,6 +563,9 @@ def forward(
hidden_states = self.norm(hidden_states)
else:
hidden_states, _ = self.norm(hidden_states, residual)

if len(aux_hidden_states) > 0:
return hidden_states, aux_hidden_states
return hidden_states

def load_weights(
Expand Down Expand Up @@ -591,7 +622,9 @@ def set_eplb_state(self, eplb_state) -> None:
moe.set_eplb_state(eplb_state)


class SarvamMLAForCausalLM(nn.Module, SupportsPP, SupportsLoRA, SarvamMixtureOfExperts):
class SarvamMLAForCausalLM(
nn.Module, SupportsPP, SupportsLoRA, SupportsEagle3, SarvamMixtureOfExperts
):
packed_modules_mapping = {
"q_proj": ["q_proj"],
"q_a_proj": ["q_a_proj"],
Expand Down Expand Up @@ -658,7 +691,8 @@ def forward(
positions: torch.Tensor,
intermediate_tensors: IntermediateTensors | None = None,
inputs_embeds: torch.Tensor | None = None,
) -> torch.Tensor | IntermediateTensors:
) -> torch.Tensor | IntermediateTensors | tuple[torch.Tensor, list[torch.Tensor]]:
"""Return backbone outputs, including auxiliary states when captured."""
return self.model(
input_ids=input_ids,
positions=positions,
Expand Down
Loading