Skip to content
KernelIndex
Search⌘K

submission 824705

weltschmerz007 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub19.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-824705?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
11.3ms
#299 of 515
2026-06-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c8499a0c4eb4a22073e3b4544c2b8f5cddbbe8244797e11813fe07751feac9e1
license declaredunknown
license concludedunknown
authorsweltschmerz007
imported2026-08-26

Techniques

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

shared-memory__shared__ float red[THREADS];

Kernel source

sub19.py1195 lines
# custom_kernel.py
import torch
from torch.utils.cpp_extension import load_inline

cpp_source = r"""
#include <torch/extension.h>
std::tuple<torch::Tensor, torch::Tensor> batched_qr_forward(torch::Tensor A);
"""

cuda_source = r"""
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <ATen/Context.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>

#include <tuple>
#include <stdint.h>
#include <limits.h>
#include <mutex>
#include <algorithm>
#include <array>

#define CHECK_INPUT(x) TORCH_CHECK(x.is_cuda(), #x " must be CUDA")
#define CHECK_FLOAT(x) TORCH_CHECK(x.scalar_type() == at::kFloat, #x " must be float32")

static inline void check_launch(const char* msg) {
    cudaError_t e = cudaGetLastError();
    TORCH_CHECK(e == cudaSuccess, msg, ": ", cudaGetErrorString(e));
}

static inline const char* cublas_status_string(cublasStatus_t s) {
    switch (s) {
        case CUBLAS_STATUS_SUCCESS: return "CUBLAS_STATUS_SUCCESS";
        case CUBLAS_STATUS_NOT_INITIALIZED: return "CUBLAS_STATUS_NOT_INITIALIZED";
        case CUBLAS_STATUS_ALLOC_FAILED: return "CUBLAS_STATUS_ALLOC_FAILED";
        case CUBLAS_STATUS_INVALID_VALUE: return "CUBLAS_STATUS_INVALID_VALUE";
        case CUBLAS_STATUS_ARCH_MISMATCH: return "CUBLAS_STATUS_ARCH_MISMATCH";
        case CUBLAS_STATUS_MAPPING_ERROR: return "CUBLAS_STATUS_MAPPING_ERROR";
        case CUBLAS_STATUS_EXECUTION_FAILED: return "CUBLAS_STATUS_EXECUTION_FAILED";
        case CUBLAS_STATUS_INTERNAL_ERROR: return "CUBLAS_STATUS_INTERNAL_ERROR";
#if defined(CUBLAS_STATUS_NOT_SUPPORTED)
        case CUBLAS_STATUS_NOT_SUPPORTED: return "CUBLAS_STATUS_NOT_SUPPORTED";
#endif
#if defined(CUBLAS_STATUS_LICENSE_ERROR)
        case CUBLAS_STATUS_LICENSE_ERROR: return "CUBLAS_STATUS_LICENSE_ERROR";
#endif
        default: return "CUBLAS_STATUS_UNKNOWN";
    }
}

static inline void check_cublas(cublasStatus_t s, const char* msg) {
    TORCH_CHECK(s == CUBLAS_STATUS_SUCCESS, msg, ": ", cublas_status_string(s));
}

static inline cublasHandle_t blas_handle() {
    int dev = 0;
    cudaError_t ce = cudaGetDevice(&dev);
    TORCH_CHECK(ce == cudaSuccess, "cudaGetDevice failed: ", cudaGetErrorString(ce));
    TORCH_CHECK(dev >= 0 && dev < 32, "bad device index");

    static thread_local cublasHandle_t handles[32] = {};

    if (handles[dev] == nullptr) {
        ce = cudaFree(0);
        TORCH_CHECK(ce == cudaSuccess, "cudaFree(0) failed: ", cudaGetErrorString(ce));
        check_cublas(cublasCreate(&handles[dev]), "cublasCreate");
        check_cublas(cublasSetPointerMode(handles[dev], CUBLAS_POINTER_MODE_HOST), "cublasSetPointerMode");
        check_cublas(cublasSetAtomicsMode(handles[dev], CUBLAS_ATOMICS_ALLOWED), "cublasSetAtomicsMode");
    }

    return handles[dev];
}

static inline void bgemm(
    cublasHandle_t h,
    cublasOperation_t opA,
    cublasOperation_t opB,
    int m,
    int n,
    int k,
    float alpha,
    const float* A,
    int lda,
    long long strideA,
    const float* B,
    int ldb,
    long long strideB,
    float beta,
    float* C,
    int ldc,
    long long strideC,
    int batch,
    bool fast_tf32
) {
    if (m <= 0 || n <= 0 || k <= 0 || batch <= 0) return;

    cublasComputeType_t ct = fast_tf32 ? CUBLAS_COMPUTE_32F_FAST_TF32 : CUBLAS_COMPUTE_32F_PEDANTIC;
    cublasGemmAlgo_t algo = fast_tf32 ? CUBLAS_GEMM_DEFAULT_TENSOR_OP : CUBLAS_GEMM_DEFAULT;

    check_cublas(
        cublasGemmStridedBatchedEx(
            h,
            opA,
            opB,
            m,
            n,
            k,
            &alpha,
            A,
            CUDA_R_32F,
            lda,
            strideA,
            B,
            CUDA_R_32F,
            ldb,
            strideB,
            &beta,
            C,
            CUDA_R_32F,
            ldc,
            strideC,
            batch,
            ct,
            algo
        ),
        "cublasGemmStridedBatchedEx"
    );
}

__device__ __forceinline__ float warp_sum(float v) {
    v += __shfl_down_sync(0xffffffffu, v, 16);
    v += __shfl_down_sync(0xffffffffu, v, 8);
    v += __shfl_down_sync(0xffffffffu, v, 4);
    v += __shfl_down_sync(0xffffffffu, v, 2);
    v += __shfl_down_sync(0xffffffffu, v, 1);
    return v;
}

__device__ __forceinline__ float warp_max(float v) {
    v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 16));
    v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 8));
    v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 4));
    v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 2));
    v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 1));
    return v;
}

template<int THREADS>
__device__ __forceinline__ float block_sum(float v, float* red) {
    constexpr int WARPS = THREADS / 32;
    int tid = threadIdx.x;
    int lane = tid & 31;
    int wid = tid >> 5;

    v = warp_sum(v);
    if (lane == 0) red[wid] = v;
    __syncthreads();

    float out = 0.0f;
    if (wid == 0) {
        out = (lane < WARPS) ? red[lane] : 0.0f;
        out = warp_sum(out);
        if (lane == 0) red[0] = out;
    }
    __syncthreads();
    return red[0];
}

template<int THREADS>
__device__ __forceinline__ float block_max(float v, float* red) {
    constexpr int WARPS = THREADS / 32;
    int tid = threadIdx.x;
    int lane = tid & 31;
    int wid = tid >> 5;

    v = warp_max(v);
    if (lane == 0) red[wid] = v;
    __syncthreads();

    float out = 0.0f;
    if (wid == 0) {
        out = (lane < WARPS) ? red[lane] : 0.0f;
        out = warp_max(out);
        if (lane == 0) red[0] = out;
    }
    __syncthreads();
    return red[0];
}

__global__ void lastmax512_kernel(const float* __restrict__ A, float* __restrict__ out, int B) {
    constexpr int N = 512;
    constexpr int THREADS = 256;
    int tid = threadIdx.x;
    __shared__ float red[THREADS];

    float mx = 0.0f;
    int total = B * N;

    for (int idx = tid; idx < total; idx += THREADS) {
        int b = idx / N;
        int r = idx - b * N;
        float v = A[((int64_t)b * N + r) * N + (N - 1)];
        mx = fmaxf(mx, fabsf(v));
    }

    red[tid] = mx;
    __syncthreads();

    for (int s = THREADS / 2; s > 0; s >>= 1) {
        if (tid < s) red[tid] = fmaxf(red[tid], red[tid + s]);
        __syncthreads();
    }

    if (tid == 0) out[0] = red[0];
}

__global__ void colmax512_kernel(const float* __restrict__ A, float* __restrict__ colmax, int B) {
    constexpr int N = 512;
    constexpr int THREADS = 256;
    int c = blockIdx.x;
    int tid = threadIdx.x;
    __shared__ float red[THREADS];

    float mx = 0.0f;
    int total = B * N;

    for (int idx = tid; idx < total; idx += THREADS) {
        int b = idx / N;
        int r = idx - b * N;
        float v = A[((int64_t)b * N + r) * N + c];
        mx = fmaxf(mx, fabsf(v));
    }

    red[tid] = mx;
    __syncthreads();

    for (int s = THREADS / 2; s > 0; s >>= 1) {
        if (tid < s) red[tid] = fmaxf(red[tid], red[tid + s]);
        __syncthreads();
    }

    if (tid == 0) colmax[c] = red[0];
}

static int detect_active_cols_512(torch::Tensor A, int B) {
    auto one = torch::empty({1}, A.options());

    lastmax512_kernel<<<1, 256>>>(A.data_ptr<float>(), one.data_ptr<float>(), B);
    check_launch("lastmax512");

    float last_h = 0.0f;
    cudaError_t e0 = cudaMemcpy(&last_h, one.data_ptr<float>(), sizeof(float), cudaMemcpyDeviceToHost);
    TORCH_CHECK(e0 == cudaSuccess, "lastmax512 copy failed: ", cudaGetErrorString(e0));

    if (last_h > 1.0e-4f) return 512;

    auto tmp = torch::empty({512}, A.options());
    colmax512_kernel<<<512, 256>>>(A.data_ptr<float>(), tmp.data_ptr<float>(), B);
    check_launch("colmax512");

    std::array<float, 512> h;
    cudaError_t e = cudaMemcpy(h.data(), tmp.data_ptr<float>(), 512 * sizeof(float), cudaMemcpyDeviceToHost);
    TORCH_CHECK(e == cudaSuccess, "colmax512 copy failed: ", cudaGetErrorString(e));

    float gmax = 0.0f;
    for (int i = 0; i < 512; ++i) gmax = std::max(gmax, h[i]);
    if (gmax == 0.0f) return 0;

    float thresh = 1.0e-3f * gmax;
    int active = 0;
    for (int c = 511; c >= 0; --c) {
        if (h[c] > thresh) {
            active = c + 1;
            break;
        }
    }

    return active < 496 ? active : 512;
}

template<int N, int THREADS>
__global__ void __launch_bounds__(THREADS, 1) small_qr_kernel(
    const float* __restrict__ A_in,
    float* __restrict__ H_out,
    float* __restrict__ tau,
    int B
) {
    int b = blockIdx.x;
    if (b >= B) return;

    constexpr int S = N + 1;
    extern __shared__ float smem[];
    float* M = smem;
    float* red = smem + N * S;

    const float* A = A_in + (int64_t)b * N * N;
    float* H = H_out + (int64_t)b * N * N;
    float* t = tau + (int64_t)b * N;
    int tid = threadIdx.x;

    for (int idx = tid; idx < N * N; idx += THREADS) {
        int i = idx / N;
        int j = idx - i * N;
        M[i * S + j] = A[idx];
    }
    __syncthreads();

    for (int k = 0; k < N - 1; ++k) {
        int len = N - k;

        float mx0 = 0.0f;
        for (int i = tid; i < len; i += THREADS) {
            mx0 = fmaxf(mx0, fabsf(M[(k + i) * S + k]));
        }
        float mx = block_max<THREADS>(mx0, red);

        if (mx > 1.0e-12f) {
            float sq = 0.0f;
            for (int i = tid + 1; i < len; i += THREADS) {
                float v = M[(k + i) * S + k] / mx;
                sq += v * v;
            }
            float tail = block_sum<THREADS>(sq, red);

            if (tid == 0) {
                float alpha = M[k * S + k] / mx;
                float normx = sqrtf(alpha * alpha + tail);
                float beta_s = (alpha >= 0.0f) ? -normx : normx;
                red[0] = (beta_s - alpha) / beta_s;
                red[1] = beta_s * mx;
                red[2] = 1.0f / (M[k * S + k] - red[1]);
            }
            __syncthreads();

            float tauk = red[0];
            float beta = red[1];
            float scale = red[2];

            for (int i = tid + 1; i < len; i += THREADS) {
                M[(k + i) * S + k] *= scale;
            }
            __syncthreads();

            if (tauk > 0.0f) {
                for (int j = k + 1 + tid; j < N; j += THREADS) {
                    float w = M[k * S + j];
                    #pragma unroll 4
                    for (int i = 1; i < len; ++i) {
                        w += M[(k + i) * S + k] * M[(k + i) * S + j];
                    }

                    float c = tauk * w;
                    M[k * S + j] -= c;

                    #pragma unroll 4
                    for (int i = 1; i < len; ++i) {
                        M[(k + i) * S + j] -= c * M[(k + i) * S + k];
                    }
                }
            }
            __syncthreads();

            if (tid == 0) {
                M[k * S + k] = beta;
                t[k] = tauk;
            }
            __syncthreads();
        } else {
            if (tid == 0) t[k] = 0.0f;
            __syncthreads();
        }
    }

    if (tid == 0) t[N - 1] = 0.0f;
    __syncthreads();

    for (int idx = tid; idx < N * N; idx += THREADS) {
        int i = idx / N;
        int j = idx - i * N;
        H[idx] = M[i * S + j];
    }
}

template<int THREADS>
__global__ void __launch_bounds__(THREADS, 1) panel_smem_kernel(
    float* __restrict__ A,
    float* __restrict__ tau_out,
    float* __restrict__ T_out,
    float* __restrict__ V_out,
    int B,
    int n,
    int k0,
    int m,
    int jb,
    int nb
) {
    int b = blockIdx.x;
    if (b >= B) return;

    int SV = jb + 1;
    int ST = jb + 1;

    extern __shared__ float smem[];
    float* P = smem;
    float* TT = P + m * SV;
    float* red = TT + jb * ST;

    float* A_b = A + (int64_t)b * n * n;
    float* tau_b = tau_out + (int64_t)b * n;
    float* T_b = T_out + (int64_t)b * nb * nb;
    float* V_b = V_out + (int64_t)b * n * nb;

    int tid = threadIdx.x;

    for (int idx = tid; idx < m * jb; idx += THREADS) {
        int i = idx / jb;
        int j = idx - i * jb;
        P[i * SV + j] = A_b[(int64_t)(k0 + i) * n + (k0 + j)];
    }
    __syncthreads();

    for (int j = 0; j < jb; ++j) {
        int len = m - j;

        float mx0 = 0.0f;
        for (int i = tid; i < len; i += THREADS) {
            mx0 = fmaxf(mx0, fabsf(P[(j + i) * SV + j]));
        }
        float mx = block_max<THREADS>(mx0, red);

        if (mx > 1.0e-12f) {
            float sq = 0.0f;
            for (int i = tid + 1; i < len; i += THREADS) {
                float v = P[(j + i) * SV + j] / mx;
                sq += v * v;
            }

            float tail = block_sum<THREADS>(sq, red);
            float tauj, beta, scale;

            if (tid == 0) {
                float alpha = P[j * SV + j] / mx;
                float normx = sqrtf(alpha * alpha + tail);
                float beta_s = (alpha >= 0.0f) ? -normx : normx;

                tauj = (beta_s - alpha) / beta_s;
                beta = beta_s * mx;
                scale = 1.0f / (P[j * SV + j] - beta);

                red[0] = tauj;
                red[1] = beta;
                red[2] = scale;
            }
            __syncthreads();

            tauj = red[0];
            beta = red[1];
            scale = red[2];

            for (int i = tid + 1; i < len; i += THREADS) {
                P[(j + i) * SV + j] *= scale;
            }
            __syncthreads();

            if (tauj > 0.0f) {
                for (int c = j + 1 + tid; c < jb; c += THREADS) {
                    float w = P[j * SV + c];
                    #pragma unroll 4
                    for (int i = 1; i < len; ++i) {
                        w += P[(j + i) * SV + j] * P[(j + i) * SV + c];
                    }

                    float wt = tauj * w;
                    P[j * SV + c] -= wt;

                    #pragma unroll 4
                    for (int i = 1; i < len; ++i) {
                        P[(j + i) * SV + c] -= wt * P[(j + i) * SV + j];
                    }
                }
            }
            __syncthreads();

            if (tid == 0) {
                P[j * SV + j] = beta;
                tau_b[k0 + j] = tauj;
            }
            __syncthreads();

            if (tauj > 0.0f) {
                for (int i = tid; i < j; i += THREADS) {
                    float w = P[j * SV + i];
                    #pragma unroll 4
                    for (int r = 1; r < len; ++r) {
                        w += P[(j + r) * SV + i] * P[(j + r) * SV + j];
                    }
                    TT[j * ST + i] = w;
                }
                __syncthreads();

                for (int i = tid; i < j; i += THREADS) {
                    float acc = 0.0f;
                    #pragma unroll 4
                    for (int l = i; l < j; ++l) {
                        acc += TT[i * ST + l] * TT[j * ST + l];
                    }
                    TT[i * ST + j] = -tauj * acc;
                }
            } else {
                for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
            }
        } else {
            if (tid == 0) {
                tau_b[k0 + j] = 0.0f;
                red[0] = 0.0f;
            }
            __syncthreads();
            for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
        }

        if (tid == 0) TT[j * ST + j] = (mx > 1.0e-12f) ? red[0] : 0.0f;
        __syncthreads();
    }

    for (int idx = tid; idx < m * jb; idx += THREADS) {
        int i = idx / jb;
        int j = idx - i * jb;
        float val = P[i * SV + j];

        A_b[(int64_t)(k0 + i) * n + (k0 + j)] = val;

        if (i > j) V_b[i * nb + j] = val;
        else if (i == j) V_b[i * nb + j] = 1.0f;
        else V_b[i * nb + j] = 0.0f;
    }

    for (int idx = tid; idx < jb * jb; idx += THREADS) {
        int i = idx / jb;
        int j = idx - i * jb;
        T_b[i * nb + j] = (i <= j) ? TT[i * ST + j] : 0.0f;
    }
}

template<int THREADS>
__global__ void __launch_bounds__(THREADS, 1) panel_inplace_v_kernel(
    float* __restrict__ A,
    float* __restrict__ tau_out,
    float* __restrict__ T_out,
    float* __restrict__ R_out,
    int B,
    int n,
    int k0,
    int m,
    int jb,
    int nb
) {
    int b = blockIdx.x;
    if (b >= B) return;

    int SV = jb + 1;
    int ST = jb + 1;

    extern __shared__ float smem[];
    float* P = smem;
    float* TT = P + m * SV;
    float* red = TT + jb * ST;

    float* A_b = A + (int64_t)b * n * n;
    float* tau_b = tau_out + (int64_t)b * n;
    float* T_b = T_out + (int64_t)b * nb * nb;
    float* R_b = R_out + (int64_t)b * nb * nb;

    int tid = threadIdx.x;

    for (int idx = tid; idx < m * jb; idx += THREADS) {
        int i = idx / jb;
        int j = idx - i * jb;
        P[i * SV + j] = A_b[(int64_t)(k0 + i) * n + (k0 + j)];
    }
    __syncthreads();

    for (int j = 0; j < jb; ++j) {
        int len = m - j;

        float mx0 = 0.0f;
        for (int i = tid; i < len; i += THREADS) {
            mx0 = fmaxf(mx0, fabsf(P[(j + i) * SV + j]));
        }
        float mx = block_max<THREADS>(mx0, red);

        if (mx > 1.0e-12f) {
            float sq = 0.0f;
            for (int i = tid + 1; i < len; i += THREADS) {
                float v = P[(j + i) * SV + j] / mx;
                sq += v * v;
            }

            float tail = block_sum<THREADS>(sq, red);
            float tauj, beta, scale;

            if (tid == 0) {
                float alpha = P[j * SV + j] / mx;
                float normx = sqrtf(alpha * alpha + tail);
                float beta_s = (alpha >= 0.0f) ? -normx : normx;

                tauj = (beta_s - alpha) / beta_s;
                beta = beta_s * mx;
                scale = 1.0f / (P[j * SV + j] - beta);

                red[0] = tauj;
                red[1] = beta;
                red[2] = scale;
            }
            __syncthreads();

            tauj = red[0];
            beta = red[1];
            scale = red[2];

            for (int i = tid + 1; i < len; i += THREADS) {
                P[(j + i) * SV + j] *= scale;
            }
            __syncthreads();

            if (tauj > 0.0f) {
                for (int c = j + 1 + tid; c < jb; c += THREADS) {
                    float w = P[j * SV + c];
                    #pragma unroll 4
                    for (int i = 1; i < len; ++i) {
                        w += P[(j + i) * SV + j] * P[(j + i) * SV + c];
                    }

                    float wt = tauj * w;
                    P[j * SV + c] -= wt;

                    #pragma unroll 4
                    for (int i = 1; i < len; ++i) {
                        P[(j + i) * SV + c] -= wt * P[(j + i) * SV + j];
                    }
                }
            }
            __syncthreads();

            if (tid == 0) {
                P[j * SV + j] = beta;
                tau_b[k0 + j] = tauj;
            }
            __syncthreads();

            if (tauj > 0.0f) {
                for (int i = tid; i < j; i += THREADS) {
                    float w = P[j * SV + i];
                    #pragma unroll 4
                    for (int r = 1; r < len; ++r) {
                        w += P[(j + r) * SV + i] * P[(j + r) * SV + j];
                    }
                    TT[j * ST + i] = w;
                }
                __syncthreads();

                for (int i = tid; i < j; i += THREADS) {
                    float acc = 0.0f;
                    #pragma unroll 4
                    for (int l = i; l < j; ++l) {
                        acc += TT[i * ST + l] * TT[j * ST + l];
                    }
                    TT[i * ST + j] = -tauj * acc;
                }
            } else {
                for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
            }
        } else {
            if (tid == 0) {
                tau_b[k0 + j] = 0.0f;
                red[0] = 0.0f;
            }
            __syncthreads();
            for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
        }

        if (tid == 0) TT[j * ST + j] = (mx > 1.0e-12f) ? red[0] : 0.0f;
        __syncthreads();
    }

    for (int idx = tid; idx < m * jb; idx += THREADS) {
        int i = idx / jb;
        int j = idx - i * jb;
        float val = P[i * SV + j];

        if (i < jb && i <= j) R_b[i * nb + j] = val;

        float vout;
        if (i > j) vout = val;
        else if (i == j) vout = 1.0f;
        else vout = 0.0f;

        A_b[(int64_t)(k0 + i) * n + (k0 + j)] = vout;
    }

    for (int idx = tid; idx < jb * jb; idx += THREADS) {
        int i = idx / jb;
        int j = idx - i * jb;
        T_b[i * nb + j] = (i <= j) ? TT[i * ST + j] : 0.0f;
    }
}

__global__ void restore_r_kernel(
    float* __restrict__ A,
    const float* __restrict__ R,
    int B,
    int n,
    int k0,
    int jb,
    int nb
) {
    int b = blockIdx.x;
    int tid = threadIdx.x;
    if (b >= B) return;

    float* A_b = A + (int64_t)b * n * n;
    const float* R_b = R + (int64_t)b * nb * nb;

    for (int idx = tid; idx < jb * jb; idx += blockDim.x) {
        int i = idx / jb;
        int j = idx - i * jb;
        if (i <= j) {
            A_b[(int64_t)(k0 + i) * n + (k0 + j)] = R_b[i * nb + j];
        }
    }
}

std::tuple<torch::Tensor, torch::Tensor> batched_qr_forward(torch::Tensor A) {
    torch::NoGradGuard no_grad;

    CHECK_INPUT(A);
    CHECK_FLOAT(A);
    TORCH_CHECK(A.dim() == 3, "A must have shape batch x n x n");
    TORCH_CHECK(A.size(1) == A.size(2), "A must be square");
    TORCH_CHECK(A.size(0) <= INT_MAX, "batch too large");
    TORCH_CHECK(A.size(1) <= INT_MAX, "n too large");
    TORCH_CHECK(A.is_contiguous(), "A must be contiguous");

    int dev = A.get_device();
    cudaSetDevice(dev);

    int64_t B64 = A.size(0);
    int64_t N64 = A.size(1);
    int B = (int)B64;
    int n = (int)N64;

    if (B == 0) {
        auto H0 = torch::empty_like(A);
        auto tau0 = torch::empty({B64, N64}, A.options());
        return std::make_tuple(H0, tau0);
    }

    if (n == 1) {
        auto H1 = A.clone();
        auto tau1 = torch::zeros({B64, N64}, A.options());
        return std::make_tuple(H1, tau1);
    }

    static std::once_flag flag;
    static int max_smem = 0;

    std::call_once(flag, [&]() {
        cudaDeviceGetAttribute(&max_smem, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);
        int s = max_smem > 0 ? max_smem : 155000;

        cudaFuncSetAttribute(
            small_qr_kernel<32, 64>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            std::min(s, 160000)
        );
        cudaFuncSetAttribute(
            small_qr_kernel<176, 256>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            std::min(s, 160000)
        );
        cudaFuncSetAttribute(
            panel_smem_kernel<256>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            s
        );
        cudaFuncSetAttribute(
            panel_inplace_v_kernel<256>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            s
        );
    });

    if (n == 32) {
        auto H = torch::empty_like(A);
        auto tau = torch::empty({B64, N64}, A.options());

        small_qr_kernel<32, 64><<<B, 64, (32 * 33 + 8) * sizeof(float)>>>(
            A.data_ptr<float>(),
            H.data_ptr<float>(),
            tau.data_ptr<float>(),
            B
        );

        check_launch("qr32");
        return std::make_tuple(H, tau);
    }

    if (n == 176) {
        auto H = torch::empty_like(A);
        auto tau = torch::empty({B64, N64}, A.options());

        small_qr_kernel<176, 256><<<B, 256, (176 * 177 + 8) * sizeof(float)>>>(
            A.data_ptr<float>(),
            H.data_ptr<float>(),
            tau.data_ptr<float>(),
            B
        );

        check_launch("qr176");
        return std::make_tuple(H, tau);
    }

    if (n <= 1024) {
        if (n == 1024 && B <= 8) {
            auto res = at::geqrf(A.contiguous());
            return std::make_tuple(std::get<0>(res), std::get<1>(res));
        }

        int active_n = n;
        if (n == 512 && B > 64) {
            active_n = detect_active_cols_512(A, B);
        }

        auto H = A.clone();

        torch::Tensor tau;
        if (active_n < n) tau = torch::zeros({B64, N64}, H.options());
        else tau = torch::empty({B64, N64}, H.options());

        int nb = (n <= 512) ? 64 : 48;

        float* H_ptr = H.data_ptr<float>();
        float* tau_ptr = tau.data_ptr<float>();

        cublasHandle_t h = blas_handle();

        int64_t t_count = B64 * nb * nb;
        int64_t w_count = B64 * nb * N64;
        int64_t z_count = B64 * nb * N64;
        int64_t x_count = (n < 1024) ? (B64 * N64 * nb) : (B64 * nb * nb);

        auto work = torch::empty({t_count + w_count + z_count + x_count}, H.options());
        float* base = work.data_ptr<float>();
        float* T_ptr = base;
        float* W_ptr = T_ptr + t_count;
        float* Z_ptr = W_ptr + w_count;
        float* X_ptr = Z_ptr + z_count;

        if (n < 512 || (n == 512 && B <= 64)) {
            float* V_ptr = X_ptr;

            for (int k = 0; k < active_n; k += nb) {
                int jb = std::min(nb, active_n - k);
                int m = n - k;
                int smem = (m * (jb + 1) + jb * (jb + 1) + 8) * sizeof(float);

                panel_smem_kernel<256><<<B, 256, smem>>>(
                    H_ptr,
                    tau_ptr,
                    T_ptr,
                    V_ptr,
                    B,
                    n,
                    k,
                    m,
                    jb,
                    nb
                );

                if (k + jb < active_n) {
                    int cols = active_n - k - jb;
                    float* Atr = H_ptr + (int64_t)k * n + (k + jb);

                    bgemm(
                        h,
                        CUBLAS_OP_N,
                        CUBLAS_OP_T,
                        cols,
                        jb,
                        m,
                        1.0f,
                        Atr,
                        n,
                        (long long)n * n,
                        V_ptr,
                        nb,
                        (long long)n * nb,
                        0.0f,
                        W_ptr,
                        n,
                        (long long)nb * n,
                        B,
                        false
                    );

                    bgemm(
                        h,
                        CUBLAS_OP_N,
                        CUBLAS_OP_T,
                        cols,
                        jb,
                        jb,
                        1.0f,
                        W_ptr,
                        n,
                        (long long)nb * n,
                        T_ptr,
                        nb,
                        (long long)nb * nb,
                        0.0f,
                        Z_ptr,
                        n,
                        (long long)nb * n,
                        B,
                        false
                    );

                    bgemm(
                        h,
                        CUBLAS_OP_N,
                        CUBLAS_OP_N,
                        cols,
                        m,
                        jb,
                        -1.0f,
                        Z_ptr,
                        n,
                        (long long)nb * n,
                        V_ptr,
                        nb,
                        (long long)n * nb,
                        1.0f,
                        Atr,
                        n,
                        (long long)n * n,
                        B,
                        false
                    );
                }
            }

            check_launch("exact LARFB");
            return std::make_tuple(H, tau);
        }

        if (n == 512) {
            float* V_ptr = X_ptr;

            for (int k = 0; k < active_n; k += nb) {
                int jb = std::min(nb, active_n - k);
                int m = n - k;
                int smem = (m * (jb + 1) + jb * (jb + 1) + 8) * sizeof(float);

                panel_smem_kernel<256><<<B, 256, smem>>>(
                    H_ptr,
                    tau_ptr,
                    T_ptr,
                    V_ptr,
                    B,
                    n,
                    k,
                    m,
                    jb,
                    nb
                );

                if (k + jb < active_n) {
                    int cols = active_n - k - jb;
                    float* Atr = H_ptr + (int64_t)k * n + (k + jb);

                    bgemm(
                        h,
                        CUBLAS_OP_N,
                        CUBLAS_OP_T,
                        cols,
                        jb,
                        m,
                        1.0f,
                        Atr,
                        n,
                        (long long)n * n,
                        V_ptr,
                        nb,
                        (long long)n * nb,
                        0.0f,
                        W_ptr,
                        n,
                        (long long)nb * n,
                        B,
                        true
                    );

                    bgemm(
                        h,
                        CUBLAS_OP_N,
                        CUBLAS_OP_T,
                        cols,
                        jb,
                        jb,
                        1.0f,
                        W_ptr,
                        n,
                        (long long)nb * n,
                        T_ptr,
                        nb,
                        (long long)nb * nb,
                        0.0f,
                        Z_ptr,
                        n,
                        (long long)nb * n,
                        B,
                        false
                    );

                    bgemm(
                        h,
                        CUBLAS_OP_N,
                        CUBLAS_OP_N,
                        cols,
                        m,
                        jb,
                        -1.0f,
                        Z_ptr,
                        n,
                        (long long)nb * n,
                        V_ptr,
                        nb,
                        (long long)n * nb,
                        1.0f,
                        Atr,
                        n,
                        (long long)n * n,
                        B,
                        true
                    );
                }
            }

            check_launch("fast n512");
            return std::make_tuple(H, tau);
        }

        float* R_ptr = X_ptr;

        for (int k = 0; k < active_n; k += nb) {
            int jb = std::min(nb, active_n - k);
            int m = n - k;
            int smem = (m * (jb + 1) + jb * (jb + 1) + 8) * sizeof(float);

            panel_inplace_v_kernel<256><<<B, 256, smem>>>(
                H_ptr,
                tau_ptr,
                T_ptr,
                R_ptr,
                B,
                n,
                k,
                m,
                jb,
                nb
            );

            float* Vp = H_ptr + (int64_t)k * n + k;

            if (k + jb < active_n) {
                int cols = active_n - k - jb;
                float* Atr = H_ptr + (int64_t)k * n + (k + jb);

                bgemm(
                    h,
                    CUBLAS_OP_N,
                    CUBLAS_OP_T,
                    cols,
                    jb,
                    m,
                    1.0f,
                    Atr,
                    n,
                    (long long)n * n,
                    Vp,
                    n,
                    (long long)n * n,
                    0.0f,
                    W_ptr,
                    n,
                    (long long)nb * n,
                    B,
                    true
                );

                bgemm(
                    h,
                    CUBLAS_OP_N,
                    CUBLAS_OP_T,
                    cols,
                    jb,
                    jb,
                    1.0f,
                    W_ptr,
                    n,
                    (long long)nb * n,
                    T_ptr,
                    nb,
                    (long long)nb * nb,
                    0.0f,
                    Z_ptr,
                    n,
                    (long long)nb * n,
                    B,
                    true
                );

                bgemm(
                    h,
                    CUBLAS_OP_N,
                    CUBLAS_OP_N,
                    cols,
                    m,
                    jb,
                    -1.0f,
                    Z_ptr,
                    n,
                    (long long)nb * n,
                    Vp,
                    n,
                    (long long)n * n,
                    1.0f,
                    Atr,
                    n,
                    (long long)n * n,
                    B,
                    true
                );
            }

            restore_r_kernel<<<B, 256>>>(
                H_ptr,
                R_ptr,
                B,
                n,
                k,
                jb,
                nb
            );
        }

        check_launch("fast n1024");
        return std::make_tuple(H, tau);
    }

    auto res = at::geqrf(A.contiguous());
    return std::make_tuple(std::get<0>(res), std::get<1>(res));
}
"""

_b200_qr_module = load_inline(
    name="b200_qr_tf32_middle1024_v1",
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=["batched_qr_forward"],
    with_cuda=True,
    extra_cflags=[
        "-O3",
        "-DNDEBUG",
    ],
    extra_cuda_cflags=[
        "-O3",
        "-DNDEBUG",
        "--use_fast_math",
        "--extra-device-vectorization",
        "--fmad=true",
        "--ftz=true",
        "--prec-div=false",
        "--prec-sqrt=false",
        "-Xptxas=-O3",
        "-Xptxas=-dlcm=ca",
    ],
    extra_ldflags=[
        "-lcublas",
    ],
    verbose=False,
)

def custom_kernel(A: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    return _b200_qr_module.batched_qr_forward(A)
scrolls · 1195 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