submission 835902
UjasShah · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 662 lines, June 9 Researcher Reciprocity License v1.0.
submission_panel_bucket_all.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-835902?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:4e0ec8a6d9909518baad7be54ff9d4dca6dd61352066499f1209fa9a126cb1c8
license declaredunknown
license concludedunknown
authorsUjasShah
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
_qr32_kernel[(data.shape[0],)](data, h, tau, num_warps=1)Kernel source
submission_panel_bucket_all.py662 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def _qr32_kernel(a_ptr, h_ptr, tau_ptr):
b = tl.program_id(0)
offs = tl.arange(0, 32)
rows = offs[:, None]
cols = offs[None, :]
base = b * 1024
mat = tl.load(a_ptr + base + rows * 32 + cols)
for k in tl.static_range(0, 32):
col_k = tl.sum(tl.where(cols == k, mat, 0.0), axis=1)
alpha = tl.sum(tl.where(offs == k, col_k, 0.0), axis=0)
below = offs > k
sigma = tl.sum(tl.where(below, col_k * col_k, 0.0), axis=0)
active = sigma > 0.0
norm = tl.sqrt(alpha * alpha + sigma)
beta_raw = tl.where(alpha <= 0.0, norm, -norm)
beta = tl.where(active, beta_raw, alpha)
safe_beta = tl.where(active, beta_raw, 1.0)
tau_k = tl.where(active, (beta_raw - alpha) / safe_beta, 0.0)
denom = tl.where(active, alpha - beta_raw, 1.0)
v_tail = tl.where(below, col_k / denom, 0.0)
v = tl.where(offs == k, 1.0, v_tail)
dots = tl.sum(v[:, None] * mat, axis=0)
updated = mat - tau_k * v[:, None] * dots[None, :]
mat = tl.where(cols > k, updated, mat)
diag_mask = (rows == k) & (cols == k)
tail_mask = (rows > k) & (cols == k)
mat = tl.where(diag_mask, beta, mat)
mat = tl.where(tail_mask, v_tail[:, None], mat)
tl.store(tau_ptr + b * 32 + k, tau_k)
tl.store(h_ptr + base + rows * 32 + cols, mat)
def _qr32(data: torch.Tensor) -> output_t:
h = torch.empty_like(data)
tau = torch.empty((data.shape[0], 32), device=data.device, dtype=data.dtype)
_qr32_kernel[(data.shape[0],)](data, h, tau, num_warps=1)
return h, tau
@triton.jit
def _qr_factor_kernel(h_ptr, tau_ptr, k, n: tl.constexpr, block_m: tl.constexpr):
b = tl.program_id(0)
rows = tl.arange(0, block_m)
base = b * n * n
col_k = tl.load(h_ptr + base + rows * n + k, mask=rows < n, other=0.0)
alpha = tl.load(h_ptr + base + k * n + k)
below = rows > k
sigma = tl.sum(tl.where(below, col_k * col_k, 0.0), axis=0)
active = sigma > 0.0
norm = tl.sqrt(alpha * alpha + sigma)
beta_raw = tl.where(alpha <= 0.0, norm, -norm)
beta = tl.where(active, beta_raw, alpha)
safe_beta = tl.where(active, beta_raw, 1.0)
tau_k = tl.where(active, (beta_raw - alpha) / safe_beta, 0.0)
denom = tl.where(active, alpha - beta_raw, 1.0)
v_tail = tl.where(below, col_k / denom, 0.0)
tl.store(tau_ptr + b * n + k, tau_k)
tl.store(h_ptr + base + k * n + k, beta)
tl.store(
h_ptr + base + rows * n + k,
v_tail,
mask=(rows > k) & (rows < n),
)
@triton.jit
def _qr_apply_kernel(h_ptr, tau_ptr, k, n: tl.constexpr, block_m: tl.constexpr, block_n: tl.constexpr):
b = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = col_block * block_n + tl.arange(0, block_n)
base = b * n * n
stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
v = tl.where(rows == k, 1.0, 0.0)
v = tl.where(rows > k, stored_tail, v)
tau_k = tl.load(tau_ptr + b * n + k)
tile = tl.load(
h_ptr + base + rows[:, None] * n + cols[None, :],
mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
other=0.0,
)
dots = tl.sum(v[:, None] * tile, axis=0)
updated = tile - tau_k * v[:, None] * dots[None, :]
tl.store(
h_ptr + base + rows[:, None] * n + cols[None, :],
updated,
mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
)
@triton.jit
def _qr_apply_range_kernel(
h_ptr,
tau_ptr,
k,
col_start: tl.constexpr,
n: tl.constexpr,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
b = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = col_start + tl.arange(0, block_n)
base = b * n * n
stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
v = tl.where(rows == k, 1.0, 0.0)
v = tl.where(rows > k, stored_tail, v)
tau_k = tl.load(tau_ptr + b * n + k)
tile = tl.load(
h_ptr + base + rows[:, None] * n + cols[None, :],
mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
other=0.0,
)
dots = tl.sum(v[:, None] * tile, axis=0)
updated = tile - tau_k * v[:, None] * dots[None, :]
tl.store(
h_ptr + base + rows[:, None] * n + cols[None, :],
updated,
mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
)
@triton.jit
def _qr_panel4_apply_kernel(
h_ptr,
tau_ptr,
panel_start,
n: tl.constexpr,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
b = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = panel_start + 4 + col_block * block_n + tl.arange(0, block_n)
base = b * n * n
tile = tl.load(
h_ptr + base + rows[:, None] * n + cols[None, :],
mask=(rows[:, None] < n) & (cols[None, :] < n),
other=0.0,
)
for j in tl.static_range(0, 4):
k = panel_start + j
stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
v = tl.where(rows == k, 1.0, 0.0)
v = tl.where(rows > k, stored_tail, v)
tau_k = tl.load(tau_ptr + b * n + k)
dots = tl.sum(v[:, None] * tile, axis=0)
tile = tile - tau_k * v[:, None] * dots[None, :]
tl.store(
h_ptr + base + rows[:, None] * n + cols[None, :],
tile,
mask=(rows[:, None] < n) & (cols[None, :] < n),
)
@triton.jit
def _qr_panel8_apply_kernel(
h_ptr,
tau_ptr,
panel_start,
n: tl.constexpr,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
b = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = panel_start + 8 + col_block * block_n + tl.arange(0, block_n)
base = b * n * n
tile = tl.load(
h_ptr + base + rows[:, None] * n + cols[None, :],
mask=(rows[:, None] < n) & (cols[None, :] < n),
other=0.0,
)
for j in tl.static_range(0, 8):
k = panel_start + j
stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
v = tl.where(rows == k, 1.0, 0.0)
v = tl.where(rows > k, stored_tail, v)
tau_k = tl.load(tau_ptr + b * n + k)
dots = tl.sum(v[:, None] * tile, axis=0)
tile = tile - tau_k * v[:, None] * dots[None, :]
tl.store(
h_ptr + base + rows[:, None] * n + cols[None, :],
tile,
mask=(rows[:, None] < n) & (cols[None, :] < n),
)
@triton.jit
def _qr_panel8_apply_tail_kernel(
h_ptr,
tau_ptr,
panel_start,
n: tl.constexpr,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
b = tl.program_id(0)
col_block = tl.program_id(1)
rows = panel_start + tl.arange(0, block_m)
cols = panel_start + 8 + col_block * block_n + tl.arange(0, block_n)
base = b * n * n
tile = tl.load(
h_ptr + base + rows[:, None] * n + cols[None, :],
mask=(rows[:, None] < n) & (cols[None, :] < n),
other=0.0,
)
for j in tl.static_range(0, 8):
k = panel_start + j
stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
v = tl.where(rows == k, 1.0, 0.0)
v = tl.where(rows > k, stored_tail, v)
tau_k = tl.load(tau_ptr + b * n + k)
dots = tl.sum(v[:, None] * tile, axis=0)
tile = tile - tau_k * v[:, None] * dots[None, :]
tl.store(
h_ptr + base + rows[:, None] * n + cols[None, :],
tile,
mask=(rows[:, None] < n) & (cols[None, :] < n),
)
@triton.jit
def _qr_panel8_factor_kernel(h_ptr, tau_ptr, panel_start, n: tl.constexpr, block_m: tl.constexpr):
b = tl.program_id(0)
rows = tl.arange(0, block_m)
pcols = tl.arange(0, 8)
cols = panel_start + pcols
base = b * n * n
panel = tl.load(
h_ptr + base + rows[:, None] * n + cols[None, :],
mask=(rows[:, None] < n) & (cols[None, :] < n),
other=0.0,
)
for j in tl.static_range(0, 8):
k = panel_start + j
col_j = tl.sum(tl.where(pcols[None, :] == j, panel, 0.0), axis=1)
alpha = tl.sum(tl.where(rows == k, col_j, 0.0), axis=0)
below = rows > k
sigma = tl.sum(tl.where(below, col_j * col_j, 0.0), axis=0)
active = sigma > 0.0
norm = tl.sqrt(alpha * alpha + sigma)
beta_raw = tl.where(alpha <= 0.0, norm, -norm)
beta = tl.where(active, beta_raw, alpha)
safe_beta = tl.where(active, beta_raw, 1.0)
tau_k = tl.where(active, (beta_raw - alpha) / safe_beta, 0.0)
denom = tl.where(active, alpha - beta_raw, 1.0)
v_tail = tl.where(below, col_j / denom, 0.0)
v = tl.where(rows == k, 1.0, 0.0)
v = tl.where(rows > k, v_tail, v)
dots = tl.sum(v[:, None] * panel, axis=0)
updated = panel - tau_k * v[:, None] * dots[None, :]
panel = tl.where(pcols[None, :] > j, updated, panel)
diag_mask = (rows[:, None] == k) & (pcols[None, :] == j)
tail_mask = (rows[:, None] > k) & (pcols[None, :] == j)
panel = tl.where(diag_mask, beta, panel)
panel = tl.where(tail_mask, v_tail[:, None], panel)
tl.store(tau_ptr + b * n + k, tau_k)
tl.store(
h_ptr + base + rows[:, None] * n + cols[None, :],
panel,
mask=(rows[:, None] < n) & (cols[None, :] < n),
)
def _qr176(data: torch.Tensor) -> output_t:
n = 176
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
apply_grid = (data.shape[0], triton.cdiv(n, 16))
for k in range(n):
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=256, num_warps=8)
_qr_apply_kernel[apply_grid](h, tau, k, n, block_m=256, block_n=16, num_warps=8)
return h, tau
def _qr176_panel8(data: torch.Tensor) -> output_t:
n = 176
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
_qr_panel8_factor_kernel[(data.shape[0],)](h, tau, panel_start, n, block_m=256, num_warps=4)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
active_rows = n - panel_start
if active_rows > 128:
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
elif active_rows > 64:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
elif active_rows > 32:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
else:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
return h, tau
def _qr352(data: torch.Tensor) -> output_t:
n = 352
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
apply_grid = (data.shape[0], triton.cdiv(n, 16))
for k in range(n):
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=512, num_warps=8)
_qr_apply_kernel[apply_grid](h, tau, k, n, block_m=512, block_n=16, num_warps=8)
return h, tau
def _qr352_panel8(data: torch.Tensor) -> output_t:
n = 352
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
_qr_panel8_factor_kernel[(data.shape[0],)](h, tau, panel_start, n, block_m=512, num_warps=4)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
active_rows = n - panel_start
if active_rows > 256:
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=512, block_n=32, num_warps=4)
elif active_rows > 128:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
elif active_rows > 64:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
elif active_rows > 32:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
else:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
return h, tau
def _qr512(data: torch.Tensor) -> output_t:
n = 512
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
apply_grid = (data.shape[0], triton.cdiv(n, 32))
for k in range(n):
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
_qr_apply_kernel[apply_grid](h, tau, k, n, block_m=1024, block_n=32, num_warps=8)
return h, tau
def _qr512_blocked4(data: torch.Tensor) -> output_t:
n = 512
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 4):
for j in range(4):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
if j < 3:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=1024,
block_n=4,
num_warps=8,
)
trailing = n - panel_start - 4
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
_qr_panel4_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
return h, tau
def _qr512_blocked8(data: torch.Tensor) -> output_t:
n = 512
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
for j in range(8):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
if j < 7:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=1024,
block_n=8,
num_warps=8,
)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
return h, tau
def _qr512_panel8(data: torch.Tensor) -> output_t:
n = 512
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
_qr_panel8_factor_kernel[(data.shape[0],)](h, tau, panel_start, n, block_m=512, num_warps=4)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
active_rows = n - panel_start
if active_rows > 256:
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=512, block_n=32, num_warps=4)
elif active_rows > 128:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
elif active_rows > 64:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
elif active_rows > 32:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
else:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
return h, tau
def _qr1024(data: torch.Tensor) -> output_t:
n = 1024
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
apply_grid = (data.shape[0], triton.cdiv(n, 32))
for k in range(n):
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
_qr_apply_kernel[apply_grid](h, tau, k, n, block_m=1024, block_n=32, num_warps=8)
return h, tau
def _qr1024_blocked4(data: torch.Tensor) -> output_t:
n = 1024
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 4):
for j in range(4):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
if j < 3:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=1024,
block_n=4,
num_warps=8,
)
trailing = n - panel_start - 4
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
_qr_panel4_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
return h, tau
def _qr1024_blocked8(data: torch.Tensor) -> output_t:
n = 1024
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
for j in range(8):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
if j < 7:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=1024,
block_n=8,
num_warps=8,
)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
return h, tau
def _qr1024_blocked8_bucket(data: torch.Tensor) -> output_t:
n = 1024
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
for j in range(8):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
if j < 7:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=1024,
block_n=8,
num_warps=8,
)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
active_rows = n - panel_start
if active_rows > 512:
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
elif active_rows > 256:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=512, block_n=32, num_warps=8)
elif active_rows > 128:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
elif active_rows > 64:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
elif active_rows > 32:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
else:
_qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
return h, tau
def _qr2048_blocked8(data: torch.Tensor) -> output_t:
n = 2048
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
for j in range(8):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=2048, num_warps=8)
if j < 7:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=2048,
block_n=8,
num_warps=8,
)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=2048, block_n=32, num_warps=8)
return h, tau
def _qr4096_blocked8(data: torch.Tensor) -> output_t:
n = 4096
h = data.clone()
tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
for panel_start in range(0, n, 8):
for j in range(8):
k = panel_start + j
_qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=4096, num_warps=8)
if j < 7:
_qr_apply_range_kernel[(data.shape[0],)](
h,
tau,
k,
panel_start,
n,
block_m=4096,
block_n=8,
num_warps=8,
)
trailing = n - panel_start - 8
if trailing:
apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
_qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=4096, block_n=32, num_warps=8)
return h, tau
def custom_kernel(data: input_t) -> output_t:
if (
data.shape[-1] == 32
and data.shape[-2] == 32
and data.dtype == torch.float32
and data.is_contiguous()
):
return _qr32(data)
if (
data.shape[-1] == 176
and data.shape[-2] == 176
and data.dtype == torch.float32
and data.is_contiguous()
):
return _qr176_panel8(data)
if (
data.shape[-1] == 352
and data.shape[-2] == 352
and data.dtype == torch.float32
and data.is_contiguous()
):
return _qr352_panel8(data)
if (
data.shape[-1] == 512
and data.shape[-2] == 512
and data.dtype == torch.float32
and data.is_contiguous()
):
return _qr512_panel8(data)
if (
data.shape[-1] == 1024
and data.shape[-2] == 1024
and data.dtype == torch.float32
and data.is_contiguous()
):
return _qr1024_blocked8_bucket(data)
return torch.geqrf(data)
scrolls · 662 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