Skip to content
KernelIndex
Search⌘K

submission 74120

irregular · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 269 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-grayscale-v2-74120?include=source"
interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RGB to grayscalesuite of 6 cases
NVIDIA H100
12.9ms
#33 of 36
2025-11-12

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d908c0b1ecc41deb07223296b3356dc50ac21a4de4e67c246098288ed2f3b1d7
license declaredunknown
license concludedunknown
authorsirregular
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4• Ranked NVFP4 GEMV: (a[M,K,L], b[1,K,L], sfa[M,K//16,L], sfb[1,K//16,L], c[M,1,L])

Kernel source

submission.py269 lines
# # !POPCORN leaderboard ranked
# import torch

# # ---------------------------------------------------------------------------
# # Batched GEMV in NVFP4(e2m1) with FP8(E4M3 fnuz) block scales using
# # Blackwell's torch._scaled_mm() fast path.
# # ---------------------------------------------------------------------------

# SF_VEC = 16  # per-16 K elements

# def ceil_div(a: int, b: int) -> int:
#     return (a + b - 1) // b

# def _to_blocked(sf_2d: torch.Tensor) -> torch.Tensor:
#     """
#     Convert FP8 scale factors from [rows, K//16] K-major to the flattened
#     CuTe/Blackwell block layout expected by torch._scaled_mm (view/permute/reshape only).
#     """
#     rows, sf_k = sf_2d.shape
#     n_row_blocks = ceil_div(rows, 128)
#     n_col_blocks = ceil_div(sf_k, 4)
#     # [nrb,128,ncb,4] -> [nrb,ncb,128,4]
#     t = sf_2d.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
#     # [*,128,4] -> [*,4,32,4] -> swap -> [*,32,4,4] -> [*,32,16] -> flat
#     t = t.reshape(-1, 128, 4).view(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
#     return t.flatten()

# # cached grayscale weights for 2-tensor practice probes
# _GRAYSCALE_W = None

# def _scaled_gemv_N1(A_mk, b1k, sfa_mk16, sfb_1k16, scale_a_blocked=None):
#     """
#     Fast path (preferred): N=1 directly with _scaled_mm. Returns [M,1] fp16.
#     """
#     scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
#     scale_b = _to_blocked(sfb_1k16)
#     return torch._scaled_mm(
#         A_mk, b1k, scale_a, scale_b, bias=None, out_dtype=torch.float16
#     )  # -> [M,1]

# def _scaled_gemv_N128(A_mk, b1k, sfa_mk16, sfb_1k16, scratch_B128K, scratch_SFB128, scale_a_blocked=None):
#     """
#     Fallback path: pad N to 128 using reusable scratch buffers.
#     Returns [M,1] fp16 (narrowed view of GEMM result).
#     """
#     # reset scratch cheaply
#     scratch_B128K.zero_()
#     scratch_SFB128.fill_(1)

#     # real data on row 0 only
#     scratch_B128K[0, :].copy_(b1k[0, :])
#     scratch_SFB128[0, :].copy_(sfb_1k16[0, :])

#     scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
#     scale_b = _to_blocked(scratch_SFB128)

#     outMN = torch._scaled_mm(
#         A_mk, scratch_B128K, scale_a, scale_b, bias=None, out_dtype=torch.float16
#     )  # [M,128]
#     return outMN[:, :1]

# def custom_kernel(data):
#     """
#     Supports:
#       • Ranked NVFP4 GEMV: (a[M,K,L], b[1,K,L], sfa[M,K//16,L], sfb[1,K//16,L], c[M,1,L])
#       • Practice probe (2-tensor): (x[H,W,3], out[H,W]) → grayscale
#     """
#     # ---- 2-tensor practice/warmup probe -----------------------------------
#     if len(data) == 2:
#         x, out = data
#         if x.ndim == 3 and x.shape[-1] == 3:
#             global _GRAYSCALE_W
#             if _GRAYSCALE_W is None or _GRAYSCALE_W.device != x.device or _GRAYSCALE_W.dtype != x.dtype:
#                 _GRAYSCALE_W = torch.tensor([0.2989, 0.5870, 0.1140], device=x.device, dtype=x.dtype)
#             out.copy_(torch.einsum("hwc,c->hw", x, _GRAYSCALE_W))
#         else:
#             out.copy_(x)
#         return out

#     # ---- ranked path: 5 tensors -------------------------------------------
#     a, b, sfa, sfb, c = data
#     M, K, L = a.shape

#     # minimal sanity checks
#     assert b.shape[0] == 1
#     assert a.dtype == torch.float4_e2m1fn_x2 and b.dtype == torch.float4_e2m1fn_x2
#     assert c.dtype == torch.float16
#     assert sfa.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))
#     assert sfb.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))

#     N_PAD = 128
#     # scratch reused across all batches for the padded fallback
#     scratch_B128K = torch.empty((N_PAD, K), device=b.device, dtype=b.dtype)
#     scratch_SFB128 = torch.empty((N_PAD, K // SF_VEC), device=sfb.device, dtype=sfb.dtype)

#     # try N=1 once; if it throws, stick to padded fallback
#     use_N1 = True

#     # precompute scale_a (blocked) per batch once, so we don't redo it on both paths
#     scale_a_blocked = [None] * L
#     for l in range(L):
#         scale_a_blocked[l] = _to_blocked(sfa[:, :, l].contiguous())

#     for l in range(L):
#         A_l   = a[:, :, l].contiguous()       # [M,K] nvfp4
#         b_l   = b[:, :, l].contiguous()       # [1,K] nvfp4
#         sfa_l = sfa[:, :, l].contiguous()     # [M,K//16] fp8
#         sfb_l = sfb[:, :, l].contiguous()     # [1,K//16] fp8

#         if use_N1:
#             try:
#                 outM1 = _scaled_gemv_N1(A_l, b_l, sfa_l, sfb_l, scale_a_blocked=scale_a_blocked[l])
#             except Exception:
#                 use_N1 = False
#             else:
#                 c[:, 0, l].copy_(outM1[:, 0])
#                 continue

#         outM1 = _scaled_gemv_N128(
#             A_l, b_l, sfa_l, sfb_l,
#             scratch_B128K, scratch_SFB128,
#             scale_a_blocked=scale_a_blocked[l]
#         )
#         c[:, 0, l].copy_(outM1[:, 0])

#     return c
# !POPCORN leaderboard ranked
import torch

# ---------------------------------------------------------------------------
# Batched GEMV in NVFP4(e2m1) with FP8(E4M3 fnuz) block scaling on B200
# via torch._scaled_mm (Blackwell tensor core path).
#
# Inputs (a, b, sfa, sfb, c):
#   a   : [M, K, L], dtype=float4_e2m1fn_x2,  K-major
#   b   : [1, K, L], dtype=float4_e2m1fn_x2,  K-major
#   sfa : [M, K//16, L], dtype=float8_e4m3fn/float8_e4m3fnuz
#   sfb : [1, K//16, L], dtype=float8_e4m3fn/float8_e4m3fnuz
#   c   : [M, 1, L], dtype=float16
#
# Computes per batch l:
#   c[:,0,l] = sum_k ( a[:,k,l]*sfa[:,k//16,l] * b[0,k,l]*sfb[0,k//16,l] )
#
# Notes:
# • Prefers true GEMV (N=1) if runner supports it; otherwise falls back to N=128
#   with reusable scratch to avoid allocator overhead.
# • Reorders FP8 scaling factors into CuTe/Blackwell blocked layout that
#   torch._scaled_mm expects.
# ---------------------------------------------------------------------------

SF_VEC = 16  # scale granularity: per 16 K elements

def ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b

def _to_blocked(sf_2d: torch.Tensor) -> torch.Tensor:
    """
    Convert FP8 scale factors from [rows, K//16] K-major to the flattened
    CuTe/Blackwell block layout expected by torch._scaled_mm (view/permute/reshape only).
    """
    rows, sf_k = sf_2d.shape
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(sf_k, 4)
    # [nrb,128,ncb,4] -> [nrb,ncb,128,4]
    t = sf_2d.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    # [*,128,4] -> [*,4,32,4] -> swap -> [*,32,4,4] -> [*,32,16] -> flat
    t = t.reshape(-1, 128, 4).view(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
    return t.flatten()

# cached weights for 2-tensor grayscale probe (if harness sanity-checks)
_GRAYSCALE_W = None

def _scaled_gemv_N1(A_mk, b1k, sfa_mk16, sfb_1k16, *, scale_a_blocked=None):
    """
    Preferred fast path: run torch._scaled_mm with N=1 directly. Returns [M,1] fp16.
    """
    scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
    scale_b = _to_blocked(sfb_1k16)
    return torch._scaled_mm(
        A_mk, b1k, scale_a, scale_b, bias=None, out_dtype=torch.float16
    )

def _scaled_gemv_N128(A_mk, b1k, sfa_mk16, sfb_1k16, *, scratch_B128K, scratch_SFB128, scale_a_blocked=None):
    """
    Fallback: pad N to 128 using reusable scratch buffers. Returns [M,1] fp16.
    """
    # reset scratch cheaply
    scratch_B128K.zero_()
    scratch_SFB128.fill_(1)

    # real data on row 0 only
    scratch_B128K[0, :].copy_(b1k[0, :])
    scratch_SFB128[0, :].copy_(sfb_1k16[0, :])

    scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
    scale_b = _to_blocked(scratch_SFB128)

    outMN = torch._scaled_mm(
        A_mk, scratch_B128K, scale_a, scale_b, bias=None, out_dtype=torch.float16
    )  # [M,128]
    return outMN[:, :1]

def custom_kernel(data):
    """
    Supports:
      • Ranked NVFP4 GEMV: (a[M,K,L], b[1,K,L], sfa[M,K//16,L], sfb[1,K//16,L], c[M,1,L])
      • Practice probe (2-tensor): (x[H,W,3], out[H,W]) → grayscale
    """
    # ---- 2-tensor practice/warmup probe -----------------------------------
    if len(data) == 2:
        x, out = data
        if x.ndim == 3 and x.shape[-1] == 3:
            global _GRAYSCALE_W
            if _GRAYSCALE_W is None or _GRAYSCALE_W.device != x.device or _GRAYSCALE_W.dtype != x.dtype:
                _GRAYSCALE_W = torch.tensor([0.2989, 0.5870, 0.1140], device=x.device, dtype=x.dtype)
            out.copy_(torch.einsum("hwc,c->hw", x, _GRAYSCALE_W))
        else:
            out.copy_(x)
        return out

    # ---- ranked path: 5 tensors -------------------------------------------
    a, b, sfa, sfb, c = data
    M, K, L = a.shape

    # minimal sanity checks
    assert b.shape[0] == 1
    assert a.dtype == torch.float4_e2m1fn_x2 and b.dtype == torch.float4_e2m1fn_x2
    assert c.dtype == torch.float16
    assert sfa.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))
    assert sfb.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))

    # reusable scratch for padded fallback (allocated once)
    N_PAD = 128
    scratch_B128K = torch.empty((N_PAD, K), device=b.device, dtype=b.dtype)
    scratch_SFB128 = torch.empty((N_PAD, K // SF_VEC), device=sfb.device, dtype=sfb.dtype)

    # try fast N=1 once; if it throws, stick to padded path thereafter
    use_N1 = True

    # precompute scale_a (blocked) per batch once (avoids redoing on fallback)
    scale_a_blocked = [None] * L
    for l in range(L):
        scale_a_blocked[l] = _to_blocked(sfa[:, :, l].contiguous())

    for l in range(L):
        A_l   = a[:, :, l].contiguous()       # [M,K] nvfp4
        b_l   = b[:, :, l].contiguous()       # [1,K] nvfp4
        sfa_l = sfa[:, :, l].contiguous()     # [M,K//16] fp8
        sfb_l = sfb[:, :, l].contiguous()     # [1,K//16] fp8

        if use_N1:
            try:
                outM1 = _scaled_gemv_N1(A_l, b_l, sfa_l, sfb_l, scale_a_blocked=scale_a_blocked[l])
            except Exception:
                use_N1 = False
            else:
                c[:, 0, l].copy_(outM1[:, 0])
                continue

        outM1 = _scaled_gemv_N128(
            A_l, b_l, sfa_l, sfb_l,
            scratch_B128K=scratch_B128K,
            scratch_SFB128=scratch_SFB128,
            scale_a_blocked=scale_a_blocked[l],
        )
        c[:, 0, l].copy_(outM1[:, 0])

    return c
scrolls · 269 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 74117.

+ # # !POPCORN leaderboard ranked
+ # import torch
+
+ # # ---------------------------------------------------------------------------
+ # # Batched GEMV in NVFP4(e2m1) with FP8(E4M3 fnuz) block scales using
+ # # Blackwell's torch._scaled_mm() fast path.
+ # # ---------------------------------------------------------------------------
+
+ # SF_VEC = 16 # per-16 K elements
+
+ # def ceil_div(a: int, b: int) -> int:
+ # return (a + b - 1) // b
+
+ # def _to_blocked(sf_2d: torch.Tensor) -> torch.Tensor:
+ # """
+ # Convert FP8 scale factors from [rows, K//16] K-major to the flattened
+ # CuTe/Blackwell block layout expected by torch._scaled_mm (view/permute/reshape only).
+ # """
+ # rows, sf_k = sf_2d.shape
+ # n_row_blocks = ceil_div(rows, 128)
+ # n_col_blocks = ceil_div(sf_k, 4)
+ # # [nrb,128,ncb,4] -> [nrb,ncb,128,4]
+ # t = sf_2d.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
+ # # [*,128,4] -> [*,4,32,4] -> swap -> [*,32,4,4] -> [*,32,16] -> flat
+ # t = t.reshape(-1, 128, 4).view(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
+ # return t.flatten()
+
+ # # cached grayscale weights for 2-tensor practice probes
+ # _GRAYSCALE_W = None
+
+ # def _scaled_gemv_N1(A_mk, b1k, sfa_mk16, sfb_1k16, scale_a_blocked=None):
+ # """
+ # Fast path (preferred): N=1 directly with _scaled_mm. Returns [M,1] fp16.
+ # """
+ # scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
+ # scale_b = _to_blocked(sfb_1k16)
+ # return torch._scaled_mm(
+ # A_mk, b1k, scale_a, scale_b, bias=None, out_dtype=torch.float16
+ # ) # -> [M,1]
+
+ # def _scaled_gemv_N128(A_mk, b1k, sfa_mk16, sfb_1k16, scratch_B128K, scratch_SFB128, scale_a_blocked=None):
+ # """
+ # Fallback path: pad N to 128 using reusable scratch buffers.
+ # Returns [M,1] fp16 (narrowed view of GEMM result).
+ # """
+ # # reset scratch cheaply
+ # scratch_B128K.zero_()
+ # scratch_SFB128.fill_(1)
+
+ # # real data on row 0 only
+ # scratch_B128K[0, :].copy_(b1k[0, :])
+ # scratch_SFB128[0, :].copy_(sfb_1k16[0, :])
+
+ # scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
+ # scale_b = _to_blocked(scratch_SFB128)
+
+ # outMN = torch._scaled_mm(
+ # A_mk, scratch_B128K, scale_a, scale_b, bias=None, out_dtype=torch.float16
+ # ) # [M,128]
+ # return outMN[:, :1]
+
+ # def custom_kernel(data):
+ # """
+ # Supports:
+ # • Ranked NVFP4 GEMV: (a[M,K,L], b[1,K,L], sfa[M,K//16,L], sfb[1,K//16,L], c[M,1,L])
+ # • Practice probe (2-tensor): (x[H,W,3], out[H,W]) → grayscale
+ # """
+ # # ---- 2-tensor practice/warmup probe -----------------------------------
+ # if len(data) == 2:
+ # x, out = data
+ # if x.ndim == 3 and x.shape[-1] == 3:
+ # global _GRAYSCALE_W
+ # if _GRAYSCALE_W is None or _GRAYSCALE_W.device != x.device or _GRAYSCALE_W.dtype != x.dtype:
+ # _GRAYSCALE_W = torch.tensor([0.2989, 0.5870, 0.1140], device=x.device, dtype=x.dtype)
+ # out.copy_(torch.einsum("hwc,c->hw", x, _GRAYSCALE_W))
+ # else:
+ # out.copy_(x)
+ # return out
+
+ # # ---- ranked path: 5 tensors -------------------------------------------
+ # a, b, sfa, sfb, c = data
+ # M, K, L = a.shape
+
+ # # minimal sanity checks
+ # assert b.shape[0] == 1
+ # assert a.dtype == torch.float4_e2m1fn_x2 and b.dtype == torch.float4_e2m1fn_x2
+ # assert c.dtype == torch.float16
+ # assert sfa.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))
+ # assert sfb.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))
+
+ # N_PAD = 128
+ # # scratch reused across all batches for the padded fallback
+ # scratch_B128K = torch.empty((N_PAD, K), device=b.device, dtype=b.dtype)
+ # scratch_SFB128 = torch.empty((N_PAD, K // SF_VEC), device=sfb.device, dtype=sfb.dtype)
+
+ # # try N=1 once; if it throws, stick to padded fallback
+ # use_N1 = True
+
+ # # precompute scale_a (blocked) per batch once, so we don't redo it on both paths
+ # scale_a_blocked = [None] * L
+ # for l in range(L):
+ # scale_a_blocked[l] = _to_blocked(sfa[:, :, l].contiguous())
+
+ # for l in range(L):
+ # A_l = a[:, :, l].contiguous() # [M,K] nvfp4
+ # b_l = b[:, :, l].contiguous() # [1,K] nvfp4
+ # sfa_l = sfa[:, :, l].contiguous() # [M,K//16] fp8
+ # sfb_l = sfb[:, :, l].contiguous() # [1,K//16] fp8
+
+ # if use_N1:
+ # try:
+ # outM1 = _scaled_gemv_N1(A_l, b_l, sfa_l, sfb_l, scale_a_blocked=scale_a_blocked[l])
+ # except Exception:
+ # use_N1 = False
+ # else:
+ # c[:, 0, l].copy_(outM1[:, 0])
+ # continue
+
+ # outM1 = _scaled_gemv_N128(
+ # A_l, b_l, sfa_l, sfb_l,
+ # scratch_B128K, scratch_SFB128,
+ # scale_a_blocked=scale_a_blocked[l]
+ # )
+ # c[:, 0, l].copy_(outM1[:, 0])
+
+ # return c
# !POPCORN leaderboard ranked
import torch
# ---------------------------------------------------------------------------
- # Batched GEMV in NVFP4(e2m1) with FP8(E4M3 fnuz) block scales using
- # Blackwell's torch._scaled_mm() fast path.
+ # Batched GEMV in NVFP4(e2m1) with FP8(E4M3 fnuz) block scaling on B200
+ # via torch._scaled_mm (Blackwell tensor core path).
+ #
+ # Inputs (a, b, sfa, sfb, c):
+ # a : [M, K, L], dtype=float4_e2m1fn_x2, K-major
+ # b : [1, K, L], dtype=float4_e2m1fn_x2, K-major
+ # sfa : [M, K//16, L], dtype=float8_e4m3fn/float8_e4m3fnuz
+ # sfb : [1, K//16, L], dtype=float8_e4m3fn/float8_e4m3fnuz
+ # c : [M, 1, L], dtype=float16
+ #
+ # Computes per batch l:
+ # c[:,0,l] = sum_k ( a[:,k,l]*sfa[:,k//16,l] * b[0,k,l]*sfb[0,k//16,l] )
+ #
+ # Notes:
+ # • Prefers true GEMV (N=1) if runner supports it; otherwise falls back to N=128
+ # with reusable scratch to avoid allocator overhead.
+ # • Reorders FP8 scaling factors into CuTe/Blackwell blocked layout that
+ # torch._scaled_mm expects.
# ---------------------------------------------------------------------------
- SF_VEC = 16 # per-16 K elements
+ SF_VEC = 16 # scale granularity: per 16 K elements
def ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
⋯ 12 unchanged lines
t = t.reshape(-1, 128, 4).view(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return t.flatten()
- # cached grayscale weights for 2-tensor practice probes
+ # cached weights for 2-tensor grayscale probe (if harness sanity-checks)
_GRAYSCALE_W = None
- def _scaled_gemv_N1(A_mk, b1k, sfa_mk16, sfb_1k16, scale_a_blocked=None):
+ def _scaled_gemv_N1(A_mk, b1k, sfa_mk16, sfb_1k16, *, scale_a_blocked=None):
"""
- Fast path (preferred): N=1 directly with _scaled_mm. Returns [M,1] fp16.
+ Preferred fast path: run torch._scaled_mm with N=1 directly. Returns [M,1] fp16.
"""
scale_a = scale_a_blocked if scale_a_blocked is not None else _to_blocked(sfa_mk16)
scale_b = _to_blocked(sfb_1k16)
return torch._scaled_mm(
A_mk, b1k, scale_a, scale_b, bias=None, out_dtype=torch.float16
- ) # -> [M,1]
+ )
- def _scaled_gemv_N128(A_mk, b1k, sfa_mk16, sfb_1k16, scratch_B128K, scratch_SFB128, scale_a_blocked=None):
+ def _scaled_gemv_N128(A_mk, b1k, sfa_mk16, sfb_1k16, *, scratch_B128K, scratch_SFB128, scale_a_blocked=None):
"""
- Fallback path: pad N to 128 using reusable scratch buffers.
- Returns [M,1] fp16 (narrowed view of GEMM result).
+ Fallback: pad N to 128 using reusable scratch buffers. Returns [M,1] fp16.
"""
# reset scratch cheaply
scratch_B128K.zero_()
⋯ 40 unchanged lines
assert sfa.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))
assert sfb.dtype in (torch.float8_e4m3fn, getattr(torch, "float8_e4m3fnuz", torch.float8_e4m3fn))
+ # reusable scratch for padded fallback (allocated once)
N_PAD = 128
- # scratch reused across all batches for the padded fallback
scratch_B128K = torch.empty((N_PAD, K), device=b.device, dtype=b.dtype)
scratch_SFB128 = torch.empty((N_PAD, K // SF_VEC), device=sfb.device, dtype=sfb.dtype)
- # try N=1 once; if it throws, stick to padded fallback
+ # try fast N=1 once; if it throws, stick to padded path thereafter
use_N1 = True
- # precompute scale_a (blocked) per batch once, so we don't redo it on both paths
+ # precompute scale_a (blocked) per batch once (avoids redoing on fallback)
scale_a_blocked = [None] * L
for l in range(L):
scale_a_blocked[l] = _to_blocked(sfa[:, :, l].contiguous())
⋯ 15 unchanged lines
outM1 = _scaled_gemv_N128(
A_l, b_l, sfa_l, sfb_l,
- scratch_B128K, scratch_SFB128,
- scale_a_blocked=scale_a_blocked[l]
+ scratch_B128K=scratch_B128K,
+ scratch_SFB128=scratch_SFB128,
+ scale_a_blocked=scale_a_blocked[l],
)
c[:, 0, l].copy_(outM1[:, 0])
scrolls · 218 diff lines total

Best evidence level for this revision: reported

JSON