8 Commits
Author SHA1 Message Date
ViperEkura 6ac3b51496 refactor: separate extension ops and backends 2026-08-16 21:15:52 +08:00
ViperEkura 3406157431 refactor: standardize packed 3d inference
- keep training attention on dense 4d tensors
- use packed 3d tensors with KV cache for inference
- extend CUDA rotary embedding to packed 3d inputs
- adapt torch, CUDA and FlashAttention backend dispatch
2026-08-16 13:24:02 +08:00
ViperEkura 0dd9a417b7 refactor: separate KV token address resolution 2026-08-15 22:59:35 +08:00
ViperEkura a01c1fd427 perf: bypass L1 for attention tile loads 2026-08-15 21:23:20 +08:00
ViperEkura f8d9ab344d refactor: remove unused streaming dataset 2026-08-15 20:55:08 +08:00
ViperEkura 3fb4b8ab13 perf: use int32 paged KV indices
- store page-table, request-row, and cache-location indices as int32
- preserve CUDA graph replay with bit-exact logits and KV cache coverage
- improve B=1 decode latency by 1-6% across 1K-32K contexts on L20
2026-08-15 13:17:06 +08:00
ViperEkura b5afe3d7a4 perf: optimize small-head causal prefill
- map D=32 and D=64 causal prefill to BC=64 tiles

- add small-head correctness and benchmark coverage
2026-08-14 23:25:49 +08:00
ViperEkura 69f35c46e0 fix: quantize amax from raw values, not scaled fp8 values
- amax for delayed scale was the quantized max (always ~448), so scale collapsed to 1
- this made fp8 gradients diverge (cosine 0.05) and training stall
- stop w/x transpose-quantize amax from polluting the grad scale
2026-08-14 14:26:16 +08:00
44 changed files with 654 additions and 533 deletions
-2
View File
@@ -14,7 +14,6 @@ from astrai.dataset.storage import (
Streamable, Streamable,
detect_format, detect_format,
) )
from astrai.dataset.streaming import StreamingSeqDataset
from astrai.serialization import ( from astrai.serialization import (
load_bin, load_bin,
save_bin, save_bin,
@@ -35,5 +34,4 @@ __all__ = [
"save_bin", "save_bin",
"load_bin", "load_bin",
"RDSampler", "RDSampler",
"StreamingSeqDataset",
] ]
-122
View File
@@ -1,122 +0,0 @@
"""Streaming IterableDataset for pre-training with shard-level shuffle.
Unlike the map-style datasets, the streaming dataset yields windows
sequentially through each data shard — no random access, no sampler.
Each DataLoader worker independently streams its assigned shard subset,
giving better OS page-cache locality for large-scale (TB+) datasets.
Key properties:
- Implements ``torch.utils.data.IterableDataset``.
- ``__len__`` returns total window count so ``compute_total_steps`` works.
- Shard-level shuffle with deterministic seed (reproducible across runs).
- Distributed: each rank gets a disjoint subset of shards.
- Multi-worker: each worker within a rank gets a disjoint subset.
"""
import random
from typing import Iterator, Optional
import torch
import torch.distributed as dist
from torch import Tensor
from torch.utils.data import IterableDataset
from astrai.dataset.storage import Store
def _resolve_rank_and_world_size() -> tuple[int, int]:
if dist.is_available() and dist.is_initialized():
return dist.get_rank(), dist.get_world_size()
return 0, 1
def _total_windows(token_count, window_size, stride):
if token_count <= window_size:
return 0
return (token_count - 1 - window_size) // stride + 1
class StreamingSeqDataset(IterableDataset):
"""Streaming next-token prediction dataset.
Yields ``{"input_ids": [L], "target_ids": [L]}`` dicts by sliding a
window sequentially through each data shard. Shards are shuffled
deterministically. Distributed and multi-worker DataLoader modes are
supported: each consumer gets a disjoint shard subset.
Args:
store: Already-loaded Store with a ``"sequence"`` key.
window_size: Context length per sample.
stride: Step between consecutive windows (default: window_size).
shuffle: Shuffle shard order.
seed: Base seed for deterministic shard shuffle.
"""
def __init__(
self,
store: Store,
window_size: int,
stride: Optional[int] = None,
shuffle: bool = True,
seed: int = 42,
rank: Optional[int] = None,
world_size: Optional[int] = None,
):
super().__init__()
if window_size <= 0:
raise ValueError("window_size must be positive")
self.store = store
self.window_size = window_size
self.stride = stride if stride is not None else window_size
self.shuffle = shuffle
self.seed = seed
self._rank, self._world_size = (
rank,
world_size if rank is not None else _resolve_rank_and_world_size(),
)
if "sequence" not in store.keys:
raise KeyError(
f"Store is missing required key 'sequence'; "
f"available keys: {sorted(store.keys)}"
)
@property
def num_samples(self) -> int:
return _total_windows(self.store.token_count, self.window_size, self.stride)
def __len__(self) -> int:
return self.num_samples
def __iter__(self) -> Iterator[dict[str, Tensor]]:
segments = self.store._data["sequence"]
n_shards = len(segments)
indices = list(range(n_shards))
if self.shuffle:
rng = random.Random(self.seed)
rng.shuffle(indices)
worker_info = torch.utils.data.get_worker_info()
if worker_info is None:
num_consumers = self._world_size
consumer_id = self._rank
else:
num_consumers = self._world_size * worker_info.num_workers
consumer_id = self._rank * worker_info.num_workers + worker_info.id
my_shards = [
i for idx, i in enumerate(indices) if idx % num_consumers == consumer_id
]
for shard_idx in my_shards:
segment = segments[shard_idx]
seq_len = segment.shape[0]
for begin in range(0, seq_len - self.window_size, self.stride):
end = begin + self.window_size
yield {
"input_ids": torch.as_tensor(segment[begin:end], dtype=torch.long),
"target_ids": torch.as_tensor(
segment[begin + 1 : end + 1], dtype=torch.long
),
}
+4 -4
View File
@@ -15,25 +15,25 @@ Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
SDPA is handled by the attention backend, not the wrapper functions. SDPA is handled by the attention backend, not the wrapper functions.
""" """
from astrai.extension.attention_backend import ( from astrai.extension.backend import (
ATTN_BACKEND, ATTN_BACKEND,
AttentionBackend, AttentionBackend,
AttentionBackendFactory, AttentionBackendFactory,
CudaBackend, CudaBackend,
FlashAttnBackend, FlashAttnBackend,
TorchNativeBackend, TorchNativeBackend,
apply_rotary_emb,
attention, attention,
attn_backend, attn_backend,
get_backend, get_backend,
) )
from astrai.extension.attention_ops import ( from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import (
TensorLayout, TensorLayout,
attn_decode, attn_decode,
attn_paged_decode, attn_paged_decode,
attn_prefill, attn_prefill,
) )
from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.rotary_backend import apply_rotary_emb
__all__ = [ __all__ = [
"ATTN_BACKEND", "ATTN_BACKEND",
+27
View File
@@ -0,0 +1,27 @@
"""Backend selection, fallbacks, and execution policies."""
from astrai.extension.backend.attention import (
ATTN_BACKEND,
AttentionBackend,
AttentionBackendFactory,
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
attention,
attn_backend,
get_backend,
)
from astrai.extension.backend.rotary import apply_rotary_emb
__all__ = [
"ATTN_BACKEND",
"AttentionBackend",
"AttentionBackendFactory",
"CudaBackend",
"FlashAttnBackend",
"TorchNativeBackend",
"apply_rotary_emb",
"attention",
"attn_backend",
"get_backend",
]
@@ -32,7 +32,6 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
import contextvars import contextvars
import enum import enum
import functools import functools
import importlib
import os import os
import threading import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
@@ -43,13 +42,18 @@ import torch
import torch.nn.functional as F import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.extension.attention_ops import ( from astrai.extension.loader import is_available
from astrai.extension.ops.attention import (
attn_paged_decode, attn_paged_decode,
attn_paged_prefill, attn_paged_prefill,
) )
from astrai.extension.loader import is_available
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
try:
import flash_attn as _flash_attn
except Exception:
_flash_attn = None
if TYPE_CHECKING: if TYPE_CHECKING:
from astrai.inference.cache import KVCache from astrai.inference.cache import KVCache
@@ -67,7 +71,7 @@ _current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
def flash_attn_available() -> bool: def flash_attn_available() -> bool:
if not torch.cuda.is_available(): if not torch.cuda.is_available():
return False return False
fa = _get_flash_attn() fa = _flash_attn
if fa is None: if fa is None:
return False return False
@@ -90,14 +94,6 @@ def flash_attn_available() -> bool:
return False return False
@functools.lru_cache(maxsize=1)
def _get_flash_attn():
try:
return importlib.import_module("flash_attn")
except Exception:
return None
class ATTN_BACKEND(enum.Enum): class ATTN_BACKEND(enum.Enum):
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``.""" """Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
@@ -123,6 +119,7 @@ def _backend_supports(
kv_cache: Optional["KVCache"], kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor], attn_mask: Optional[Tensor],
is_causal: bool, is_causal: bool,
fwd: Optional[str],
) -> bool: ) -> bool:
"""Whether ``backend`` can run this attention call. """Whether ``backend`` can run this attention call.
@@ -131,17 +128,20 @@ def _backend_supports(
""" """
if isinstance(backend, CudaBackend): if isinstance(backend, CudaBackend):
return ( return (
kv_cache is not None fwd in ("prefill", "decode")
and kv_cache is not None
and q.ndim == 3
and q.dtype == torch.bfloat16 and q.dtype == torch.bfloat16
and q.size(-1) in (32, 64, 128, 256) and q.size(-1) in (32, 64, 128, 256)
and is_available(f"attn_paged_{fwd}")
) )
if isinstance(backend, FlashAttnBackend): if isinstance(backend, FlashAttnBackend):
if not flash_attn_available(): if not flash_attn_available():
return False return False
if q.dtype not in (torch.float16, torch.bfloat16): if q.dtype not in (torch.float16, torch.bfloat16):
return False return False
if q.size(1) == 1 and kv_cache is not None: if fwd is not None:
return True return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
if attn_mask is None or is_causal: if attn_mask is None or is_causal:
return True return True
return attn_mask.dim() == 4 return attn_mask.dim() == 4
@@ -243,13 +243,13 @@ def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
def repeat_kv(x: Tensor, n_rep: int) -> Tensor: def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""Expand KV heads to match Q heads for GQA.""" """Expand KV heads to match Q heads for GQA."""
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1: if n_rep == 1:
return x return x
n_heads, head_dim = x.shape[-2:]
return ( return (
x[:, :, :, None, :] x.unsqueeze(-2)
.expand(bs, slen, n_heads, n_rep, head_dim) .expand(*x.shape[:-2], n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim) .reshape(*x.shape[:-2], n_heads * n_rep, head_dim)
) )
@@ -283,6 +283,7 @@ def attention(
layer_id: int = 0, layer_id: int = 0,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``. """Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
@@ -302,9 +303,11 @@ def attention(
Returns: Returns:
[batch, q_len, n_heads * head_dim] [batch, q_len, n_heads * head_dim]
""" """
explicit = get_backend(use_default=False)
backend = get_backend() backend = get_backend()
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal): if fwd is None and explicit is None:
explicit = get_backend(use_default=False) backend = TorchNativeBackend()
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal, fwd):
if explicit is not None: if explicit is not None:
raise RuntimeError( raise RuntimeError(
f"Explicitly-set backend {type(backend).__name__} cannot " f"Explicitly-set backend {type(backend).__name__} cannot "
@@ -316,10 +319,10 @@ def attention(
for candidate in _priority_backends(): for candidate in _priority_backends():
if isinstance(candidate, type(backend)): if isinstance(candidate, type(backend)):
continue continue
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal): if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal, fwd):
backend = candidate backend = candidate
break break
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal) return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
class AttentionBackend(ABC): class AttentionBackend(ABC):
@@ -355,6 +358,7 @@ class AttentionBackend(ABC):
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
"""Dispatch to decode or extend based on q_len. """Dispatch to decode or extend based on q_len.
@@ -370,9 +374,11 @@ class AttentionBackend(ABC):
Returns: Returns:
[batch, q_len, n_heads * head_dim] [batch, q_len, n_heads * head_dim]
""" """
if kv_cache is not None and q.size(1) == 1: if fwd == "decode":
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal) return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal) if fwd == "prefill" or fwd is None:
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
raise ValueError(f"unsupported attention forward mode: {fwd}")
@abstractmethod @abstractmethod
def fwd_decode( def fwd_decode(
@@ -466,23 +472,52 @@ class TorchNativeBackend(AttentionBackend):
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
if kv_cache is not None: if q.ndim == 4:
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask) n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
return (
F.scaled_dot_product_attention(
q.permute(0, 2, 1, 3),
k.permute(0, 2, 1, 3),
v.permute(0, 2, 1, 3),
attn_mask,
is_causal=is_causal,
)
.permute(0, 2, 1, 3)
.contiguous()
)
n_rep = q.size(2) // k.size(2) if kv_cache is None or kv_cache.qo_indptr is None:
if n_rep > 1: raise ValueError("packed attention requires KV cache metadata")
k = repeat_kv(k, n_rep) kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
v = repeat_kv(v, n_rep) kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
outputs = []
out = F.scaled_dot_product_attention( n_rep = q.size(1) // k.size(1)
q.permute(0, 2, 1, 3), for i in range(kv_cache.req_pool_indices.numel()):
k.permute(0, 2, 1, 3), q_start = int(kv_cache.qo_indptr[i])
v.permute(0, 2, 1, 3), q_end = int(kv_cache.qo_indptr[i + 1])
attn_mask, indices = kv_cache.req_to_token[
is_causal=is_causal, kv_cache.req_pool_indices[i], : kv_cache.seq_lens[i]
) ]
out = out.permute(0, 2, 1, 3).contiguous().flatten(2) k_i = kv_cache.k_buffer[layer_id, indices]
return out v_i = kv_cache.v_buffer[layer_id, indices]
if n_rep > 1:
k_i = repeat_kv(k_i, n_rep)
v_i = repeat_kv(v_i, n_rep)
q_len = q_end - q_start
kv_len = k_i.size(0)
q_pos = torch.arange(kv_len - q_len, kv_len, device=q.device)
causal_mask = q_pos[:, None] >= torch.arange(kv_len, device=q.device)
out = F.scaled_dot_product_attention(
q[q_start:q_end].transpose(0, 1).unsqueeze(0),
k_i.transpose(0, 1).unsqueeze(0),
v_i.transpose(0, 1).unsqueeze(0),
attn_mask=causal_mask,
)
outputs.append(out.squeeze(0).transpose(0, 1))
return torch.cat(outputs)
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value) @AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
@@ -530,16 +565,14 @@ class CudaBackend(AttentionBackend):
if kv_cache is None: if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)") raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
loc = kv_cache.out_cache_loc[:, 0] loc = kv_cache.out_cache_loc
kv_cache.k_buffer[layer_id].index_copy_(0, loc, k[:, 0]) kv_cache.k_buffer[layer_id, loc] = k
kv_cache.v_buffer[layer_id].index_copy_(0, loc, v[:, 0]) kv_cache.v_buffer[layer_id, loc] = v
q_3d = q.squeeze(1)
kv_indptr = kv_cache.kv_indptr kv_indptr = kv_cache.kv_indptr
out = attn_paged_decode( out = attn_paged_decode(
q_3d, q,
kv_cache.k_buffer[layer_id], kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id], kv_cache.v_buffer[layer_id],
kv_cache.req_to_token, kv_cache.req_to_token,
@@ -550,7 +583,7 @@ class CudaBackend(AttentionBackend):
ml_part_buf=kv_cache.decode_ml_part, ml_part_buf=kv_cache.decode_ml_part,
out_buf=kv_cache.decode_out, out_buf=kv_cache.decode_out,
) )
return out.unsqueeze(1).flatten(2) return out
def fwd_prefill( def fwd_prefill(
self, self,
@@ -565,34 +598,22 @@ class CudaBackend(AttentionBackend):
if kv_cache is None: if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)") raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
loc = kv_cache.out_cache_loc.reshape(-1) loc = kv_cache.out_cache_loc
kv_cache.k_buffer[layer_id].index_copy_( kv_cache.k_buffer[layer_id, loc] = k
0, loc, k.reshape(-1, k.size(2), k.size(3)) kv_cache.v_buffer[layer_id, loc] = v
)
kv_cache.v_buffer[layer_id].index_copy_(
0, loc, v.reshape(-1, v.size(2), v.size(3))
)
b = q.size(0)
q_len = q.size(1)
kv_indptr = kv_cache.kv_indptr
qo_indptr = kv_cache.qo_indptr
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
out = attn_paged_prefill( out = attn_paged_prefill(
q_flat, q,
kv_cache.k_buffer[layer_id], kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id], kv_cache.v_buffer[layer_id],
kv_cache.req_to_token, kv_cache.req_to_token,
kv_cache.req_pool_indices, kv_cache.req_pool_indices,
kv_indptr, kv_cache.kv_indptr,
qo_indptr, kv_cache.qo_indptr,
attn_mask, attn_mask,
is_causal=is_causal, is_causal=is_causal,
) )
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2) return out
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value) @AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
@@ -621,7 +642,7 @@ class FlashAttnBackend(AttentionBackend):
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal) return self._forward_packed(q, k, v, kv_cache, layer_id)
def fwd_prefill( def fwd_prefill(
self, self,
@@ -633,25 +654,18 @@ class FlashAttnBackend(AttentionBackend):
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal) if q.ndim == 3:
return self._forward_packed(q, k, v, kv_cache, layer_id)
return self._forward_dense(q, k, v, attn_mask, is_causal)
def _forward( def _forward_dense(
self, self,
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
if kv_cache is not None:
if q.size(1) == 1 and kv_cache.k_buffer.size(
1
) == kv_cache.req_to_token.size(0) * kv_cache.req_to_token.size(1):
return self._decode_with_kvcache(q, k, v, kv_cache, layer_id)
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask)
n_rep = q.size(2) // k.size(2) n_rep = q.size(2) // k.size(2)
if n_rep > 1: if n_rep > 1:
k = repeat_kv(k, n_rep) k = repeat_kv(k, n_rep)
@@ -662,7 +676,7 @@ class FlashAttnBackend(AttentionBackend):
"FlashAttnBackend does not support a custom attention mask; " "FlashAttnBackend does not support a custom attention mask; "
"use a causal mask or select TorchNativeBackend." "use a causal mask or select TorchNativeBackend."
) )
fa = _get_flash_attn() fa = _flash_attn
if fa is None: if fa is None:
raise RuntimeError( raise RuntimeError(
"FlashAttnBackend requires the optional 'flash-attn' package. " "FlashAttnBackend requires the optional 'flash-attn' package. "
@@ -674,9 +688,9 @@ class FlashAttnBackend(AttentionBackend):
v.contiguous(), v.contiguous(),
causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4), causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4),
) )
return out.contiguous().flatten(2) return out.contiguous()
def _decode_with_kvcache( def _forward_packed(
self, self,
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
@@ -684,22 +698,27 @@ class FlashAttnBackend(AttentionBackend):
kv_cache: "KVCache", kv_cache: "KVCache",
layer_id: int, layer_id: int,
) -> Tensor: ) -> Tensor:
max_batch = kv_cache.req_to_token.size(0) fa = _flash_attn
max_seq = kv_cache.req_to_token.size(1) if fa is None or not hasattr(fa, "flash_attn_varlen_func"):
n_kv = k.size(2) raise RuntimeError("packed inference requires flash_attn_varlen_func")
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
k_cache = kv_cache.k_buffer[layer_id].view(max_batch, max_seq, n_kv, k.size(3)) kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
v_cache = kv_cache.v_buffer[layer_id].view(max_batch, max_seq, n_kv, v.size(3)) page_table = kv_cache.req_to_token[
kv_cache.req_pool_indices, : kv_cache.max_len
fa = _get_flash_attn() ]
out = fa.flash_attn_with_kvcache( positions = torch.arange(kv_cache.max_len, device=q.device)
q=q, indices = page_table[positions.unsqueeze(0) < kv_cache.seq_lens.unsqueeze(1)]
k_cache=k_cache, k_flat = kv_cache.k_buffer[layer_id, indices].contiguous()
v_cache=v_cache, v_flat = kv_cache.v_buffer[layer_id, indices].contiguous()
k=k, out = fa.flash_attn_varlen_func(
v=v, q.contiguous(),
cache_seqlens=(kv_cache.seq_lens - 1).to(torch.int32), k_flat,
cache_batch_idx=kv_cache.req_pool_indices.to(torch.int32), v_flat,
kv_cache.qo_indptr,
kv_cache.kv_indptr,
int((kv_cache.qo_indptr[1:] - kv_cache.qo_indptr[:-1]).max()),
int(kv_cache.seq_lens.max()),
dropout_p=0.0,
causal=True, causal=True,
) )
return out.flatten(2) return out
@@ -11,6 +11,7 @@ import torch
from torch import Tensor from torch import Tensor
from astrai.extension.loader import is_available from astrai.extension.loader import is_available
from astrai.extension.ops.rotary import rotary_emb as _cuda_rotary
_cache = {"available": None} _cache = {"available": None}
@@ -26,7 +27,7 @@ def _torch_apply(x: Tensor, freqs_cis: Tensor) -> Tensor:
dtype = x.dtype dtype = x.dtype
x_ = x.float().reshape(*x.shape[:-1], -1, 2) x_ = x.float().reshape(*x.shape[:-1], -1, 2)
x_complex = torch.view_as_complex(x_) x_complex = torch.view_as_complex(x_)
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(2) freqs_cis_complex = torch.complex(cos, sin).unsqueeze(-2)
x_rotated = x_complex * freqs_cis_complex x_rotated = x_complex * freqs_cis_complex
x_out = torch.view_as_real(x_rotated).flatten(-2) x_out = torch.view_as_real(x_rotated).flatten(-2)
return x_out.to(dtype) return x_out.to(dtype)
@@ -48,7 +49,5 @@ def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
and x.is_cuda and x.is_cuda
and x.dtype == torch.bfloat16 and x.dtype == torch.bfloat16
): ):
from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary
return _cuda_rotary(x, freqs_cis) return _cuda_rotary(x, freqs_cis)
return _torch_apply(x, freqs_cis) return _torch_apply(x, freqs_cis)
+62 -12
View File
@@ -1,8 +1,8 @@
"""FP8 training: scaling state and aten::linear dispatch. """FP8 training: scaling state and aten::linear dispatch.
Layered (see also ``fp8_ops.py`` for the CUDA interface adapter): Layered (see also ``ops/fp8.py`` for the CUDA interface adapter):
1. Kernel interface: "fp8_ops" the only module touching the pybind. 1. Kernel interface: ``ops.fp8`` - the only module touching the pybind.
2. Training state (this module): per-tensor scales, amax history, delayed 2. Training state (this module): per-tensor scales, amax history, delayed
scaling, and the ``fp8_autocast`` context (TE-style, like scaling, and the ``fp8_autocast`` context (TE-style, like
``torch.autocast``). ``torch.autocast``).
@@ -25,7 +25,7 @@ from contextlib import contextmanager
import torch import torch
from torch.library import Library from torch.library import Library
from astrai.extension.fp8_ops import ( from astrai.extension.ops.fp8 import (
linear_backward_scaled, linear_backward_scaled,
linear_forward_scaled, linear_forward_scaled,
) )
@@ -52,8 +52,15 @@ class FP8TensorMeta:
"idx", "idx",
"x_scale", "x_scale",
"x_scale_inv", "x_scale_inv",
"x_history",
"x_idx",
"g_scale", "g_scale",
"g_scale_inv", "g_scale_inv",
"g_history",
"g_idx",
"w_init",
"x_init",
"g_init",
) )
def __init__(self, device: torch.device, update_interval: int): def __init__(self, device: torch.device, update_interval: int):
@@ -65,8 +72,43 @@ class FP8TensorMeta:
self.idx = 0 self.idx = 0
self.x_scale = torch.ones(1, device=device, dtype=torch.float32) self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32) self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.x_history = torch.ones(update_interval, device=device, dtype=torch.float32)
self.x_idx = 0
self.g_scale = torch.ones(1, device=device, dtype=torch.float32) self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32) self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.g_history = torch.ones(update_interval, device=device, dtype=torch.float32)
self.g_idx = 0
self.w_init = False
self.x_init = False
self.g_init = False
def init_scale(self, t: torch.Tensor) -> None:
"""Immediate scale from the current amax; used on the first call.
A scale of 1 would underflow small activations/gradients (e4m3 min
normal is 2^-6); initialize from the actual amax once, then delayed
updates take over.
"""
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
self.scale.copy_(amax / E4M3_MAX)
self.scale_inv.copy_(E4M3_MAX / amax)
self.record(amax)
def push_x_scale(self, amax: torch.Tensor) -> None:
"""Window update for the activation scale (delayed, TE style)."""
self.x_history[self.x_idx] = amax.reshape(())
self.x_idx = (self.x_idx + 1) % self.x_history.numel()
m = self.x_history.max()
self.x_scale.copy_(m / E4M3_MAX)
self.x_scale_inv.copy_(E4M3_MAX / m)
def push_g_scale(self, amax: torch.Tensor) -> None:
"""Window update for the gradient scale (delayed, TE style)."""
self.g_history[self.g_idx] = amax.reshape(())
self.g_idx = (self.g_idx + 1) % self.g_history.numel()
m = self.g_history.max()
self.g_scale.copy_(m / E4M3_MAX)
self.g_scale_inv.copy_(E4M3_MAX / m)
def record(self, amax: torch.Tensor) -> None: def record(self, amax: torch.Tensor) -> None:
"""Push the latest amax into the ring buffer (device-side copy, no sync).""" """Push the latest amax into the ring buffer (device-side copy, no sync)."""
@@ -155,13 +197,6 @@ def fp8_autocast(enabled: bool = True, update_interval: int = 16):
state.update_interval = prev_interval state.update_interval = prev_interval
def _update_delayed_scale(scale, scale_inv, amax) -> None:
"""scale = amax / 448 for the *next* call (device-side, no sync)."""
amax_f = amax.reshape(()).to(torch.float32).clamp_min(1e-12)
scale.copy_(amax_f / E4M3_MAX)
scale_inv.copy_(E4M3_MAX / amax_f)
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None): def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
"""TE-style scaled fp8 linear forward (called from the aten::linear impl). """TE-style scaled fp8 linear forward (called from the aten::linear impl).
@@ -173,6 +208,15 @@ def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
bias = torch.empty(0, device=x.device, dtype=x.dtype) bias = torch.empty(0, device=x.device, dtype=x.dtype)
state = fp8_state() state = fp8_state()
meta = state.get_weight_meta(w) meta = state.get_weight_meta(w)
if not meta.w_init:
meta.init_scale(w)
meta.w_init = True
if not meta.x_init:
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
meta.x_history.fill_(amax)
meta.x_scale.copy_(amax / E4M3_MAX)
meta.x_scale_inv.copy_(E4M3_MAX / amax)
meta.x_init = True
amax_x = torch.empty(1, device=x.device, dtype=torch.float32) amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
amax_w = torch.empty(1, device=x.device, dtype=torch.float32) amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
out = linear_forward_scaled( out = linear_forward_scaled(
@@ -187,7 +231,7 @@ def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
amax_w, amax_w,
) )
meta.record(amax_w) meta.record(amax_w)
_update_delayed_scale(meta.x_scale, meta.x_scale_inv, amax_x) meta.push_x_scale(amax_x)
return out return out
@@ -195,6 +239,12 @@ def fp8_linear_backward(g, x, w, masks):
"""TE-style scaled fp8 linear backward (called from aten::linear_backward).""" """TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
state = fp8_state() state = fp8_state()
meta = state.get_weight_meta(w) meta = state.get_weight_meta(w)
if not meta.g_init:
amax = g.abs().amax().to(torch.float32).clamp_min(1e-12)
meta.g_history.fill_(amax)
meta.g_scale.copy_(amax / E4M3_MAX)
meta.g_scale_inv.copy_(E4M3_MAX / amax)
meta.g_init = True
amax_g = torch.empty(1, device=g.device, dtype=torch.float32) amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
out = linear_backward_scaled( out = linear_backward_scaled(
g, g,
@@ -209,7 +259,7 @@ def fp8_linear_backward(g, x, w, masks):
meta.x_scale_inv, meta.x_scale_inv,
amax_g, amax_g,
) )
_update_delayed_scale(meta.g_scale, meta.g_scale_inv, amax_g) meta.push_g_scale(amax_g)
return out return out
+19
View File
@@ -0,0 +1,19 @@
"""Stateless wrappers around compiled extension kernels."""
from astrai.extension.ops.attention import (
TensorLayout,
attn_decode,
attn_paged_decode,
attn_paged_prefill,
attn_prefill,
)
from astrai.extension.ops.rotary import rotary_emb
__all__ = [
"TensorLayout",
"attn_decode",
"attn_paged_decode",
"attn_paged_prefill",
"attn_prefill",
"rotary_emb",
]
@@ -1,4 +1,4 @@
"""Attention kernel wrapper functions one entry point per compiled kernel. """Attention kernel wrapper functions - one entry point per compiled kernel.
Each wrapper calls its CUDA kernel directly. If the kernel is not Each wrapper calls its CUDA kernel directly. If the kernel is not
available, raises ``RuntimeError``. Fallback to torch SDPA is the available, raises ``RuntimeError``. Fallback to torch SDPA is the
@@ -113,8 +113,8 @@ def attn_paged_decode(
q: [batch, n_heads, head_dim] (bf16, 3D no seq dim) q: [batch, n_heads, head_dim] (bf16, 3D no seq dim)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat) k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int64) token -> slot req_to_token: [num_reqs, max_context_len] (int32) token -> slot
req_pool_indices: [batch] (int64) rows into req_to_token req_pool_indices: [batch] (int32) rows into req_to_token
kv_indptr: [batch+1] (int32) prefix sum of per-request seq_lens kv_indptr: [batch+1] (int32) prefix sum of per-request seq_lens
mask: 2D [batch, max_context_len] (bool, True=keep) or None mask: 2D [batch, max_context_len] (bool, True=keep) or None
is_causal: apply causal mask is_causal: apply causal mask
@@ -163,8 +163,8 @@ def attn_paged_prefill(
q: [total_q, n_heads, head_dim] (bf16, 3D flattened across requests) q: [total_q, n_heads, head_dim] (bf16, 3D flattened across requests)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat) k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int64) req_to_token: [num_reqs, max_context_len] (int32)
req_pool_indices: [batch] (int64) req_pool_indices: [batch] (int32)
kv_indptr: [batch+1] (int32) prefix sum of per-request kv_lens kv_indptr: [batch+1] (int32) prefix sum of per-request kv_lens
qo_indptr: [batch+1] (int32) prefix sum of per-request q_lens qo_indptr: [batch+1] (int32) prefix sum of per-request q_lens
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
@@ -2,10 +2,10 @@
Calls the compiled CUDA kernel directly. If the kernel is not available, Calls the compiled CUDA kernel directly. If the kernel is not available,
raises ``RuntimeError``. Fallback to torch complex multiply is the raises ``RuntimeError``. Fallback to torch complex multiply is the
responsibility of ``astrai.extension.rotary_backend.apply_rotary_emb``. responsibility of ``astrai.extension.backend.rotary.apply_rotary_emb``.
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16, contiguous). Layout: x is packed [tokens, n_heads, head_dim] or dense
freqs_cis is [batch, seq_len, head_dim/2, 2] (f32, contiguous) [cos, sin] pairs. [batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes.
""" """
import torch import torch
@@ -25,11 +25,11 @@ def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
"""Fused rotary embedding kernel. """Fused rotary embedding kernel.
Args: Args:
x: [batch, seq_len, n_heads, head_dim] (bf16, contiguous) x: packed 3D or dense 4D bf16 tensor.
freqs_cis: [batch, seq_len, head_dim/2, 2] (f32, contiguous) [cos, sin] pairs freqs_cis: matching token axes followed by [head_dim/2, 2].
Returns: Returns:
[batch, seq_len, n_heads, head_dim] (bf16) Tensor with the same shape as ``x``.
""" """
_check_available() _check_available()
if not x.is_contiguous(): if not x.is_contiguous():
+1 -1
View File
@@ -27,7 +27,7 @@ class ReqToTokenPool:
self.size = size self.size = size
self.max_context_len = max_context_len self.max_context_len = max_context_len
self.req_to_token = torch.zeros( self.req_to_token = torch.zeros(
(size, max_context_len), dtype=torch.long, device=device (size, max_context_len), dtype=torch.int32, device=device
) )
self.free_slots = list(range(size)) self.free_slots = list(range(size))
self._lock = threading.Lock() self._lock = threading.Lock()
+25 -12
View File
@@ -115,6 +115,8 @@ class PagePool:
self.contiguous = n_tokens is None self.contiguous = n_tokens is None
self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens
if self.n_tokens > torch.iinfo(torch.int32).max:
raise ValueError("KV cache token count exceeds the int32 slot index limit")
self._storage = KVStorage( self._storage = KVStorage(
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
@@ -124,7 +126,10 @@ class PagePool:
if self.contiguous: if self.contiguous:
for i in range(max_batch_size): for i in range(max_batch_size):
self._req_pool.req_to_token[i] = torch.arange( self._req_pool.req_to_token[i] = torch.arange(
i * max_seq_len, (i + 1) * max_seq_len, device=device i * max_seq_len,
(i + 1) * max_seq_len,
dtype=torch.int32,
device=device,
) )
self._strategy: AllocationStrategy = ContiguousStrategy() self._strategy: AllocationStrategy = ContiguousStrategy()
else: else:
@@ -184,7 +189,7 @@ class PagePool:
kvp_buf[: b + 1] += inc_buf[: b + 1] kvp_buf[: b + 1] += inc_buf[: b + 1]
else: else:
rpi_buf[:b].copy_( rpi_buf[:b].copy_(
torch.tensor(req_indices, dtype=torch.long, device=device) torch.tensor(req_indices, dtype=torch.int32, device=device)
) )
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device)) sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
kvp_buf[: b + 1].zero_() kvp_buf[: b + 1].zero_()
@@ -195,14 +200,21 @@ class PagePool:
kv_indptr = kvp_buf[: b + 1] kv_indptr = kvp_buf[: b + 1]
if start_pos is not None: if start_pos is not None:
# ---- prefill: out_cache_loc covers prefix range [start_pos:seq_len] ---- # Packed prefill concatenates each request's query tokens.
seq_len = seq_lens[0] q_lens = [seq_len - start_pos for seq_len in seq_lens]
out_cache_loc = self._req_pool.req_to_token[ if any(q_len <= 0 for q_len in q_lens):
req_pool_indices, start_pos:seq_len raise ValueError("prefill sequence lengths must exceed start_pos")
] out_cache_loc = torch.cat(
q_len = seq_len - start_pos [
workspace.qo_indptr[: b + 1].copy_( self._req_pool.req_to_token[
torch.arange(b + 1, dtype=torch.int32, device=device) * q_len req_pool_indices[i], start_pos : seq_lens[i]
]
for i in range(b)
]
)
workspace.qo_indptr[: b + 1].zero_()
workspace.qo_indptr[1 : b + 1].copy_(
torch.tensor(q_lens, dtype=torch.int32, device=device).cumsum(0)
) )
qo_indptr = workspace.qo_indptr[: b + 1] qo_indptr = workspace.qo_indptr[: b + 1]
decode_o_part = decode_ml_part = decode_out = None decode_o_part = decode_ml_part = decode_out = None
@@ -211,8 +223,9 @@ class PagePool:
write_pos = seq_lens_t - 1 write_pos = seq_lens_t - 1
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1) loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
ocl_buf[:b].copy_(loc) ocl_buf[:b].copy_(loc)
out_cache_loc = ocl_buf[:b] out_cache_loc = ocl_buf[:b].reshape(-1)
qo_indptr = None workspace.qo_indptr[: b + 1].copy_(inc_buf[: b + 1])
qo_indptr = workspace.qo_indptr[: b + 1]
decode_o_part = getattr(workspace, "decode_o_part", None) decode_o_part = getattr(workspace, "decode_o_part", None)
decode_ml_part = getattr(workspace, "decode_ml_part", None) decode_ml_part = getattr(workspace, "decode_ml_part", None)
decode_out = getattr(workspace, "decode_out", None) decode_out = getattr(workspace, "decode_out", None)
+22 -31
View File
@@ -7,7 +7,7 @@ from typing import List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.extension.attention_backend import ( from astrai.extension.backend.attention import (
CudaBackend, CudaBackend,
get_backend, get_backend,
) )
@@ -118,13 +118,13 @@ def _warmup_cuda_graphs(
timed("warmup prefill", logger), timed("warmup prefill", logger),
): ):
kv = task_cache.bind([tid], ws, start_pos=0) kv = task_cache.bind([tid], ws, start_pos=0)
ids_in = torch.arange(warmup_len, device=dev).unsqueeze(0) ids_in = torch.arange(warmup_len, device=dev)
pos_in = ids_in pos_in = ids_in
model( model(
ids_in, ids_in,
input_mask=pos_in.unsqueeze(-1) >= torch.arange(warmup_len, device=dev),
kv_cache=kv, kv_cache=kv,
position_ids=pos_in, position_ids=pos_in,
fwd="prefill",
) )
task_cache.task_free(tid) task_cache.task_free(tid)
@@ -159,15 +159,14 @@ def _warmup_cuda_graphs(
for tid in task_ids: for tid in task_ids:
task_cache.task_extend(tid, seq_pos) task_cache.task_extend(tid, seq_pos)
kv = task_cache.bind(task_ids, ws) kv = task_cache.bind(task_ids, ws)
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
ids_buf = ws.fill_input_ids([step] * b) ids_buf = ws.fill_input_ids([step] * b)
gctx.forward( gctx.forward(
model, model,
key=(b,), key=(b,),
input_ids=ids_buf.unsqueeze(1), input_ids=ids_buf,
input_mask=input_mask,
kv_cache=kv, kv_cache=kv,
position_ids=ws.position_ids[:b].unsqueeze(1), position_ids=ws.position_ids[:b],
fwd="decode",
) )
for tid in task_ids: for tid in task_ids:
@@ -308,20 +307,15 @@ class Executor:
batch_sz = len(tasks) batch_sz = len(tasks)
input_ids = torch.tensor( input_ids = torch.tensor(
[t.prompt_ids[start_pos:prompt_len] for t in tasks], [token for t in tasks for token in t.prompt_ids[start_pos:prompt_len]],
dtype=torch.long, dtype=torch.long,
device=self.device, device=self.device,
) )
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
position_ids = ( position_ids = torch.arange(
torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device) start_pos, prompt_len, dtype=torch.long, device=self.device
.unsqueeze(0) ).repeat(batch_sz)
.expand(batch_sz, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_len, device=self.device
)
with ( with (
torch.inference_mode(), torch.inference_mode(),
@@ -329,15 +323,18 @@ class Executor:
): ):
outputs = self.model( outputs = self.model(
input_ids, input_ids,
input_mask=input_mask,
position_ids=position_ids, position_ids=position_ids,
kv_cache=self.task_cache.bind( kv_cache=self.task_cache.bind(
task_ids, task_ids,
self._workspace, self._workspace,
start_pos=start_pos, start_pos=start_pos,
), ),
fwd="prefill",
) )
logits = outputs["logits"][:, -1, :] q_len = prompt_len - start_pos
logits = outputs["logits"][
torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1
]
return tasks, self._sample_logits(logits, tasks, return_logprobs) return tasks, self._sample_logits(logits, tasks, return_logprobs)
@@ -391,9 +388,6 @@ class Executor:
) )
self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info) self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info)
total_len = max(cur_positions) + 1
input_mask = ws.decode_mask(ws.position_ids[:b], total_len)
# ---- forward (graph replay or live run + capture) ---- # ---- forward (graph replay or live run + capture) ----
use_graph = ( use_graph = (
@@ -402,9 +396,6 @@ class Executor:
and get_backend().supports_graph() and get_backend().supports_graph()
) )
key = (b,) key = (b,)
if use_graph:
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
with ( with (
torch.inference_mode(), torch.inference_mode(),
timed(f"execute_decode forward b={b}", logger), timed(f"execute_decode forward b={b}", logger),
@@ -413,18 +404,18 @@ class Executor:
outputs = self._graph_ctx.forward( outputs = self._graph_ctx.forward(
self.model, self.model,
key=key, key=key,
input_ids=input_ids.unsqueeze(1), input_ids=input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=ws.position_ids[:b].unsqueeze(1), position_ids=ws.position_ids[:b],
fwd="decode",
) )
else: else:
outputs = self.model( outputs = self.model(
input_ids.unsqueeze(1), input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=ws.position_ids[:b].unsqueeze(1), position_ids=ws.position_ids[:b],
fwd="decode",
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"]
return self._sample_logits(logits, tasks, return_logprobs, info=info) return self._sample_logits(logits, tasks, return_logprobs, info=info)
+2 -2
View File
@@ -74,7 +74,7 @@ class InferenceWorkspace:
# when the Executor passes this workspace). Stable addresses make the # when the Executor passes this workspace). Stable addresses make the
# decode forward CUDA-graph capturable. # decode forward CUDA-graph capturable.
self.req_pool_indices = torch.empty( self.req_pool_indices = torch.empty(
(max_batch_size,), dtype=torch.long, device=device (max_batch_size,), dtype=torch.int32, device=device
) )
self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device) self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device)
self.kv_indptr = torch.empty( self.kv_indptr = torch.empty(
@@ -85,7 +85,7 @@ class InferenceWorkspace:
) )
self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device) self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device)
self.out_cache_loc = torch.empty( self.out_cache_loc = torch.empty(
(max_batch_size, 1), dtype=torch.long, device=device (max_batch_size, 1), dtype=torch.int32, device=device
) )
# Per-step position IDs (must be at a fixed address for CUDA-graph capture). # Per-step position IDs (must be at a fixed address for CUDA-graph capture).
+1 -1
View File
@@ -1,4 +1,4 @@
from astrai.extension.rotary_backend import apply_rotary_emb from astrai.extension.backend.rotary import apply_rotary_emb
from astrai.model.components.attention import GQA, MLA from astrai.model.components.attention import GQA, MLA
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding from astrai.model.components.embedding import Embedding
+12 -11
View File
@@ -5,8 +5,7 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.extension import attention from astrai.extension.backend import apply_rotary_emb, attention
from astrai.extension.rotary_backend import apply_rotary_emb
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.inference.cache import KVCache from astrai.inference.cache import KVCache
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
@@ -56,9 +55,7 @@ class GQA(nn.Module):
self.gate = Linear(dim, dim) self.gate = Linear(dim, dim)
def _split_heads(self, x: Tensor, n_heads) -> Tensor: def _split_heads(self, x: Tensor, n_heads) -> Tensor:
batch_size, seq_len, _ = x.shape return x.reshape(*x.shape[:-1], n_heads, self.head_dim)
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
return x
def forward( def forward(
self, self,
@@ -67,6 +64,7 @@ class GQA(nn.Module):
attn_mask: Tensor = None, attn_mask: Tensor = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
q = self._split_heads(self.q_proj(x), self.n_heads) q = self._split_heads(self.q_proj(x), self.n_heads)
k = self._split_heads(self.k_proj(x), self.n_kv_heads) k = self._split_heads(self.k_proj(x), self.n_kv_heads)
@@ -76,7 +74,9 @@ class GQA(nn.Module):
if self.use_qk_norm: if self.use_qk_norm:
q, k = self.q_norm(q), self.k_norm(k) q, k = self.q_norm(q), self.k_norm(k)
sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal) sdqa_out = attention(
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
).reshape(*x.shape[:-1], self.dim)
if self.use_gated_attention: if self.use_gated_attention:
sdqa_out = sdqa_out * F.sigmoid(self.gate(x)) sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
@@ -141,17 +141,16 @@ class MLA(nn.Module):
attn_mask: Tensor = None, attn_mask: Tensor = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size()
q = self.q_proj(x) q = self.q_proj(x)
q = q.view(bsz, seq_len, self.n_heads, self.head_dim) q = q.reshape(*x.shape[:-1], self.n_heads, self.head_dim)
kv_compressed = self.kv_a_proj(x) kv_compressed = self.kv_a_proj(x)
kv_compressed = self.kv_norm(kv_compressed) kv_compressed = self.kv_norm(kv_compressed)
kv = self.kv_b_proj(kv_compressed) kv = self.kv_b_proj(kv_compressed)
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1) kv = kv.reshape(*x.shape[:-1], self.n_kv_heads, -1)
k_nope, k_rope, v = torch.split( k_nope, k_rope, v = torch.split(
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1 kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
@@ -171,7 +170,9 @@ class MLA(nn.Module):
q = self.q_norm(q) q = self.q_norm(q)
k = self.k_norm(k) k = self.k_norm(k)
attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal) attn_out = attention(
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
).reshape(*x.shape[:-1], self.dim)
if self.use_gated_attention: if self.use_gated_attention:
attn_out = attn_out * F.sigmoid(self.gate(x)) attn_out = attn_out * F.sigmoid(self.gate(x))
+2
View File
@@ -54,6 +54,7 @@ class DecoderBlock(nn.Module):
attention_mask: Optional[Tensor] = None, attention_mask: Optional[Tensor] = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> DecoderOutput: ) -> DecoderOutput:
attn_output = self.attention( attn_output = self.attention(
self.input_norm(x), self.input_norm(x),
@@ -61,6 +62,7 @@ class DecoderBlock(nn.Module):
attention_mask, attention_mask,
kv_cache, kv_cache,
is_causal, is_causal,
fwd,
) )
x = attn_output + x x = attn_output + x
normalized = self.post_attention_norm(x) normalized = self.post_attention_norm(x)
+3 -2
View File
@@ -100,13 +100,14 @@ class DeepSeekMoE(nn.Module):
def forward(self, x: Tensor) -> FFNOutput: def forward(self, x: Tensor) -> FFNOutput:
include_aux_loss = self.training and torch.is_grad_enabled() include_aux_loss = self.training and torch.is_grad_enabled()
bsz, seq_len, dim = x.shape shape = x.shape
dim = shape[-1]
x_flat = x.view(-1, dim) x_flat = x.view(-1, dim)
shared_out = self._shared_forward(x_flat) shared_out = self._shared_forward(x_flat)
routed_output = self._routed_forward(x_flat, include_aux_loss) routed_output = self._routed_forward(x_flat, include_aux_loss)
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim) out = (shared_out + routed_output["hidden_states"]).view(shape)
return { return {
"hidden_states": out, "hidden_states": out,
"aux_loss": routed_output["aux_loss"], "aux_loss": routed_output["aux_loss"],
+8 -5
View File
@@ -65,9 +65,12 @@ class RotaryEmbedding(nn.Module):
[batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs. [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
""" """
if position_ids is None: if position_ids is None:
position_ids = ( if x.ndim == 2:
torch.arange(x.size(1), device=x.device) position_ids = torch.arange(x.size(0), device=x.device)
.unsqueeze(0) else:
.expand(x.size(0), -1) position_ids = (
) torch.arange(x.size(1), device=x.device)
.unsqueeze(0)
.expand(x.size(0), -1)
)
return self.freqs_cis[position_ids].float() return self.freqs_cis[position_ids].float()
+14 -1
View File
@@ -105,8 +105,20 @@ class AutoRegressiveLM(AutoModel):
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
position_ids: Optional[Tensor] = None, position_ids: Optional[Tensor] = None,
fwd: Optional[str] = None,
) -> Dict[str, Tensor]: ) -> Dict[str, Tensor]:
assert input_ids.ndim == 2 if fwd is None:
if input_ids.ndim != 2:
raise ValueError("training input_ids must be [batch, seq_len]")
if kv_cache is not None:
raise ValueError("training forward does not accept a KV cache")
elif fwd in ("prefill", "decode"):
if input_ids.ndim != 1:
raise ValueError("inference input_ids must be packed [tokens]")
if kv_cache is None:
raise ValueError("inference forward requires a KV cache")
else:
raise ValueError(f"unsupported forward mode: {fwd}")
x = self.embed_tokens(input_ids) x = self.embed_tokens(input_ids)
rotary_emb = self.rotary_embedding(x, position_ids) rotary_emb = self.rotary_embedding(x, position_ids)
@@ -122,6 +134,7 @@ class AutoRegressiveLM(AutoModel):
attn_mask, attn_mask,
kv_cache, kv_cache,
use_sdpa_causal_mask, use_sdpa_causal_mask,
fwd,
) )
x = layer_output["hidden_states"] x = layer_output["hidden_states"]
stats = layer_output.get("router_stats") stats = layer_output.get("router_stats")
+2 -2
View File
@@ -55,8 +55,8 @@ struct AttentionParams {
int mask_l_stride; int mask_l_stride;
// Paged K/V addressing // Paged K/V addressing
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len] const int* __restrict__ req_to_token; // [num_reqs, max_context_len]
const int64_t* __restrict__ req_pool_indices; // [batch] const int* __restrict__ req_pool_indices; // [batch]
const int* __restrict__ kv_indptr; // [batch + 1] const int* __restrict__ kv_indptr; // [batch + 1]
const int* __restrict__ qo_indptr; // [batch + 1] or nullptr for decode const int* __restrict__ qo_indptr; // [batch + 1] or nullptr for decode
int max_context_len; // req_to_token stride (dim 1) int max_context_len; // req_to_token stride (dim 1)
+2 -1
View File
@@ -57,7 +57,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int s = i / p.head_dim; int s = i / p.head_dim;
int d_dim = i % p.head_dim; int d_dim = i % p.head_dim;
int kc = chunk_start + s; int kc = chunk_start + s;
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); int token = KV::resolve_token(p, kctx, kc, true);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f; k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f; v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
} }
+2 -1
View File
@@ -73,7 +73,8 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < seq_len; bool valid = kc < seq_len;
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); int token = KV::resolve_token(p, kctx, kc, valid);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid); cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid); cp_async_16_pred(&dV[off], a.v, a.valid);
+24 -4
View File
@@ -59,14 +59,34 @@ inline int compute_num_splits(int base_blocks, int tiles_total,
// ====================================================================== // ======================================================================
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
template <int BC_>
struct PrefillKernelConfig {
static constexpr int BC = BC_;
static constexpr int WARPS = 4;
static constexpr int STAGES = 2;
};
// Compile-time configuration map shared by contiguous and paged prefill.
// Unsupported head dimensions intentionally have no mapping.
template <int HEAD_DIM, bool IsCausal>
struct PrefillConfigMap;
template <> struct PrefillConfigMap<32, false> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<32, true> : PrefillKernelConfig<64> {};
template <> struct PrefillConfigMap<64, false> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<64, true> : PrefillKernelConfig<64> {};
template <> struct PrefillConfigMap<128, false> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<128, true> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<256, false> : PrefillKernelConfig<16> {};
template <> struct PrefillConfigMap<256, true> : PrefillKernelConfig<16> {};
template <typename KV> template <typename KV>
struct PrefillLauncherMMA { struct PrefillLauncherMMA {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) { static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
constexpr int WARPS = 4; using Config = PrefillConfigMap<HEAD_DIM, IsCausal>;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; using Traits = KernelTraits<HEAD_DIM, Config::BC, Config::WARPS, Config::STAGES>;
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>; constexpr int ROWS = Traits::BR * Config::WARPS;
constexpr int ROWS = Traits::BR * WARPS;
dim3 grid(KV::host_q_blocks(p, ROWS), p.q_head, dim3 grid(KV::host_q_blocks(p, ROWS), p.q_head,
KV::kPaged ? 1 : p.batch); KV::kPaged ? 1 : p.batch);
dim3 block(Traits::NUM_THREADS); dim3 block(Traits::NUM_THREADS);
+10 -8
View File
@@ -160,8 +160,9 @@ inline void attn_pack_paged_decode_params(
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16"); TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16"); TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16"); TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64"); TORCH_CHECK(req_to_token.dtype() == torch::kInt32, "req_to_token must be int32");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64"); TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
"req_pool_indices must be int32");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32"); TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match"); TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]"); TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
@@ -184,8 +185,8 @@ inline void attn_pack_paged_decode_params(
p.k_ptr = (const T*)k_cache.data_ptr(); p.k_ptr = (const T*)k_cache.data_ptr();
p.v_ptr = (const T*)v_cache.data_ptr(); p.v_ptr = (const T*)v_cache.data_ptr();
p.q_ptr = (const T*)q.data_ptr(); p.q_ptr = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>(); p.req_to_token = req_to_token.data_ptr<int>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>(); p.req_pool_indices = req_pool_indices.data_ptr<int>();
p.kv_indptr = kv_indptr.data_ptr<int>(); p.kv_indptr = kv_indptr.data_ptr<int>();
p.qo_indptr = nullptr; p.qo_indptr = nullptr;
p.max_context_len = (int)req_to_token.size(1); p.max_context_len = (int)req_to_token.size(1);
@@ -239,8 +240,9 @@ inline void attn_pack_paged_prefill_params(
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16"); TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16"); TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16"); TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64"); TORCH_CHECK(req_to_token.dtype() == torch::kInt32, "req_to_token must be int32");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64"); TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
"req_pool_indices must be int32");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32"); TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32"); TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match"); TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
@@ -267,8 +269,8 @@ inline void attn_pack_paged_prefill_params(
p.k_ptr = (const T*)k_cache.data_ptr(); p.k_ptr = (const T*)k_cache.data_ptr();
p.v_ptr = (const T*)v_cache.data_ptr(); p.v_ptr = (const T*)v_cache.data_ptr();
p.q_ptr = (const T*)q.data_ptr(); p.q_ptr = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>(); p.req_to_token = req_to_token.data_ptr<int>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>(); p.req_pool_indices = req_pool_indices.data_ptr<int>();
p.kv_indptr = kv_indptr.data_ptr<int>(); p.kv_indptr = kv_indptr.data_ptr<int>();
p.qo_indptr = qo_indptr.data_ptr<int>(); p.qo_indptr = qo_indptr.data_ptr<int>();
p.max_context_len = (int)req_to_token.size(1); p.max_context_len = (int)req_to_token.size(1);
+23 -11
View File
@@ -33,7 +33,7 @@ using bf16 = __nv_bfloat16;
// Hoisted per-(batch, kv_head) addressing context. // Hoisted per-(batch, kv_head) addressing context.
struct KVContext { struct KVContext {
int kv_base; // contig: batch*kv_b_stride + kv_head*kv_h_stride int kv_base; // contig: batch*kv_b_stride + kv_head*kv_h_stride
int64_t req_idx; // paged: req_pool_indices[batch] int req_idx; // paged: req_pool_indices[batch]
int64_t rtt_stride; // paged: max_context_len int64_t rtt_stride; // paged: max_context_len
int64_t pool_stride; // paged: kv_head * HEAD_DIM int64_t pool_stride; // paged: kv_head * HEAD_DIM
int64_t head_off; // paged: kv_head * HEAD_DIM int64_t head_off; // paged: kv_head * HEAD_DIM
@@ -104,10 +104,18 @@ struct ContigKV {
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride; c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
return c; return c;
} }
HOST_DEV_FORCEINLINE KVAddr kv_addr( HOST_DEV_FORCEINLINE int resolve_token(
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) { const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
const int g_off = c.kv_base + kc * p.kv_l_stride + d * p.kv_d_stride; return valid ? kc : -1;
return {&p.k_ptr[g_off], &p.v_ptr[g_off], valid}; }
HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token(
const AttentionParams<bf16>& p, const KVContext& c, int token, int d) {
const bool valid = token >= 0;
const int safe_token = valid ? token : 0;
const int64_t gmem_off = (int64_t)c.kv_base
+ (int64_t)safe_token * p.kv_l_stride
+ (int64_t)d * p.kv_d_stride;
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
} }
}; };
@@ -175,12 +183,16 @@ struct PagedKV {
c.head_off = (int64_t)kv_head * HEAD_DIM; c.head_off = (int64_t)kv_head * HEAD_DIM;
return c; return c;
} }
HOST_DEV_FORCEINLINE KVAddr kv_addr( HOST_DEV_FORCEINLINE int resolve_token(
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) { const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
const int64_t slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0; return valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : -1;
const bool ok = valid && (slot >= 0); }
const int64_t gmem_off = slot * c.pool_stride + c.head_off + d; HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token(
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], ok}; const AttentionParams<bf16>& p, const KVContext& c, int slot, int d) {
const bool valid = slot >= 0;
const int safe_slot = valid ? slot : 0;
const int64_t gmem_off = (int64_t)safe_slot * c.pool_stride + c.head_off + d;
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
} }
}; };
+10 -10
View File
@@ -99,22 +99,22 @@ __device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7); return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
} }
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly.
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
:: "r"(smem_addr), "l"(gmem_ptr));
}
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill. // Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// src_size=0 → no bytes read from src, so out-of-bounds src address is safe. // BypassL1 defaults to .cg (L2 only); false selects .ca (L1 + L2).
// src_size=0 means no bytes are read, so an out-of-bounds address is safe.
template <bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr, __device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr, const void* gmem_ptr,
bool pred) { bool pred) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr); unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
int src_size = pred ? 16 : 0; int src_size = pred ? 16 : 0;
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;" if constexpr (BypassL1) {
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size)); asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
} else {
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
} }
__device__ __forceinline__ void cp_async_commit() { __device__ __forceinline__ void cp_async_commit() {
+2 -1
View File
@@ -90,7 +90,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
int s = i / HEAD_DIM; int s = i / HEAD_DIM;
int d_dim = i % HEAD_DIM; int d_dim = i % HEAD_DIM;
int kc = kv0 + s; int kc = kv0 + s;
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); int token = KV::resolve_token(p, kctx, kc, true);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f; sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f; sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
} }
+2 -1
View File
@@ -83,7 +83,8 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < seq_len; bool valid = kc < seq_len;
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); int token = KV::resolve_token(p, kctx, kc, valid);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid); cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid); cp_async_16_pred(&dV[off], a.v, a.valid);
+15 -12
View File
@@ -154,9 +154,9 @@ __global__ void quantize_kernel(const __nv_bfloat16* __restrict__ src,
int64_t i = blockIdx.x * (int64_t)blockDim.x + threadIdx.x; int64_t i = blockIdx.x * (int64_t)blockDim.x + threadIdx.x;
float amax = 0.f; float amax = 0.f;
if (i < n) { if (i < n) {
float v = __bfloat162float(src[i]) * *scale_inv; float raw = __bfloat162float(src[i]);
dst[i] = cast_fp8<T8>(v); dst[i] = cast_fp8<T8>(raw * *scale_inv);
amax = fabsf(v); amax = fabsf(raw);
} }
for (int off = 16; off; off >>= 1) for (int off = 16; off; off >>= 1)
amax = fmaxf(amax, __shfl_xor_sync(0xffffffffu, amax, off)); amax = fmaxf(amax, __shfl_xor_sync(0xffffffffu, amax, off));
@@ -182,9 +182,9 @@ __global__ void transpose_quantize_kernel(
float amax = 0.f; float amax = 0.f;
for (int j = 0; j < 32; j += 8) { for (int j = 0; j < 32; j += 8) {
if (x < cols && y + j < rows) { if (x < cols && y + j < rows) {
float v = __bfloat162float(src[(y + j) * cols + x]) * *scale_inv; float raw = __bfloat162float(src[(y + j) * cols + x]);
tile[threadIdx.y + j][threadIdx.x] = cast_fp8<T8>(v); tile[threadIdx.y + j][threadIdx.x] = cast_fp8<T8>(raw * *scale_inv);
amax = fmaxf(amax, fabsf(v)); amax = fmaxf(amax, fabsf(raw));
} }
} }
__syncthreads(); __syncthreads();
@@ -382,6 +382,9 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scal
auto gt8 = masks[1] ? torch::empty({n, m}, fp8_options) : torch::Tensor(); auto gt8 = masks[1] ? torch::empty({n, m}, fp8_options) : torch::Tensor();
auto wt8 = masks[0] ? torch::empty({k, n}, fp8_options) : torch::Tensor(); auto wt8 = masks[0] ? torch::empty({k, n}, fp8_options) : torch::Tensor();
auto xt8 = masks[1] ? torch::empty({k, m}, fp8_options) : torch::Tensor(); auto xt8 = masks[1] ? torch::empty({k, m}, fp8_options) : torch::Tensor();
// w/x transpose-quantize amax goes to a scratch buffer, NOT amax_g: the
// gradient scale must only see the gradient's own max-abs.
auto amax_t = torch::zeros({1}, g_c.options().dtype(torch::kFloat32));
int64_t block = 256; int64_t block = 256;
quantize_kernel<__nv_fp8_e4m3> quantize_kernel<__nv_fp8_e4m3>
@@ -394,8 +397,8 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scal
transpose_quantize_kernel<__nv_fp8_e4m3> transpose_quantize_kernel<__nv_fp8_e4m3>
<<<blocks, threads, 0, stream.stream()>>>( <<<blocks, threads, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(w_c.data_ptr()), swi_ptr, reinterpret_cast<const __nv_bfloat16*>(w_c.data_ptr()), swi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(wt8.data_ptr()), amax_g_ptr, reinterpret_cast<__nv_fp8_e4m3*>(wt8.data_ptr()),
n, k); amax_t.data_ptr<float>(), n, k);
fp8_gemm_into(g8, wt8, grad_input.reshape({m, k}), m, n, k, sg_ptr, fp8_gemm_into(g8, wt8, grad_input.reshape({m, k}), m, n, k, sg_ptr,
sw_ptr, stream.stream()); sw_ptr, stream.stream());
} }
@@ -405,13 +408,13 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scal
transpose_quantize_kernel<__nv_fp8_e4m3> transpose_quantize_kernel<__nv_fp8_e4m3>
<<<g_blocks, threads, 0, stream.stream()>>>( <<<g_blocks, threads, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()), sgi_ptr, reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()), sgi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()), amax_g_ptr, reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()),
m, n); amax_t.data_ptr<float>(), m, n);
transpose_quantize_kernel<__nv_fp8_e4m3> transpose_quantize_kernel<__nv_fp8_e4m3>
<<<x_blocks, threads, 0, stream.stream()>>>( <<<x_blocks, threads, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(x_c.data_ptr()), sxi_ptr, reinterpret_cast<const __nv_bfloat16*>(x_c.data_ptr()), sxi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(xt8.data_ptr()), amax_g_ptr, reinterpret_cast<__nv_fp8_e4m3*>(xt8.data_ptr()),
m, k); amax_t.data_ptr<float>(), m, k);
fp8_gemm_into(gt8, xt8, grad_weight, n, m, k, sg_ptr, sx_ptr, fp8_gemm_into(gt8, xt8, grad_weight, n, m, k, sg_ptr, sx_ptr,
stream.stream()); stream.stream());
} }
+19 -20
View File
@@ -7,13 +7,12 @@ __global__ void rotary_emb_kernel(
const __nv_bfloat16* __restrict__ x, const __nv_bfloat16* __restrict__ x,
const float* __restrict__ freqs_cis, const float* __restrict__ freqs_cis,
__nv_bfloat16* __restrict__ out, __nv_bfloat16* __restrict__ out,
int batch, int n_tokens,
int seq_len,
int n_heads, int n_heads,
int head_dim int head_dim
) { ) {
const int half_dim = head_dim >> 1; const int half_dim = head_dim >> 1;
const int total = batch * seq_len * n_heads * half_dim; const int total = n_tokens * n_heads * half_dim;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; for (int idx = blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx < total;
@@ -23,11 +22,10 @@ __global__ void rotary_emb_kernel(
int tmp = idx / half_dim; int tmp = idx / half_dim;
int head = tmp % n_heads; int head = tmp % n_heads;
tmp /= n_heads; tmp /= n_heads;
int seq = tmp % seq_len; int token = tmp;
int b = tmp / seq_len;
int x_offset = ((b * seq_len + seq) * n_heads + head) * head_dim + (pair << 1); int x_offset = (token * n_heads + head) * head_dim + (pair << 1);
int cs_offset = ((b * seq_len + seq) * half_dim + pair) * 2; int cs_offset = (token * half_dim + pair) * 2;
__nv_bfloat162 x_pair = *reinterpret_cast<const __nv_bfloat162*>(x + x_offset); __nv_bfloat162 x_pair = *reinterpret_cast<const __nv_bfloat162*>(x + x_offset);
float x_even = __bfloat162float(__low2bfloat16(x_pair)); float x_even = __bfloat162float(__low2bfloat16(x_pair));
@@ -54,27 +52,28 @@ torch::Tensor rotary_emb(
TORCH_CHECK(x.is_cuda(), "x must be on CUDA"); TORCH_CHECK(x.is_cuda(), "x must be on CUDA");
TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA"); TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16"); TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
TORCH_CHECK(x.dim() == 4, "x must be 4D [batch, seq_len, n_heads, head_dim]"); TORCH_CHECK(x.dim() == 3 || x.dim() == 4,
"x must be [tokens, n_heads, head_dim] or "
"[batch, seq_len, n_heads, head_dim]");
TORCH_CHECK(x.is_contiguous(), "x must be contiguous"); TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
TORCH_CHECK(freqs_cis.dim() == 4, "freqs_cis must be 4D [batch, seq_len, dim/2, 2]"); TORCH_CHECK(freqs_cis.dim() == x.dim(), "freqs_cis rank must match x rank");
TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous"); TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous");
TORCH_CHECK(freqs_cis.scalar_type() == torch::kFloat32, "freqs_cis must be f32"); TORCH_CHECK(freqs_cis.scalar_type() == torch::kFloat32, "freqs_cis must be f32");
int batch = x.size(0); int n_tokens = x.dim() == 3 ? x.size(0) : x.size(0) * x.size(1);
int seq_len = x.size(1); int n_heads = x.size(x.dim() - 2);
int n_heads = x.size(2); int head_dim = x.size(x.dim() - 1);
int head_dim = x.size(3);
TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even"); TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even");
TORCH_CHECK(freqs_cis.size(0) == batch, "freqs_cis batch mismatch"); TORCH_CHECK(freqs_cis.numel() == (int64_t)n_tokens * head_dim,
TORCH_CHECK(freqs_cis.size(1) == seq_len, "freqs_cis seq_len mismatch"); "freqs_cis token or rotary dimension mismatch");
TORCH_CHECK(freqs_cis.size(2) == head_dim / 2, "freqs_cis dim/2 mismatch"); TORCH_CHECK(freqs_cis.size(-2) == head_dim / 2, "freqs_cis dim/2 mismatch");
TORCH_CHECK(freqs_cis.size(3) == 2, "freqs_cis last dim must be 2 [cos, sin]"); TORCH_CHECK(freqs_cis.size(-1) == 2, "freqs_cis last dim must be 2 [cos, sin]");
auto out = torch::empty_like(x); auto out = torch::empty_like(x);
int half_dim = head_dim / 2; int half_dim = head_dim / 2;
int total = batch * seq_len * n_heads * half_dim; int total = n_tokens * n_heads * half_dim;
int block = 256; int block = 256;
int grid = std::min((total + block - 1) / block, 1024); int grid = std::min((total + block - 1) / block, 1024);
@@ -82,7 +81,7 @@ torch::Tensor rotary_emb(
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()), reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
freqs_cis.data_ptr<float>(), freqs_cis.data_ptr<float>(),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
batch, seq_len, n_heads, head_dim n_tokens, n_heads, head_dim
); );
C10_CUDA_CHECK(cudaGetLastError()); C10_CUDA_CHECK(cudaGetLastError());
@@ -93,6 +92,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("rotary_emb", &rotary_emb, m.def("rotary_emb", &rotary_emb,
py::arg("x"), py::arg("x"),
py::arg("freqs_cis"), py::arg("freqs_cis"),
"Fused rotary embedding (bf16 x, f32 freqs_cis [b,s,d/2,2], bf16 out)" "Fused rotary embedding for packed 3D or dense 4D tensors"
); );
} }
+39 -39
View File
@@ -18,7 +18,7 @@ struct PagedPrefillDispatch { AttentionParams<bf16>& p; template<int H> void ope
// kv_indptr: [B+1]. mask: [B, max_seq_len] bool (True=keep) or NULL. // kv_indptr: [B+1]. mask: [B, max_seq_len] bool (True=keep) or NULL.
static void cpu_paged_decode_ref( static void cpu_paged_decode_ref(
const float* Q, const float* K_pool, const float* V_pool, const float* Q, const float* K_pool, const float* V_pool,
const int64_t* req_to_token, const int64_t* req_pool_indices, const int* req_to_token, const int* req_pool_indices,
const int* kv_indptr, const bool* mask, int mask_b_stride, const int* kv_indptr, const bool* mask, int mask_b_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len, int B, int Hq, int Hkv, int D, int max_ctx_len,
float* O) float* O)
@@ -27,7 +27,7 @@ static void cpu_paged_decode_ref(
int n_rep = Hq / Hkv; int n_rep = Hq / Hkv;
for (int b = 0; b < B; b++) { for (int b = 0; b < B; b++) {
int seq_len = kv_indptr[b + 1] - kv_indptr[b]; int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int64_t req_idx = req_pool_indices[b]; int req_idx = req_pool_indices[b];
#pragma omp parallel for schedule(dynamic) #pragma omp parallel for schedule(dynamic)
for (int h = 0; h < Hq; h++) { for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep; int kv_h = h / n_rep;
@@ -35,7 +35,7 @@ static void cpu_paged_decode_ref(
float accum[256] = {0.0f}; float accum[256] = {0.0f};
for (int kj = 0; kj < seq_len; kj++) { for (int kj = 0; kj < seq_len; kj++) {
if (mask && !mask[b * mask_b_stride + kj]) continue; if (mask && !mask[b * mask_b_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj]; int slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f; float dot = 0.0f;
for (int d = 0; d < D; d++) for (int d = 0; d < D; d++)
dot += Q[(b * Hq + h) * D + d] * dot += Q[(b * Hq + h) * D + d] *
@@ -66,7 +66,7 @@ static void cpu_paged_decode_ref(
// attention mask on top of the (unused) causal logic. // attention mask on top of the (unused) causal logic.
static void cpu_paged_prefill_ref( static void cpu_paged_prefill_ref(
const float* Q, const float* K_pool, const float* V_pool, const float* Q, const float* K_pool, const float* V_pool,
const int64_t* req_to_token, const int64_t* req_pool_indices, const int* req_to_token, const int* req_pool_indices,
const int* kv_indptr, const int* qo_indptr, const int* kv_indptr, const int* qo_indptr,
const bool* mask, int mask_l_stride, int mask_kv_stride, const bool* mask, int mask_l_stride, int mask_kv_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal, int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
@@ -78,7 +78,7 @@ static void cpu_paged_prefill_ref(
int seq_len = kv_indptr[b + 1] - kv_indptr[b]; int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int q_len = qo_indptr[b + 1] - qo_indptr[b]; int q_len = qo_indptr[b + 1] - qo_indptr[b];
int causal_off = seq_len - q_len; int causal_off = seq_len - q_len;
int64_t req_idx = req_pool_indices[b]; int req_idx = req_pool_indices[b];
#pragma omp parallel for collapse(2) schedule(dynamic) #pragma omp parallel for collapse(2) schedule(dynamic)
for (int h = 0; h < Hq; h++) { for (int h = 0; h < Hq; h++) {
for (int qi = 0; qi < q_len; qi++) { for (int qi = 0; qi < q_len; qi++) {
@@ -89,7 +89,7 @@ static void cpu_paged_prefill_ref(
for (int kj = 0; kj < lim; kj++) { for (int kj = 0; kj < lim; kj++) {
if (mask && !mask[b * mask_l_stride * mask_kv_stride if (mask && !mask[b * mask_l_stride * mask_kv_stride
+ qi * mask_kv_stride + kj]) continue; + qi * mask_kv_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj]; int slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f; float dot = 0.0f;
for (int d = 0; d < D; d++) for (int d = 0; d < D; d++)
dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] * dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] *
@@ -149,14 +149,14 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi; int *d_kvi;
float *d_op, *d_ml; float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
@@ -181,7 +181,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
// req_to_token: assign unique slots per request (scattered, not contiguous) // req_to_token: assign unique slots per request (scattered, not contiguous)
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -191,7 +191,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
// req_pool_indices: pick B random request rows // req_pool_indices: pick B random request rows
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -278,15 +278,15 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_mask = (size_t)B * max_sl * sizeof(bool); size_t sz_mask = (size_t)B * max_sl * sizeof(bool);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi; int *d_kvi;
bool *d_mask; bool *d_mask;
float *d_op, *d_ml; float *d_op, *d_ml;
@@ -312,7 +312,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -321,7 +321,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
} }
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -417,13 +417,13 @@ static int run_prefill_test(int B, int Hq, int Hkv,
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int); size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi, *d_qoi; int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv); cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
@@ -446,7 +446,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -455,7 +455,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
} }
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -546,14 +546,14 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int); size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
size_t sz_mask = (size_t)B * q_len * q_len * sizeof(bool); size_t sz_mask = (size_t)B * q_len * q_len * sizeof(bool);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi, *d_qoi; int *d_kvi, *d_qoi;
bool *d_mask; bool *d_mask;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
@@ -577,7 +577,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -586,7 +586,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
} }
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
h_rpi[0] = 0; h_rpi[0] = 0;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -667,20 +667,20 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
// ====================================================================== // ======================================================================
template <int HEAD_DIM> template <int HEAD_DIM>
static void bench_decode(int B, int Hq, int Hkv, int seq_len) { static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
int max_ctx = seq_len + 16; int max_ctx = max(16384, seq_len + 16);
int pool_size = B * max_ctx; int pool_size = B * (seq_len + 16);
int num_reqs = B; int num_reqs = B;
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi; int *d_kvi;
float *d_op, *d_ml; float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
@@ -696,12 +696,12 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size; h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi); int* h_kvi = (int*)malloc(sz_kvi);
@@ -749,13 +749,13 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int); size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi, *d_qoi; int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv); cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
@@ -769,12 +769,12 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size; h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi); int* h_kvi = (int*)malloc(sz_kvi);
@@ -933,9 +933,9 @@ int main() {
bench_decode<128>(1, 32, 4, 1024); bench_decode<128>(1, 32, 4, 1024);
bench_decode<128>(1, 32, 4, 2048); bench_decode<128>(1, 32, 4, 2048);
bench_decode<128>(1, 32, 4, 4096); bench_decode<128>(1, 32, 4, 4096);
bench_decode<128>(1, 32, 4, 16384);
bench_decode<128>(4, 32, 4, 2048); bench_decode<128>(4, 32, 4, 2048);
bench_decode<128>(16, 32, 4, 2048); bench_decode<128>(16, 32, 4, 2048);
bench_decode<128>(32, 32, 4, 1024);
printf("\n===== PAGED PREFILL BENCH =====\n"); printf("\n===== PAGED PREFILL BENCH =====\n");
print_bench_header(); print_bench_header();
+10 -1
View File
@@ -118,7 +118,8 @@ static void bench_decode() {
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS); printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
print_bench_header(); print_bench_header();
for (int ci = 0; ci < 6; ci++) { int n = sizeof(cfgs) / sizeof(cfgs[0]);
for (int ci = 0; ci < n; ci++) {
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2]; int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
int sl = cfgs[ci][3], D = cfgs[ci][4]; int sl = cfgs[ci][3], D = cfgs[ci][4];
size_t nQ = (size_t)B * Hq * D; size_t nQ = (size_t)B * Hq * D;
@@ -229,6 +230,12 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
static void bench_prefill() { static void bench_prefill() {
const int cfgs[][7] = { const int cfgs[][7] = {
{1,32,4,1024,1024,32,0},
{1,32,4,1024,1024,32,1},
{1,32,4,4096,4096,32,1},
{1,32,4,1024,1024,64,0},
{1,32,4,1024,1024,64,1},
{1,32,4,4096,4096,64,1},
{1,32,4,512,512,128,0}, {1,32,4,512,512,128,0},
{1,32,4,1024,1024,128,0}, {1,32,4,1024,1024,128,0},
{1,32,4,2048,2048,128,0}, {1,32,4,2048,2048,128,0},
@@ -324,7 +331,9 @@ int main() {
{ {
const int configs[][7] = { const int configs[][7] = {
{1,2,1,64,128,32,0}, // scalar fallback D=32 {1,2,1,64,128,32,0}, // scalar fallback D=32
{1,4,2,256,256,32,1}, // causal D=32 dispatch
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal {1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,4,2,256,256,64,1}, // causal D=64 dispatch
{1,32,4,512,512,128,0}, // standard {1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium {1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal {1,4,2,256,256,128,1}, // causal
+1 -1
View File
@@ -1456,7 +1456,7 @@ classDiagram
| **Context** | `TrainContext` | Unified training state bag | | **Context** | `TrainContext` | Unified training state bag |
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction | | **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
| **Strategy (Attention)** | `AttentionBackend`, `CudaBackend`, `FlashAttnBackend`, `TorchNativeBackend` | Attention computation backend switching via context manager | | **Strategy (Attention)** | `AttentionBackend`, `CudaBackend`, `FlashAttnBackend`, `TorchNativeBackend` | Attention computation backend switching via context manager |
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback | | **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `backend/rotary.py`, `ops/rotary.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution | | **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support | | **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
+5 -5
View File
@@ -30,7 +30,7 @@ The `rotary_emb` kernel (`csrc/kernels/rotary_emb.cu`) fuses cos/sin lookup and
- One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store - One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store
- f32 cos/sin input, bf16 compute and output - f32 cos/sin input, bf16 compute and output
- 256-thread blocks, grid-stride loop - 256-thread blocks, grid-stride loop
- Auto-dispatched via `apply_rotary_emb` in `astrai/extension/rotary_backend.py` (CUDA when available + inference mode, else torch complex-multiply fallback) - Auto-dispatched via `apply_rotary_emb` in `astrai/extension/backend/rotary.py` (CUDA when available + inference mode, else torch complex-multiply fallback)
- No context-manager backend needed — rotary is backend-agnostic, both attention backends benefit - No context-manager backend needed — rotary is backend-agnostic, both attention backends benefit
Standalone benchmark vs torch complex-multiply (48 calls = 24 layers × q+k): 6-9x faster, max diff 0 (decode) to 3e-2 (large prefill, bf16). Standalone benchmark vs torch complex-multiply (48 calls = 24 layers × q+k): 6-9x faster, max diff 0 (decode) to 3e-2 (large prefill, bf16).
@@ -83,7 +83,7 @@ Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 mod
## Attention Backend ## Attention Backend
`astrai/extension/attention_backend.py` provides the backend abstraction: `astrai/extension/backend/attention.py` provides the backend abstraction:
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len - **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU. - **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU.
@@ -106,7 +106,7 @@ with attn_backend(ATTN_BACKEND.CUDA):
### Rotary Backend ### Rotary Backend
`astrai/extension/rotary_backend.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch: `astrai/extension/backend/rotary.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch:
- **CUDA path**: calls `rotary_emb` kernel directly when available, input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference) - **CUDA path**: calls `rotary_emb` kernel directly when available, input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference)
- **Torch fallback**: complex multiply (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable - **Torch fallback**: complex multiply (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable
@@ -115,9 +115,9 @@ No context-manager switching needed — the dispatch is automatic per call.
## Python Wrappers ## Python Wrappers
`astrai/extension/attention_ops.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions. `astrai/extension/ops/attention.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions.
`astrai/extension/rotary_ops.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `rotary_backend.py`. `astrai/extension/ops/rotary.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `backend/rotary.py`.
Interface (all functions): Interface (all functions):
``` ```
+2 -2
View File
@@ -176,14 +176,14 @@ Three-layer separation (SGLang-inspired):
### Attention Backend ### Attention Backend
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`): Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/backend/attention.py`):
- **`CudaBackend`** (default): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). Falls back to `FlashAttnBackend` when dtype unsupported. - **`CudaBackend`** (default): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). Falls back to `FlashAttnBackend` when dtype unsupported.
- **`FlashAttnBackend`**: optional flash-attn dispatch with `flash_attn_with_kvcache` fast path for contiguous cache; falls back to KV gather + `flash_attn_func`. - **`FlashAttnBackend`**: optional flash-attn dispatch with `flash_attn_with_kvcache` fast path for contiguous cache; falls back to KV gather + `flash_attn_func`.
- **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`. - **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
- Default priority: cuda > flash > torch. Set `ASTR_BACKEND=cuda|torch_native|flash` to override. - Default priority: cuda > flash > torch. Set `ASTR_BACKEND=cuda|torch_native|flash` to override.
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch. Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`: Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`:
+1 -1
View File
@@ -90,7 +90,7 @@ Fallback: when `CudaBackend` cannot handle an input (wrong dtype or head_dim), `
### Rotary Embedding Backend ### Rotary Embedding Backend
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches: Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches:
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, the input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode) - **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, the input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
- **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available - **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
+21 -40
View File
@@ -118,16 +118,11 @@ class GenerationBenchmark:
workspace: InferenceWorkspace, workspace: InferenceWorkspace,
) -> list: ) -> list:
input_ids = torch.randint( input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, prompt_len), device=self.device 0, self.config.vocab_size, (batch_size * prompt_len,), device=self.device
)
position_ids = (
torch.arange(0, prompt_len, dtype=torch.long, device=self.device)
.unsqueeze(0)
.expand(batch_size, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_len, device=self.device
) )
position_ids = torch.arange(
prompt_len, dtype=torch.long, device=self.device
).repeat(batch_size)
task_ids = [f"bench_{i}" for i in range(batch_size)] task_ids = [f"bench_{i}" for i in range(batch_size)]
for tid in task_ids: for tid in task_ids:
@@ -137,9 +132,9 @@ class GenerationBenchmark:
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
torch.cuda.synchronize() torch.cuda.synchronize()
return task_ids return task_ids
@@ -154,24 +149,20 @@ class GenerationBenchmark:
): ):
batch_size = len(task_ids) batch_size = len(task_ids)
input_ids = torch.randint( input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, 1), device=self.device 0, self.config.vocab_size, (batch_size,), device=self.device
) )
position_ids = torch.tensor( position_ids = torch.tensor(
[[seq_len] for _ in range(batch_size)], dtype=torch.long, device=self.device [seq_len] * batch_size, dtype=torch.long, device=self.device
) )
total_len = seq_len + 1
for tid in task_ids: for tid in task_ids:
task_cache.task_extend(tid, seq_len) task_cache.task_extend(tid, seq_len)
input_mask = position_ids[:, :, None] >= torch.arange(
total_len, device=self.device
)
kv_cache = task_cache.bind(task_ids, workspace, self.device) kv_cache = task_cache.bind(task_ids, workspace, self.device)
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="decode",
) )
def run_prefill_benchmark( def run_prefill_benchmark(
@@ -188,25 +179,23 @@ class GenerationBenchmark:
task_cache.task_alloc(tid, list(range(prompt_length))) task_cache.task_alloc(tid, list(range(prompt_length)))
input_ids = torch.randint( input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, prompt_length), device=self.device 0,
) self.config.vocab_size,
position_ids = ( (batch_size * prompt_length,),
torch.arange(0, prompt_length, dtype=torch.long, device=self.device) device=self.device,
.unsqueeze(0)
.expand(batch_size, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_length, device=self.device
) )
position_ids = torch.arange(
prompt_length, dtype=torch.long, device=self.device
).repeat(batch_size)
kv_cache = task_cache.bind(task_ids, workspace, self.device, start_pos=0) kv_cache = task_cache.bind(task_ids, workspace, self.device, start_pos=0)
for _ in range(3): for _ in range(3):
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -215,9 +204,9 @@ class GenerationBenchmark:
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
torch.cuda.synchronize() torch.cuda.synchronize()
elapsed = time.perf_counter() - t0 elapsed = time.perf_counter() - t0
@@ -311,37 +300,29 @@ class GenerationBenchmark:
) )
b = batch_size b = batch_size
input_ids_buf = torch.zeros(b, 1, dtype=torch.long, device=self.device) input_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device)
position_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device) position_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device)
arange = torch.arange(max_seq_len, device=self.device)
gctx = CudaGraphContext(enabled=True) gctx = CudaGraphContext(enabled=True)
graph_key = (b,) graph_key = (b,)
def _decode_graph_step(seq_len): def _decode_graph_step(seq_len):
input_ids_buf.copy_( input_ids_buf.copy_(
torch.randint(0, self.config.vocab_size, (b, 1), device=self.device) torch.randint(0, self.config.vocab_size, (b,), device=self.device)
) )
position_ids_buf[:] = seq_len position_ids_buf[:] = seq_len
for tid in task_ids: for tid in task_ids:
task_cache.task_extend(tid, seq_len) task_cache.task_extend(tid, seq_len)
kv_cache = task_cache.bind(task_ids, workspace, self.device) kv_cache = task_cache.bind(task_ids, workspace, self.device)
input_mask = torch.ge(
position_ids_buf[:, None],
arange,
out=workspace.input_mask[:b, 0, :max_seq_len],
)
input_mask = input_mask.unsqueeze(1)
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
return gctx.forward( return gctx.forward(
self.model, self.model,
key=graph_key, key=graph_key,
input_ids=input_ids_buf, input_ids=input_ids_buf,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids_buf.unsqueeze(1), position_ids=position_ids_buf,
fwd="decode",
) )
for i in range(5): for i in range(5):
+2 -7
View File
@@ -10,6 +10,8 @@ from astrai.extension import (
ATTN_BACKEND, ATTN_BACKEND,
AttentionBackendFactory, AttentionBackendFactory,
CudaBackend, CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
attn_backend, attn_backend,
get_backend, get_backend,
) )
@@ -17,13 +19,6 @@ from astrai.extension import (
def test_default_backend_resolves_to_available(): def test_default_backend_resolves_to_available():
"""Default backend is the first available in cuda > flash > torch order.""" """Default backend is the first available in cuda > flash > torch order."""
from astrai.extension.attention_backend import (
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
_resolve_default_backend,
)
backend = get_backend() backend = get_backend()
assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend)) assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend))
+77 -37
View File
@@ -8,8 +8,11 @@ import torch
from astrai.extension import ATTN_BACKEND, attn_backend from astrai.extension import ATTN_BACKEND, attn_backend
from astrai.inference.cache import PagePool, TaskCacheManager from astrai.inference.cache import PagePool, TaskCacheManager
from astrai.inference.runtime.graph import CudaGraphContext
from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.workspace import InferenceWorkspace from astrai.inference.workspace import InferenceWorkspace
from tests.extension.conftest import D, skip_no_kernel from tests.extension.conftest import D, skip_no_kernel
from tests.helpers import FakeTokenizer
def _mk_task_cache(pool: PagePool) -> TaskCacheManager: def _mk_task_cache(pool: PagePool) -> TaskCacheManager:
@@ -56,17 +59,9 @@ def test_prefill_with_kv_cache_matches_torch(cuda_model):
"""Inference prefill with KV cache should match torch backend.""" """Inference prefill with KV cache should match torch backend."""
model, _ = cuda_model model, _ = cuda_model
prompt_ids = [[1, 2, 3, 4, 5, 6, 7, 8], [10, 11, 12, 13, 14, 15]] prompt_ids = [[1, 2, 3, 4, 5, 6, 7, 8], [10, 11, 12, 13, 14, 15]]
max_len = max(len(p) for p in prompt_ids)
batch = len(prompt_ids)
device = "cuda" device = "cuda"
input_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device) input_ids = torch.tensor(sum(prompt_ids, []), dtype=torch.long, device=device)
input_mask = torch.zeros(batch, max_len, dtype=torch.bool, device=device) position_ids = torch.cat([torch.arange(len(p), device=device) for p in prompt_ids])
position_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
for i, p in enumerate(prompt_ids):
input_ids[i, : len(p)] = torch.tensor(p, device=device)
input_mask[i, : len(p)] = True
position_ids[i, : len(p)] = torch.arange(len(p), device=device)
cache = PagePool( cache = PagePool(
n_layers=2, n_layers=2,
@@ -85,7 +80,7 @@ def test_prefill_with_kv_cache_matches_torch(cuda_model):
kv1 = task_cache.bind(["t1", "t2"], ws, start_pos=0) kv1 = task_cache.bind(["t1", "t2"], ws, start_pos=0)
with torch.inference_mode(): with torch.inference_mode():
out_torch = model( out_torch = model(
input_ids, input_mask=input_mask, kv_cache=kv1, position_ids=position_ids input_ids, kv_cache=kv1, position_ids=position_ids, fwd="prefill"
) )
task_cache.task_free("t1") task_cache.task_free("t1")
@@ -97,22 +92,24 @@ def test_prefill_with_kv_cache_matches_torch(cuda_model):
with torch.inference_mode(): with torch.inference_mode():
out_cuda = model( out_cuda = model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv2, kv_cache=kv2,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
offset = 0
for i, p in enumerate(prompt_ids): for i, p in enumerate(prompt_ids):
d = ( d = (
( (
out_torch["logits"][i, : len(p)].float() out_torch["logits"][offset : offset + len(p)].float()
- out_cuda["logits"][i, : len(p)].float() - out_cuda["logits"][offset : offset + len(p)].float()
) )
.abs() .abs()
.max() .max()
.item() .item()
) )
assert d == 0.0, f"Prefill diff for sample {i}: {d}" assert d == 0.0, f"Prefill diff for sample {i}: {d}"
offset += len(p)
@skip_no_kernel @skip_no_kernel
@@ -133,15 +130,8 @@ def test_decode_mixed_seq_lens_matches_torch(cuda_model):
) )
# Prefill to populate cache # Prefill to populate cache
max_len = max(len(p) for p in prompt_ids) input_ids = torch.tensor(sum(prompt_ids, []), dtype=torch.long, device=device)
batch = len(prompt_ids) position_ids = torch.cat([torch.arange(len(p), device=device) for p in prompt_ids])
input_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
input_mask = torch.zeros(batch, max_len, dtype=torch.bool, device=device)
position_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
for i, p in enumerate(prompt_ids):
input_ids[i, : len(p)] = torch.tensor(p, device=device)
input_mask[i, : len(p)] = True
position_ids[i, : len(p)] = torch.arange(len(p), device=device)
task_cache = _mk_task_cache(cache) task_cache = _mk_task_cache(cache)
ws = _ws(cache) ws = _ws(cache)
@@ -149,39 +139,89 @@ def test_decode_mixed_seq_lens_matches_torch(cuda_model):
task_cache.task_alloc("t2", prompt_ids[1]) task_cache.task_alloc("t2", prompt_ids[1])
kv = task_cache.bind(["t1", "t2"], ws, start_pos=0) kv = task_cache.bind(["t1", "t2"], ws, start_pos=0)
with torch.inference_mode(): with torch.inference_mode():
model(input_ids, input_mask=input_mask, kv_cache=kv, position_ids=position_ids) model(input_ids, kv_cache=kv, position_ids=position_ids, fwd="prefill")
# Decode step — seq_lens are 9 and 7 (after extending) # Decode step — seq_lens are 9 and 7 (after extending)
dec_ids = torch.tensor([[99], [98]], dtype=torch.long, device=device) dec_ids = torch.tensor([99, 98], dtype=torch.long, device=device)
dec_pos = torch.tensor([[8], [6]], dtype=torch.long, device=device) dec_pos = torch.tensor([8, 6], dtype=torch.long, device=device)
total_len = 9
dec_mask = dec_pos[:, None, None] >= torch.arange(total_len, device=device)
task_cache.task_extend("t1", 8) task_cache.task_extend("t1", 8)
task_cache.task_extend("t2", 6) task_cache.task_extend("t2", 6)
kv_t = task_cache.bind(["t1", "t2"], ws) kv_t = task_cache.bind(["t1", "t2"], ws)
with torch.inference_mode(): with torch.inference_mode():
out_torch = model( out_torch = model(dec_ids, kv_cache=kv_t, position_ids=dec_pos, fwd="decode")
dec_ids, input_mask=dec_mask, kv_cache=kv_t, position_ids=dec_pos
)
kv_c = task_cache.bind(["t1", "t2"], ws) kv_c = task_cache.bind(["t1", "t2"], ws)
with attn_backend(ATTN_BACKEND.CUDA): with attn_backend(ATTN_BACKEND.CUDA):
with torch.inference_mode(): with torch.inference_mode():
out_cuda = model( out_cuda = model(dec_ids, kv_cache=kv_c, position_ids=dec_pos, fwd="decode")
dec_ids, input_mask=dec_mask, kv_cache=kv_c, position_ids=dec_pos
)
diff = (out_torch["logits"].float() - out_cuda["logits"].float()).abs().max().item() diff = (out_torch["logits"].float() - out_cuda["logits"].float()).abs().max().item()
assert diff < 0.05, f"Decode diff (mixed seq_lens): {diff}" assert diff < 0.05, f"Decode diff (mixed seq_lens): {diff}"
@skip_no_kernel
def test_decode_cuda_graph_replay_is_exact(cuda_model):
"""INT32 cache indices must remain graph-capturable and replay exactly."""
model, _ = cuda_model
device = "cuda"
prompt_ids = [1, 2, 3, 4, 5, 6, 7, 8]
cache = PagePool(
n_layers=2,
n_kv_heads=1,
head_dim=D,
max_batch_size=1,
max_seq_len=64,
device=device,
dtype=torch.bfloat16,
)
task_cache = _mk_task_cache(cache)
ws = _ws(cache)
task_cache.task_alloc("t1", prompt_ids)
input_ids = torch.tensor(prompt_ids, dtype=torch.long, device=device)
position_ids = torch.arange(len(prompt_ids), device=device)
with attn_backend(ATTN_BACKEND.CUDA), torch.inference_mode():
model(
input_ids,
position_ids=position_ids,
kv_cache=task_cache.bind(["t1"], ws, start_pos=0),
fwd="prefill",
)
task_cache.task_extend("t1", len(prompt_ids))
kv_cache = task_cache.bind(["t1"], ws)
assert kv_cache.req_to_token.dtype == torch.int32
assert kv_cache.req_pool_indices.dtype == torch.int32
assert kv_cache.out_cache_loc.dtype == torch.int32
decode_args = {
"input_ids": torch.tensor([9], dtype=torch.long, device=device),
"position_ids": torch.tensor([len(prompt_ids)], device=device),
"kv_cache": kv_cache,
"fwd": "decode",
}
graph = CudaGraphContext(enabled=True)
graph.forward(model, key=(1,), **decode_args)
graph.forward(model, key=(1,), **decode_args)
first = graph.forward(model, key=(1,), **decode_args)["logits"].clone()
slot = kv_cache.out_cache_loc[0]
first_k = kv_cache.k_buffer[:, slot].clone()
first_v = kv_cache.v_buffer[:, slot].clone()
second = graph.forward(model, key=(1,), **decode_args)["logits"].clone()
torch.cuda.synchronize()
assert graph.has_graph((1,))
torch.testing.assert_close(second, first, rtol=0, atol=0)
torch.testing.assert_close(kv_cache.k_buffer[:, slot], first_k, rtol=0, atol=0)
torch.testing.assert_close(kv_cache.v_buffer[:, slot], first_v, rtol=0, atol=0)
@skip_no_kernel @skip_no_kernel
def test_run_batch_cuda_matches_torch_greedy(cuda_model): def test_run_batch_cuda_matches_torch_greedy(cuda_model):
"""Greedy decode (temperature=0) should produce identical tokens.""" """Greedy decode (temperature=0) should produce identical tokens."""
from astrai.inference.scheduler import InferenceScheduler
from tests.helpers import FakeTokenizer
model, _ = cuda_model model, _ = cuda_model
tokenizer = FakeTokenizer() tokenizer = FakeTokenizer()
+1 -8
View File
@@ -2,14 +2,13 @@
import torch import torch
from astrai.extension.ops.attention import attn_prefill
from tests.extension.conftest import D, skip_no_kernel from tests.extension.conftest import D, skip_no_kernel
@skip_no_kernel @skip_no_kernel
def test_kernel_accepts_2d_mask(): def test_kernel_accepts_2d_mask():
"""Kernel should accept 2D mask [batch, kv_len].""" """Kernel should accept 2D mask [batch, kv_len]."""
from astrai.extension.attention_ops import attn_prefill
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
kv_len = 8 kv_len = 8
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
@@ -25,8 +24,6 @@ def test_kernel_accepts_2d_mask():
@skip_no_kernel @skip_no_kernel
def test_kernel_accepts_3d_mask(): def test_kernel_accepts_3d_mask():
"""Kernel should accept 3D mask [batch, q_len, kv_len].""" """Kernel should accept 3D mask [batch, q_len, kv_len]."""
from astrai.extension.attention_ops import attn_prefill
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
kv_len = 8 kv_len = 8
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
@@ -41,8 +38,6 @@ def test_kernel_accepts_3d_mask():
@skip_no_kernel @skip_no_kernel
def test_kernel_accepts_4d_mask(): def test_kernel_accepts_4d_mask():
"""Kernel should accept 4D mask [batch, n_heads, q_len, kv_len].""" """Kernel should accept 4D mask [batch, n_heads, q_len, kv_len]."""
from astrai.extension.attention_ops import attn_prefill
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
kv_len = 8 kv_len = 8
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
@@ -58,8 +53,6 @@ def test_kernel_accepts_4d_mask():
@skip_no_kernel @skip_no_kernel
def test_4d_mask_matches_no_mask_when_all_true(): def test_4d_mask_matches_no_mask_when_all_true():
"""A 4D all-True mask should produce the same output as no mask.""" """A 4D all-True mask should produce the same output as no mask."""
from astrai.extension.attention_ops import attn_prefill
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
kv_len = 8 kv_len = 8
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
+5 -2
View File
@@ -176,6 +176,7 @@ def test_page_pool_task_cacheable_ids_excludes_unmaterialized_tail():
def test_req_to_token_pool_alloc_free(): def test_req_to_token_pool_alloc_free():
pool = ReqToTokenPool(4, 128, torch.device("cpu")) pool = ReqToTokenPool(4, 128, torch.device("cpu"))
assert pool.req_to_token.dtype == torch.int32
slots = pool.alloc(2) slots = pool.alloc(2)
assert len(slots) == 2 assert len(slots) == 2
assert len(pool.free_slots) == 2 assert len(pool.free_slots) == 2
@@ -278,9 +279,11 @@ def test_page_pool_contiguous_bind_tasks_prefill():
task_cache.task_alloc("t1", list(range(10))) task_cache.task_alloc("t1", list(range(10)))
task_cache.task_alloc("t2", list(range(10))) task_cache.task_alloc("t2", list(range(10)))
kv = task_cache.bind(["t1", "t2"], _ws(pool), start_pos=0) kv = task_cache.bind(["t1", "t2"], _ws(pool), start_pos=0)
assert kv.out_cache_loc.shape == (2, 10) assert kv.out_cache_loc.shape == (20,)
assert kv.out_cache_loc.dtype == torch.int32
assert kv.seq_lens.tolist() == [10, 10] assert kv.seq_lens.tolist() == [10, 10]
assert kv.req_pool_indices.shape == (2,) assert kv.req_pool_indices.shape == (2,)
assert kv.req_pool_indices.dtype == torch.int32
def test_page_pool_contiguous_bind_tasks_decode(): def test_page_pool_contiguous_bind_tasks_decode():
@@ -292,7 +295,7 @@ def test_page_pool_contiguous_bind_tasks_decode():
assert task_cache.task_extend("t1", 10) assert task_cache.task_extend("t1", 10)
assert task_cache.task_extend("t2", 8) assert task_cache.task_extend("t2", 8)
kv = task_cache.bind(["t1", "t2"], _ws(pool)) kv = task_cache.bind(["t1", "t2"], _ws(pool))
assert kv.out_cache_loc.shape == (2, 1) assert kv.out_cache_loc.shape == (2,)
assert kv.seq_lens.tolist() == [11, 9] assert kv.seq_lens.tolist() == [11, 9]
+46
View File
@@ -39,6 +39,52 @@ def _make_model(config=None) -> AutoRegressiveLM:
return AutoRegressiveLM(config) return AutoRegressiveLM(config)
def test_model_forward_contract_uses_dense_training_and_packed_inference():
from astrai.inference.cache import PagePool, TaskCacheManager
from astrai.inference.workspace import InferenceWorkspace
config = AutoRegressiveLMConfig(**TINY_CONFIG)
model = AutoRegressiveLM(config).eval()
dense = model(torch.tensor([[1, 2, 3]]))
assert dense["logits"].shape == (1, 3, config.vocab_size)
pool = PagePool(
n_layers=config.num_hidden_layers,
n_kv_heads=config.num_key_value_heads,
head_dim=config.hidden_size // config.num_attention_heads,
max_batch_size=1,
max_seq_len=config.max_position_embeddings,
device="cpu",
dtype=torch.float32,
)
cache = TaskCacheManager(pool)
workspace = InferenceWorkspace(
1,
config.max_position_embeddings,
config.num_attention_heads,
config.hidden_size // config.num_attention_heads,
torch.device("cpu"),
torch.float32,
)
assert cache.task_alloc("t", [1, 2, 3])
packed = model(
torch.tensor([1, 2, 3]),
position_ids=torch.arange(3),
kv_cache=cache.bind(["t"], workspace, start_pos=0),
fwd="prefill",
)
assert packed["logits"].shape == (3, config.vocab_size)
with pytest.raises(ValueError, match="training input_ids"):
model(torch.tensor([1, 2, 3]))
with pytest.raises(ValueError, match="inference input_ids"):
model(
torch.tensor([[1, 2, 3]]),
kv_cache=cache.bind(["t"], workspace, start_pos=0),
fwd="prefill",
)
def _router_stats(probs, topk_indices): def _router_stats(probs, topk_indices):
return {"probs": probs, "topk_indices": topk_indices} return {"probs": probs, "topk_indices": topk_indices}