feat: add timed() context manager and backend supports()
- Each backend exposes static supports(**kwargs) for capability query - CudaBackend.supports checks head_dim + kernel availability - FlashAttnBackend/TorchNativeBackend always return True - timed() context manager gated by ASTRAI_TIMED=1 env var, logs via logger.info - Wraps warmup prefill/decode, execute_prefill, and execute_decode
This commit is contained in:
@@ -388,6 +388,10 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
runs SDPA directly on the projected q/k/v.
|
runs SDPA directly on the projected q/k/v.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports(**kwargs) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def fwd_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
@@ -479,6 +483,15 @@ class CudaBackend(AttentionBackend):
|
|||||||
Raises ``RuntimeError`` if the required kernel is not available.
|
Raises ``RuntimeError`` if the required kernel is not available.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports(**kwargs) -> bool:
|
||||||
|
head_dim = kwargs.get("head_dim", -1)
|
||||||
|
return (
|
||||||
|
head_dim in (32, 64, 128, 256)
|
||||||
|
and is_available("attn_paged_decode")
|
||||||
|
and is_available("attn_paged_prefill")
|
||||||
|
)
|
||||||
|
|
||||||
def fwd_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
@@ -575,15 +588,12 @@ class FlashAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
Prefill / non-contiguous decode: falls back to KV gather +
|
Prefill / non-contiguous decode: falls back to KV gather +
|
||||||
``flash_attn_func``.
|
``flash_attn_func``.
|
||||||
|
|
||||||
This backend only does flash attention — inputs ``flash-attn`` cannot
|
|
||||||
express (missing package, custom attention mask on prefill, fp32,
|
|
||||||
unsupported head_dim) raise a clear error instead of silently falling
|
|
||||||
back to torch.
|
|
||||||
|
|
||||||
For a torch fallback, select ``TorchNativeBackend`` instead.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports(**kwargs) -> bool:
|
||||||
|
return flash_attn_available()
|
||||||
|
|
||||||
def fwd_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
|
|||||||
@@ -1,4 +1,7 @@
|
|||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
@@ -20,6 +23,19 @@ from astrai.model.automodel import AutoModel
|
|||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
_TIMED = os.environ.get("ASTRAI_TIMED", "") == "1"
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def timed(label: str, log: Optional[logging.Logger] = None):
|
||||||
|
"""Wall-clock debug timer, enabled via ``ASTRAI_TIMED=1``."""
|
||||||
|
if not _TIMED:
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
tic = time.perf_counter()
|
||||||
|
yield
|
||||||
|
elapsed_ms = (time.perf_counter() - tic) * 1000
|
||||||
|
(log or logger).info("%s %.1fms", label, elapsed_ms)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -64,7 +80,7 @@ def _warmup_cuda_graphs(
|
|||||||
ws: InferenceWorkspace,
|
ws: InferenceWorkspace,
|
||||||
gctx: "CudaGraphContext",
|
gctx: "CudaGraphContext",
|
||||||
max_batch_size: int,
|
max_batch_size: int,
|
||||||
prompt_len: int = 32,
|
prompt_len: int = 1,
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
):
|
):
|
||||||
batch_sizes = [1]
|
batch_sizes = [1]
|
||||||
@@ -90,7 +106,11 @@ def _warmup_cuda_graphs(
|
|||||||
pool.task_free(tid)
|
pool.task_free(tid)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
attn_backend(ATTN_BACKEND.CUDA),
|
||||||
|
timed(f"warmup prefill b={b}", logger),
|
||||||
|
):
|
||||||
kv_cache = pool.bind_tasks(task_ids, ws, start_pos=0)
|
kv_cache = pool.bind_tasks(task_ids, ws, start_pos=0)
|
||||||
ids_in = torch.tensor(prompt_tokens, dtype=torch.long, device=dev)
|
ids_in = torch.tensor(prompt_tokens, dtype=torch.long, device=dev)
|
||||||
pos_in = torch.arange(prompt_len, device=dev).unsqueeze(0).expand(b, -1)
|
pos_in = torch.arange(prompt_len, device=dev).unsqueeze(0).expand(b, -1)
|
||||||
@@ -101,6 +121,11 @@ def _warmup_cuda_graphs(
|
|||||||
position_ids=pos_in,
|
position_ids=pos_in,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
attn_backend(ATTN_BACKEND.CUDA),
|
||||||
|
timed(f"warmup decode b={b}", logger),
|
||||||
|
):
|
||||||
for step in range(2):
|
for step in range(2):
|
||||||
seq_pos = prompt_len + step
|
seq_pos = prompt_len + step
|
||||||
ws.position_ids[:b] = seq_pos
|
ws.position_ids[:b] = seq_pos
|
||||||
@@ -153,6 +178,10 @@ class Executor:
|
|||||||
config = model.config
|
config = model.config
|
||||||
max_q_heads = config.num_attention_heads
|
max_q_heads = config.num_attention_heads
|
||||||
head_dim = config.hidden_size // config.num_attention_heads
|
head_dim = config.hidden_size // config.num_attention_heads
|
||||||
|
self._head_dim = head_dim
|
||||||
|
self._graph_supported = CudaBackend.supports(
|
||||||
|
head_dim=head_dim
|
||||||
|
) and "cuda" in str(self.device)
|
||||||
self._workspace = InferenceWorkspace(
|
self._workspace = InferenceWorkspace(
|
||||||
max_batch_size=kv_cache.max_batch_size,
|
max_batch_size=kv_cache.max_batch_size,
|
||||||
max_seq_len=kv_cache.max_seq_len,
|
max_seq_len=kv_cache.max_seq_len,
|
||||||
@@ -169,11 +198,7 @@ class Executor:
|
|||||||
self._try_enable_cuda_graph()
|
self._try_enable_cuda_graph()
|
||||||
|
|
||||||
def _try_enable_cuda_graph(self):
|
def _try_enable_cuda_graph(self):
|
||||||
on_cuda = "cuda" in str(self.device)
|
if not self._graph_supported:
|
||||||
head_dim = (
|
|
||||||
self.model.config.hidden_size // self.model.config.num_attention_heads
|
|
||||||
)
|
|
||||||
if not on_cuda or head_dim not in (32, 64, 128, 256):
|
|
||||||
return
|
return
|
||||||
|
|
||||||
self._graph_ctx.set_enabled(True)
|
self._graph_ctx.set_enabled(True)
|
||||||
@@ -183,7 +208,6 @@ class Executor:
|
|||||||
self._workspace,
|
self._workspace,
|
||||||
self._graph_ctx,
|
self._graph_ctx,
|
||||||
max_batch_size=self.kv_cache.max_batch_size,
|
max_batch_size=self.kv_cache.max_batch_size,
|
||||||
prompt_len=32,
|
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -266,7 +290,10 @@ class Executor:
|
|||||||
prompt_len, device=self.device
|
prompt_len, device=self.device
|
||||||
)
|
)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
timed(f"execute_prefill b={batch_sz} prompt_len={prompt_len}", logger),
|
||||||
|
):
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
@@ -339,19 +366,17 @@ class Executor:
|
|||||||
|
|
||||||
use_graph = (
|
use_graph = (
|
||||||
self._graph_ctx.enabled
|
self._graph_ctx.enabled
|
||||||
and "cuda" in str(self.device)
|
and self._graph_supported
|
||||||
and isinstance(get_backend(), CudaBackend)
|
and isinstance(get_backend(), CudaBackend)
|
||||||
)
|
)
|
||||||
if use_graph:
|
|
||||||
head_dim = (
|
|
||||||
self.model.config.hidden_size // self.model.config.num_attention_heads
|
|
||||||
)
|
|
||||||
use_graph = use_graph and head_dim in (32, 64, 128, 256)
|
|
||||||
key = (b,)
|
key = (b,)
|
||||||
if use_graph:
|
if use_graph:
|
||||||
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
|
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
timed(f"execute_decode forward b={b}", logger),
|
||||||
|
):
|
||||||
if use_graph:
|
if use_graph:
|
||||||
outputs = self._graph_ctx.forward(
|
outputs = self._graph_ctx.forward(
|
||||||
self.model,
|
self.model,
|
||||||
|
|||||||
Reference in New Issue
Block a user