Files
sglang/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py
T

63 lines
1.5 KiB
Python

from abc import ABC, abstractmethod
import torch
class LinearAttnKernelBase(ABC):
"""Abstract base class for linear attention kernel implementations.
Each concrete implementation wraps a specific kernel (Triton, CuTe DSL, etc.)
and provides decode/extend/target_verify methods with a unified interface.
"""
@abstractmethod
def decode(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
*,
A_log: torch.Tensor,
dt_bias: torch.Tensor,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
query_start_loc: torch.Tensor,
**kwargs,
) -> torch.Tensor: ...
@abstractmethod
def extend(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
*,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
query_start_loc: torch.Tensor,
**kwargs,
) -> tuple: ...
def target_verify(
self,
A_log: torch.Tensor,
dt_bias: torch.Tensor,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
*,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
query_start_loc: torch.Tensor,
**kwargs,
) -> torch.Tensor:
raise NotImplementedError(
f"{self.__class__.__name__} does not support target_verify"
)