Skip to content
KernelIndex
Search⌘K

submission 838628

Dortamac · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-838628?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
14.1ms
#318 of 515
2026-06-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f0355594d23bb1fc453e16920c8c12cd5f765a41cff249320b92b195dd4b9ec8
license declaredunknown
license concludedunknown
authorsDortamac
imported2026-08-26

Techniques

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

shared-memory__shared__ double scratch[THREADS];

Kernel source

submission.py730 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t


_PANEL = 8
_PANEL_32 = 32
_PANEL_352 = 8
_PANEL_2048 = 4
_PANEL_4096 = 4
_SUPERPANEL_512 = 32
_SUPERPANEL_NS = ()
_PANEL_T_NS = (176,)
_USE_PRECOMPUTE_U_NS = ()
_BLOCKED_NS = (32, 176, 352, 512, 1024, 2048)
_NATIVE_PANEL_NS = (32, 176, 352, 512, 1024, 2048)


CPP_SRC = r"""
#include <torch/extension.h>

void geqrf_32_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_176_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_352_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_512_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_1024_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_2048_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_4096_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_176_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_352_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_512_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_1024_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_2048_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void pack_v_panel(torch::Tensor h, torch::Tensor v, int64_t k, int64_t width);
void make_t_panel(torch::Tensor v, torch::Tensor tau, torch::Tensor t, int64_t width);
"""


CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <stdexcept>
#include <string>

namespace {

constexpr int PANEL = 64;
constexpr int MAX_PANEL = 128;
constexpr int THREADS = 256;

inline void check_cuda(cudaError_t status, const char* what) {
    if (status != cudaSuccess) {
        throw std::runtime_error(std::string(what) + ": " + cudaGetErrorString(status));
    }
}

__device__ double shfl_down_double(double value, int offset) {
    int2 words = *reinterpret_cast<int2*>(&value);
    words.x = __shfl_down_sync(0xffffffffu, words.x, offset);
    words.y = __shfl_down_sync(0xffffffffu, words.y, offset);
    return *reinterpret_cast<double*>(&words);
}

__device__ double warp_sum(double value) {
    for (int offset = 16; offset > 0; offset >>= 1) {
        value += shfl_down_double(value, offset);
    }
    return value;
}

__device__ double block_sum(double value, double* scratch) {
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;

    value = warp_sum(value);
    if (lane == 0) {
        scratch[warp] = value;
    }
    __syncthreads();

    value = 0.0;
    const int num_warps = (blockDim.x + 31) >> 5;
    if (warp == 0 && lane < num_warps) {
        value = scratch[lane];
    }
    value = warp_sum(value);
    if (tid == 0) {
        scratch[0] = value;
    }
    __syncthreads();
    return scratch[0];
}

template <int N, bool BUILD_T>
__global__ void geqrf_panel_kernel(float* __restrict__ h,
                                   float* __restrict__ tau,
                                   float* __restrict__ t_out,
                                   int batch,
                                   int k,
                                   int width) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    if (b >= batch) {
        return;
    }

    __shared__ double scratch[THREADS];
    __shared__ float beta_s;
    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float update_s;

    float* a = h + static_cast<long long>(b) * N * N;
    float* t = tau + static_cast<long long>(b) * N;

    for (int jj = 0; jj < width; ++jj) {
        const int col = k + jj;
        double tail_local = 0.0;
        for (int row = col + 1 + tid; row < N; row += blockDim.x) {
            const float x = a[row * N + col];
            tail_local += static_cast<double>(x) * static_cast<double>(x);
        }
        const double tail_norm_sq = block_sum(tail_local, scratch);

        if (tid == 0) {
            const float alpha = a[col * N + col];
            if (tail_norm_sq == 0.0) {
                beta_s = alpha;
                tau_s = 0.0f;
                inv_s = 0.0f;
            } else {
                const double norm = sqrt(static_cast<double>(alpha) * static_cast<double>(alpha) + tail_norm_sq);
                const double beta = (alpha >= 0.0f) ? -norm : norm;
                beta_s = static_cast<float>(beta);
                tau_s = static_cast<float>((beta - static_cast<double>(alpha)) / beta);
                inv_s = static_cast<float>(1.0 / (static_cast<double>(alpha) - beta));
            }
            t[col] = tau_s;
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            for (int row = col + 1 + tid; row < N; row += blockDim.x) {
                a[row * N + col] *= inv_s;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            for (int j2 = jj + 1; j2 < width; ++j2) {
                const int update_col = k + j2;
                double dot_local = (tid == 0) ? static_cast<double>(a[col * N + update_col]) : 0.0;
                for (int row = col + 1 + tid; row < N; row += blockDim.x) {
                    dot_local += static_cast<double>(a[row * N + col]) *
                                 static_cast<double>(a[row * N + update_col]);
                }
                const double dot = block_sum(dot_local, scratch);

                if (tid == 0) {
                    update_s = tau_s * static_cast<float>(dot);
                    a[col * N + update_col] -= update_s;
                }
                __syncthreads();

                for (int row = col + 1 + tid; row < N; row += blockDim.x) {
                    a[row * N + update_col] -= a[row * N + col] * update_s;
                }
            }
            __syncthreads();
        }

        if (tid == 0) {
            a[col * N + col] = beta_s;
        }
        __syncthreads();
    }

    if constexpr (BUILD_T) {
        __shared__ float t_col[MAX_PANEL];
        __shared__ float tau_j_s;
        float* tb = t_out + static_cast<long long>(b) * width * width;
        const int rows = N - k;
        for (int idx = tid; idx < width * width; idx += blockDim.x) {
            tb[idx] = 0.0f;
        }
        __syncthreads();

        for (int j = 0; j < width; ++j) {
            if (tid == 0) {
                tau_j_s = t[k + j];
            }
            __syncthreads();

            if (j > 0 && tau_j_s != 0.0f) {
                for (int i = 0; i < j; ++i) {
                    double local = 0.0;
                    for (int row = j + tid; row < rows; row += blockDim.x) {
                        float vi = 0.0f;
                        if (row == i) {
                            vi = 1.0f;
                        } else if (row > i) {
                            vi = a[(k + row) * N + (k + i)];
                        }

                        float vj = 0.0f;
                        if (row == j) {
                            vj = 1.0f;
                        } else if (row > j) {
                            vj = a[(k + row) * N + (k + j)];
                        }
                        local += static_cast<double>(vi) * static_cast<double>(vj);
                    }
                    const double dot = block_sum(local, scratch);
                    if (tid == 0) {
                        t_col[i] = -tau_j_s * static_cast<float>(dot);
                    }
                }
                __syncthreads();

                for (int i = tid; i < j; i += blockDim.x) {
                    float acc = 0.0f;
                    for (int q = 0; q < j; ++q) {
                        acc += tb[i * width + q] * t_col[q];
                    }
                    tb[i * width + j] = acc;
                }
                __syncthreads();
            }

            if (tid == 0) {
                tb[j * width + j] = tau_j_s;
            }
            __syncthreads();
        }
    }
}

__global__ void make_t_panel_kernel(const float* __restrict__ v,
                                    const float* __restrict__ tau,
                                    float* __restrict__ t,
                                    int batch,
                                    int rows,
                                    int width,
                                    long long tau_s0,
                                    long long tau_s1) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    if (b >= batch) {
        return;
    }

    __shared__ double scratch[THREADS];
    __shared__ float col[MAX_PANEL];
    __shared__ float tau_j_s;

    const float* vb = v + static_cast<long long>(b) * rows * width;
    const float* taub = tau + static_cast<long long>(b) * tau_s0;
    float* tb = t + static_cast<long long>(b) * width * width;

    for (int idx = tid; idx < width * width; idx += blockDim.x) {
        tb[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < width; ++j) {
        if (tid == 0) {
            tau_j_s = taub[static_cast<long long>(j) * tau_s1];
        }
        __syncthreads();

        if (j > 0 && tau_j_s != 0.0f) {
            for (int i = 0; i < j; ++i) {
                double local = 0.0;
                for (int row = j + tid; row < rows; row += blockDim.x) {
                    local += static_cast<double>(vb[row * width + i]) *
                             static_cast<double>(vb[row * width + j]);
                }
                const double dot = block_sum(local, scratch);
                if (tid == 0) {
                    col[i] = -tau_j_s * static_cast<float>(dot);
                }
            }
            __syncthreads();

            for (int i = tid; i < j; i += blockDim.x) {
                float acc = 0.0f;
                for (int q = 0; q < j; ++q) {
                    acc += tb[i * width + q] * col[q];
                }
                tb[i * width + j] = acc;
            }
            __syncthreads();
        }

        if (tid == 0) {
            tb[j * width + j] = tau_j_s;
        }
        __syncthreads();
    }
}

__global__ void pack_v_panel_kernel(const float* __restrict__ h,
                                    float* __restrict__ v,
                                    int batch,
                                    int n,
                                    int rows,
                                    int k,
                                    int width) {
    const int b = blockIdx.z;
    const int row = blockIdx.y * blockDim.y + threadIdx.y;
    const int col = blockIdx.x * blockDim.x + threadIdx.x;
    if (b >= batch || row >= rows || col >= width) {
        return;
    }

    float value = 0.0f;
    if (row == col) {
        value = 1.0f;
    } else if (row > col) {
        value = h[(static_cast<long long>(b) * n + (k + row)) * n + (k + col)];
    }
    v[(static_cast<long long>(b) * rows + row) * width + col] = value;
}

} // namespace

template <int N>
void geqrf_panel(torch::Tensor h, torch::Tensor tau, torch::Tensor t_panel, int64_t k, int64_t width, const char* name) {
    TORCH_CHECK(h.is_cuda() && tau.is_cuda(), name, " expects CUDA tensors");
    TORCH_CHECK(h.scalar_type() == torch::kFloat32, "h must be float32");
    TORCH_CHECK(tau.scalar_type() == torch::kFloat32, "tau must be float32");
    TORCH_CHECK(h.is_contiguous(), "h must be contiguous");
    TORCH_CHECK(tau.is_contiguous(), "tau must be contiguous");
    TORCH_CHECK(h.dim() == 3 && h.size(1) == N && h.size(2) == N, "h has wrong shape");
    TORCH_CHECK(tau.dim() == 2 && tau.size(1) == N, "tau has wrong shape");
    float* t_ptr = nullptr;
    if (t_panel.defined()) {
        TORCH_CHECK(t_panel.is_cuda(), name, " t expects CUDA tensor");
        TORCH_CHECK(t_panel.scalar_type() == torch::kFloat32, "t must be float32");
        TORCH_CHECK(t_panel.is_contiguous(), "t must be contiguous");
        TORCH_CHECK(t_panel.dim() == 3, "t must be batch x width x width");
        t_ptr = t_panel.data_ptr<float>();
    }

    const int kk = static_cast<int>(k);
    const int ww = static_cast<int>(width);
    TORCH_CHECK(ww > 0 && ww <= MAX_PANEL, "invalid panel width");
    TORCH_CHECK(kk >= 0 && kk + ww <= N, "invalid panel offset");
    const int batch = static_cast<int>(h.size(0));
    if (t_panel.defined()) {
        TORCH_CHECK(t_panel.size(0) == batch && t_panel.size(1) == ww && t_panel.size(2) == ww, "t has wrong shape");
    }
    if (batch == 0) {
        return;
    }
    if (t_panel.defined()) {
        geqrf_panel_kernel<N, true><<<batch, THREADS>>>(h.data_ptr<float>(), tau.data_ptr<float>(), t_ptr, batch, kk, ww);
    } else {
        geqrf_panel_kernel<N, false><<<batch, THREADS>>>(h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, batch, kk, ww);
    }
    check_cuda(cudaGetLastError(), name);
}

void geqrf_32_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<32>(h, tau, torch::Tensor(), k, width, "geqrf_32_panel");
}

void geqrf_176_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<176>(h, tau, torch::Tensor(), k, width, "geqrf_176_panel");
}

void geqrf_352_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<352>(h, tau, torch::Tensor(), k, width, "geqrf_352_panel");
}

void geqrf_512_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<512>(h, tau, torch::Tensor(), k, width, "geqrf_512_panel");
}

void geqrf_1024_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<1024>(h, tau, torch::Tensor(), k, width, "geqrf_1024_panel");
}

void geqrf_2048_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<2048>(h, tau, torch::Tensor(), k, width, "geqrf_2048_panel");
}

void geqrf_4096_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
    geqrf_panel<4096>(h, tau, torch::Tensor(), k, width, "geqrf_4096_panel");
}

void geqrf_176_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
    geqrf_panel<176>(h, tau, t, k, width, "geqrf_176_panel_t");
}

void geqrf_352_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
    geqrf_panel<352>(h, tau, t, k, width, "geqrf_352_panel_t");
}

void geqrf_512_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
    geqrf_panel<512>(h, tau, t, k, width, "geqrf_512_panel_t");
}

void geqrf_1024_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
    geqrf_panel<1024>(h, tau, t, k, width, "geqrf_1024_panel_t");
}

void geqrf_2048_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
    geqrf_panel<2048>(h, tau, t, k, width, "geqrf_2048_panel_t");
}

void pack_v_panel(torch::Tensor h, torch::Tensor v, int64_t k, int64_t width) {
    TORCH_CHECK(h.is_cuda() && v.is_cuda(), "pack_v_panel expects CUDA tensors");
    TORCH_CHECK(h.scalar_type() == torch::kFloat32, "h must be float32");
    TORCH_CHECK(v.scalar_type() == torch::kFloat32, "v must be float32");
    TORCH_CHECK(h.is_contiguous(), "h must be contiguous");
    TORCH_CHECK(v.is_contiguous(), "v must be contiguous");
    TORCH_CHECK(h.dim() == 3 && h.size(1) == h.size(2), "h must be batch x n x n");
    TORCH_CHECK(v.dim() == 3, "v must be batch x rows x width");

    const int batch = static_cast<int>(h.size(0));
    const int n = static_cast<int>(h.size(1));
    const int rows = static_cast<int>(v.size(1));
    const int kk = static_cast<int>(k);
    const int ww = static_cast<int>(width);
    TORCH_CHECK(v.size(0) == batch, "v batch mismatch");
    TORCH_CHECK(ww > 0 && ww <= MAX_PANEL && v.size(2) == ww, "invalid pack_v_panel width");
    TORCH_CHECK(kk >= 0 && kk + ww <= n && rows == n - kk, "invalid pack_v_panel shape");
    if (batch == 0) {
        return;
    }

    const dim3 block(16, 16, 1);
    const dim3 grid((ww + block.x - 1) / block.x, (rows + block.y - 1) / block.y, batch);
    pack_v_panel_kernel<<<grid, block>>>(h.data_ptr<float>(), v.data_ptr<float>(), batch, n, rows, kk, ww);
    check_cuda(cudaGetLastError(), "pack_v_panel");
}

void make_t_panel(torch::Tensor v, torch::Tensor tau, torch::Tensor t, int64_t width) {
    TORCH_CHECK(v.is_cuda() && tau.is_cuda() && t.is_cuda(), "make_t_panel expects CUDA tensors");
    TORCH_CHECK(v.scalar_type() == torch::kFloat32, "v must be float32");
    TORCH_CHECK(tau.scalar_type() == torch::kFloat32, "tau must be float32");
    TORCH_CHECK(t.scalar_type() == torch::kFloat32, "t must be float32");
    TORCH_CHECK(v.is_contiguous(), "v must be contiguous");
    TORCH_CHECK(t.is_contiguous(), "t must be contiguous");
    TORCH_CHECK(v.dim() == 3, "v must be batch x rows x width");
    TORCH_CHECK(tau.dim() == 2, "tau must be batch x width");
    TORCH_CHECK(t.dim() == 3, "t must be batch x width x width");

    const int batch = static_cast<int>(v.size(0));
    const int rows = static_cast<int>(v.size(1));
    const int ww = static_cast<int>(width);
    TORCH_CHECK(ww > 0 && ww <= MAX_PANEL, "invalid make_t_panel width");
    TORCH_CHECK(v.size(2) == ww && tau.size(1) == ww && t.size(1) == ww && t.size(2) == ww, "make_t_panel shape mismatch");
    if (batch == 0) {
        return;
    }
    make_t_panel_kernel<<<batch, THREADS>>>(
        v.data_ptr<float>(),
        tau.data_ptr<float>(),
        t.data_ptr<float>(),
        batch,
        rows,
        ww,
        static_cast<long long>(tau.stride(0)),
        static_cast<long long>(tau.stride(1)));
    check_cuda(cudaGetLastError(), "make_t_panel");
}

"""


if torch.cuda.is_available():
    _native_module = load_inline(
        name="qr_v2_native_panel_p8_n32_n4096_v1",
        cpp_sources=[CPP_SRC],
        cuda_sources=[CUDA_SRC],
        functions=[
            "geqrf_32_panel",
            "geqrf_176_panel",
            "geqrf_352_panel",
            "geqrf_512_panel",
            "geqrf_1024_panel",
            "geqrf_2048_panel",
            "geqrf_4096_panel",
            "geqrf_176_panel_t",
            "geqrf_352_panel_t",
            "geqrf_512_panel_t",
            "geqrf_1024_panel_t",
            "geqrf_2048_panel_t",
            "pack_v_panel",
            "make_t_panel",
        ],
        extra_cflags=["-O3"],
        extra_cuda_cflags=["-O3"],
        verbose=False,
    )
else:
    _native_module = None


def _make_v_inplace(panel_h: torch.Tensor, width: int) -> torch.Tensor:
    v = panel_h[:, :, :width]
    v.tril_(-1)
    diag = torch.arange(width, device=panel_h.device)
    v[:, diag, diag] = 1.0
    return v


def _make_t(
    v: torch.Tensor,
    tau: torch.Tensor,
    width: int,
    use_native_t: bool,
) -> torch.Tensor:
    batch = v.shape[0]
    if use_native_t and width <= 128 and v.is_cuda and _native_module is not None:
        t = torch.empty((batch, width, width), device=v.device, dtype=v.dtype)
        _native_module.make_t_panel(v, tau, t, width)
        return t

    t = torch.zeros((batch, width, width), device=v.device, dtype=v.dtype)
    for j in range(width):
        tau_j = tau[:, j]
        if j > 0:
            col = -tau_j[:, None] * torch.bmm(
                v[:, j:, :j].transpose(1, 2),
                v[:, j:, j : j + 1],
            ).squeeze(-1)
            t[:, :j, j] = torch.bmm(t[:, :j, :j], col.unsqueeze(-1)).squeeze(-1)
        t[:, j, j] = tau_j
    return t


def _panel_for_n(n: int) -> int:
    if n == 32:
        return _PANEL_32
    if n == 2048:
        return _PANEL_2048
    if n == 4096:
        return _PANEL_4096
    if n == 352:
        return _PANEL_352
    return _PANEL


def _apply_trailing_update(
    trailing: torch.Tensor,
    v: torch.Tensor,
    t: torch.Tensor,
    precompute_u: bool,
) -> None:
    work = torch.bmm(v.transpose(1, 2), trailing)
    if precompute_u:
        u = torch.bmm(v, t.transpose(1, 2))
        torch.baddbmm(trailing, u, work, beta=1.0, alpha=-1.0, out=trailing)
    else:
        work = torch.bmm(t.transpose(1, 2), work)
        torch.baddbmm(trailing, v, work, beta=1.0, alpha=-1.0, out=trailing)


def _native_geqrf_panel(
    h: torch.Tensor,
    tau: torch.Tensor,
    n: int,
    k: int,
    width: int,
    t: torch.Tensor | None = None,
) -> None:
    if n == 32:
        _native_module.geqrf_32_panel(h, tau, k, width)
    elif n == 176:
        if t is not None:
            _native_module.geqrf_176_panel_t(h, tau, t, k, width)
        else:
            _native_module.geqrf_176_panel(h, tau, k, width)
    elif n == 352:
        if t is not None:
            _native_module.geqrf_352_panel_t(h, tau, t, k, width)
        else:
            _native_module.geqrf_352_panel(h, tau, k, width)
    elif n == 512:
        if t is not None:
            _native_module.geqrf_512_panel_t(h, tau, t, k, width)
        else:
            _native_module.geqrf_512_panel(h, tau, k, width)
    elif n == 1024:
        if t is not None:
            _native_module.geqrf_1024_panel_t(h, tau, t, k, width)
        else:
            _native_module.geqrf_1024_panel(h, tau, k, width)
    elif n == 2048:
        if t is not None:
            _native_module.geqrf_2048_panel_t(h, tau, t, k, width)
        else:
            _native_module.geqrf_2048_panel(h, tau, k, width)


def _blocked_qr_superpanel_512(data: torch.Tensor) -> output_t:
    h = data.clone()
    batch, n, _ = h.shape
    tau = torch.empty((batch, n), device=h.device, dtype=h.dtype)
    panel = _PANEL
    superpanel = _SUPERPANEL_512
    precompute_u = n in _USE_PRECOMPUTE_U_NS

    for k in range(0, n, superpanel):
        block_width = min(superpanel, n - k)
        block_end = k + block_width

        for kk in range(k, block_end, panel):
            width = min(panel, block_end - kk)
            next_col = kk + width

            _native_geqrf_panel(h, tau, n, kk, width)

            if next_col >= block_end:
                continue

            v = torch.empty((batch, n - kk, width), device=h.device, dtype=h.dtype)
            _native_module.pack_v_panel(h, v, kk, width)
            panel_tau = tau[:, kk : kk + width]
            t = _make_t(v, panel_tau, width, True)
            local_trailing = h[:, kk:, next_col:block_end]
            _apply_trailing_update(local_trailing, v, t, precompute_u)

        if block_end >= n:
            continue

        v_block = torch.empty((batch, n - k, block_width), device=h.device, dtype=h.dtype)
        _native_module.pack_v_panel(h, v_block, k, block_width)
        t_block = _make_t(v_block, tau[:, k:block_end], block_width, True)
        trailing = h[:, k:, block_end:]
        _apply_trailing_update(trailing, v_block, t_block, precompute_u)

    return h, tau


def _blocked_qr(data: torch.Tensor, panel: int) -> output_t:
    h = data.clone()
    batch, n, _ = h.shape
    tau = torch.empty((batch, n), device=h.device, dtype=h.dtype)
    panel = _panel_for_n(n)
    use_native_panel = n in _NATIVE_PANEL_NS and _native_module is not None
    use_native_t = use_native_panel
    precompute_u = n in _USE_PRECOMPUTE_U_NS

    for k in range(0, n, panel):
        width = min(panel, n - k)
        next_col = k + width
        has_trailing = next_col < n
        use_panel_t = use_native_panel and n in _PANEL_T_NS and has_trailing
        t = None
        if use_native_panel:
            if use_panel_t:
                t = torch.empty((batch, width, width), device=h.device, dtype=h.dtype)
            if n == 32:
                _native_module.geqrf_32_panel(h, tau, k, width)
            elif n == 176:
                if use_panel_t:
                    _native_module.geqrf_176_panel_t(h, tau, t, k, width)
                else:
                    _native_module.geqrf_176_panel(h, tau, k, width)
            elif n == 352:
                if use_panel_t:
                    _native_module.geqrf_352_panel_t(h, tau, t, k, width)
                else:
                    _native_module.geqrf_352_panel(h, tau, k, width)
            elif n == 512:
                if use_panel_t:
                    _native_module.geqrf_512_panel_t(h, tau, t, k, width)
                else:
                    _native_module.geqrf_512_panel(h, tau, k, width)
            elif n == 1024:
                if use_panel_t:
                    _native_module.geqrf_1024_panel_t(h, tau, t, k, width)
                else:
                    _native_module.geqrf_1024_panel(h, tau, k, width)
            elif n == 2048:
                if use_panel_t:
                    _native_module.geqrf_2048_panel_t(h, tau, t, k, width)
                else:
                    _native_module.geqrf_2048_panel(h, tau, k, width)
            elif n == 4096:
                _native_module.geqrf_4096_panel(h, tau, k, width)
            panel_tau = tau[:, k : k + width]
            panel_h = None
        else:
            panel_h, panel_tau = torch.geqrf(h[:, k:, k : k + width].contiguous())
            h[:, k:, k : k + width] = panel_h
            tau[:, k : k + width] = panel_tau

        if not has_trailing:
            continue

        if use_native_panel:
            v = torch.empty((batch, n - k, width), device=h.device, dtype=h.dtype)
            _native_module.pack_v_panel(h, v, k, width)
            if t is None:
                t = _make_t(v, panel_tau, width, use_native_t)
        else:
            v = _make_v_inplace(panel_h, width)
            t = _make_t(v, panel_tau, width, use_native_t)
        trailing = h[:, k:, next_col:]
        _apply_trailing_update(trailing, v, t, precompute_u)

    return h, tau

def custom_kernel(data: input_t) -> output_t:
    if (
        data.is_cuda
        and data.dtype == torch.float32
        and data.dim() == 3
        and data.shape[-1] in _BLOCKED_NS
    ):
        if data.shape[-1] in _SUPERPANEL_NS and _native_module is not None:
            return _blocked_qr_superpanel_512(
                data.contiguous() if not data.is_contiguous() else data,
            )
        return _blocked_qr(
            data.contiguous() if not data.is_contiguous() else data,
            _panel_for_n(data.shape[-1]),
        )
    return torch.geqrf(data)
scrolls · 730 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