fix: forward update_mamba_state_after_mtp_verify in HybridAttnBackend (#25883)

This commit is contained in:
fatSheep
2026-06-09 20:06:50 -07:00
committed by GitHub
parent f101b287ef
commit d21c31f681
@@ -139,6 +139,16 @@ class HybridAttnBackend(AttentionBackend):
backend = self._select_backend(forward_batch.forward_mode)
return backend.get_indexer_metadata(layer_id, forward_batch)
def update_mamba_state_after_mtp_verify(self, *args, **kwargs):
# Forward to whichever sub-backend handled target_verify, since its inner
# linear_attn_backend.forward_metadata holds the mamba_cache_indices the
# method consumes. Mirrors _select_backend's target_verify branch.
if self.model_runner.server_args.speculative_attention_mode == "decode":
backend = self.decode_backend
else:
backend = self.prefill_backend
return backend.update_mamba_state_after_mtp_verify(*args, **kwargs)
def forward(
self,
q: torch.Tensor = None,