63 lines
1.5 KiB
Python
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"
|
|
)
|