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)
|
||||
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,
|
||||
|
||||
Reference in New Issue
Block a user