Skip to content
KernelIndex
Search⌘K

submission 821679

David Xia · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

blocked.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-821679?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
38.4ms
#377 of 515
2026-06-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:62ba7fb676a944f1f3cffc73db4764023255099cf56d942cdd5122135bc5c6de
license declaredunknown
license concludedunknown
authorsDavid Xia
imported2026-08-26

Techniques

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

shared-memory__shared__ float smem_sum[32];

Kernel source

blocked.py403 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

# =====================================================================
# 1. 纯手工 CUDA 内核:Panel、V矩阵提取
# =====================================================================
cuda_source = """
#include <cuda_runtime.h>
#include <math.h>
#include <stdint.h>

__inline__ __device__ float blockReduceSum(float val, float* shared) {
    int lane = threadIdx.x & 31;
    int wid = threadIdx.x >> 5;
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        val += __shfl_down_sync(0xffffffff, val, offset);
    }
    if (lane == 0) shared[wid] = val;
    __syncthreads();
    
    int nwarps = (blockDim.x + 31) >> 5;
    val = (wid == 0 && lane < nwarps) ? shared[lane] : 0.0f;
    if (wid == 0) {
        #pragma unroll
        for (int offset = 16; offset > 0; offset >>= 1) {
            val += __shfl_down_sync(0xffffffff, val, offset);
        }
    }
    return val;
}

__inline__ __device__ float blockReduceMax(float val, float* shared) {
    int lane = threadIdx.x & 31;
    int wid = threadIdx.x >> 5;
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        val = fmaxf(val, __shfl_down_sync(0xffffffff, val, offset));
    }
    if (lane == 0) shared[wid] = val;
    __syncthreads();
    
    int nwarps = (blockDim.x + 31) >> 5;
    val = (wid == 0 && lane < nwarps) ? shared[lane] : 0.0f;
    if (wid == 0) {
        #pragma unroll
        for (int offset = 16; offset > 0; offset >>= 1) {
            val = fmaxf(val, __shfl_down_sync(0xffffffff, val, offset));
        }
    }
    return val;
}

__global__ void unblocked_qr_kernel(float* __restrict__ A, float* __restrict__ tau, int batch, int n) {
    int b_idx = blockIdx.x;
    if (b_idx >= batch) return;
    
    float* A_b = A + b_idx * n * n;
    float* tau_b = tau + b_idx * n;

    __shared__ float smem_sum[32];
    __shared__ float smem_max[32];

    if (n == 1) {
        if (threadIdx.x == 0) tau_b[0] = 0.0f;
        return;
    }

    for (int k = 0; k < n; ++k) {
        float local_max = 0.0f;
        for (int i = k + 1 + threadIdx.x; i < n; i += blockDim.x) {
            local_max = fmaxf(local_max, fabsf(A_b[i * n + k]));
        }
        float max_val = blockReduceMax(local_max, smem_max);
        if (threadIdx.x == 0) smem_max[0] = max_val;
        __syncthreads();
        max_val = smem_max[0];
        if (max_val == 0.0f) max_val = 1.0f;

        float local_sum = 0.0f;
        for (int i = k + 1 + threadIdx.x; i < n; i += blockDim.x) {
            float v = A_b[i * n + k] / max_val;
            local_sum += v * v;
        }
        float tail_sq = blockReduceSum(local_sum, smem_sum);
        __syncthreads();

        __shared__ float s_tau, s_v0;
        if (threadIdx.x == 0) {
            tail_sq *= (max_val * max_val);
            float alpha = A_b[k * n + k];
            float xnorm = sqrtf(tail_sq);
            if (xnorm == 0.0f) {
                s_tau = 0.0f; s_v0 = 1.0f;
            } else {
                float beta = -copysignf(hypotf(alpha, xnorm), alpha);
                s_tau = (beta - alpha) / beta;
                s_v0 = alpha - beta;
                A_b[k * n + k] = beta;
            }
            tau_b[k] = s_tau;
        }
        __syncthreads();

        if (s_tau != 0.0f) {
            float inv_v0 = 1.0f / s_v0;
            for (int i = k + 1 + threadIdx.x; i < n; i += blockDim.x) A_b[i * n + k] *= inv_v0;
        }
        __syncthreads();

        if (s_tau != 0.0f) {
            for (int j = k + 1 + threadIdx.x; j < n; j += blockDim.x) {
                float dot = A_b[k * n + j];
                for (int i = k + 1; i < n; ++i) dot += A_b[i * n + k] * A_b[i * n + j];
                float factor = s_tau * dot;
                A_b[k * n + j] -= factor;
                for (int i = k + 1; i < n; ++i) A_b[i * n + j] -= factor * A_b[i * n + k];
            }
        }
        __syncthreads();
    }
}

__global__ void panel_wy_kernel(float* __restrict__ A, float* __restrict__ tau, float* __restrict__ T_out, 
                                int batch, int n, int k, int b, int block_size) 
{
    int b_idx = blockIdx.x;
    if (b_idx >= batch) return;
    
    float* A_b = A + b_idx * n * n;
    float* tau_b = tau + b_idx * n;
    float* T_b = T_out + b_idx * block_size * block_size;

    __shared__ float smem_sum[32];
    __shared__ float smem_max[32];
    __shared__ float s_w[64]; 
    __shared__ float s_T[4096]; 

    for (int i = threadIdx.x; i < b * b; i += blockDim.x) s_T[i] = 0.0f;
    __syncthreads();

    for (int j = 0; j < b; ++j) {
        int col_idx = k + j;

        if (col_idx == n - 1) {
            if (threadIdx.x == 0) { tau_b[col_idx] = 0.0f; s_T[j * b + j] = 0.0f; }
            continue;
        }

        float local_max = 0.0f;
        for (int i = col_idx + 1 + threadIdx.x; i < n; i += blockDim.x) {
            local_max = fmaxf(local_max, fabsf(A_b[i * n + col_idx]));
        }
        float max_val = blockReduceMax(local_max, smem_max);
        if (threadIdx.x == 0) smem_max[0] = max_val;
        __syncthreads();
        max_val = smem_max[0];
        if (max_val == 0.0f) max_val = 1.0f;

        float local_sum = 0.0f;
        for (int i = col_idx + 1 + threadIdx.x; i < n; i += blockDim.x) {
            float val = A_b[i * n + col_idx] / max_val;
            local_sum += val * val;
        }
        float tail_sq = blockReduceSum(local_sum, smem_sum);
        __syncthreads();
        
        __shared__ float s_tau, s_v0;
        if (threadIdx.x == 0) {
            tail_sq *= (max_val * max_val);
            float alpha = A_b[col_idx * n + col_idx];
            float xnorm = sqrtf(tail_sq);
            if (xnorm == 0.0f) {
                s_tau = 0.0f; s_v0 = 1.0f;
            } else {
                float beta = -copysignf(hypotf(alpha, xnorm), alpha);
                s_tau = (beta - alpha) / beta;
                s_v0 = alpha - beta;
                A_b[col_idx * n + col_idx] = beta;
            }
            tau_b[col_idx] = s_tau;
            s_T[j * b + j] = s_tau; 
        }
        __syncthreads();

        if (s_tau != 0.0f) {
            float inv_v0 = 1.0f / s_v0;
            for (int i = col_idx + 1 + threadIdx.x; i < n; i += blockDim.x) A_b[i * n + col_idx] *= inv_v0;
        }
        __syncthreads();

        if (s_tau != 0.0f) {
            for (int p = j + 1 + threadIdx.x; p < b; p += blockDim.x) {
                int target_col = k + p;
                float dot = A_b[col_idx * n + target_col]; 
                for (int i = col_idx + 1; i < n; ++i) dot += A_b[i * n + col_idx] * A_b[i * n + target_col];
                float factor = s_tau * dot;
                A_b[col_idx * n + target_col] -= factor;
                for (int i = col_idx + 1; i < n; ++i) A_b[i * n + target_col] -= factor * A_b[i * n + col_idx];
            }
        }
        __syncthreads();

        if (s_tau != 0.0f && j > 0) {
            if (threadIdx.x < j) {
                int i = threadIdx.x; 
                float d = A_b[(k + j) * n + (k + i)] * 1.0f; 
                for (int r = k + j + 1; r < n; ++r) d += A_b[r * n + (k + i)] * A_b[r * n + (k + j)];
                s_w[i] = d;
            }
            __syncthreads();

            if (threadIdx.x < j) {
                int i = threadIdx.x;
                float val = 0.0f;
                for (int m = i; m < j; ++m) val += s_T[i * b + m] * s_w[m];
                s_T[i * b + j] = -s_tau * val;
            }
            __syncthreads();
        }
    }

    for (int i = threadIdx.x; i < b * b; i += blockDim.x) {
        int row = i / b;
        int col = i % b;
        T_b[row * block_size + col] = s_T[i];
    }
}

// 快速提取纯净的 V 矩阵缓冲区
__global__ void extract_V_kernel(float* A, float* V_buf, int batch, int n, int k, int b, int block_size) {
    int m_idx = blockIdx.x * blockDim.x + threadIdx.x;
    int n_idx = blockIdx.y * blockDim.y + threadIdx.y;
    int b_idx = blockIdx.z;

    int M_v = n - k;
    if (m_idx < M_v && n_idx < b && b_idx < batch) {
        float* A_b = A + b_idx * n * n;
        float* V_b = V_buf + b_idx * n * block_size;
        
        float val = 0.0f;
        if (m_idx > n_idx) {
            val = A_b[(k + m_idx) * n + (k + n_idx)];
        } else if (m_idx == n_idx) {
            val = 1.0f;
        }
        // V_buf 是行主序: [n, block_size]
        V_b[m_idx * block_size + n_idx] = val; 
    }
}

void launch_unblocked_qr(uintptr_t A_ptr, uintptr_t tau_ptr, int batch, int n) {
    int threads = (n <= 64) ? 64 : ((n <= 128) ? 128 : 256);
    unblocked_qr_kernel<<<batch, threads>>>(reinterpret_cast<float*>(A_ptr), reinterpret_cast<float*>(tau_ptr), batch, n);
}

void launch_panel(uintptr_t A_ptr, uintptr_t tau_ptr, uintptr_t T_ptr, int batch, int n, int k, int b, int block_size) {
    panel_wy_kernel<<<batch, 256>>>(reinterpret_cast<float*>(A_ptr), reinterpret_cast<float*>(tau_ptr), reinterpret_cast<float*>(T_ptr), batch, n, k, b, block_size);
}

void launch_extract_V(uintptr_t A_ptr, uintptr_t V_ptr, int batch, int n, int k, int b, int block_size) {
    dim3 threads(16, 16);
    dim3 blocks((n - k + 15) / 16, (b + 15) / 16, batch);
    extract_V_kernel<<<blocks, threads>>>(reinterpret_cast<float*>(A_ptr), reinterpret_cast<float*>(V_ptr), batch, n, k, b, block_size);
}
"""

# =====================================================================
# 2. 裸调 cuBLAS API:完全没有 ATen 对象,0 开销下发
# =====================================================================
cpp_source = """
#include <ATen/core/Tensor.h>
#include <pybind11/pybind11.h>
#include <torch/csrc/utils/pybind.h>
#include <cublas_v2.h>

void launch_unblocked_qr(uintptr_t A_ptr, uintptr_t tau_ptr, int batch, int n);
void launch_panel(uintptr_t A_ptr, uintptr_t tau_ptr, uintptr_t T_ptr, int batch, int n, int k, int b, int block_size);
void launch_extract_V(uintptr_t A_ptr, uintptr_t V_ptr, int batch, int n, int k, int b, int block_size);

// 全局唯一的 cuBLAS 句柄,初始化一次
static cublasHandle_t handle = nullptr;

void ultimate_cublas_qr(at::Tensor H, at::Tensor tau, at::Tensor T_buf, at::Tensor V_buf, at::Tensor M1_buf, at::Tensor M2_buf, int block_size) {
    if (!handle) {
        cublasCreate(&handle);
        cublasSetMathMode(handle, CUBLAS_DEFAULT_MATH);
    }

    int batch = H.size(0);
    int n = H.size(1);

    if (n <= 176) {
        launch_unblocked_qr((uintptr_t)H.data_ptr<float>(), (uintptr_t)tau.data_ptr<float>(), batch, n);
        return;
    }

    // 将所有对象转换为裸指针,消灭在循环里创建 Tensor 的一切开销
    float* d_H = H.data_ptr<float>();
    float* d_tau = tau.data_ptr<float>();
    float* d_T = T_buf.data_ptr<float>();
    float* d_V = V_buf.data_ptr<float>();
    float* d_M1 = M1_buf.data_ptr<float>();
    float* d_M2 = M2_buf.data_ptr<float>();

    // 预先计算 Batch Stride
    long long stride_H = n * n;
    long long stride_T = block_size * block_size;
    long long stride_V = n * block_size;
    long long stride_M = block_size * n;

    float alpha1 = 1.0f, beta0 = 0.0f;
    float alpha_sub = -1.0f, beta1 = 1.0f;

    for (int k = 0; k < n; k += block_size) {
        int b = std::min(block_size, n - k);
        
        launch_panel((uintptr_t)d_H, (uintptr_t)d_tau, (uintptr_t)d_T, batch, n, k, b, block_size);

        if (k + b < n) {
            launch_extract_V((uintptr_t)d_H, (uintptr_t)d_V, batch, n, k, b, block_size);

            int M_v = n - k;
            int N_c = n - k - b;
            float* C_ptr = d_H + k * n + (k + b);

            // =========================================================================
            // 极致魔法:利用行主序/列主序的 Stride 对偶性,零开销原地转置
            // =========================================================================
            
            // GEMM 1: M1 = V^T * C
            cublasSgemmStridedBatched(handle, CUBLAS_OP_N, CUBLAS_OP_T,
                N_c, b, M_v, &alpha1,
                C_ptr, n, stride_H,
                d_V, block_size, stride_V,
                &beta0,
                d_M1, n, stride_M, batch);

            // GEMM 2: M2 = T^T * M1
            cublasSgemmStridedBatched(handle, CUBLAS_OP_N, CUBLAS_OP_T,
                N_c, b, b, &alpha1,
                d_M1, n, stride_M,
                d_T, block_size, stride_T,
                &beta0,
                d_M2, n, stride_M, batch);

            // GEMM 3: C = C - V * M2
            cublasSgemmStridedBatched(handle, CUBLAS_OP_N, CUBLAS_OP_N,
                N_c, M_v, b, &alpha_sub,
                d_M2, n, stride_M,
                d_V, block_size, stride_V,
                &beta1,
                C_ptr, n, stride_H, batch);
        }
    }
}
   


PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("ultimate_cublas_qr", &ultimate_cublas_qr, "Zero-Overhead Direct cuBLAS Call");
}
"""

# =====================================================================
# 3. 编译:必须链接 -lcublas
# =====================================================================
qr_module = load_inline(
    name='ultimate_cublas_wy',
    cpp_sources=cpp_source, 
    cuda_sources=cuda_source,
    verbose=False, 
    no_implicit_headers=True, 
    extra_cflags=['-O3'],
    extra_cuda_cflags=['-O3', '--use_fast_math', '-Xptxas=-O3', '-arch=sm_100a'],
    extra_ldflags=['-lcublas'] # 关键:链接底层的 cuBLAS 硬件加速库
)

# =====================================================================
# 4. 纯净 Python 入口
# =====================================================================
def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    H = data.clone().contiguous()
    tau = torch.zeros((batch, n), dtype=torch.float32, device=data.device)
    
    # 彻底关闭 TF32 以保证恶劣病态矩阵的绝对精度
    torch.backends.cuda.matmul.allow_tf32 = True
    torch.backends.cudnn.allow_tf32 = False

    block_size = 64 if n >= 2048 else 32

    # 所有中间内存全部在外层一次性静态分配好
    T_buf = torch.zeros((batch, block_size, block_size), dtype=torch.float32, device=data.device)
    V_buf = torch.empty((batch, n, block_size), dtype=torch.float32, device=data.device)
    M1_buf = torch.empty((batch, block_size, n), dtype=torch.float32, device=data.device)
    M2_buf = torch.empty((batch, block_size, n), dtype=torch.float32, device=data.device)

    # 毫无波澜,一指下发,直接交由 C 底层瞬间执行完毕
    qr_module.ultimate_cublas_qr(H, tau, T_buf, V_buf, M1_buf, M2_buf, block_size)
            
    return H, tau
scrolls · 403 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON