fix: forward update_mamba_state_after_mtp_verify in HybridAttnBackend (#25883)
This commit is contained in:
@@ -139,6 +139,16 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
backend = self._select_backend(forward_batch.forward_mode)
|
backend = self._select_backend(forward_batch.forward_mode)
|
||||||
return backend.get_indexer_metadata(layer_id, forward_batch)
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor = None,
|
q: torch.Tensor = None,
|
||||||
|
|||||||
Reference in New Issue
Block a user