From bdca8ee2d9c587b14c82abd4cf31cfe5b56318b4 Mon Sep 17 00:00:00 2001 From: mohit-sarvam Date: Thu, 20 Aug 2026 05:25:07 +0000 Subject: [PATCH 1/6] [Model] Support EAGLE3 for Sarvam Enable EAGLE-3 speculative decoding for the Sarvam MLA architecture by adopting the consolidated SupportsEagle interface from #36063. SarvamMLAModel mixes in EagleModelMixin and collects auxiliary hidden states, using absolute layer indices so the capture points stay correct under pipeline parallelism. SarvamMLAForCausalLM declares SupportsEagle3 and relies on the protocol defaults for set_aux_hidden_state_layers and get_eagle3_default_aux_hidden_state_layers. Signed-off-by: mohit-sarvam Co-authored-by: Cursor --- vllm/model_executor/models/sarvam.py | 33 +++++++++++++++++++++++----- 1 file changed, 27 insertions(+), 6 deletions(-) diff --git a/vllm/model_executor/models/sarvam.py b/vllm/model_executor/models/sarvam.py index 04590a2a913a..9d63eded78d9 100644 --- a/vllm/model_executor/models/sarvam.py +++ b/vllm/model_executor/models/sarvam.py @@ -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, @@ -442,7 +448,7 @@ def forward( return hidden_states, residual -class SarvamMLAModel(nn.Module): +class SarvamMLAModel(nn.Module, EagleModelMixin): hf_to_vllm_mapper = WeightsMapper( orig_to_new_stacked={ # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP @@ -508,7 +514,7 @@ 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]]: if get_pp_group().is_first_rank: if inputs_embeds is not None: hidden_states = inputs_embeds @@ -521,12 +527,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 + ) + 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} @@ -535,6 +551,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( @@ -591,7 +610,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"], @@ -658,7 +679,7 @@ 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 self.model( input_ids=input_ids, positions=positions, From 1c5b972867dec399c8cca6931ef4a329f0b38fb8 Mon Sep 17 00:00:00 2001 From: mohit-sarvam Date: Thu, 20 Aug 2026 05:39:48 +0000 Subject: [PATCH 2/6] [Model] Add EAGLE3 unit tests for Sarvam Cover the four things that can regress silently: the SupportsEagle3 declaration, the protocol hooks reaching the inner SarvamMLAModel, the captured tensors being complete layer outputs, and forward still returning a bare tensor when no aux layers are configured. The tests build SarvamMLAModel via object.__new__ so they need neither weights nor an initialized distributed environment. Signed-off-by: mohit-sarvam Co-authored-by: Cursor --- tests/model_executor/test_sarvam_eagle3.py | 152 +++++++++++++++++++++ 1 file changed, 152 insertions(+) create mode 100644 tests/model_executor/test_sarvam_eagle3.py diff --git a/tests/model_executor/test_sarvam_eagle3.py b/tests/model_executor/test_sarvam_eagle3.py new file mode 100644 index 000000000000..fa1fbdbdd310 --- /dev/null +++ b/tests/model_executor/test_sarvam_eagle3.py @@ -0,0 +1,152 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace +from unittest.mock import Mock + +import torch + +from vllm.model_executor.models import sarvam as sarvam_mod +from vllm.model_executor.models.interfaces import supports_eagle3 +from vllm.model_executor.models.sarvam import SarvamMLAForCausalLM, SarvamMLAModel +from vllm.sequence import IntermediateTensors + + +def _make_model( + *, + start_layer: int = 0, + end_layer: int = 1, + layers: list | None = None, + aux_hidden_state_layers: tuple[int, ...] = (), +) -> SarvamMLAModel: + """Build a SarvamMLAModel without running __init__. + + Constructing the real module needs a full VllmConfig plus an initialized + distributed environment, neither of which the aux hidden state plumbing + depends on. + """ + model = object.__new__(SarvamMLAModel) + object.__setattr__(model, "start_layer", start_layer) + object.__setattr__(model, "end_layer", end_layer) + object.__setattr__(model, "layers", layers if layers is not None else []) + object.__setattr__(model, "aux_hidden_state_layers", aux_hidden_state_layers) + object.__setattr__(model, "embedding_dropout", lambda x: x) + # RMSNorm in fused-residual mode returns (norm(h + r), h + r); the scaling + # is irrelevant here, so pass the sum through unchanged. + object.__setattr__( + model, "norm", lambda h, r=None: (h + r, h + r) if r is not None else h + ) + return model + + +def _layer(hidden_states: torch.Tensor, residual: torch.Tensor) -> Mock: + return Mock(return_value=(hidden_states, residual)) + + +def _patch_pp_group(monkeypatch, *, is_first_rank=True, is_last_rank=True) -> None: + monkeypatch.setattr( + sarvam_mod, + "get_pp_group", + lambda: SimpleNamespace(is_first_rank=is_first_rank, is_last_rank=is_last_rank), + ) + + +def test_sarvam_mla_advertises_eagle3_support(): + assert supports_eagle3(SarvamMLAForCausalLM) + + +def test_sarvam_mla_configures_aux_hidden_state_layers(): + # SarvamMLAForCausalLM inherits both hooks from the SupportsEagle3 + # protocol, so this pins that the protocol defaults actually reach the + # inner SarvamMLAModel rather than silently doing nothing. + target = object.__new__(SarvamMLAForCausalLM) + torch.nn.Module.__init__(target) + model = _make_model(layers=[None] * 12) + object.__setattr__(target, "model", model) + + target.set_aux_hidden_state_layers((2, 6, 9)) + + assert model.aux_hidden_state_layers == (2, 6, 9) + assert target.get_eagle3_default_aux_hidden_state_layers() == (2, 6, 9) + + +def test_sarvam_mla_forward_captures_aux_hidden_states(monkeypatch): + inputs_embeds = torch.tensor([[1.0, 2.0]]) + layer_hidden_states = torch.tensor([[3.0, 4.0]]) + layer_residual = torch.tensor([[5.0, 6.0]]) + model = _make_model( + layers=[_layer(layer_hidden_states, layer_residual)], + aux_hidden_state_layers=(0, 1), + ) + _patch_pp_group(monkeypatch) + + output, aux_hidden_states = model.forward( + input_ids=None, + positions=torch.tensor([0]), + intermediate_tensors=None, + inputs_embeds=inputs_embeds, + ) + + # Index 0 is the embedding output; index 1 is the output of layer 0, which + # is only complete once the pending residual is added back. + expected_layer_output = layer_hidden_states + layer_residual + torch.testing.assert_close(output, expected_layer_output) + assert len(aux_hidden_states) == 2 + torch.testing.assert_close(aux_hidden_states[0], inputs_embeds) + torch.testing.assert_close(aux_hidden_states[1], expected_layer_output) + + +def test_sarvam_mla_forward_returns_bare_tensor_without_eagle3(monkeypatch): + # Without a drafter the runner unpacks a single tensor, so returning a + # tuple unconditionally would break every non-speculative request. + model = _make_model( + layers=[_layer(torch.tensor([[3.0, 4.0]]), torch.tensor([[5.0, 6.0]]))], + ) + _patch_pp_group(monkeypatch) + + output = model.forward( + input_ids=None, + positions=torch.tensor([0]), + intermediate_tensors=None, + inputs_embeds=torch.tensor([[1.0, 2.0]]), + ) + + assert isinstance(output, torch.Tensor) + + +def test_sarvam_mla_forward_uses_absolute_layer_indices(monkeypatch): + # On a pipeline stage that does not start at layer 0, the capture points + # are numbered by absolute layer index. With relative numbering the + # requested layers would silently resolve to different tensors. + stage_hidden_states = torch.tensor([[1.0, 2.0]]) + stage_residual = torch.tensor([[0.5, 0.5]]) + second_hidden_states = torch.tensor([[3.0, 4.0]]) + second_residual = torch.tensor([[5.0, 6.0]]) + model = _make_model( + start_layer=1, + end_layer=3, + layers=[ + None, + _layer(torch.tensor([[7.0, 8.0]]), torch.tensor([[9.0, 10.0]])), + _layer(second_hidden_states, second_residual), + ], + aux_hidden_state_layers=(1, 3), + ) + _patch_pp_group(monkeypatch, is_first_rank=False) + + _, aux_hidden_states = model.forward( + input_ids=None, + positions=torch.tensor([0]), + intermediate_tensors=IntermediateTensors( + {"hidden_states": stage_hidden_states, "residual": stage_residual} + ), + inputs_embeds=None, + ) + + assert len(aux_hidden_states) == 2 + torch.testing.assert_close( + aux_hidden_states[0], stage_hidden_states + stage_residual + ) + torch.testing.assert_close( + aux_hidden_states[1], second_hidden_states + second_residual + ) From 6c64b610b67ed3c2168ab7ee9ef7294db44baa76 Mon Sep 17 00:00:00 2001 From: mohit-sarvam Date: Sat, 5 Sep 2026 08:11:10 +0000 Subject: [PATCH 3/6] Document Sarvam EAGLE3 pipeline scope and test ordinary PP Co-authored-by: Codex Signed-off-by: mohit-sarvam --- tests/model_executor/test_sarvam_eagle3.py | 43 ++++++++++++++++++++++ vllm/model_executor/models/sarvam.py | 13 +++++++ 2 files changed, 56 insertions(+) diff --git a/tests/model_executor/test_sarvam_eagle3.py b/tests/model_executor/test_sarvam_eagle3.py index fa1fbdbdd310..0afe8f4962ac 100644 --- a/tests/model_executor/test_sarvam_eagle3.py +++ b/tests/model_executor/test_sarvam_eagle3.py @@ -40,10 +40,12 @@ def _make_model( def _layer(hidden_states: torch.Tensor, residual: torch.Tensor) -> Mock: + """Return a decoder stub with fixed hidden states and residual.""" return Mock(return_value=(hidden_states, residual)) def _patch_pp_group(monkeypatch, *, is_first_rank=True, is_last_rank=True) -> None: + """Select the pipeline stage simulated by the forward call.""" monkeypatch.setattr( sarvam_mod, "get_pp_group", @@ -52,10 +54,12 @@ def _patch_pp_group(monkeypatch, *, is_first_rank=True, is_last_rank=True) -> No def test_sarvam_mla_advertises_eagle3_support(): + """Allow the runner to detect the EAGLE3 interface.""" assert supports_eagle3(SarvamMLAForCausalLM) def test_sarvam_mla_configures_aux_hidden_state_layers(): + """Propagate auxiliary capture configuration to the inner model.""" # SarvamMLAForCausalLM inherits both hooks from the SupportsEagle3 # protocol, so this pins that the protocol defaults actually reach the # inner SarvamMLAModel rather than silently doing nothing. @@ -71,6 +75,7 @@ def test_sarvam_mla_configures_aux_hidden_state_layers(): def test_sarvam_mla_forward_captures_aux_hidden_states(monkeypatch): + """Capture embeddings and complete layer outputs for the drafter.""" inputs_embeds = torch.tensor([[1.0, 2.0]]) layer_hidden_states = torch.tensor([[3.0, 4.0]]) layer_residual = torch.tensor([[5.0, 6.0]]) @@ -97,6 +102,7 @@ def test_sarvam_mla_forward_captures_aux_hidden_states(monkeypatch): def test_sarvam_mla_forward_returns_bare_tensor_without_eagle3(monkeypatch): + """Preserve the tensor return contract when auxiliary capture is disabled.""" # Without a drafter the runner unpacks a single tensor, so returning a # tuple unconditionally would break every non-speculative request. model = _make_model( @@ -115,6 +121,10 @@ def test_sarvam_mla_forward_returns_bare_tensor_without_eagle3(monkeypatch): def test_sarvam_mla_forward_uses_absolute_layer_indices(monkeypatch): + """Use global capture indices within one simulated pipeline stage. + + This checks local indexing, not cross-stage EAGLE3 state transport. + """ # On a pipeline stage that does not start at layer 0, the capture points # are numbered by absolute layer index. With relative numbering the # requested layers would silently resolve to different tensors. @@ -150,3 +160,36 @@ def test_sarvam_mla_forward_uses_absolute_layer_indices(monkeypatch): torch.testing.assert_close( aux_hidden_states[1], second_hidden_states + second_residual ) + + +def test_sarvam_mla_pipeline_forward_without_eagle3(monkeypatch): + """Preserve hidden states and residual across two stages without EAGLE3.""" + inputs_embeds = torch.tensor([[1.0, 2.0]]) + positions = torch.tensor([0]) + hidden_states = torch.tensor([[3.0, 4.0]]) + residual = torch.tensor([[5.0, 6.0]]) + first = _make_model(layers=[_layer(hidden_states, residual)]) + final_layer = Mock(side_effect=lambda h, p, r: (h * 2, r + 1)) + last = _make_model(start_layer=1, end_layer=2, layers=[None, final_layer]) + _patch_pp_group(monkeypatch, is_last_rank=False) + + intermediate = first.forward( + input_ids=None, + positions=positions, + intermediate_tensors=None, + inputs_embeds=inputs_embeds, + ) + + assert isinstance(intermediate, IntermediateTensors) + torch.testing.assert_close(intermediate["hidden_states"], hidden_states) + torch.testing.assert_close(intermediate["residual"], residual) + _patch_pp_group(monkeypatch, is_first_rank=False) + + output = last.forward( + input_ids=None, + positions=positions, + intermediate_tensors=intermediate, + ) + + assert isinstance(output, torch.Tensor) + torch.testing.assert_close(output, hidden_states * 2 + residual + 1) diff --git a/vllm/model_executor/models/sarvam.py b/vllm/model_executor/models/sarvam.py index eefc8133b5d0..3ce94b3f3297 100644 --- a/vllm/model_executor/models/sarvam.py +++ b/vllm/model_executor/models/sarvam.py @@ -449,6 +449,8 @@ def forward( 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 @@ -515,6 +517,16 @@ def forward( intermediate_tensors: IntermediateTensors | None, inputs_embeds: torch.Tensor | None = None, ) -> 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 @@ -680,6 +692,7 @@ def forward( intermediate_tensors: IntermediateTensors | None = None, inputs_embeds: torch.Tensor | None = None, ) -> 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, From 176b76216e9ce29e4ca57b64745781d6bd044b78 Mon Sep 17 00:00:00 2001 From: mohit-sarvam Date: Sat, 5 Sep 2026 08:20:27 +0000 Subject: [PATCH 4/6] Document Sarvam EAGLE3 pipeline requirement in README Co-authored-by: Codex Signed-off-by: mohit-sarvam --- README.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/README.md b/README.md index ca31a8cce94d..bfd4ed979480 100644 --- a/README.md +++ b/README.md @@ -61,6 +61,8 @@ vLLM seamlessly supports 200+ model architectures on Hugging Face, including: Find the full list of supported models [here](https://docs.vllm.ai/en/latest/models/supported_models.html). +Sarvam MLA (`SarvamMLAForCausalLM`) supports [EAGLE3 speculative decoding](docs/features/speculative_decoding/eagle.md) with a compatible draft checkpoint and a single pipeline stage (`--pipeline-parallel-size 1`). Auxiliary hidden states are not transported across pipeline stages. Ordinary pipeline parallelism remains supported when EAGLE3 is disabled. + ## Getting Started Install vLLM with [`uv`](https://docs.astral.sh/uv/) (recommended) or `pip`: From 4e1c874a3594b512c633842cfa6dde3c4c8550b7 Mon Sep 17 00:00:00 2001 From: mohit-sarvam Date: Sat, 5 Sep 2026 08:28:24 +0000 Subject: [PATCH 5/6] Revert Sarvam-specific README note Keep the EAGLE3 pipeline scope in the PR description and model docstrings. Co-authored-by: Codex Signed-off-by: mohit-sarvam --- README.md | 2 -- 1 file changed, 2 deletions(-) diff --git a/README.md b/README.md index bfd4ed979480..ca31a8cce94d 100644 --- a/README.md +++ b/README.md @@ -61,8 +61,6 @@ vLLM seamlessly supports 200+ model architectures on Hugging Face, including: Find the full list of supported models [here](https://docs.vllm.ai/en/latest/models/supported_models.html). -Sarvam MLA (`SarvamMLAForCausalLM`) supports [EAGLE3 speculative decoding](docs/features/speculative_decoding/eagle.md) with a compatible draft checkpoint and a single pipeline stage (`--pipeline-parallel-size 1`). Auxiliary hidden states are not transported across pipeline stages. Ordinary pipeline parallelism remains supported when EAGLE3 is disabled. - ## Getting Started Install vLLM with [`uv`](https://docs.astral.sh/uv/) (recommended) or `pip`: From e336e1e007c78c9444c4cba1ee00ef2b98554771 Mon Sep 17 00:00:00 2001 From: mohit-sarvam Date: Mon, 7 Sep 2026 15:39:09 +0000 Subject: [PATCH 6/6] Remove Sarvam EAGLE3 unit test file The PR author tested the Sarvam MLA model end to end with DSpark and confirmed that inference works. Co-authored-by: Codex Signed-off-by: mohit-sarvam --- tests/model_executor/test_sarvam_eagle3.py | 195 --------------------- 1 file changed, 195 deletions(-) delete mode 100644 tests/model_executor/test_sarvam_eagle3.py diff --git a/tests/model_executor/test_sarvam_eagle3.py b/tests/model_executor/test_sarvam_eagle3.py deleted file mode 100644 index 0afe8f4962ac..000000000000 --- a/tests/model_executor/test_sarvam_eagle3.py +++ /dev/null @@ -1,195 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -from types import SimpleNamespace -from unittest.mock import Mock - -import torch - -from vllm.model_executor.models import sarvam as sarvam_mod -from vllm.model_executor.models.interfaces import supports_eagle3 -from vllm.model_executor.models.sarvam import SarvamMLAForCausalLM, SarvamMLAModel -from vllm.sequence import IntermediateTensors - - -def _make_model( - *, - start_layer: int = 0, - end_layer: int = 1, - layers: list | None = None, - aux_hidden_state_layers: tuple[int, ...] = (), -) -> SarvamMLAModel: - """Build a SarvamMLAModel without running __init__. - - Constructing the real module needs a full VllmConfig plus an initialized - distributed environment, neither of which the aux hidden state plumbing - depends on. - """ - model = object.__new__(SarvamMLAModel) - object.__setattr__(model, "start_layer", start_layer) - object.__setattr__(model, "end_layer", end_layer) - object.__setattr__(model, "layers", layers if layers is not None else []) - object.__setattr__(model, "aux_hidden_state_layers", aux_hidden_state_layers) - object.__setattr__(model, "embedding_dropout", lambda x: x) - # RMSNorm in fused-residual mode returns (norm(h + r), h + r); the scaling - # is irrelevant here, so pass the sum through unchanged. - object.__setattr__( - model, "norm", lambda h, r=None: (h + r, h + r) if r is not None else h - ) - return model - - -def _layer(hidden_states: torch.Tensor, residual: torch.Tensor) -> Mock: - """Return a decoder stub with fixed hidden states and residual.""" - return Mock(return_value=(hidden_states, residual)) - - -def _patch_pp_group(monkeypatch, *, is_first_rank=True, is_last_rank=True) -> None: - """Select the pipeline stage simulated by the forward call.""" - monkeypatch.setattr( - sarvam_mod, - "get_pp_group", - lambda: SimpleNamespace(is_first_rank=is_first_rank, is_last_rank=is_last_rank), - ) - - -def test_sarvam_mla_advertises_eagle3_support(): - """Allow the runner to detect the EAGLE3 interface.""" - assert supports_eagle3(SarvamMLAForCausalLM) - - -def test_sarvam_mla_configures_aux_hidden_state_layers(): - """Propagate auxiliary capture configuration to the inner model.""" - # SarvamMLAForCausalLM inherits both hooks from the SupportsEagle3 - # protocol, so this pins that the protocol defaults actually reach the - # inner SarvamMLAModel rather than silently doing nothing. - target = object.__new__(SarvamMLAForCausalLM) - torch.nn.Module.__init__(target) - model = _make_model(layers=[None] * 12) - object.__setattr__(target, "model", model) - - target.set_aux_hidden_state_layers((2, 6, 9)) - - assert model.aux_hidden_state_layers == (2, 6, 9) - assert target.get_eagle3_default_aux_hidden_state_layers() == (2, 6, 9) - - -def test_sarvam_mla_forward_captures_aux_hidden_states(monkeypatch): - """Capture embeddings and complete layer outputs for the drafter.""" - inputs_embeds = torch.tensor([[1.0, 2.0]]) - layer_hidden_states = torch.tensor([[3.0, 4.0]]) - layer_residual = torch.tensor([[5.0, 6.0]]) - model = _make_model( - layers=[_layer(layer_hidden_states, layer_residual)], - aux_hidden_state_layers=(0, 1), - ) - _patch_pp_group(monkeypatch) - - output, aux_hidden_states = model.forward( - input_ids=None, - positions=torch.tensor([0]), - intermediate_tensors=None, - inputs_embeds=inputs_embeds, - ) - - # Index 0 is the embedding output; index 1 is the output of layer 0, which - # is only complete once the pending residual is added back. - expected_layer_output = layer_hidden_states + layer_residual - torch.testing.assert_close(output, expected_layer_output) - assert len(aux_hidden_states) == 2 - torch.testing.assert_close(aux_hidden_states[0], inputs_embeds) - torch.testing.assert_close(aux_hidden_states[1], expected_layer_output) - - -def test_sarvam_mla_forward_returns_bare_tensor_without_eagle3(monkeypatch): - """Preserve the tensor return contract when auxiliary capture is disabled.""" - # Without a drafter the runner unpacks a single tensor, so returning a - # tuple unconditionally would break every non-speculative request. - model = _make_model( - layers=[_layer(torch.tensor([[3.0, 4.0]]), torch.tensor([[5.0, 6.0]]))], - ) - _patch_pp_group(monkeypatch) - - output = model.forward( - input_ids=None, - positions=torch.tensor([0]), - intermediate_tensors=None, - inputs_embeds=torch.tensor([[1.0, 2.0]]), - ) - - assert isinstance(output, torch.Tensor) - - -def test_sarvam_mla_forward_uses_absolute_layer_indices(monkeypatch): - """Use global capture indices within one simulated pipeline stage. - - This checks local indexing, not cross-stage EAGLE3 state transport. - """ - # On a pipeline stage that does not start at layer 0, the capture points - # are numbered by absolute layer index. With relative numbering the - # requested layers would silently resolve to different tensors. - stage_hidden_states = torch.tensor([[1.0, 2.0]]) - stage_residual = torch.tensor([[0.5, 0.5]]) - second_hidden_states = torch.tensor([[3.0, 4.0]]) - second_residual = torch.tensor([[5.0, 6.0]]) - model = _make_model( - start_layer=1, - end_layer=3, - layers=[ - None, - _layer(torch.tensor([[7.0, 8.0]]), torch.tensor([[9.0, 10.0]])), - _layer(second_hidden_states, second_residual), - ], - aux_hidden_state_layers=(1, 3), - ) - _patch_pp_group(monkeypatch, is_first_rank=False) - - _, aux_hidden_states = model.forward( - input_ids=None, - positions=torch.tensor([0]), - intermediate_tensors=IntermediateTensors( - {"hidden_states": stage_hidden_states, "residual": stage_residual} - ), - inputs_embeds=None, - ) - - assert len(aux_hidden_states) == 2 - torch.testing.assert_close( - aux_hidden_states[0], stage_hidden_states + stage_residual - ) - torch.testing.assert_close( - aux_hidden_states[1], second_hidden_states + second_residual - ) - - -def test_sarvam_mla_pipeline_forward_without_eagle3(monkeypatch): - """Preserve hidden states and residual across two stages without EAGLE3.""" - inputs_embeds = torch.tensor([[1.0, 2.0]]) - positions = torch.tensor([0]) - hidden_states = torch.tensor([[3.0, 4.0]]) - residual = torch.tensor([[5.0, 6.0]]) - first = _make_model(layers=[_layer(hidden_states, residual)]) - final_layer = Mock(side_effect=lambda h, p, r: (h * 2, r + 1)) - last = _make_model(start_layer=1, end_layer=2, layers=[None, final_layer]) - _patch_pp_group(monkeypatch, is_last_rank=False) - - intermediate = first.forward( - input_ids=None, - positions=positions, - intermediate_tensors=None, - inputs_embeds=inputs_embeds, - ) - - assert isinstance(intermediate, IntermediateTensors) - torch.testing.assert_close(intermediate["hidden_states"], hidden_states) - torch.testing.assert_close(intermediate["residual"], residual) - _patch_pp_group(monkeypatch, is_first_rank=False) - - output = last.forward( - input_ids=None, - positions=positions, - intermediate_tensors=intermediate, - ) - - assert isinstance(output, torch.Tensor) - torch.testing.assert_close(output, hidden_states * 2 + residual + 1)