refactor: remove bf16 gemm and swiglu kernels and rebuild csrc benchmarks

- delete csrc/kernels/gemm.cu and swiglu.cu and drop their CMake and setup.py registration
- remove the ops wrappers plus backend/linear.py and backend/swiglu.py so Linear and MLP call F.linear directly
- drop the four gemm and swiglu kernel test files and prune the stale cuda_kernels.md sections
- add csrc/bench benchmarks for the remaining kernels: attention decode prefill paged decode paged prefill versus single-launch SDPA references, rotary versus the torch fallback, fp8 quantize and mm_fp8 versus torch baselines
- attention, rotary_emb, and fp8_ops kernels are unchanged
This commit is contained in:
2026-09-05 01:38:10 +08:00
parent a77e35dd51
commit 6709534d64
24 changed files with 1306 additions and 3394 deletions
+2 -3
View File
@@ -1,9 +1,8 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from astrai.extension.backend.linear import linear
class Linear(nn.Module):
def __init__(
@@ -22,4 +21,4 @@ class Linear(nn.Module):
nn.init.uniform_(self.bias, -bound, bound)
def forward(self, x: Tensor) -> Tensor:
return linear(x, self.weight, self.bias)
return F.linear(x, self.weight, self.bias)
+1 -2
View File
@@ -5,7 +5,6 @@ import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from astrai.extension.backend.swiglu import swiglu
from astrai.factory import BaseFactory
from astrai.model.components.linear import Linear
@@ -39,7 +38,7 @@ class MLP(nn.Module):
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
def forward(self, x: Tensor) -> FFNOutput:
gated = swiglu(x, self.up.weight, self.gate.weight)
gated = self.up(x) * F.silu(self.gate(x))
out = self.down(gated)
return {"hidden_states": out, "aux_loss": None, "router_stats": None}