submission 418609
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1994 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-418609?include=source"interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:a26f0794f0f56fa06af6bf00f528fd05986afc6e073fb1a3e6d9233c74242574
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)Kernel source
submission.py1994 lines
from __future__ import annotations
from typing import Any, Dict, Tuple
import torch
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
from cutlass.cutlass_dsl import for_generate, if_generate, yield_out, range_constexpr
class _LayerNormLastDimF32ToF16:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
self.warps = self.threads // 32
@cute.jit
def __call__(
self,
x_ptr: "cute.Pointer",
w_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
y_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, d = problem
stride_bs = n * n * d
stride_i = n * d
stride_j = d
x = cute.make_tensor(
x_ptr,
cute.make_layout((bs, n, n, d), stride=(stride_bs, stride_i, stride_j, 1)),
)
w = cute.make_tensor(w_ptr, cute.make_layout((d,), stride=(1,)))
b = cute.make_tensor(b_ptr, cute.make_layout((d,), stride=(1,)))
y = cute.make_tensor(
y_ptr,
cute.make_layout((bs, n, n, d), stride=(stride_bs, stride_i, stride_j, 1)),
)
total = bs * n * n
grid = (total + self.warps - 1) // self.warps
self.kernel(x, w, b, y, bs, n, d).launch(
grid=[grid, 1, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
x: "cute.Tensor",
w: "cute.Tensor",
b: "cute.Tensor",
y: "cute.Tensor",
bs: int,
n: int,
d: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, _, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
idx = bx * self.warps + warp_id
def _do_one():
j = idx % n
t0 = idx // n
i = t0 % n
bb = t0 // n
sum0_init = cutlass.Float32(0.0)
sumsq0_init = cutlass.Float32(0.0)
for dd, acc, acc_out in for_generate(lane, d, 32, iter_args=[sum0_init, sumsq0_init]):
sum0_it = acc[0]
sumsq0_it = acc[1]
v = x[bb, i, j, dd]
sum0_it = sum0_it + v
sumsq0_it = sumsq0_it + v * v
yield_out([sum0_it, sumsq0_it])
sum0 = cute.arch.warp_reduction_sum(acc_out[0])
sumsq0 = cute.arch.warp_reduction_sum(acc_out[1])
inv_d = cutlass.Float32(1.0) / cutlass.Float32(d)
mean = sum0 * inv_d
var = sumsq0 * inv_d - mean * mean
inv_std = cute.rsqrt(var + cutlass.Float32(1e-5), fastmath=True)
for dd in for_generate(lane, d, 32):
v = x[bb, i, j, dd]
nrm = (v - mean) * inv_std
out = nrm * w[dd] + b[dd]
y[bb, i, j, dd] = out.to(cutlass.Float16)
yield_out()
if_generate(idx < bs * n * n, _do_one)
class _GemmF16F16ToF16:
def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
self.threads = int(threads)
self.tile_m = int(tile_m)
self.tile_n = int(tile_n)
self.tile_k = int(tile_k)
self.warps = self.threads // 32
self.rows_per_warp = self.tile_m // self.warps
@cute.jit
def __call__(
self,
a_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
c_ptr: "cute.Pointer",
problem: tuple,
):
m, n, k = problem
a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))
b = cute.make_tensor(b_ptr, cute.make_layout((n, k), stride=(k, 1)))
c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))
grid_n = (n + self.tile_n - 1) // self.tile_n
grid_m = (m + self.tile_m - 1) // self.tile_m
self.kernel(a, b, c, m, n, k).launch(
grid=[grid_n, grid_m, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
a: "cute.Tensor",
b: "cute.Tensor",
c: "cute.Tensor",
m: int,
n: int,
k: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
base_m = by * self.tile_m
base_n = bx * self.tile_n
row0 = base_m + warp_id * self.rows_per_warp
col0 = base_n + lane
col1 = col0 + 32
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc02 = cutlass.Float32(0.0)
acc03 = cutlass.Float32(0.0)
acc04 = cutlass.Float32(0.0)
acc05 = cutlass.Float32(0.0)
acc06 = cutlass.Float32(0.0)
acc07 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
acc12 = cutlass.Float32(0.0)
acc13 = cutlass.Float32(0.0)
acc14 = cutlass.Float32(0.0)
acc15 = cutlass.Float32(0.0)
acc16 = cutlass.Float32(0.0)
acc17 = cutlass.Float32(0.0)
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_n * (self.tile_k + 2), alignment=16)
smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_n, self.tile_k + 2), stride=(self.tile_k + 2, 1)))
for k0, acc, acc_out in for_generate(
0,
k,
self.tile_k,
iter_args=[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
],
):
acc00 = acc[0]
acc01 = acc[1]
acc02 = acc[2]
acc03 = acc[3]
acc04 = acc[4]
acc05 = acc[5]
acc06 = acc[6]
acc07 = acc[7]
acc10 = acc[8]
acc11 = acc[9]
acc12 = acc[10]
acc13 = acc[11]
acc14 = acc[12]
acc15 = acc[13]
acc16 = acc[14]
acc17 = acc[15]
for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
mm = idx // self.tile_k
kk = idx - mm * self.tile_k
gm = base_m + mm
gk = k0 + kk
def _ld():
smem_a[mm, kk] = a[gm, gk]
def _stz():
smem_a[mm, kk] = cutlass.Float16(0.0)
if_generate((gm < m) & (gk < k), _ld, _stz)
yield_out()
for idx in for_generate(tx, self.tile_n * self.tile_k, self.threads):
nn = idx // self.tile_k
kk = idx - nn * self.tile_k
gn = base_n + nn
gk = k0 + kk
def _ld():
smem_b[nn, kk] = b[gn, gk]
def _stz():
smem_b[nn, kk] = cutlass.Float16(0.0)
if_generate((gn < n) & (gk < k), _ld, _stz)
yield_out()
cute.arch.sync_threads()
for kk in range_constexpr(128):
bv0 = smem_b[lane, kk].to(cutlass.Float32)
bv1 = smem_b[lane + 32, kk].to(cutlass.Float32)
a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)
acc00 = acc00 + a0 * bv0
acc01 = acc01 + a1 * bv0
acc02 = acc02 + a2 * bv0
acc03 = acc03 + a3 * bv0
acc04 = acc04 + a4 * bv0
acc05 = acc05 + a5 * bv0
acc06 = acc06 + a6 * bv0
acc07 = acc07 + a7 * bv0
acc10 = acc10 + a0 * bv1
acc11 = acc11 + a1 * bv1
acc12 = acc12 + a2 * bv1
acc13 = acc13 + a3 * bv1
acc14 = acc14 + a4 * bv1
acc15 = acc15 + a5 * bv1
acc16 = acc16 + a6 * bv1
acc17 = acc17 + a7 * bv1
def _need_sync():
cute.arch.sync_threads()
if_generate(k0 + self.tile_k < k, _need_sync)
yield_out(
[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
]
)
def _st_row(r: int, col: int, val: cutlass.Float32):
gm = row0 + r
if_generate((gm < m) & (col < n), lambda: c.__setitem__((gm, col), val.to(cutlass.Float16)))
_st_row(0, col0, acc_out[0])
_st_row(1, col0, acc_out[1])
_st_row(2, col0, acc_out[2])
_st_row(3, col0, acc_out[3])
_st_row(4, col0, acc_out[4])
_st_row(5, col0, acc_out[5])
_st_row(6, col0, acc_out[6])
_st_row(7, col0, acc_out[7])
_st_row(0, col1, acc_out[8])
_st_row(1, col1, acc_out[9])
_st_row(2, col1, acc_out[10])
_st_row(3, col1, acc_out[11])
_st_row(4, col1, acc_out[12])
_st_row(5, col1, acc_out[13])
_st_row(6, col1, acc_out[14])
_st_row(7, col1, acc_out[15])
class _GemmF16F16ToF16_Full:
def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
self.threads = int(threads)
self.tile_m = int(tile_m)
self.tile_n = int(tile_n)
self.tile_k = int(tile_k)
self.warps = self.threads // 32
self.rows_per_warp = self.tile_m // self.warps
@cute.jit
def __call__(
self,
a_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
c_ptr: "cute.Pointer",
problem: tuple,
):
m, n, k = problem
a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))
b = cute.make_tensor(b_ptr, cute.make_layout((n, k), stride=(k, 1)))
c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))
grid_n = (n + self.tile_n - 1) // self.tile_n
grid_m = (m + self.tile_m - 1) // self.tile_m
self.kernel(a, b, c, m, n, k).launch(
grid=[grid_n, grid_m, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
a: "cute.Tensor",
b: "cute.Tensor",
c: "cute.Tensor",
m: int,
n: int,
k: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
base_m = by * self.tile_m
base_n = bx * self.tile_n
row0 = base_m + warp_id * self.rows_per_warp
col0 = base_n + lane
col1 = col0 + 32
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc02 = cutlass.Float32(0.0)
acc03 = cutlass.Float32(0.0)
acc04 = cutlass.Float32(0.0)
acc05 = cutlass.Float32(0.0)
acc06 = cutlass.Float32(0.0)
acc07 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
acc12 = cutlass.Float32(0.0)
acc13 = cutlass.Float32(0.0)
acc14 = cutlass.Float32(0.0)
acc15 = cutlass.Float32(0.0)
acc16 = cutlass.Float32(0.0)
acc17 = cutlass.Float32(0.0)
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_n * (self.tile_k + 2), alignment=16)
smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_n, self.tile_k + 2), stride=(self.tile_k + 2, 1)))
for k0, acc, acc_out in for_generate(
0,
k,
self.tile_k,
iter_args=[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
],
):
acc00 = acc[0]
acc01 = acc[1]
acc02 = acc[2]
acc03 = acc[3]
acc04 = acc[4]
acc05 = acc[5]
acc06 = acc[6]
acc07 = acc[7]
acc10 = acc[8]
acc11 = acc[9]
acc12 = acc[10]
acc13 = acc[11]
acc14 = acc[12]
acc15 = acc[13]
acc16 = acc[14]
acc17 = acc[15]
for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
mm = idx // self.tile_k
kk = idx - mm * self.tile_k
smem_a[mm, kk] = a[base_m + mm, k0 + kk]
yield_out()
for idx in for_generate(tx, self.tile_n * self.tile_k, self.threads):
nn = idx // self.tile_k
kk = idx - nn * self.tile_k
smem_b[nn, kk] = b[base_n + nn, k0 + kk]
yield_out()
cute.arch.sync_threads()
for kk in range_constexpr(128):
bv0 = smem_b[lane, kk].to(cutlass.Float32)
bv1 = smem_b[lane + 32, kk].to(cutlass.Float32)
a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)
acc00 = acc00 + a0 * bv0
acc01 = acc01 + a1 * bv0
acc02 = acc02 + a2 * bv0
acc03 = acc03 + a3 * bv0
acc04 = acc04 + a4 * bv0
acc05 = acc05 + a5 * bv0
acc06 = acc06 + a6 * bv0
acc07 = acc07 + a7 * bv0
acc10 = acc10 + a0 * bv1
acc11 = acc11 + a1 * bv1
acc12 = acc12 + a2 * bv1
acc13 = acc13 + a3 * bv1
acc14 = acc14 + a4 * bv1
acc15 = acc15 + a5 * bv1
acc16 = acc16 + a6 * bv1
acc17 = acc17 + a7 * bv1
def _need_sync():
cute.arch.sync_threads()
if_generate(k0 + self.tile_k < k, _need_sync)
yield_out(
[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
]
)
gm0 = row0 + 0
gm1 = row0 + 1
gm2 = row0 + 2
gm3 = row0 + 3
gm4 = row0 + 4
gm5 = row0 + 5
gm6 = row0 + 6
gm7 = row0 + 7
c[gm0, col0] = acc_out[0].to(cutlass.Float16)
c[gm1, col0] = acc_out[1].to(cutlass.Float16)
c[gm2, col0] = acc_out[2].to(cutlass.Float16)
c[gm3, col0] = acc_out[3].to(cutlass.Float16)
c[gm4, col0] = acc_out[4].to(cutlass.Float16)
c[gm5, col0] = acc_out[5].to(cutlass.Float16)
c[gm6, col0] = acc_out[6].to(cutlass.Float16)
c[gm7, col0] = acc_out[7].to(cutlass.Float16)
c[gm0, col1] = acc_out[8].to(cutlass.Float16)
c[gm1, col1] = acc_out[9].to(cutlass.Float16)
c[gm2, col1] = acc_out[10].to(cutlass.Float16)
c[gm3, col1] = acc_out[11].to(cutlass.Float16)
c[gm4, col1] = acc_out[12].to(cutlass.Float16)
c[gm5, col1] = acc_out[13].to(cutlass.Float16)
c[gm6, col1] = acc_out[14].to(cutlass.Float16)
c[gm7, col1] = acc_out[15].to(cutlass.Float16)
class _GemmF16F16ToF32:
def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
self.threads = int(threads)
self.tile_m = int(tile_m)
self.tile_n = int(tile_n)
self.tile_k = int(tile_k)
self.warps = self.threads // 32
self.rows_per_warp = self.tile_m // self.warps
@cute.jit
def __call__(
self,
a_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
c_ptr: "cute.Pointer",
problem: tuple,
):
m, n, k = problem
a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))
b = cute.make_tensor(b_ptr, cute.make_layout((k, n), stride=(n, 1)))
c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))
grid_n = (n + self.tile_n - 1) // self.tile_n
grid_m = (m + self.tile_m - 1) // self.tile_m
self.kernel(a, b, c, m, n, k).launch(
grid=[grid_n, grid_m, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
a: "cute.Tensor",
b: "cute.Tensor",
c: "cute.Tensor",
m: int,
n: int,
k: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
base_m = by * self.tile_m
base_n = bx * self.tile_n
row0 = base_m + warp_id * self.rows_per_warp
col0 = base_n + lane
col1 = col0 + 32
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc02 = cutlass.Float32(0.0)
acc03 = cutlass.Float32(0.0)
acc04 = cutlass.Float32(0.0)
acc05 = cutlass.Float32(0.0)
acc06 = cutlass.Float32(0.0)
acc07 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
acc12 = cutlass.Float32(0.0)
acc13 = cutlass.Float32(0.0)
acc14 = cutlass.Float32(0.0)
acc15 = cutlass.Float32(0.0)
acc16 = cutlass.Float32(0.0)
acc17 = cutlass.Float32(0.0)
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)
smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))
for k0, acc, acc_out in for_generate(
0,
k,
self.tile_k,
iter_args=[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
],
):
acc00 = acc[0]
acc01 = acc[1]
acc02 = acc[2]
acc03 = acc[3]
acc04 = acc[4]
acc05 = acc[5]
acc06 = acc[6]
acc07 = acc[7]
acc10 = acc[8]
acc11 = acc[9]
acc12 = acc[10]
acc13 = acc[11]
acc14 = acc[12]
acc15 = acc[13]
acc16 = acc[14]
acc17 = acc[15]
for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
mm = idx // self.tile_k
kk = idx - mm * self.tile_k
gm = base_m + mm
gk = k0 + kk
def _ld():
smem_a[mm, kk] = a[gm, gk]
def _stz():
smem_a[mm, kk] = cutlass.Float16(0.0)
if_generate((gm < m) & (gk < k), _ld, _stz)
yield_out()
for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):
kk = idx // self.tile_n
nn = idx - kk * self.tile_n
gk = k0 + kk
gn = base_n + nn
def _ld():
smem_b[kk, nn] = b[gk, gn]
def _stz():
smem_b[kk, nn] = cutlass.Float16(0.0)
if_generate((gk < k) & (gn < n), _ld, _stz)
yield_out()
cute.arch.sync_threads()
for kk in range_constexpr(128):
bv0 = smem_b[kk, lane].to(cutlass.Float32)
bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)
a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)
acc00 = acc00 + a0 * bv0
acc01 = acc01 + a1 * bv0
acc02 = acc02 + a2 * bv0
acc03 = acc03 + a3 * bv0
acc04 = acc04 + a4 * bv0
acc05 = acc05 + a5 * bv0
acc06 = acc06 + a6 * bv0
acc07 = acc07 + a7 * bv0
acc10 = acc10 + a0 * bv1
acc11 = acc11 + a1 * bv1
acc12 = acc12 + a2 * bv1
acc13 = acc13 + a3 * bv1
acc14 = acc14 + a4 * bv1
acc15 = acc15 + a5 * bv1
acc16 = acc16 + a6 * bv1
acc17 = acc17 + a7 * bv1
def _need_sync():
cute.arch.sync_threads()
if_generate(k0 + self.tile_k < k, _need_sync)
yield_out(
[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
]
)
def _st_row(r: int, col: int, val: cutlass.Float32):
gm = row0 + r
if_generate((gm < m) & (col < n), lambda: c.__setitem__((gm, col), val))
_st_row(0, col0, acc_out[0])
_st_row(1, col0, acc_out[1])
_st_row(2, col0, acc_out[2])
_st_row(3, col0, acc_out[3])
_st_row(4, col0, acc_out[4])
_st_row(5, col0, acc_out[5])
_st_row(6, col0, acc_out[6])
_st_row(7, col0, acc_out[7])
_st_row(0, col1, acc_out[8])
_st_row(1, col1, acc_out[9])
_st_row(2, col1, acc_out[10])
_st_row(3, col1, acc_out[11])
_st_row(4, col1, acc_out[12])
_st_row(5, col1, acc_out[13])
_st_row(6, col1, acc_out[14])
_st_row(7, col1, acc_out[15])
class _GemmF16F16ToF32_Full:
def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
self.threads = int(threads)
self.tile_m = int(tile_m)
self.tile_n = int(tile_n)
self.tile_k = int(tile_k)
self.warps = self.threads // 32
self.rows_per_warp = self.tile_m // self.warps
@cute.jit
def __call__(
self,
a_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
c_ptr: "cute.Pointer",
problem: tuple,
):
m, n, k = problem
a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))
b = cute.make_tensor(b_ptr, cute.make_layout((k, n), stride=(n, 1)))
c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))
grid_n = (n + self.tile_n - 1) // self.tile_n
grid_m = (m + self.tile_m - 1) // self.tile_m
self.kernel(a, b, c, m, n, k).launch(
grid=[grid_n, grid_m, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
a: "cute.Tensor",
b: "cute.Tensor",
c: "cute.Tensor",
m: int,
n: int,
k: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
base_m = by * self.tile_m
base_n = bx * self.tile_n
row0 = base_m + warp_id * self.rows_per_warp
col0 = base_n + lane
col1 = col0 + 32
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc02 = cutlass.Float32(0.0)
acc03 = cutlass.Float32(0.0)
acc04 = cutlass.Float32(0.0)
acc05 = cutlass.Float32(0.0)
acc06 = cutlass.Float32(0.0)
acc07 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
acc12 = cutlass.Float32(0.0)
acc13 = cutlass.Float32(0.0)
acc14 = cutlass.Float32(0.0)
acc15 = cutlass.Float32(0.0)
acc16 = cutlass.Float32(0.0)
acc17 = cutlass.Float32(0.0)
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)
smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))
for k0, acc, acc_out in for_generate(
0,
k,
self.tile_k,
iter_args=[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
],
):
acc00 = acc[0]
acc01 = acc[1]
acc02 = acc[2]
acc03 = acc[3]
acc04 = acc[4]
acc05 = acc[5]
acc06 = acc[6]
acc07 = acc[7]
acc10 = acc[8]
acc11 = acc[9]
acc12 = acc[10]
acc13 = acc[11]
acc14 = acc[12]
acc15 = acc[13]
acc16 = acc[14]
acc17 = acc[15]
for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
mm = idx // self.tile_k
kk = idx - mm * self.tile_k
smem_a[mm, kk] = a[base_m + mm, k0 + kk]
yield_out()
for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):
kk = idx // self.tile_n
nn = idx - kk * self.tile_n
smem_b[kk, nn] = b[k0 + kk, base_n + nn]
yield_out()
cute.arch.sync_threads()
for kk in range_constexpr(128):
bv0 = smem_b[kk, lane].to(cutlass.Float32)
bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)
a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)
acc00 = acc00 + a0 * bv0
acc01 = acc01 + a1 * bv0
acc02 = acc02 + a2 * bv0
acc03 = acc03 + a3 * bv0
acc04 = acc04 + a4 * bv0
acc05 = acc05 + a5 * bv0
acc06 = acc06 + a6 * bv0
acc07 = acc07 + a7 * bv0
acc10 = acc10 + a0 * bv1
acc11 = acc11 + a1 * bv1
acc12 = acc12 + a2 * bv1
acc13 = acc13 + a3 * bv1
acc14 = acc14 + a4 * bv1
acc15 = acc15 + a5 * bv1
acc16 = acc16 + a6 * bv1
acc17 = acc17 + a7 * bv1
def _need_sync():
cute.arch.sync_threads()
if_generate(k0 + self.tile_k < k, _need_sync)
yield_out(
[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
]
)
gm0 = row0 + 0
gm1 = row0 + 1
gm2 = row0 + 2
gm3 = row0 + 3
gm4 = row0 + 4
gm5 = row0 + 5
gm6 = row0 + 6
gm7 = row0 + 7
c[gm0, col0] = acc_out[0]
c[gm1, col0] = acc_out[1]
c[gm2, col0] = acc_out[2]
c[gm3, col0] = acc_out[3]
c[gm4, col0] = acc_out[4]
c[gm5, col0] = acc_out[5]
c[gm6, col0] = acc_out[6]
c[gm7, col0] = acc_out[7]
c[gm0, col1] = acc_out[8]
c[gm1, col1] = acc_out[9]
c[gm2, col1] = acc_out[10]
c[gm3, col1] = acc_out[11]
c[gm4, col1] = acc_out[12]
c[gm5, col1] = acc_out[13]
c[gm6, col1] = acc_out[14]
c[gm7, col1] = acc_out[15]
_LOG2E = 1.4426950408889634
def _sigmoid_f16(x: cutlass.Float16) -> cutlass.Float16:
xx = x.to(cutlass.Float32)
t = (cutlass.Float32(0.0) - xx) * cutlass.Float32(_LOG2E)
ee = cute.exp2(t, fastmath=True)
denom = cutlass.Float32(1.0) + ee
y = cute.arch.rcp_approx(denom)
y = y * (cutlass.Float32(2.0) - denom * y)
return y.to(cutlass.Float16)
class _ProcessProj:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
@cute.jit
def __call__(
self,
proj_ptr: "cute.Pointer",
mask_ptr: "cute.Pointer",
left_ptr: "cute.Pointer",
right_ptr: "cute.Pointer",
gate_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, h = problem
stride_proj_bs = n * n * (5 * h)
stride_proj_i = n * (5 * h)
stride_proj_j = 5 * h
proj = cute.make_tensor(
proj_ptr,
cute.make_layout(
(bs, n, n, 5 * h),
stride=(stride_proj_bs, stride_proj_i, stride_proj_j, 1),
),
)
mask = cute.make_tensor(mask_ptr, cute.make_layout((bs, n, n), stride=(n * n, n, 1)))
left = cute.make_tensor(left_ptr, cute.make_layout((bs, n, h, n), stride=(n * h * n, h * n, n, 1)))
right = cute.make_tensor(right_ptr, cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)))
gate = cute.make_tensor(gate_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
total = bs * n * n * h
grid = (total + self.threads - 1) // self.threads
self.kernel(proj, mask, left, right, gate, bs, n, h).launch(
grid=[grid, 1, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
proj: "cute.Tensor",
mask: "cute.Tensor",
left: "cute.Tensor",
right: "cute.Tensor",
gate: "cute.Tensor",
bs: int,
n: int,
h: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, _, _ = cute.arch.block_idx()
bdx, _, _ = cute.arch.block_dim()
idx = bx * bdx + tx
total = bs * n * n * h
def _do_one():
hh = idx % h
t0 = idx // h
j = t0 % n
t1 = t0 // n
i = t1 % n
bb = t1 // n
m = mask[bb, i, j].to(cutlass.Float16)
lp = proj[bb, i, j, hh]
rp = proj[bb, i, j, hh + h]
lg = proj[bb, i, j, hh + 2 * h]
rg = proj[bb, i, j, hh + 3 * h]
og = proj[bb, i, j, hh + 4 * h]
gl = _sigmoid_f16(lg)
gr = _sigmoid_f16(rg)
go = _sigmoid_f16(og)
left[bb, i, hh, j] = (lp * gl * m).to(cutlass.Float16)
right[bb, hh, j, i] = (rp * gr * m).to(cutlass.Float16)
gate[bb, i, j, hh] = go
if_generate(idx < total, _do_one)
class _ContractHiddenGemm:
def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
self.threads = int(threads)
self.tile_m = int(tile_m)
self.tile_n = int(tile_n)
self.tile_k = int(tile_k)
self.warps = self.threads // 32
self.rows_per_warp = self.tile_m // self.warps
@cute.jit
def __call__(
self,
left_ptr: "cute.Pointer",
right_ptr: "cute.Pointer",
out_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, h = problem
left = cute.make_tensor(
left_ptr,
cute.make_layout((bs, n, h, n), stride=(n * h * n, h * n, n, 1)),
)
right = cute.make_tensor(
right_ptr,
cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),
)
out = cute.make_tensor(
out_ptr,
cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),
)
grid_n = (n + self.tile_n - 1) // self.tile_n
grid_m = (n + self.tile_m - 1) // self.tile_m
grid_z = bs * h
self.kernel(left, right, out, bs, n, h).launch(
grid=[grid_n, grid_m, grid_z],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
left: "cute.Tensor",
right: "cute.Tensor",
out: "cute.Tensor",
bs: int,
n: int,
h: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, bz = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
bb = bz // h
hh = bz - bb * h
base_m = by * self.tile_m
base_n = bx * self.tile_n
row0 = base_m + warp_id * self.rows_per_warp
col0 = base_n + lane
col1 = col0 + 32
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc02 = cutlass.Float32(0.0)
acc03 = cutlass.Float32(0.0)
acc04 = cutlass.Float32(0.0)
acc05 = cutlass.Float32(0.0)
acc06 = cutlass.Float32(0.0)
acc07 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
acc12 = cutlass.Float32(0.0)
acc13 = cutlass.Float32(0.0)
acc14 = cutlass.Float32(0.0)
acc15 = cutlass.Float32(0.0)
acc16 = cutlass.Float32(0.0)
acc17 = cutlass.Float32(0.0)
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)
smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))
for k0, acc, acc_out in for_generate(
0,
n,
self.tile_k,
iter_args=[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
],
):
acc00 = acc[0]
acc01 = acc[1]
acc02 = acc[2]
acc03 = acc[3]
acc04 = acc[4]
acc05 = acc[5]
acc06 = acc[6]
acc07 = acc[7]
acc10 = acc[8]
acc11 = acc[9]
acc12 = acc[10]
acc13 = acc[11]
acc14 = acc[12]
acc15 = acc[13]
acc16 = acc[14]
acc17 = acc[15]
for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
mm = idx // self.tile_k
kk = idx - mm * self.tile_k
gi = base_m + mm
gk = k0 + kk
def _ld():
smem_a[mm, kk] = left[bb, gi, hh, gk]
def _stz():
smem_a[mm, kk] = cutlass.Float16(0.0)
if_generate((gi < n) & (gk < n), _ld, _stz)
yield_out()
for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):
kk = idx // self.tile_n
nn = idx - kk * self.tile_n
gk = k0 + kk
gj = base_n + nn
def _ld():
smem_b[kk, nn] = right[bb, hh, gk, gj]
def _stz():
smem_b[kk, nn] = cutlass.Float16(0.0)
if_generate((gk < n) & (gj < n), _ld, _stz)
yield_out()
cute.arch.sync_threads()
for kk in range_constexpr(128):
bv0 = smem_b[kk, lane].to(cutlass.Float32)
bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)
a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)
acc00 = acc00 + a0 * bv0
acc01 = acc01 + a1 * bv0
acc02 = acc02 + a2 * bv0
acc03 = acc03 + a3 * bv0
acc04 = acc04 + a4 * bv0
acc05 = acc05 + a5 * bv0
acc06 = acc06 + a6 * bv0
acc07 = acc07 + a7 * bv0
acc10 = acc10 + a0 * bv1
acc11 = acc11 + a1 * bv1
acc12 = acc12 + a2 * bv1
acc13 = acc13 + a3 * bv1
acc14 = acc14 + a4 * bv1
acc15 = acc15 + a5 * bv1
acc16 = acc16 + a6 * bv1
acc17 = acc17 + a7 * bv1
def _need_sync():
cute.arch.sync_threads()
if_generate(k0 + self.tile_k < n, _need_sync)
yield_out(
[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
]
)
def _st_row(r: int, col: int, val: cutlass.Float32):
gi = row0 + r
if_generate((gi < n) & (col < n), lambda: out.__setitem__((bb, hh, gi, col), val))
_st_row(0, col0, acc_out[0])
_st_row(1, col0, acc_out[1])
_st_row(2, col0, acc_out[2])
_st_row(3, col0, acc_out[3])
_st_row(4, col0, acc_out[4])
_st_row(5, col0, acc_out[5])
_st_row(6, col0, acc_out[6])
_st_row(7, col0, acc_out[7])
_st_row(0, col1, acc_out[8])
_st_row(1, col1, acc_out[9])
_st_row(2, col1, acc_out[10])
_st_row(3, col1, acc_out[11])
_st_row(4, col1, acc_out[12])
_st_row(5, col1, acc_out[13])
_st_row(6, col1, acc_out[14])
_st_row(7, col1, acc_out[15])
class _ContractHiddenGemm_Full:
def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
self.threads = int(threads)
self.tile_m = int(tile_m)
self.tile_n = int(tile_n)
self.tile_k = int(tile_k)
self.warps = self.threads // 32
self.rows_per_warp = self.tile_m // self.warps
@cute.jit
def __call__(
self,
left_ptr: "cute.Pointer",
right_ptr: "cute.Pointer",
out_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, h = problem
left = cute.make_tensor(
left_ptr,
cute.make_layout((bs, n, h, n), stride=(n * h * n, h * n, n, 1)),
)
right = cute.make_tensor(
right_ptr,
cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),
)
out = cute.make_tensor(
out_ptr,
cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),
)
grid_n = (n + self.tile_n - 1) // self.tile_n
grid_m = (n + self.tile_m - 1) // self.tile_m
grid_z = bs * h
self.kernel(left, right, out, bs, n, h).launch(
grid=[grid_n, grid_m, grid_z],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
left: "cute.Tensor",
right: "cute.Tensor",
out: "cute.Tensor",
bs: int,
n: int,
h: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, by, bz = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
bb = bz // h
hh = bz - bb * h
base_m = by * self.tile_m
base_n = bx * self.tile_n
row0 = base_m + warp_id * self.rows_per_warp
col0 = base_n + lane
col1 = col0 + 32
acc00 = cutlass.Float32(0.0)
acc01 = cutlass.Float32(0.0)
acc02 = cutlass.Float32(0.0)
acc03 = cutlass.Float32(0.0)
acc04 = cutlass.Float32(0.0)
acc05 = cutlass.Float32(0.0)
acc06 = cutlass.Float32(0.0)
acc07 = cutlass.Float32(0.0)
acc10 = cutlass.Float32(0.0)
acc11 = cutlass.Float32(0.0)
acc12 = cutlass.Float32(0.0)
acc13 = cutlass.Float32(0.0)
acc14 = cutlass.Float32(0.0)
acc15 = cutlass.Float32(0.0)
acc16 = cutlass.Float32(0.0)
acc17 = cutlass.Float32(0.0)
smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)
smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))
for k0, acc, acc_out in for_generate(
0,
n,
self.tile_k,
iter_args=[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
],
):
acc00 = acc[0]
acc01 = acc[1]
acc02 = acc[2]
acc03 = acc[3]
acc04 = acc[4]
acc05 = acc[5]
acc06 = acc[6]
acc07 = acc[7]
acc10 = acc[8]
acc11 = acc[9]
acc12 = acc[10]
acc13 = acc[11]
acc14 = acc[12]
acc15 = acc[13]
acc16 = acc[14]
acc17 = acc[15]
for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
mm = idx // self.tile_k
kk = idx - mm * self.tile_k
smem_a[mm, kk] = left[bb, base_m + mm, hh, k0 + kk]
yield_out()
for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):
kk = idx // self.tile_n
nn = idx - kk * self.tile_n
smem_b[kk, nn] = right[bb, hh, k0 + kk, base_n + nn]
yield_out()
cute.arch.sync_threads()
for kk in range_constexpr(128):
bv0 = smem_b[kk, lane].to(cutlass.Float32)
bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)
a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)
acc00 = acc00 + a0 * bv0
acc01 = acc01 + a1 * bv0
acc02 = acc02 + a2 * bv0
acc03 = acc03 + a3 * bv0
acc04 = acc04 + a4 * bv0
acc05 = acc05 + a5 * bv0
acc06 = acc06 + a6 * bv0
acc07 = acc07 + a7 * bv0
acc10 = acc10 + a0 * bv1
acc11 = acc11 + a1 * bv1
acc12 = acc12 + a2 * bv1
acc13 = acc13 + a3 * bv1
acc14 = acc14 + a4 * bv1
acc15 = acc15 + a5 * bv1
acc16 = acc16 + a6 * bv1
acc17 = acc17 + a7 * bv1
def _need_sync():
cute.arch.sync_threads()
if_generate(k0 + self.tile_k < n, _need_sync)
yield_out(
[
acc00,
acc01,
acc02,
acc03,
acc04,
acc05,
acc06,
acc07,
acc10,
acc11,
acc12,
acc13,
acc14,
acc15,
acc16,
acc17,
]
)
gi0 = row0 + 0
gi1 = row0 + 1
gi2 = row0 + 2
gi3 = row0 + 3
gi4 = row0 + 4
gi5 = row0 + 5
gi6 = row0 + 6
gi7 = row0 + 7
out[bb, hh, gi0, col0] = acc_out[0]
out[bb, hh, gi1, col0] = acc_out[1]
out[bb, hh, gi2, col0] = acc_out[2]
out[bb, hh, gi3, col0] = acc_out[3]
out[bb, hh, gi4, col0] = acc_out[4]
out[bb, hh, gi5, col0] = acc_out[5]
out[bb, hh, gi6, col0] = acc_out[6]
out[bb, hh, gi7, col0] = acc_out[7]
out[bb, hh, gi0, col1] = acc_out[8]
out[bb, hh, gi1, col1] = acc_out[9]
out[bb, hh, gi2, col1] = acc_out[10]
out[bb, hh, gi3, col1] = acc_out[11]
out[bb, hh, gi4, col1] = acc_out[12]
out[bb, hh, gi5, col1] = acc_out[13]
out[bb, hh, gi6, col1] = acc_out[14]
out[bb, hh, gi7, col1] = acc_out[15]
class _LayerNormHiddenF32ToF16:
def __init__(self, threads: int = 256) -> None:
self.threads = int(threads)
self.warps = self.threads // 32
@cute.jit
def __call__(
self,
x_ptr: "cute.Pointer",
w_ptr: "cute.Pointer",
b_ptr: "cute.Pointer",
g_ptr: "cute.Pointer",
y_ptr: "cute.Pointer",
problem: tuple,
):
bs, n, h = problem
stride_bs = h * n * n
stride_h = n * n
stride_i = n
x = cute.make_tensor(x_ptr, cute.make_layout((bs, h, n, n), stride=(stride_bs, stride_h, stride_i, 1)))
w = cute.make_tensor(w_ptr, cute.make_layout((h,), stride=(1,)))
b = cute.make_tensor(b_ptr, cute.make_layout((h,), stride=(1,)))
g = cute.make_tensor(g_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
y = cute.make_tensor(y_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
total = bs * n * n
grid = (total + self.warps - 1) // self.warps
self.kernel(x, w, b, g, y, bs, n, h).launch(
grid=[grid, 1, 1],
block=[self.threads, 1, 1],
)
return
@cute.kernel
def kernel(
self,
x: "cute.Tensor",
w: "cute.Tensor",
b: "cute.Tensor",
g: "cute.Tensor",
y: "cute.Tensor",
bs: int,
n: int,
h: int,
):
tx, _, _ = cute.arch.thread_idx()
bx, _, _ = cute.arch.block_idx()
warp_id = tx >> 5
lane = tx & 31
idx = bx * self.warps + warp_id
def _do_one():
j = idx % n
t0 = idx // n
i = t0 % n
bb = t0 // n
sum0_init = cutlass.Float32(0.0)
sumsq0_init = cutlass.Float32(0.0)
for hh, acc, acc_out in for_generate(lane, h, 32, iter_args=[sum0_init, sumsq0_init]):
sum0_it = acc[0]
sumsq0_it = acc[1]
v = x[bb, hh, i, j]
sum0_it = sum0_it + v
sumsq0_it = sumsq0_it + v * v
yield_out([sum0_it, sumsq0_it])
sum0 = cute.arch.warp_reduction_sum(acc_out[0])
sumsq0 = cute.arch.warp_reduction_sum(acc_out[1])
inv_h = cutlass.Float32(1.0) / cutlass.Float32(h)
mean = sum0 * inv_h
var = sumsq0 * inv_h - mean * mean
inv_std = cute.rsqrt(var + cutlass.Float32(1e-5), fastmath=True)
for hh in for_generate(lane, h, 32):
v = x[bb, hh, i, j]
nrm = (v - mean) * inv_std
out = (nrm * w[hh] + b[hh]).to(cutlass.Float16)
y[bb, i, j, hh] = (out * g[bb, i, j, hh]).to(cutlass.Float16)
yield_out()
if_generate(idx < bs * n * n, _do_one)
_LN_X = _LayerNormLastDimF32ToF16()
_LN_X_C = None
_GEMM_PROJ = _GemmF16F16ToF16()
_GEMM_PROJ_C = None
_GEMM_PROJ_FULL = _GemmF16F16ToF16_Full()
_GEMM_PROJ_FULL_C = None
_PROC = _ProcessProj()
_PROC_C = None
_CONTRACT = _ContractHiddenGemm()
_CONTRACT_C = None
_CONTRACT_FULL = _ContractHiddenGemm_Full()
_CONTRACT_FULL_C = None
_LN_H = _LayerNormHiddenF32ToF16()
_LN_H_C = None
_GEMM_OUT = _GemmF16F16ToF32()
_GEMM_OUT_C = None
_GEMM_OUT_FULL = _GemmF16F16ToF32_Full()
_GEMM_OUT_FULL_C = None
def _compile_once():
global _LN_X_C, _GEMM_PROJ_C, _GEMM_PROJ_FULL_C, _PROC_C, _CONTRACT_C, _CONTRACT_FULL_C, _LN_H_C, _GEMM_OUT_C, _GEMM_OUT_FULL_C
if _LN_X_C is None:
x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
w_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
y_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_LN_X_C = cute.compile(_LN_X, x_ptr, w_ptr, b_ptr, y_ptr, (0, 0, 0), options="--opt-level 3")
if _GEMM_PROJ_C is None:
a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_GEMM_PROJ_C = cute.compile(_GEMM_PROJ, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")
if _GEMM_PROJ_FULL_C is None:
a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_GEMM_PROJ_FULL_C = cute.compile(_GEMM_PROJ_FULL, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")
if _PROC_C is None:
proj_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
mask_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
gate_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_PROC_C = cute.compile(
_PROC,
proj_ptr,
mask_ptr,
left_ptr,
right_ptr,
gate_ptr,
(0, 0, 0),
options="--opt-level 3",
)
if _CONTRACT_C is None:
left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
out_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_CONTRACT_C = cute.compile(_CONTRACT, left_ptr, right_ptr, out_ptr, (0, 0, 0), options="--opt-level 3")
if _CONTRACT_FULL_C is None:
left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
out_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_CONTRACT_FULL_C = cute.compile(_CONTRACT_FULL, left_ptr, right_ptr, out_ptr, (0, 0, 0), options="--opt-level 3")
if _LN_H_C is None:
x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
w_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
g_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
y_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
_LN_H_C = cute.compile(_LN_H, x_ptr, w_ptr, b_ptr, g_ptr, y_ptr, (0, 0, 0), options="--opt-level 3")
if _GEMM_OUT_C is None:
a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_GEMM_OUT_C = cute.compile(_GEMM_OUT, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")
if _GEMM_OUT_FULL_C is None:
a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_GEMM_OUT_FULL_C = cute.compile(_GEMM_OUT_FULL, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")
def _as_ptr(ty, t: torch.Tensor):
return make_ptr(ty, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
_W_KEY = None
_W_PACK16 = None
_W_OUT_T16 = None
_W_NORM_W = None
_W_NORM_B = None
_W_TO_OUT_NORM_W = None
_W_TO_OUT_NORM_B = None
_BUF_KEY = None
_BUF = None
def _prep_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int, device: torch.device):
global _W_KEY, _W_PACK16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B
w_lp = weights["left_proj.weight"]
w_rp = weights["right_proj.weight"]
w_lg = weights["left_gate.weight"]
w_rg = weights["right_gate.weight"]
w_og = weights["out_gate.weight"]
w_no_w = weights["norm.weight"]
w_no_b = weights["norm.bias"]
w_o_nw = weights["to_out_norm.weight"]
w_o_nb = weights["to_out_norm.bias"]
w_out = weights["to_out.weight"]
key = (
device,
dim,
hidden,
int(w_lp.data_ptr()),
int(w_rp.data_ptr()),
int(w_lg.data_ptr()),
int(w_rg.data_ptr()),
int(w_og.data_ptr()),
int(w_no_w.data_ptr()),
int(w_no_b.data_ptr()),
int(w_o_nw.data_ptr()),
int(w_o_nb.data_ptr()),
int(w_out.data_ptr()),
)
if key == _W_KEY:
return _W_PACK16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B
w_pack16 = torch.cat([w_lp, w_rp, w_lg, w_rg, w_og], dim=0).contiguous().to(torch.float16)
w_out_t16 = w_out.contiguous().to(torch.float16).transpose(0, 1).contiguous()
_W_PACK16 = w_pack16
_W_OUT_T16 = w_out_t16
_W_NORM_W = w_no_w.contiguous()
_W_NORM_B = w_no_b.contiguous()
_W_TO_OUT_NORM_W = w_o_nw.contiguous()
_W_TO_OUT_NORM_B = w_o_nb.contiguous()
_W_KEY = key
return w_pack16, w_out_t16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B
def _get_buf(device: torch.device, bs: int, n: int, dim: int, hidden: int):
global _BUF_KEY, _BUF
key = (device, bs, n, dim, hidden)
if key == _BUF_KEY and _BUF is not None:
return _BUF
m = bs * n * n
buf = {
"x_norm": torch.empty((bs, n, n, dim), device=device, dtype=torch.float16),
"proj": torch.empty((m, 5 * hidden), device=device, dtype=torch.float16),
"mask16": torch.empty((bs, n, n), device=device, dtype=torch.float16),
"left_t": torch.empty((bs, n, hidden, n), device=device, dtype=torch.float16),
"right_t": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float16),
"out_gate": torch.empty((bs, n, n, hidden), device=device, dtype=torch.float16),
"out_tmp": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float32),
"out_norm": torch.empty((bs, n, n, hidden), device=device, dtype=torch.float16),
}
_BUF_KEY = key
_BUF = buf
return buf
@torch.inference_mode()
def custom_kernel(data: Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict[str, Any]]) -> torch.Tensor:
x, mask, weights, config = data
if not x.is_cuda:
raise RuntimeError("仅支持 CUDA 张量。")
if x.dtype != torch.float32:
x = x.to(torch.float32)
x = x.contiguous()
bs, n, n2, dim = x.shape
if n != n2:
raise RuntimeError("输入必须是 [bs, N, N, dim] 的方阵。")
dim_cfg = int(config["dim"])
hidden = int(config["hidden_dim"])
if dim_cfg != dim:
raise RuntimeError("config['dim'] 与 x.shape[-1] 不一致。")
_compile_once()
w_pack16, w_out_t16, w_norm_w, w_norm_b, w_out_norm_w, w_out_norm_b = _prep_weights(
weights, dim, hidden, x.device
)
buf = _get_buf(x.device, bs, n, dim, hidden)
x_norm = buf["x_norm"]
proj = buf["proj"]
mask16 = buf["mask16"]
left_t = buf["left_t"]
right_t = buf["right_t"]
out_gate = buf["out_gate"]
out_tmp = buf["out_tmp"]
out_norm = buf["out_norm"]
_LN_X_C(
_as_ptr(cutlass.Float32, x),
_as_ptr(cutlass.Float32, w_norm_w),
_as_ptr(cutlass.Float32, w_norm_b),
_as_ptr(cutlass.Float16, x_norm),
(bs, n, dim),
)
m = bs * n * n
use_proj_full = (m % 64 == 0) & ((5 * hidden) % 64 == 0) & (dim % 128 == 0)
if use_proj_full:
_GEMM_PROJ_FULL_C(
_as_ptr(cutlass.Float16, x_norm.view(m, dim)),
_as_ptr(cutlass.Float16, w_pack16),
_as_ptr(cutlass.Float16, proj),
(m, 5 * hidden, dim),
)
else:
_GEMM_PROJ_C(
_as_ptr(cutlass.Float16, x_norm.view(m, dim)),
_as_ptr(cutlass.Float16, w_pack16),
_as_ptr(cutlass.Float16, proj),
(m, 5 * hidden, dim),
)
mask16.copy_(mask)
_PROC_C(
_as_ptr(cutlass.Float16, proj.view(bs, n, n, 5 * hidden)),
_as_ptr(cutlass.Float16, mask16),
_as_ptr(cutlass.Float16, left_t),
_as_ptr(cutlass.Float16, right_t),
_as_ptr(cutlass.Float16, out_gate),
(bs, n, hidden),
)
use_contract_full = (n % 64 == 0) & (n % 128 == 0)
if use_contract_full:
_CONTRACT_FULL_C(
_as_ptr(cutlass.Float16, left_t),
_as_ptr(cutlass.Float16, right_t),
_as_ptr(cutlass.Float32, out_tmp),
(bs, n, hidden),
)
else:
_CONTRACT_C(
_as_ptr(cutlass.Float16, left_t),
_as_ptr(cutlass.Float16, right_t),
_as_ptr(cutlass.Float32, out_tmp),
(bs, n, hidden),
)
_LN_H_C(
_as_ptr(cutlass.Float32, out_tmp),
_as_ptr(cutlass.Float32, w_out_norm_w),
_as_ptr(cutlass.Float32, w_out_norm_b),
_as_ptr(cutlass.Float16, out_gate),
_as_ptr(cutlass.Float16, out_norm),
(bs, n, hidden),
)
y = torch.empty((m, dim), device=x.device, dtype=torch.float32)
use_out_full = (m % 64 == 0) & (dim % 64 == 0) & (hidden % 128 == 0)
if use_out_full:
_GEMM_OUT_FULL_C(
_as_ptr(cutlass.Float16, out_norm.view(m, hidden)),
_as_ptr(cutlass.Float16, w_out_t16),
_as_ptr(cutlass.Float32, y),
(m, dim, hidden),
)
else:
_GEMM_OUT_C(
_as_ptr(cutlass.Float16, out_norm.view(m, hidden)),
_as_ptr(cutlass.Float16, w_out_t16),
_as_ptr(cutlass.Float32, y),
(m, dim, hidden),
)
return y.view(bs, n, n, dim)
__all__ = ["custom_kernel"]
scrolls · 1994 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 418500.
⋯ 295 unchanged linesacc16 = acc16 + a6 * bv1acc17 = acc17 + a7 * bv1- cute.arch.sync_threads()+ def _need_sync():+ cute.arch.sync_threads()++ if_generate(k0 + self.tile_k < k, _need_sync)yield_out([acc00,⋯ 38 unchanged lines_st_row(7, col1, acc_out[15])+ class _GemmF16F16ToF16_Full:+ def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:+ self.threads = int(threads)+ self.tile_m = int(tile_m)+ self.tile_n = int(tile_n)+ self.tile_k = int(tile_k)+ self.warps = self.threads // 32+ self.rows_per_warp = self.tile_m // self.warps++ @cute.jit+ def __call__(+ self,+ a_ptr: "cute.Pointer",+ b_ptr: "cute.Pointer",+ c_ptr: "cute.Pointer",+ problem: tuple,+ ):+ m, n, k = problem++ a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))+ b = cute.make_tensor(b_ptr, cute.make_layout((n, k), stride=(k, 1)))+ c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))++ grid_n = (n + self.tile_n - 1) // self.tile_n+ grid_m = (m + self.tile_m - 1) // self.tile_m++ self.kernel(a, b, c, m, n, k).launch(+ grid=[grid_n, grid_m, 1],+ block=[self.threads, 1, 1],+ )+ return++ @cute.kernel+ def kernel(+ self,+ a: "cute.Tensor",+ b: "cute.Tensor",+ c: "cute.Tensor",+ m: int,+ n: int,+ k: int,+ ):+ tx, _, _ = cute.arch.thread_idx()+ bx, by, _ = cute.arch.block_idx()++ warp_id = tx >> 5+ lane = tx & 31++ base_m = by * self.tile_m+ base_n = bx * self.tile_n++ row0 = base_m + warp_id * self.rows_per_warp+ col0 = base_n + lane+ col1 = col0 + 32++ acc00 = cutlass.Float32(0.0)+ acc01 = cutlass.Float32(0.0)+ acc02 = cutlass.Float32(0.0)+ acc03 = cutlass.Float32(0.0)+ acc04 = cutlass.Float32(0.0)+ acc05 = cutlass.Float32(0.0)+ acc06 = cutlass.Float32(0.0)+ acc07 = cutlass.Float32(0.0)+ acc10 = cutlass.Float32(0.0)+ acc11 = cutlass.Float32(0.0)+ acc12 = cutlass.Float32(0.0)+ acc13 = cutlass.Float32(0.0)+ acc14 = cutlass.Float32(0.0)+ acc15 = cutlass.Float32(0.0)+ acc16 = cutlass.Float32(0.0)+ acc17 = cutlass.Float32(0.0)++ smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)+ smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_n * (self.tile_k + 2), alignment=16)+ smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))+ smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_n, self.tile_k + 2), stride=(self.tile_k + 2, 1)))++ for k0, acc, acc_out in for_generate(+ 0,+ k,+ self.tile_k,+ iter_args=[+ acc00,+ acc01,+ acc02,+ acc03,+ acc04,+ acc05,+ acc06,+ acc07,+ acc10,+ acc11,+ acc12,+ acc13,+ acc14,+ acc15,+ acc16,+ acc17,+ ],+ ):+ acc00 = acc[0]+ acc01 = acc[1]+ acc02 = acc[2]+ acc03 = acc[3]+ acc04 = acc[4]+ acc05 = acc[5]+ acc06 = acc[6]+ acc07 = acc[7]+ acc10 = acc[8]+ acc11 = acc[9]+ acc12 = acc[10]+ acc13 = acc[11]+ acc14 = acc[12]+ acc15 = acc[13]+ acc16 = acc[14]+ acc17 = acc[15]++ for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):+ mm = idx // self.tile_k+ kk = idx - mm * self.tile_k+ smem_a[mm, kk] = a[base_m + mm, k0 + kk]+ yield_out()++ for idx in for_generate(tx, self.tile_n * self.tile_k, self.threads):+ nn = idx // self.tile_k+ kk = idx - nn * self.tile_k+ smem_b[nn, kk] = b[base_n + nn, k0 + kk]+ yield_out()++ cute.arch.sync_threads()++ for kk in range_constexpr(128):+ bv0 = smem_b[lane, kk].to(cutlass.Float32)+ bv1 = smem_b[lane + 32, kk].to(cutlass.Float32)++ a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)+ a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)+ a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)+ a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)+ a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)+ a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)+ a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)+ a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)++ acc00 = acc00 + a0 * bv0+ acc01 = acc01 + a1 * bv0+ acc02 = acc02 + a2 * bv0+ acc03 = acc03 + a3 * bv0+ acc04 = acc04 + a4 * bv0+ acc05 = acc05 + a5 * bv0+ acc06 = acc06 + a6 * bv0+ acc07 = acc07 + a7 * bv0++ acc10 = acc10 + a0 * bv1+ acc11 = acc11 + a1 * bv1+ acc12 = acc12 + a2 * bv1+ acc13 = acc13 + a3 * bv1+ acc14 = acc14 + a4 * bv1+ acc15 = acc15 + a5 * bv1+ acc16 = acc16 + a6 * bv1+ acc17 = acc17 + a7 * bv1++ def _need_sync():+ cute.arch.sync_threads()++ if_generate(k0 + self.tile_k < k, _need_sync)+ yield_out(+ [+ acc00,+ acc01,+ acc02,+ acc03,+ acc04,+ acc05,+ acc06,+ acc07,+ acc10,+ acc11,+ acc12,+ acc13,+ acc14,+ acc15,+ acc16,+ acc17,+ ]+ )++ gm0 = row0 + 0+ gm1 = row0 + 1+ gm2 = row0 + 2+ gm3 = row0 + 3+ gm4 = row0 + 4+ gm5 = row0 + 5+ gm6 = row0 + 6+ gm7 = row0 + 7++ c[gm0, col0] = acc_out[0].to(cutlass.Float16)+ c[gm1, col0] = acc_out[1].to(cutlass.Float16)+ c[gm2, col0] = acc_out[2].to(cutlass.Float16)+ c[gm3, col0] = acc_out[3].to(cutlass.Float16)+ c[gm4, col0] = acc_out[4].to(cutlass.Float16)+ c[gm5, col0] = acc_out[5].to(cutlass.Float16)+ c[gm6, col0] = acc_out[6].to(cutlass.Float16)+ c[gm7, col0] = acc_out[7].to(cutlass.Float16)++ c[gm0, col1] = acc_out[8].to(cutlass.Float16)+ c[gm1, col1] = acc_out[9].to(cutlass.Float16)+ c[gm2, col1] = acc_out[10].to(cutlass.Float16)+ c[gm3, col1] = acc_out[11].to(cutlass.Float16)+ c[gm4, col1] = acc_out[12].to(cutlass.Float16)+ c[gm5, col1] = acc_out[13].to(cutlass.Float16)+ c[gm6, col1] = acc_out[14].to(cutlass.Float16)+ c[gm7, col1] = acc_out[15].to(cutlass.Float16)++class _GemmF16F16ToF32:def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:self.threads = int(threads)⋯ 174 unchanged linesacc16 = acc16 + a6 * bv1acc17 = acc17 + a7 * bv1- cute.arch.sync_threads()+ def _need_sync():+ cute.arch.sync_threads()++ if_generate(k0 + self.tile_k < k, _need_sync)yield_out([acc00,⋯ 38 unchanged lines_st_row(7, col1, acc_out[15])+ class _GemmF16F16ToF32_Full:+ def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:+ self.threads = int(threads)+ self.tile_m = int(tile_m)+ self.tile_n = int(tile_n)+ self.tile_k = int(tile_k)+ self.warps = self.threads // 32+ self.rows_per_warp = self.tile_m // self.warps+ @cute.jit+ def __call__(+ self,+ a_ptr: "cute.Pointer",+ b_ptr: "cute.Pointer",+ c_ptr: "cute.Pointer",+ problem: tuple,+ ):+ m, n, k = problem+ a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))+ b = cute.make_tensor(b_ptr, cute.make_layout((k, n), stride=(n, 1)))+ c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))+ grid_n = (n + self.tile_n - 1) // self.tile_n+ grid_m = (m + self.tile_m - 1) // self.tile_m+ self.kernel(a, b, c, m, n, k).launch(+ grid=[grid_n, grid_m, 1],+ block=[self.threads, 1, 1],+ )+ return+ @cute.kernel+ def kernel(+ self,+ a: "cute.Tensor",+ b: "cute.Tensor",+ c: "cute.Tensor",+ m: int,+ n: int,+ k: int,+ ):+ tx, _, _ = cute.arch.thread_idx()+ bx, by, _ = cute.arch.block_idx()+ warp_id = tx >> 5+ lane = tx & 31++ base_m = by * self.tile_m+ base_n = bx * self.tile_n++ row0 = base_m + warp_id * self.rows_per_warp+ col0 = base_n + lane+ col1 = col0 + 32++ acc00 = cutlass.Float32(0.0)+ acc01 = cutlass.Float32(0.0)+ acc02 = cutlass.Float32(0.0)+ acc03 = cutlass.Float32(0.0)+ acc04 = cutlass.Float32(0.0)+ acc05 = cutlass.Float32(0.0)+ acc06 = cutlass.Float32(0.0)+ acc07 = cutlass.Float32(0.0)+ acc10 = cutlass.Float32(0.0)+ acc11 = cutlass.Float32(0.0)+ acc12 = cutlass.Float32(0.0)+ acc13 = cutlass.Float32(0.0)+ acc14 = cutlass.Float32(0.0)+ acc15 = cutlass.Float32(0.0)+ acc16 = cutlass.Float32(0.0)+ acc17 = cutlass.Float32(0.0)++ smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)+ smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)+ smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))+ smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))++ for k0, acc, acc_out in for_generate(+ 0,+ k,+ self.tile_k,+ iter_args=[+ acc00,+ acc01,+ acc02,+ acc03,+ acc04,+ acc05,+ acc06,+ acc07,+ acc10,+ acc11,+ acc12,+ acc13,+ acc14,+ acc15,+ acc16,+ acc17,+ ],+ ):+ acc00 = acc[0]+ acc01 = acc[1]+ acc02 = acc[2]+ acc03 = acc[3]+ acc04 = acc[4]+ acc05 = acc[5]+ acc06 = acc[6]+ acc07 = acc[7]+ acc10 = acc[8]+ acc11 = acc[9]+ acc12 = acc[10]+ acc13 = acc[11]+ acc14 = acc[12]+ acc15 = acc[13]+ acc16 = acc[14]+ acc17 = acc[15]++ for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):+ mm = idx // self.tile_k+ kk = idx - mm * self.tile_k+ smem_a[mm, kk] = a[base_m + mm, k0 + kk]+ yield_out()++ for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):+ kk = idx // self.tile_n+ nn = idx - kk * self.tile_n+ smem_b[kk, nn] = b[k0 + kk, base_n + nn]+ yield_out()++ cute.arch.sync_threads()++ for kk in range_constexpr(128):+ bv0 = smem_b[kk, lane].to(cutlass.Float32)+ bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)++ a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)+ a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)+ a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)+ a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)+ a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)+ a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)+ a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)+ a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)++ acc00 = acc00 + a0 * bv0+ acc01 = acc01 + a1 * bv0+ acc02 = acc02 + a2 * bv0+ acc03 = acc03 + a3 * bv0+ acc04 = acc04 + a4 * bv0+ acc05 = acc05 + a5 * bv0+ acc06 = acc06 + a6 * bv0+ acc07 = acc07 + a7 * bv0++ acc10 = acc10 + a0 * bv1+ acc11 = acc11 + a1 * bv1+ acc12 = acc12 + a2 * bv1+ acc13 = acc13 + a3 * bv1+ acc14 = acc14 + a4 * bv1+ acc15 = acc15 + a5 * bv1+ acc16 = acc16 + a6 * bv1+ acc17 = acc17 + a7 * bv1++ def _need_sync():+ cute.arch.sync_threads()++ if_generate(k0 + self.tile_k < k, _need_sync)+ yield_out(+ [+ acc00,+ acc01,+ acc02,+ acc03,+ acc04,+ acc05,+ acc06,+ acc07,+ acc10,+ acc11,+ acc12,+ acc13,+ acc14,+ acc15,+ acc16,+ acc17,+ ]+ )++ gm0 = row0 + 0+ gm1 = row0 + 1+ gm2 = row0 + 2+ gm3 = row0 + 3+ gm4 = row0 + 4+ gm5 = row0 + 5+ gm6 = row0 + 6+ gm7 = row0 + 7++ c[gm0, col0] = acc_out[0]+ c[gm1, col0] = acc_out[1]+ c[gm2, col0] = acc_out[2]+ c[gm3, col0] = acc_out[3]+ c[gm4, col0] = acc_out[4]+ c[gm5, col0] = acc_out[5]+ c[gm6, col0] = acc_out[6]+ c[gm7, col0] = acc_out[7]++ c[gm0, col1] = acc_out[8]+ c[gm1, col1] = acc_out[9]+ c[gm2, col1] = acc_out[10]+ c[gm3, col1] = acc_out[11]+ c[gm4, col1] = acc_out[12]+ c[gm5, col1] = acc_out[13]+ c[gm6, col1] = acc_out[14]+ c[gm7, col1] = acc_out[15]+++++++++ _LOG2E = 1.4426950408889634++def _sigmoid_f16(x: cutlass.Float16) -> cutlass.Float16:++xx = x.to(cutlass.Float32)- ee = cute.exp(cutlass.Float32(0.0) - xx, fastmath=True)- yy = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + ee)- return yy.to(cutlass.Float16)+ t = (cutlass.Float32(0.0) - xx) * cutlass.Float32(_LOG2E)+ ee = cute.exp2(t, fastmath=True)+ denom = cutlass.Float32(1.0) + ee+ y = cute.arch.rcp_approx(denom)+ y = y * (cutlass.Float32(2.0) - denom * y)+ return y.to(cutlass.Float16)class _ProcessProj:⋯ 88 unchanged lines-class _ContractHiddenGemm:def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:self.threads = int(threads)⋯ 187 unchanged linesacc16 = acc16 + a6 * bv1acc17 = acc17 + a7 * bv1- cute.arch.sync_threads()+ def _need_sync():+ cute.arch.sync_threads()++ if_generate(k0 + self.tile_k < n, _need_sync)yield_out([acc00,⋯ 38 unchanged lines_st_row(7, col1, acc_out[15])+ class _ContractHiddenGemm_Full:+ def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:+ self.threads = int(threads)+ self.tile_m = int(tile_m)+ self.tile_n = int(tile_n)+ self.tile_k = int(tile_k)+ self.warps = self.threads // 32+ self.rows_per_warp = self.tile_m // self.warps+ @cute.jit+ def __call__(+ self,+ left_ptr: "cute.Pointer",+ right_ptr: "cute.Pointer",+ out_ptr: "cute.Pointer",+ problem: tuple,+ ):+ bs, n, h = problem+ left = cute.make_tensor(+ left_ptr,+ cute.make_layout((bs, n, h, n), stride=(n * h * n, h * n, n, 1)),+ )+ right = cute.make_tensor(+ right_ptr,+ cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),+ )+ out = cute.make_tensor(+ out_ptr,+ cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),+ )+ grid_n = (n + self.tile_n - 1) // self.tile_n+ grid_m = (n + self.tile_m - 1) // self.tile_m+ grid_z = bs * h+ self.kernel(left, right, out, bs, n, h).launch(+ grid=[grid_n, grid_m, grid_z],+ block=[self.threads, 1, 1],+ )+ return+ @cute.kernel+ def kernel(+ self,+ left: "cute.Tensor",+ right: "cute.Tensor",+ out: "cute.Tensor",+ bs: int,+ n: int,+ h: int,+ ):+ tx, _, _ = cute.arch.thread_idx()+ bx, by, bz = cute.arch.block_idx()+ warp_id = tx >> 5+ lane = tx & 31++ bb = bz // h+ hh = bz - bb * h++ base_m = by * self.tile_m+ base_n = bx * self.tile_n++ row0 = base_m + warp_id * self.rows_per_warp+ col0 = base_n + lane+ col1 = col0 + 32++ acc00 = cutlass.Float32(0.0)+ acc01 = cutlass.Float32(0.0)+ acc02 = cutlass.Float32(0.0)+ acc03 = cutlass.Float32(0.0)+ acc04 = cutlass.Float32(0.0)+ acc05 = cutlass.Float32(0.0)+ acc06 = cutlass.Float32(0.0)+ acc07 = cutlass.Float32(0.0)+ acc10 = cutlass.Float32(0.0)+ acc11 = cutlass.Float32(0.0)+ acc12 = cutlass.Float32(0.0)+ acc13 = cutlass.Float32(0.0)+ acc14 = cutlass.Float32(0.0)+ acc15 = cutlass.Float32(0.0)+ acc16 = cutlass.Float32(0.0)+ acc17 = cutlass.Float32(0.0)++ smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)+ smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)+ smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))+ smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))++ for k0, acc, acc_out in for_generate(+ 0,+ n,+ self.tile_k,+ iter_args=[+ acc00,+ acc01,+ acc02,+ acc03,+ acc04,+ acc05,+ acc06,+ acc07,+ acc10,+ acc11,+ acc12,+ acc13,+ acc14,+ acc15,+ acc16,+ acc17,+ ],+ ):+ acc00 = acc[0]+ acc01 = acc[1]+ acc02 = acc[2]+ acc03 = acc[3]+ acc04 = acc[4]+ acc05 = acc[5]+ acc06 = acc[6]+ acc07 = acc[7]+ acc10 = acc[8]+ acc11 = acc[9]+ acc12 = acc[10]+ acc13 = acc[11]+ acc14 = acc[12]+ acc15 = acc[13]+ acc16 = acc[14]+ acc17 = acc[15]++ for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):+ mm = idx // self.tile_k+ kk = idx - mm * self.tile_k+ smem_a[mm, kk] = left[bb, base_m + mm, hh, k0 + kk]+ yield_out()++ for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):+ kk = idx // self.tile_n+ nn = idx - kk * self.tile_n+ smem_b[kk, nn] = right[bb, hh, k0 + kk, base_n + nn]+ yield_out()++ cute.arch.sync_threads()++ for kk in range_constexpr(128):+ bv0 = smem_b[kk, lane].to(cutlass.Float32)+ bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)++ a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)+ a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)+ a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)+ a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)+ a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)+ a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)+ a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)+ a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)++ acc00 = acc00 + a0 * bv0+ acc01 = acc01 + a1 * bv0+ acc02 = acc02 + a2 * bv0+ acc03 = acc03 + a3 * bv0+ acc04 = acc04 + a4 * bv0+ acc05 = acc05 + a5 * bv0+ acc06 = acc06 + a6 * bv0+ acc07 = acc07 + a7 * bv0++ acc10 = acc10 + a0 * bv1+ acc11 = acc11 + a1 * bv1+ acc12 = acc12 + a2 * bv1+ acc13 = acc13 + a3 * bv1+ acc14 = acc14 + a4 * bv1+ acc15 = acc15 + a5 * bv1+ acc16 = acc16 + a6 * bv1+ acc17 = acc17 + a7 * bv1++ def _need_sync():+ cute.arch.sync_threads()++ if_generate(k0 + self.tile_k < n, _need_sync)+ yield_out(+ [+ acc00,+ acc01,+ acc02,+ acc03,+ acc04,+ acc05,+ acc06,+ acc07,+ acc10,+ acc11,+ acc12,+ acc13,+ acc14,+ acc15,+ acc16,+ acc17,+ ]+ )++ gi0 = row0 + 0+ gi1 = row0 + 1+ gi2 = row0 + 2+ gi3 = row0 + 3+ gi4 = row0 + 4+ gi5 = row0 + 5+ gi6 = row0 + 6+ gi7 = row0 + 7++ out[bb, hh, gi0, col0] = acc_out[0]+ out[bb, hh, gi1, col0] = acc_out[1]+ out[bb, hh, gi2, col0] = acc_out[2]+ out[bb, hh, gi3, col0] = acc_out[3]+ out[bb, hh, gi4, col0] = acc_out[4]+ out[bb, hh, gi5, col0] = acc_out[5]+ out[bb, hh, gi6, col0] = acc_out[6]+ out[bb, hh, gi7, col0] = acc_out[7]++ out[bb, hh, gi0, col1] = acc_out[8]+ out[bb, hh, gi1, col1] = acc_out[9]+ out[bb, hh, gi2, col1] = acc_out[10]+ out[bb, hh, gi3, col1] = acc_out[11]+ out[bb, hh, gi4, col1] = acc_out[12]+ out[bb, hh, gi5, col1] = acc_out[13]+ out[bb, hh, gi6, col1] = acc_out[14]+ out[bb, hh, gi7, col1] = acc_out[15]++++++++class _LayerNormHiddenF32ToF16:def __init__(self, threads: int = 256) -> None:self.threads = int(threads)⋯ 84 unchanged linesif_generate(idx < bs * n * n, _do_one)- class _FinalLinear:- def __init__(self) -> None:- self.gemm = _GemmF16F16ToF32()- def __call__(self, a: torch.Tensor, w: torch.Tensor, y: torch.Tensor) -> None:- m, k = a.shape- n = y.shape[1]- self.gemm(a, w, y, (m, n, k))---_LN_X = _LayerNormLastDimF32ToF16()_LN_X_C = None_GEMM_PROJ = _GemmF16F16ToF16()_GEMM_PROJ_C = None+ _GEMM_PROJ_FULL = _GemmF16F16ToF16_Full()+ _GEMM_PROJ_FULL_C = None_PROC = _ProcessProj()_PROC_C = None_CONTRACT = _ContractHiddenGemm()_CONTRACT_C = None+ _CONTRACT_FULL = _ContractHiddenGemm_Full()+ _CONTRACT_FULL_C = None_LN_H = _LayerNormHiddenF32ToF16()_LN_H_C = None_GEMM_OUT = _GemmF16F16ToF32()_GEMM_OUT_C = None+ _GEMM_OUT_FULL = _GemmF16F16ToF32_Full()+ _GEMM_OUT_FULL_C = Nonedef _compile_once():- global _LN_X_C, _GEMM_PROJ_C, _PROC_C, _CONTRACT_C, _LN_H_C, _GEMM_OUT_C+ global _LN_X_C, _GEMM_PROJ_C, _GEMM_PROJ_FULL_C, _PROC_C, _CONTRACT_C, _CONTRACT_FULL_C, _LN_H_C, _GEMM_OUT_C, _GEMM_OUT_FULL_Cif _LN_X_C is None:x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)⋯ 8 unchanged linesc_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)_GEMM_PROJ_C = cute.compile(_GEMM_PROJ, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")+ if _GEMM_PROJ_FULL_C is None:+ a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ c_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ _GEMM_PROJ_FULL_C = cute.compile(_GEMM_PROJ_FULL, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")+if _PROC_C is None:proj_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)mask_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)⋯ 17 unchanged linesout_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)_CONTRACT_C = cute.compile(_CONTRACT, left_ptr, right_ptr, out_ptr, (0, 0, 0), options="--opt-level 3")+ if _CONTRACT_FULL_C is None:+ left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ out_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)+ _CONTRACT_FULL_C = cute.compile(_CONTRACT_FULL, left_ptr, right_ptr, out_ptr, (0, 0, 0), options="--opt-level 3")+if _LN_H_C is None:x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)w_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)⋯ 8 unchanged linesc_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)_GEMM_OUT_C = cute.compile(_GEMM_OUT, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")+ if _GEMM_OUT_FULL_C is None:+ a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)+ c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)+ _GEMM_OUT_FULL_C = cute.compile(_GEMM_OUT_FULL, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")+def _as_ptr(ty, t: torch.Tensor):return make_ptr(ty, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)⋯ 129 unchanged lines)m = bs * n * n- _GEMM_PROJ_C(- _as_ptr(cutlass.Float16, x_norm.view(m, dim)),- _as_ptr(cutlass.Float16, w_pack16),- _as_ptr(cutlass.Float16, proj),- (m, 5 * hidden, dim),- )+ use_proj_full = (m % 64 == 0) & ((5 * hidden) % 64 == 0) & (dim % 128 == 0)+ if use_proj_full:+ _GEMM_PROJ_FULL_C(+ _as_ptr(cutlass.Float16, x_norm.view(m, dim)),+ _as_ptr(cutlass.Float16, w_pack16),+ _as_ptr(cutlass.Float16, proj),+ (m, 5 * hidden, dim),+ )+ else:+ _GEMM_PROJ_C(+ _as_ptr(cutlass.Float16, x_norm.view(m, dim)),+ _as_ptr(cutlass.Float16, w_pack16),+ _as_ptr(cutlass.Float16, proj),+ (m, 5 * hidden, dim),+ )+mask16.copy_(mask)_PROC_C(_as_ptr(cutlass.Float16, proj.view(bs, n, n, 5 * hidden)),⋯ 4 unchanged lines(bs, n, hidden),)- _CONTRACT_C(- _as_ptr(cutlass.Float16, left_t),- _as_ptr(cutlass.Float16, right_t),- _as_ptr(cutlass.Float32, out_tmp),- (bs, n, hidden),- )+ use_contract_full = (n % 64 == 0) & (n % 128 == 0)+ if use_contract_full:+ _CONTRACT_FULL_C(+ _as_ptr(cutlass.Float16, left_t),+ _as_ptr(cutlass.Float16, right_t),+ _as_ptr(cutlass.Float32, out_tmp),+ (bs, n, hidden),+ )+ else:+ _CONTRACT_C(+ _as_ptr(cutlass.Float16, left_t),+ _as_ptr(cutlass.Float16, right_t),+ _as_ptr(cutlass.Float32, out_tmp),+ (bs, n, hidden),+ )_LN_H_C(_as_ptr(cutlass.Float32, out_tmp),⋯ 5 unchanged lines)y = torch.empty((m, dim), device=x.device, dtype=torch.float32)- _GEMM_OUT_C(- _as_ptr(cutlass.Float16, out_norm.view(m, hidden)),- _as_ptr(cutlass.Float16, w_out_t16),- _as_ptr(cutlass.Float32, y),- (m, dim, hidden),- )+ use_out_full = (m % 64 == 0) & (dim % 64 == 0) & (hidden % 128 == 0)+ if use_out_full:+ _GEMM_OUT_FULL_C(+ _as_ptr(cutlass.Float16, out_norm.view(m, hidden)),+ _as_ptr(cutlass.Float16, w_out_t16),+ _as_ptr(cutlass.Float32, y),+ (m, dim, hidden),+ )+ else:+ _GEMM_OUT_C(+ _as_ptr(cutlass.Float16, out_norm.view(m, hidden)),+ _as_ptr(cutlass.Float16, w_out_t16),+ _as_ptr(cutlass.Float32, y),+ (m, dim, hidden),+ )return y.view(bs, n, n, dim)__all__ = ["custom_kernel"]+
scrolls · 930 diff lines total
Best evidence level for this revision: reported
JSON