submission 930424
gpuseed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2629 lines, June 9 Researcher Reciprocity License v1.0.
submission_1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-930424?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:daac65de2f36db1d5584f3a55d0982d7c4366c53e48ae2bf330c50140e985c0f
license declaredunknown
license concludedunknown
authorsgpuseed
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4,shared-memory
__shared__ float panels[1][2][32][32];Kernel source
submission_1.py2629 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True
torch.set_float32_matmul_precision("high")
_cuda_n32_source = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
__global__ void cholesky32_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
constexpr unsigned mask = 0xffffffffu;
const int lane = threadIdx.x & 31;
const int matrix = blockIdx.x * 4 + (threadIdx.x >> 5);
if (matrix >= batch) {
return;
}
const float* matrix_input = input + matrix * 32 * 32;
float* matrix_output = output + matrix * 32 * 32;
float row[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
row[col] = matrix_input[lane * 32 + col];
}
#pragma unroll
for (int pivot = 0; pivot < 32; ++pivot) {
float value = row[pivot];
#pragma unroll
for (int col = 0; col < pivot; ++col) {
value -= row[col] * __shfl_sync(mask, row[col], pivot);
}
if (lane == pivot) {
row[pivot] = sqrtf(value);
}
__syncwarp(mask);
const float diagonal = __shfl_sync(mask, row[pivot], pivot);
if (lane > pivot) {
row[pivot] = value / diagonal;
}
__syncwarp(mask);
}
#pragma unroll
for (int col = 0; col < 32; ++col) {
matrix_output[lane * 32 + col] = lane >= col ? row[col] : 0.0f;
}
}
torch::Tensor cholesky32_cuda(torch::Tensor input, torch::Tensor output) {
const int batch = static_cast<int>(input.size(0));
const int blocks = (batch + 3) / 4;
cholesky32_kernel<<<blocks, 128>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
batch
);
return output;
}
__global__ void cholesky64_recursive_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr unsigned mask = 0xffffffffu;
const int lane = threadIdx.x & 31;
const int warp_in_matrix = (threadIdx.x >> 5) & 1;
const int local_matrix = threadIdx.x >> 6;
const int matrix = blockIdx.x + local_matrix;
const float* matrix_input = input + matrix * 64 * 64;
float* matrix_output = output + matrix * 64 * 64;
// [column][row] keeps simultaneous row accesses bank-conflict free.
__shared__ float panels[1][2][32][32];
if (warp_in_matrix == 0) {
float row[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
row[col] = matrix_input[lane * 64 + col];
}
#pragma unroll
for (int pivot = 0; pivot < 32; ++pivot) {
float value = row[pivot];
#pragma unroll
for (int col = 0; col < pivot; ++col) {
value -= row[col] * __shfl_sync(mask, row[col], pivot);
}
if (lane == pivot) {
row[pivot] = sqrtf(value);
}
__syncwarp(mask);
const float diagonal = __shfl_sync(mask, row[pivot], pivot);
if (lane > pivot) {
row[pivot] = value / diagonal;
}
__syncwarp(mask);
}
#pragma unroll
for (int col = 0; col < 32; ++col) {
const float value = lane >= col ? row[col] : 0.0f;
panels[local_matrix][0][col][lane] = value;
matrix_output[lane * 64 + col] = value;
matrix_output[lane * 64 + col + 32] = 0.0f;
}
}
__syncthreads();
if (warp_in_matrix == 1) {
float panel[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
float value = matrix_input[(lane + 32) * 64 + col];
#pragma unroll
for (int inner = 0; inner < col; ++inner) {
value -= (
panel[inner]
* panels[local_matrix][0][inner][col]
);
}
panel[col] = (
value / panels[local_matrix][0][col][col]
);
panels[local_matrix][1][col][lane] = panel[col];
matrix_output[(lane + 32) * 64 + col] = panel[col];
}
}
__syncthreads();
if (warp_in_matrix == 1) {
float row[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
float value = matrix_input[(lane + 32) * 64 + col + 32];
#pragma unroll
for (int inner = 0; inner < 32; ++inner) {
value -= (
panels[local_matrix][1][inner][lane]
* panels[local_matrix][1][inner][col]
);
}
row[col] = lane >= col ? value : 0.0f;
}
#pragma unroll
for (int pivot = 0; pivot < 32; ++pivot) {
float value = row[pivot];
#pragma unroll
for (int col = 0; col < pivot; ++col) {
value -= row[col] * __shfl_sync(mask, row[col], pivot);
}
if (lane == pivot) {
row[pivot] = sqrtf(value);
}
__syncwarp(mask);
const float diagonal = __shfl_sync(mask, row[pivot], pivot);
if (lane > pivot) {
row[pivot] = value / diagonal;
}
__syncwarp(mask);
}
#pragma unroll
for (int col = 0; col < 32; ++col) {
matrix_output[(lane + 32) * 64 + col + 32] = (
lane >= col ? row[col] : 0.0f
);
}
}
}
torch::Tensor cholesky64_recursive_cuda(
torch::Tensor input,
torch::Tensor output
) {
const int batch = static_cast<int>(input.size(0));
cholesky64_recursive_kernel<<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>()
);
return output;
}
template <int stride>
__global__ void factor64_strided_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int tile_start
) {
constexpr unsigned mask = 0xffffffffu;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int matrix = blockIdx.x;
const float* matrix_input = input + matrix * stride * stride;
float* matrix_output = output + matrix * stride * stride;
__shared__ float panels[2][32][32];
if (warp == 0) {
// Lane is the column during global transfer, so each row is one
// coalesced transaction. Shared storage is [column][row].
#pragma unroll
for (int row_idx = 0; row_idx < 32; ++row_idx) {
panels[0][lane][row_idx] = matrix_input[
(tile_start + row_idx) * stride + tile_start + lane
];
}
} else {
#pragma unroll
for (int row_idx = 0; row_idx < 32; ++row_idx) {
panels[1][lane][row_idx] = matrix_input[
(tile_start + 32 + row_idx) * stride
+ tile_start + lane
];
}
}
__syncthreads();
if (warp == 0) {
float factor_row[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
factor_row[col] = panels[0][col][lane];
}
#pragma unroll
for (int pivot = 0; pivot < 32; ++pivot) {
float value = factor_row[pivot];
#pragma unroll
for (int col = 0; col < pivot; ++col) {
value -= (
factor_row[col]
* __shfl_sync(mask, factor_row[col], pivot)
);
}
if (lane == pivot) {
factor_row[pivot] = sqrtf(value);
}
__syncwarp(mask);
const float diagonal = __shfl_sync(
mask, factor_row[pivot], pivot
);
if (lane > pivot) {
factor_row[pivot] = value / diagonal;
}
__syncwarp(mask);
}
#pragma unroll
for (int col = 0; col < 32; ++col) {
panels[0][col][lane] = (
lane >= col ? factor_row[col] : 0.0f
);
}
}
__syncthreads();
if (warp == 1) {
float panel[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
float value = panels[1][col][lane];
#pragma unroll
for (int inner = 0; inner < col; ++inner) {
value -= panel[inner] * panels[0][inner][col];
}
panel[col] = value / panels[0][col][col];
panels[1][col][lane] = panel[col];
}
} else {
// Store L00 and its upper-right zero quadrant coalescently.
#pragma unroll
for (int row_idx = 0; row_idx < 32; ++row_idx) {
matrix_output[
(tile_start + row_idx) * stride + tile_start + lane
] = panels[0][lane][row_idx];
matrix_output[
(tile_start + row_idx) * stride
+ tile_start + 32 + lane
] = 0.0f;
if (tile_start == 0) {
matrix_output[row_idx * stride + 64 + lane] = 0.0f;
matrix_output[row_idx * stride + 96 + lane] = 0.0f;
}
}
}
__syncthreads();
if (warp == 1) {
// Store L10 and reuse the first tile buffer for A11.
#pragma unroll
for (int row_idx = 0; row_idx < 32; ++row_idx) {
matrix_output[
(tile_start + 32 + row_idx) * stride
+ tile_start + lane
] = panels[1][lane][row_idx];
panels[0][lane][row_idx] = matrix_input[
(tile_start + 32 + row_idx) * stride
+ tile_start + 32 + lane
];
}
}
__syncthreads();
if (warp == 1) {
float factor_row[32];
#pragma unroll
for (int col = 0; col < 32; ++col) {
float value = panels[0][col][lane];
#pragma unroll
for (int inner = 0; inner < 32; ++inner) {
value -= panels[1][inner][lane] * panels[1][inner][col];
}
factor_row[col] = lane >= col ? value : 0.0f;
}
#pragma unroll
for (int pivot = 0; pivot < 32; ++pivot) {
float value = factor_row[pivot];
#pragma unroll
for (int col = 0; col < pivot; ++col) {
value -= (
factor_row[col]
* __shfl_sync(mask, factor_row[col], pivot)
);
}
if (lane == pivot) {
factor_row[pivot] = sqrtf(value);
}
__syncwarp(mask);
const float diagonal = __shfl_sync(
mask, factor_row[pivot], pivot
);
if (lane > pivot) {
factor_row[pivot] = value / diagonal;
}
__syncwarp(mask);
}
#pragma unroll
for (int col = 0; col < 32; ++col) {
panels[0][col][lane] = (
lane >= col ? factor_row[col] : 0.0f
);
}
__syncwarp(mask);
#pragma unroll
for (int row_idx = 0; row_idx < 32; ++row_idx) {
matrix_output[
(tile_start + 32 + row_idx) * stride
+ tile_start + 32 + lane
] = panels[0][lane][row_idx];
if (tile_start == 0) {
matrix_output[
(row_idx + 32) * stride + 64 + lane
] = 0.0f;
matrix_output[
(row_idx + 32) * stride + 96 + lane
] = 0.0f;
}
}
}
}
torch::Tensor factor64_strided128_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start
) {
const int batch = static_cast<int>(input.size(0));
const int stride = static_cast<int>(input.size(1));
if (stride == 128) {
factor64_strided_kernel<128><<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start)
);
} else if (stride == 256) {
factor64_strided_kernel<256><<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start)
);
} else if (stride == 512) {
factor64_strided_kernel<512><<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start)
);
} else if (stride == 1024) {
factor64_strided_kernel<1024><<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start)
);
} else {
factor64_strided_kernel<2048><<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start)
);
}
return output;
}
template <int stride>
__global__ void trsm_panel64_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int tile_start,
int panel_count,
int panel_base
) {
const int thread = threadIdx.x;
const int matrix = blockIdx.x / panel_count;
const int panel_id = blockIdx.x - matrix * panel_count;
const int panel_start = panel_base + panel_id * 64;
const float* matrix_input = input + matrix * stride * stride;
float* matrix_output = output + matrix * stride * stride;
__shared__ float lower[64][64];
__shared__ float panel[64][64];
// Thread is the column during transfer, giving coalesced rows.
#pragma unroll
for (int row = 0; row < 64; ++row) {
lower[thread][row] = matrix_output[
(tile_start + row) * stride + tile_start + thread
];
panel[thread][row] = matrix_input[
(panel_start + row) * stride + tile_start + thread
];
}
__syncthreads();
// Thread becomes the panel row. [column][row] keeps row-parallel
// accesses conflict free throughout the forward substitution.
#pragma unroll
for (int col = 0; col < 64; ++col) {
float value = panel[col][thread];
#pragma unroll
for (int inner = 0; inner < col; ++inner) {
value -= panel[inner][thread] * lower[inner][col];
}
panel[col][thread] = value / lower[col][col];
}
__syncthreads();
// Reinterpret thread as the column for a coalesced global store.
#pragma unroll
for (int row = 0; row < 64; ++row) {
matrix_output[
(panel_start + row) * stride + tile_start + thread
] = panel[thread][row];
matrix_output[
(tile_start + row) * stride + panel_start + thread
] = 0.0f;
}
}
torch::Tensor trsm128_panel64_cuda(
torch::Tensor input,
torch::Tensor output
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<128><<<batch, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
0,
1,
64
);
return output;
}
torch::Tensor trsm256_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<256><<<batch * panel_count, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start),
static_cast<int>(panel_count),
static_cast<int>(tile_start + 64)
);
return output;
}
torch::Tensor trsm512_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<512><<<batch * panel_count, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start),
static_cast<int>(panel_count),
static_cast<int>(tile_start + 64)
);
return output;
}
torch::Tensor trsm1024_external_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<1024><<<batch * 8, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start),
8,
512
);
return output;
}
torch::Tensor trsm1024_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<1024><<<batch * panel_count, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start),
static_cast<int>(panel_count),
static_cast<int>(tile_start + 64)
);
return output;
}
torch::Tensor trsm2048_external_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<2048><<<batch * 16, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start),
16,
1024
);
return output;
}
torch::Tensor trsm2048_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
) {
const int batch = static_cast<int>(input.size(0));
trsm_panel64_kernel<2048><<<batch * panel_count, 64>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(tile_start),
static_cast<int>(panel_count),
static_cast<int>(tile_start + 64)
);
return output;
}
"""
_cuda_n32_cpp = r"""
#include <torch/extension.h>
torch::Tensor cholesky32_cuda(torch::Tensor input, torch::Tensor output);
torch::Tensor cholesky64_recursive_cuda(
torch::Tensor input,
torch::Tensor output
);
torch::Tensor factor64_strided128_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start
);
torch::Tensor trsm128_panel64_cuda(
torch::Tensor input,
torch::Tensor output
);
torch::Tensor trsm256_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
);
torch::Tensor trsm512_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
);
torch::Tensor trsm1024_external_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start
);
torch::Tensor trsm1024_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
);
torch::Tensor trsm2048_external_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start
);
torch::Tensor trsm2048_panel64_cuda(
torch::Tensor input,
torch::Tensor output,
int64_t tile_start,
int64_t panel_count
);
"""
_cuda_n32_module = load_inline(
name="cholesky_cuda_direct2048_v5",
cpp_sources=_cuda_n32_cpp,
cuda_sources=_cuda_n32_source,
functions=[
"cholesky32_cuda",
"cholesky64_recursive_cuda",
"factor64_strided128_cuda",
"trsm128_panel64_cuda",
"trsm256_panel64_cuda",
"trsm512_panel64_cuda",
"trsm1024_external_panel64_cuda",
"trsm1024_panel64_cuda",
"trsm2048_external_panel64_cuda",
"trsm2048_panel64_cuda",
],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
@triton.jit
def _cholesky_banachiewicz_32(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
):
"""Factor one 32x32 SPD matrix per Triton program."""
matrix_id = tl.program_id(0)
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
offsets = matrix_id * matrix_stride + rows * 32 + cols
# Keep the lower triangle in registers. At step k, previously computed
# columns contain L and the remaining columns still contain A.
values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)
for k in range(32):
pivot_row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(
tl.where(col_ids == k, pivot_row, 0.0),
axis=0,
)
diagonal -= tl.sum(
tl.where(col_ids < k, pivot_row * pivot_row, 0.0),
axis=0,
)
diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
current_column = tl.sum(
tl.where(cols == k, values, 0.0),
axis=1,
)
dot_products = tl.sum(
tl.where(cols < k, values * pivot_row[None, :], 0.0),
axis=1,
)
current_column = (current_column - dot_products) / diagonal
values = tl.where(
(rows == k) & (cols == k),
diagonal,
values,
)
values = tl.where(
(rows > k) & (cols == k),
current_column[:, None],
values,
)
tl.store(output_ptr + offsets, values)
@triton.jit
def _factor_lower_tile32(values):
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
for k in range(32):
pivot_row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(
tl.where(col_ids == k, pivot_row, 0.0),
axis=0,
)
diagonal -= tl.sum(
tl.where(col_ids < k, pivot_row * pivot_row, 0.0),
axis=0,
)
diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
current_column = tl.sum(
tl.where(cols == k, values, 0.0),
axis=1,
)
dot_products = tl.sum(
tl.where(cols < k, values * pivot_row[None, :], 0.0),
axis=1,
)
current_column = (current_column - dot_products) / diagonal
values = tl.where(
(rows == k) & (cols == k),
diagonal,
values,
)
values = tl.where(
(rows > k) & (cols == k),
current_column[:, None],
values,
)
return values
@triton.jit
def _potrf64_diagonal_tile(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
):
matrix_id = tl.program_id(0)
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
base = matrix_id * matrix_stride
offsets = (
base
+ (rows + tile_start) * 64
+ (cols + tile_start)
)
values = tl.where(
rows >= cols,
tl.load(input_ptr + offsets),
0.0,
)
values = _factor_lower_tile32(values)
tl.store(output_ptr + offsets, values)
if tile_start == 0:
upper_offsets = base + rows * 64 + (cols + 32)
tl.store(output_ptr + upper_offsets, 0.0)
@triton.jit
def _trsm64_panel(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
):
matrix_id = tl.program_id(0)
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
base = matrix_id * matrix_stride
offsets_00 = base + rows * 64 + cols
offsets_10 = base + (rows + 32) * 64 + cols
l00 = tl.load(output_ptr + offsets_00)
l10 = tl.load(input_ptr + offsets_10)
for k in range(32):
pivot_row = tl.sum(tl.where(rows == k, l00, 0.0), axis=0)
diagonal = tl.sum(
tl.where(col_ids == k, pivot_row, 0.0),
axis=0,
)
current_column = tl.sum(
tl.where(cols == k, l10, 0.0),
axis=1,
)
dot_products = tl.sum(
tl.where(cols < k, l10 * pivot_row[None, :], 0.0),
axis=1,
)
solved_column = (current_column - dot_products) / diagonal
l10 = tl.where(cols == k, solved_column[:, None], l10)
tl.store(output_ptr + offsets_10, l10)
@triton.jit
def _syrk64_trailing_update(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
):
matrix_id = tl.program_id(0)
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
base = matrix_id * matrix_stride
offsets_10 = base + (rows + 32) * 64 + cols
offsets_11 = base + (rows + 32) * 64 + (cols + 32)
l10 = tl.load(output_ptr + offsets_10)
a11 = tl.load(input_ptr + offsets_11)
update = tl.dot(
l10,
tl.trans(l10),
input_precision="ieee",
out_dtype=tl.float32,
)
schur = tl.where(rows >= cols, a11 - update, 0.0)
tl.store(output_ptr + offsets_11, schur)
@triton.jit
def _right_trsm32(panel, lower):
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
for k in range(32):
pivot_row = tl.sum(tl.where(rows == k, lower, 0.0), axis=0)
diagonal = tl.sum(
tl.where(col_ids == k, pivot_row, 0.0),
axis=0,
)
current_column = tl.sum(
tl.where(cols == k, panel, 0.0),
axis=1,
)
dot_products = tl.sum(
tl.where(cols < k, panel * pivot_row[None, :], 0.0),
axis=1,
)
solved_column = (current_column - dot_products) / diagonal
panel = tl.where(cols == k, solved_column[:, None], panel)
return panel
@triton.jit
def _cholesky64_fused(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
):
matrix_id = tl.program_id(0)
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
lower_mask = rows >= cols
base = matrix_id * matrix_stride
offsets_00 = base + rows * 64 + cols
offsets_01 = offsets_00 + 32
offsets_10 = offsets_00 + 32 * 64
offsets_11 = offsets_10 + 32
l00 = _factor_lower_tile32(
tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
)
l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
a11 = tl.where(
lower_mask,
tl.load(input_ptr + offsets_11),
0.0,
)
update = tl.dot(
l10,
tl.trans(l10),
input_precision="tf32x3",
out_dtype=tl.float32,
)
l11 = _factor_lower_tile32(
tl.where(lower_mask, a11 - update, 0.0)
)
tl.store(output_ptr + offsets_00, l00)
tl.store(output_ptr + offsets_01, 0.0)
tl.store(output_ptr + offsets_10, l10)
tl.store(output_ptr + offsets_11, l11)
@triton.jit
def _potrf128_tile64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
):
matrix_id = tl.program_id(0)
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
lower_mask = rows >= cols
base = matrix_id * matrix_stride
offsets_00 = base + (rows + tile_start) * 128 + cols + tile_start
offsets_01 = offsets_00 + 32
offsets_10 = offsets_00 + 32 * 128
offsets_11 = offsets_10 + 32
l00 = tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
l00 = _factor_lower_tile32(l00)
l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
a11 = tl.where(lower_mask, tl.load(input_ptr + offsets_11), 0.0)
update = tl.dot(
l10,
tl.trans(l10),
input_precision="tf32x3",
out_dtype=tl.float32,
)
l11 = _factor_lower_tile32(
tl.where(lower_mask, a11 - update, 0.0)
)
tl.store(output_ptr + offsets_00, l00)
tl.store(output_ptr + offsets_01, 0.0)
tl.store(output_ptr + offsets_10, l10)
tl.store(output_ptr + offsets_11, l11)
if tile_start == 0:
top_right = base + rows * 128 + cols + 64
tl.store(output_ptr + top_right, 0.0)
tl.store(output_ptr + top_right + 32, 0.0)
tl.store(output_ptr + top_right + 32 * 128, 0.0)
tl.store(output_ptr + top_right + 32 * 128 + 32, 0.0)
@triton.jit
def _trsm128_panel64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // 2
row_tile = program_id % 2
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
row_start = 64 + row_tile * 32
diagonal_00 = base + rows * 128 + cols
diagonal_10 = diagonal_00 + 32 * 128
diagonal_11 = diagonal_10 + 32
panel_0 = base + (rows + row_start) * 128 + cols
panel_1 = panel_0 + 32
l00 = tl.load(output_ptr + diagonal_00)
l10 = tl.load(output_ptr + diagonal_10)
l11 = tl.load(output_ptr + diagonal_11)
solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
rhs_1 = tl.load(input_ptr + panel_1)
rhs_1 -= tl.dot(
solved_0,
tl.trans(l10),
input_precision="tf32x3",
out_dtype=tl.float32,
)
solved_1 = _right_trsm32(rhs_1, l11)
tl.store(output_ptr + panel_0, solved_0)
tl.store(output_ptr + panel_1, solved_1)
@triton.jit
def _syrk128_trailing64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // 3
tile_id = program_id % 3
row_tile = tl.where(tile_id == 0, 0, 1)
col_tile = tl.where(tile_id == 2, 1, 0)
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
row_start = 64 + row_tile * 32
col_start = 64 + col_tile * 32
left_0 = base + (rows + row_start) * 128 + cols
left_1 = left_0 + 32
right_0 = base + (rows + col_start) * 128 + cols
right_1 = right_0 + 32
trailing = base + (rows + row_start) * 128 + cols + col_start
update = tl.dot(
tl.load(output_ptr + left_0),
tl.trans(tl.load(output_ptr + right_0)),
input_precision="tf32x3",
out_dtype=tl.float32,
)
update += tl.dot(
tl.load(output_ptr + left_1),
tl.trans(tl.load(output_ptr + right_1)),
input_precision="tf32x3",
out_dtype=tl.float32,
)
values = tl.load(input_ptr + trailing) - update
values = tl.where(
(row_tile > col_tile) | (rows >= cols),
values,
0.0,
)
tl.store(output_ptr + trailing, values)
@triton.jit
def _potrf256_tile64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
):
matrix_id = tl.program_id(0)
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
lower_mask = rows >= cols
base = matrix_id * matrix_stride
offsets_00 = base + (rows + tile_start) * 256 + cols + tile_start
offsets_01 = offsets_00 + 32
offsets_10 = offsets_00 + 32 * 256
offsets_11 = offsets_10 + 32
l00 = _factor_lower_tile32(
tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
)
l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
a11 = tl.where(lower_mask, tl.load(input_ptr + offsets_11), 0.0)
update = tl.dot(
l10,
tl.trans(l10),
input_precision="tf32x3",
out_dtype=tl.float32,
)
l11 = _factor_lower_tile32(
tl.where(lower_mask, a11 - update, 0.0)
)
tl.store(output_ptr + offsets_00, l00)
tl.store(output_ptr + offsets_01, 0.0)
tl.store(output_ptr + offsets_10, l10)
tl.store(output_ptr + offsets_11, l11)
@triton.jit
def _trsm256_panel64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
panel_count: tl.constexpr,
):
rows_per_matrix: tl.constexpr = panel_count * 2
program_id = tl.program_id(0)
matrix_id = program_id // rows_per_matrix
local_id = program_id % rows_per_matrix
panel_id = local_id // 2
row_tile = local_id % 2
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
row_start = tile_start + 64 + panel_id * 64 + row_tile * 32
diagonal_00 = (
base + (rows + tile_start) * 256 + cols + tile_start
)
diagonal_10 = diagonal_00 + 32 * 256
diagonal_11 = diagonal_10 + 32
panel_0 = base + (rows + row_start) * 256 + cols + tile_start
panel_1 = panel_0 + 32
l00 = tl.load(output_ptr + diagonal_00)
l10 = tl.load(output_ptr + diagonal_10)
l11 = tl.load(output_ptr + diagonal_11)
solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
rhs_1 = tl.load(input_ptr + panel_1)
rhs_1 -= tl.dot(
solved_0,
tl.trans(l10),
input_precision="tf32x3",
out_dtype=tl.float32,
)
solved_1 = _right_trsm32(rhs_1, l11)
tl.store(output_ptr + panel_0, solved_0)
tl.store(output_ptr + panel_1, solved_1)
upper_0 = base + (rows + tile_start) * 256 + cols + row_start
upper_1 = upper_0 + 32 * 256
tl.store(output_ptr + upper_0, 0.0)
tl.store(output_ptr + upper_1, 0.0)
@triton.jit
def _update256_tiles64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
remaining_tiles: tl.constexpr,
):
pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
programs_per_matrix: tl.constexpr = pair_count
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
local_id = program_id % programs_per_matrix
pair_id = local_id
if remaining_tiles == 3:
row_tile = tl.where(pair_id == 0, 0, tl.where(pair_id < 3, 1, 2))
col_tile = tl.where(
pair_id == 0,
0,
tl.where(
pair_id == 1,
0,
tl.where(pair_id == 2, 1, pair_id - 3),
),
)
elif remaining_tiles == 2:
row_tile = tl.where(pair_id == 0, 0, 1)
col_tile = tl.where(pair_id == 2, 1, 0)
else:
row_tile = 0
col_tile = 0
row_start = tile_start + 64 + row_tile * 64
col_start = tile_start + 64 + col_tile * 64
rows = tl.arange(0, 64)[:, None]
cols = tl.arange(0, 64)[None, :]
base = matrix_id * matrix_stride
left_0 = base + (rows + row_start) * 256 + cols + tile_start
right_0 = base + (rows + col_start) * 256 + cols + tile_start
trailing = base + (rows + row_start) * 256 + cols + col_start
update = tl.dot(
tl.load(output_ptr + left_0),
tl.trans(tl.load(output_ptr + right_0)),
input_precision="tf32x3",
out_dtype=tl.float32,
)
values = tl.load(input_ptr + trailing) - update
lower_tile = row_tile > col_tile
tl.store(
output_ptr + trailing,
tl.where(lower_tile | (rows >= cols), values, 0.0),
)
@triton.jit
def _zero256_upper(output_ptr, matrix_stride: tl.constexpr):
program_id = tl.program_id(0)
matrix_id = program_id // 24
local_id = program_id % 24
pair_id = local_id // 4
block_id = local_id % 4
row_tile = tl.where(pair_id < 3, 0, tl.where(pair_id < 5, 1, 2))
col_tile = tl.where(
pair_id == 0,
1,
tl.where(
pair_id == 1,
2,
tl.where(pair_id == 2, 3, tl.where(pair_id == 3, 2, 3)),
),
)
row_subtile = block_id // 2
col_subtile = block_id % 2
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
offsets = (
base
+ (rows + row_tile * 64 + row_subtile * 32) * 256
+ cols
+ col_tile * 64
+ col_subtile * 32
)
tl.store(output_ptr + offsets, 0.0)
@triton.jit
def _potrf512_tile64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
dot_precision: tl.constexpr,
):
matrix_id = tl.program_id(0)
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
lower_mask = rows >= cols
base = matrix_id * matrix_stride
offsets_00 = base + (rows + tile_start) * 512 + cols + tile_start
offsets_01 = offsets_00 + 32
offsets_10 = offsets_00 + 32 * 512
offsets_11 = offsets_10 + 32
l00 = _factor_lower_tile32(
tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
)
l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
a11 = tl.where(lower_mask, tl.load(input_ptr + offsets_11), 0.0)
update = tl.dot(
l10,
tl.trans(l10),
input_precision=dot_precision,
out_dtype=tl.float32,
)
l11 = _factor_lower_tile32(
tl.where(lower_mask, a11 - update, 0.0)
)
tl.store(output_ptr + offsets_00, l00)
tl.store(output_ptr + offsets_01, 0.0)
tl.store(output_ptr + offsets_10, l10)
tl.store(output_ptr + offsets_11, l11)
@triton.jit
def _trsm512_panel64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
panel_count: tl.constexpr,
dot_precision: tl.constexpr,
use_2d_grid: tl.constexpr,
):
rows_per_matrix: tl.constexpr = panel_count * 2
if use_2d_grid:
matrix_id = tl.program_id(1)
local_id = tl.program_id(0)
else:
program_id = tl.program_id(0)
matrix_id = program_id // rows_per_matrix
local_id = program_id % rows_per_matrix
panel_id = local_id // 2
row_tile = local_id % 2
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
row_start = tile_start + 64 + panel_id * 64 + row_tile * 32
diagonal_00 = (
base + (rows + tile_start) * 512 + cols + tile_start
)
diagonal_10 = diagonal_00 + 32 * 512
diagonal_11 = diagonal_10 + 32
panel_0 = base + (rows + row_start) * 512 + cols + tile_start
panel_1 = panel_0 + 32
l00 = tl.load(output_ptr + diagonal_00)
l10 = tl.load(output_ptr + diagonal_10)
l11 = tl.load(output_ptr + diagonal_11)
solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
rhs_1 = tl.load(input_ptr + panel_1)
rhs_1 -= tl.dot(
solved_0,
tl.trans(l10),
input_precision=dot_precision,
out_dtype=tl.float32,
)
solved_1 = _right_trsm32(rhs_1, l11)
tl.store(output_ptr + panel_0, solved_0)
tl.store(output_ptr + panel_1, solved_1)
upper_0 = base + (rows + tile_start) * 512 + cols + row_start
upper_1 = upper_0 + 32 * 512
tl.store(output_ptr + upper_0, 0.0)
tl.store(output_ptr + upper_1, 0.0)
@triton.jit
def _update512_tiles32(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
remaining_tiles: tl.constexpr,
dot_precision: tl.constexpr,
use_2d_grid: tl.constexpr,
):
pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
programs_per_matrix: tl.constexpr = pair_count * 4
if use_2d_grid:
matrix_id = tl.program_id(1)
local_id = tl.program_id(0)
else:
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
local_id = program_id % programs_per_matrix
pair_id = local_id // 4
block_id = local_id % 4
row_tile = tl.cast(
(tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = pair_id - row_tile * (row_tile + 1) // 2
row_subtile = block_id // 2
col_subtile = block_id % 2
row_start = tile_start + 64 + row_tile * 64 + row_subtile * 32
col_start = tile_start + 64 + col_tile * 64 + col_subtile * 32
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
left_0 = base + (rows + row_start) * 512 + cols + tile_start
left_1 = left_0 + 32
right_0 = base + (rows + col_start) * 512 + cols + tile_start
right_1 = right_0 + 32
trailing = base + (rows + row_start) * 512 + cols + col_start
update = tl.dot(
tl.load(output_ptr + left_0),
tl.trans(tl.load(output_ptr + right_0)),
input_precision=dot_precision,
out_dtype=tl.float32,
)
update += tl.dot(
tl.load(output_ptr + left_1),
tl.trans(tl.load(output_ptr + right_1)),
input_precision=dot_precision,
out_dtype=tl.float32,
)
values = tl.load(input_ptr + trailing) - update
lower_tile = row_tile > col_tile
lower_block = (row_subtile > col_subtile) | (
(row_subtile == col_subtile) & (rows >= cols)
)
tl.store(
output_ptr + trailing,
tl.where(lower_tile | lower_block, values, 0.0),
)
@triton.jit
def _update512_tiles64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
remaining_tiles: tl.constexpr,
dot_precision: tl.constexpr,
use_2d_grid: tl.constexpr,
):
pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
programs_per_matrix: tl.constexpr = pair_count
if use_2d_grid:
matrix_id = tl.program_id(1)
local_id = tl.program_id(0)
else:
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
local_id = program_id % programs_per_matrix
pair_id = local_id
row_tile = tl.cast(
(tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = pair_id - row_tile * (row_tile + 1) // 2
row_start = tile_start + 64 + row_tile * 64
col_start = tile_start + 64 + col_tile * 64
rows = tl.arange(0, 64)[:, None]
cols = tl.arange(0, 64)[None, :]
base = matrix_id * matrix_stride
left_0 = base + (rows + row_start) * 512 + cols + tile_start
right_0 = base + (rows + col_start) * 512 + cols + tile_start
trailing = base + (rows + row_start) * 512 + cols + col_start
update = tl.dot(
tl.load(output_ptr + left_0),
tl.trans(tl.load(output_ptr + right_0)),
input_precision=dot_precision,
out_dtype=tl.float32,
)
values = tl.load(input_ptr + trailing) - update
lower_tile = row_tile > col_tile
tl.store(
output_ptr + trailing,
tl.where(lower_tile | (rows >= cols), values, 0.0),
)
@triton.jit
def _trsm512_second64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
panel_count: tl.constexpr,
):
programs_per_matrix: tl.constexpr = panel_count * 4
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
row_block = program_id % programs_per_matrix
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
row_start = tile_start + 128 + row_block * 32
diagonal_00 = (
base + (rows + tile_start) * 512 + cols + tile_start
)
diagonal_20 = diagonal_00 + 64 * 512
diagonal_21 = diagonal_20 + 32
diagonal_22 = diagonal_20 + 64
diagonal_30 = diagonal_20 + 32 * 512
diagonal_31 = diagonal_30 + 32
diagonal_32 = diagonal_30 + 64
diagonal_33 = diagonal_30 + 96
panel_0 = base + (rows + row_start) * 512 + cols + tile_start
panel_1 = panel_0 + 32
panel_2 = panel_0 + 64
panel_3 = panel_0 + 96
solved_0 = tl.load(output_ptr + panel_0)
solved_1 = tl.load(output_ptr + panel_1)
l20 = tl.load(output_ptr + diagonal_20)
l21 = tl.load(output_ptr + diagonal_21)
l22 = tl.load(output_ptr + diagonal_22)
l30 = tl.load(output_ptr + diagonal_30)
l31 = tl.load(output_ptr + diagonal_31)
l32 = tl.load(output_ptr + diagonal_32)
l33 = tl.load(output_ptr + diagonal_33)
rhs_2 = tl.load(input_ptr + panel_2)
rhs_2 -= tl.dot(
solved_0,
tl.trans(l20),
input_precision="tf32",
out_dtype=tl.float32,
)
rhs_2 -= tl.dot(
solved_1,
tl.trans(l21),
input_precision="tf32",
out_dtype=tl.float32,
)
solved_2 = _right_trsm32(rhs_2, l22)
rhs_3 = tl.load(input_ptr + panel_3)
rhs_3 -= tl.dot(
solved_0,
tl.trans(l30),
input_precision="tf32",
out_dtype=tl.float32,
)
rhs_3 -= tl.dot(
solved_1,
tl.trans(l31),
input_precision="tf32",
out_dtype=tl.float32,
)
rhs_3 -= tl.dot(
solved_2,
tl.trans(l32),
input_precision="tf32",
out_dtype=tl.float32,
)
solved_3 = _right_trsm32(rhs_3, l33)
tl.store(output_ptr + panel_2, solved_2)
tl.store(output_ptr + panel_3, solved_3)
upper_2 = (
base + (rows + tile_start + 64) * 512 + cols + row_start
)
tl.store(output_ptr + upper_2, 0.0)
tl.store(output_ptr + upper_2 + 32 * 512, 0.0)
@triton.jit
def _trsm512_panel128(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
panel_count: tl.constexpr,
):
programs_per_matrix: tl.constexpr = panel_count * 4
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
local_id = program_id % programs_per_matrix
panel_id = local_id // 4
row_block = local_id % 4
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
row_start = tile_start + 128 + panel_id * 128 + row_block * 32
diagonal_00 = (
base + (rows + tile_start) * 512 + cols + tile_start
)
diagonal_10 = diagonal_00 + 32 * 512
diagonal_11 = diagonal_10 + 32
diagonal_20 = diagonal_00 + 64 * 512
diagonal_21 = diagonal_20 + 32
diagonal_22 = diagonal_20 + 64
diagonal_30 = diagonal_20 + 32 * 512
diagonal_31 = diagonal_30 + 32
diagonal_32 = diagonal_30 + 64
diagonal_33 = diagonal_30 + 96
panel_0 = base + (rows + row_start) * 512 + cols + tile_start
panel_1 = panel_0 + 32
panel_2 = panel_0 + 64
panel_3 = panel_0 + 96
l00 = tl.load(output_ptr + diagonal_00)
l10 = tl.load(output_ptr + diagonal_10)
l11 = tl.load(output_ptr + diagonal_11)
l20 = tl.load(output_ptr + diagonal_20)
l21 = tl.load(output_ptr + diagonal_21)
l22 = tl.load(output_ptr + diagonal_22)
l30 = tl.load(output_ptr + diagonal_30)
l31 = tl.load(output_ptr + diagonal_31)
l32 = tl.load(output_ptr + diagonal_32)
l33 = tl.load(output_ptr + diagonal_33)
solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
rhs_1 = tl.load(input_ptr + panel_1)
rhs_1 -= tl.dot(
solved_0,
tl.trans(l10),
input_precision="tf32",
out_dtype=tl.float32,
)
solved_1 = _right_trsm32(rhs_1, l11)
rhs_2 = tl.load(input_ptr + panel_2)
rhs_2 -= tl.dot(
solved_0,
tl.trans(l20),
input_precision="tf32",
out_dtype=tl.float32,
)
rhs_2 -= tl.dot(
solved_1,
tl.trans(l21),
input_precision="tf32",
out_dtype=tl.float32,
)
solved_2 = _right_trsm32(rhs_2, l22)
rhs_3 = tl.load(input_ptr + panel_3)
rhs_3 -= tl.dot(
solved_0,
tl.trans(l30),
input_precision="tf32",
out_dtype=tl.float32,
)
rhs_3 -= tl.dot(
solved_1,
tl.trans(l31),
input_precision="tf32",
out_dtype=tl.float32,
)
rhs_3 -= tl.dot(
solved_2,
tl.trans(l32),
input_precision="tf32",
out_dtype=tl.float32,
)
solved_3 = _right_trsm32(rhs_3, l33)
tl.store(output_ptr + panel_0, solved_0)
tl.store(output_ptr + panel_1, solved_1)
tl.store(output_ptr + panel_2, solved_2)
tl.store(output_ptr + panel_3, solved_3)
upper_0 = base + (rows + tile_start) * 512 + cols + row_start
tl.store(output_ptr + upper_0, 0.0)
tl.store(output_ptr + upper_0 + 32 * 512, 0.0)
tl.store(output_ptr + upper_0 + 64 * 512, 0.0)
tl.store(output_ptr + upper_0 + 96 * 512, 0.0)
@triton.jit
def _update512_second_panel64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
panel_count: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // panel_count
row_tile = program_id % panel_count
rows = tl.arange(0, 64)
cols = tl.arange(0, 64)
k_ids = tl.arange(0, 64)
row_start = tile_start + 64 + row_tile * 64
col_start = tile_start + 64
base = matrix_id * matrix_stride
left = tl.load(
output_ptr
+ base
+ (rows[:, None] + row_start) * 512
+ k_ids[None, :]
+ tile_start
)
right = tl.load(
output_ptr
+ base
+ (cols[None, :] + col_start) * 512
+ k_ids[:, None]
+ tile_start
)
update = tl.dot(
left,
right,
input_precision="tf32",
out_dtype=tl.float32,
)
panel = (
base
+ (rows[:, None] + row_start) * 512
+ cols[None, :]
+ col_start
)
values = tl.load(input_ptr + panel) - update
lower_mask = (row_tile > 0) | (
rows[:, None] >= cols[None, :]
)
tl.store(
output_ptr + panel,
tl.where(lower_mask, values, 0.0),
)
@triton.jit
def _update512_tiles128(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
remaining_tiles: tl.constexpr,
):
pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
program_id = tl.program_id(0)
matrix_id = program_id // pair_count
pair_id = program_id % pair_count
row_tile = tl.cast(
(tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = pair_id - row_tile * (row_tile + 1) // 2
rows = tl.arange(0, 64)
cols = tl.arange(0, 64)
k_ids = tl.arange(0, 128)
row_start = tile_start + 128 + row_tile * 64
col_start = tile_start + 128 + col_tile * 64
base = matrix_id * matrix_stride
left = tl.load(
output_ptr
+ base
+ (rows[:, None] + row_start) * 512
+ k_ids[None, :]
+ tile_start
)
right = tl.load(
output_ptr
+ base
+ (cols[None, :] + col_start) * 512
+ k_ids[:, None]
+ tile_start
)
update = tl.dot(
left,
right,
input_precision="tf32",
out_dtype=tl.float32,
)
trailing = (
base
+ (rows[:, None] + row_start) * 512
+ cols[None, :]
+ col_start
)
values = tl.load(input_ptr + trailing) - update
lower_mask = (row_tile > col_tile) | (
rows[:, None] >= cols[None, :]
)
tl.store(
output_ptr + trailing,
tl.where(lower_mask, values, 0.0),
)
@triton.jit
def _update1024_second_panel64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start,
panel_count,
dot_precision: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // panel_count
row_tile = program_id % panel_count
rows = tl.arange(0, 64)
cols = tl.arange(0, 64)
k_ids = tl.arange(0, 64)
row_start = tile_start + 64 + row_tile * 64
col_start = tile_start + 64
base = matrix_id * matrix_stride
left = tl.load(
output_ptr
+ base
+ (rows[:, None] + row_start) * 1024
+ k_ids[None, :]
+ tile_start
)
right = tl.load(
output_ptr
+ base
+ (cols[None, :] + col_start) * 1024
+ k_ids[:, None]
+ tile_start
)
update = tl.dot(
left,
right,
input_precision=dot_precision,
out_dtype=tl.float32,
)
panel = (
base
+ (rows[:, None] + row_start) * 1024
+ cols[None, :]
+ col_start
)
values = tl.load(input_ptr + panel) - update
lower_mask = (row_tile > 0) | (rows[:, None] >= cols[None, :])
tl.store(
output_ptr + panel,
tl.where(lower_mask, values, 0.0),
)
@triton.jit
def _update1024_tiles128(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start,
remaining_tiles,
dot_precision: tl.constexpr,
):
pair_count = remaining_tiles * (remaining_tiles + 1) // 2
program_id = tl.program_id(0)
matrix_id = program_id // pair_count
pair_id = program_id % pair_count
row_tile = tl.cast(
(tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = pair_id - row_tile * (row_tile + 1) // 2
rows = tl.arange(0, 64)
cols = tl.arange(0, 64)
k_ids = tl.arange(0, 128)
row_start = tile_start + 128 + row_tile * 64
col_start = tile_start + 128 + col_tile * 64
base = matrix_id * matrix_stride
left = tl.load(
output_ptr
+ base
+ (rows[:, None] + row_start) * 1024
+ k_ids[None, :]
+ tile_start
)
right = tl.load(
output_ptr
+ base
+ (cols[None, :] + col_start) * 1024
+ k_ids[:, None]
+ tile_start
)
update = tl.dot(
left,
right,
input_precision=dot_precision,
out_dtype=tl.float32,
)
trailing = (
base
+ (rows[:, None] + row_start) * 1024
+ cols[None, :]
+ col_start
)
values = tl.load(input_ptr + trailing) - update
lower_mask = (row_tile > col_tile) | (
rows[:, None] >= cols[None, :]
)
tl.store(
output_ptr + trailing,
tl.where(lower_mask, values, 0.0),
)
@triton.jit
def _update2048_second_panel64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start,
panel_count,
dot_precision: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // panel_count
row_tile = program_id % panel_count
rows = tl.arange(0, 64)
cols = tl.arange(0, 64)
k_ids = tl.arange(0, 64)
row_start = tile_start + 64 + row_tile * 64
col_start = tile_start + 64
base = matrix_id * matrix_stride
left = tl.load(
output_ptr
+ base
+ (rows[:, None] + row_start) * 2048
+ k_ids[None, :]
+ tile_start
)
right = tl.load(
output_ptr
+ base
+ (cols[None, :] + col_start) * 2048
+ k_ids[:, None]
+ tile_start
)
update = tl.dot(
left,
right,
input_precision=dot_precision,
out_dtype=tl.float32,
)
panel = (
base
+ (rows[:, None] + row_start) * 2048
+ cols[None, :]
+ col_start
)
values = tl.load(input_ptr + panel) - update
lower_mask = (row_tile > 0) | (rows[:, None] >= cols[None, :])
tl.store(
output_ptr + panel,
tl.where(lower_mask, values, 0.0),
)
@triton.jit
def _update2048_tiles128(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start,
remaining_tiles,
dot_precision: tl.constexpr,
):
pair_count = remaining_tiles * (remaining_tiles + 1) // 2
program_id = tl.program_id(0)
matrix_id = program_id // pair_count
pair_id = program_id % pair_count
row_tile = tl.cast(
(tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = pair_id - row_tile * (row_tile + 1) // 2
rows = tl.arange(0, 64)
cols = tl.arange(0, 64)
k_ids = tl.arange(0, 128)
row_start = tile_start + 128 + row_tile * 64
col_start = tile_start + 128 + col_tile * 64
base = matrix_id * matrix_stride
left = tl.load(
output_ptr
+ base
+ (rows[:, None] + row_start) * 2048
+ k_ids[None, :]
+ tile_start
)
right = tl.load(
output_ptr
+ base
+ (cols[None, :] + col_start) * 2048
+ k_ids[:, None]
+ tile_start
)
update = tl.dot(
left,
right,
input_precision=dot_precision,
out_dtype=tl.float32,
)
trailing = (
base
+ (rows[:, None] + row_start) * 2048
+ cols[None, :]
+ col_start
)
values = tl.load(input_ptr + trailing) - update
lower_mask = (row_tile > col_tile) | (
rows[:, None] >= cols[None, :]
)
tl.store(
output_ptr + trailing,
tl.where(lower_mask, values, 0.0),
)
@triton.jit
def _zero512_upper(output_ptr, matrix_stride: tl.constexpr):
program_id = tl.program_id(0)
matrix_id = program_id // 120
pair_id = program_id % 120
col_block = tl.cast(
(1.0 + tl.sqrt(1.0 + 8.0 * pair_id)) * 0.5,
tl.int32,
)
row_block = pair_id - col_block * (col_block - 1) // 2
rows = tl.arange(0, 32)[:, None]
cols = tl.arange(0, 32)[None, :]
base = matrix_id * matrix_stride
offsets = (
base + (rows + row_block * 32) * 512 + cols + col_block * 32
)
tl.store(output_ptr + offsets, 0.0)
@triton.jit
def _gather_lower_tiles(
input_ptr,
output_ptr,
input_batch_stride: tl.constexpr,
input_ld: tl.constexpr,
output_n: tl.constexpr,
tile_count: tl.constexpr,
):
pair_count: tl.constexpr = tile_count * (tile_count + 1) // 2
program_id = tl.program_id(0)
matrix_id = program_id // pair_count
tile_id = program_id % pair_count
row_tile = tl.cast(
(tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = tile_id - row_tile * (row_tile + 1) // 2
rows = row_tile * 64 + tl.arange(0, 64)[:, None]
cols = col_tile * 64 + tl.arange(0, 64)[None, :]
input_offsets = (
matrix_id * input_batch_stride
+ rows * input_ld
+ cols
)
output_offsets = (
matrix_id * output_n * output_n
+ rows * output_n
+ cols
)
tl.store(output_ptr + output_offsets, tl.load(input_ptr + input_offsets))
@triton.jit
def _update1024_trsm_tiles64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
remaining_tiles: tl.constexpr,
):
programs_per_matrix: tl.constexpr = 8 * remaining_tiles
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
local_id = program_id % programs_per_matrix
row_tile = local_id // remaining_tiles
col_tile = local_id % remaining_tiles
row_start = 512 + row_tile * 64
col_start = tile_start + 64 + col_tile * 64
rows = tl.arange(0, 64)[:, None]
cols = tl.arange(0, 64)[None, :]
base = matrix_id * matrix_stride
solved = base + (rows + row_start) * 1024 + cols + tile_start
factor = base + (rows + col_start) * 1024 + cols + tile_start
trailing = base + (rows + row_start) * 1024 + cols + col_start
update = tl.dot(
tl.load(output_ptr + solved),
tl.trans(tl.load(output_ptr + factor)),
input_precision="tf32x3",
out_dtype=tl.float32,
)
values = tl.load(input_ptr + trailing) - update
tl.store(output_ptr + trailing, values)
@triton.jit
def _update2048_trsm_tiles64(
input_ptr,
output_ptr,
matrix_stride: tl.constexpr,
tile_start: tl.constexpr,
remaining_tiles: tl.constexpr,
):
programs_per_matrix: tl.constexpr = 16 * remaining_tiles
program_id = tl.program_id(0)
matrix_id = program_id // programs_per_matrix
local_id = program_id % programs_per_matrix
row_tile = local_id // remaining_tiles
col_tile = local_id % remaining_tiles
row_start = 1024 + row_tile * 64
col_start = tile_start + 64 + col_tile * 64
rows = tl.arange(0, 64)[:, None]
cols = tl.arange(0, 64)[None, :]
base = matrix_id * matrix_stride
solved = base + (rows + row_start) * 2048 + cols + tile_start
factor = base + (rows + col_start) * 2048 + cols + tile_start
trailing = base + (rows + row_start) * 2048 + cols + col_start
update = tl.dot(
tl.load(output_ptr + solved),
tl.trans(tl.load(output_ptr + factor)),
input_precision="tf32x3",
out_dtype=tl.float32,
)
values = tl.load(input_ptr + trailing) - update
tl.store(output_ptr + trailing, values)
@triton.jit
def _schur1024_split512(
input_ptr,
l10_ptr,
schur_ptr,
matrix_stride: tl.constexpr,
l10_batch_stride: tl.constexpr,
l10_row_stride: tl.constexpr,
l10_col_stride: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // 36
tile_id = program_id % 36
row_tile = tl.cast(
(tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = tile_id - row_tile * (row_tile + 1) // 2
rows = row_tile * 64 + tl.arange(0, 64)
cols = col_tile * 64 + tl.arange(0, 64)
accumulator = tl.zeros((64, 64), tl.float32)
l10_base = matrix_id * l10_batch_stride
for k_start in range(0, 512, 64):
k_ids = k_start + tl.arange(0, 64)
left = tl.load(
l10_ptr
+ l10_base
+ rows[:, None] * l10_row_stride
+ k_ids[None, :] * l10_col_stride
)
right = tl.load(
l10_ptr
+ l10_base
+ k_ids[:, None] * l10_col_stride
+ cols[None, :] * l10_row_stride
)
accumulator += tl.dot(
left,
right,
input_precision="tf32x3",
out_dtype=tl.float32,
)
input_offsets = (
matrix_id * matrix_stride
+ (rows[:, None] + 512) * 1024
+ cols[None, :]
+ 512
)
schur_offsets = (
matrix_id * 512 * 512 + rows[:, None] * 512 + cols[None, :]
)
values = tl.load(input_ptr + input_offsets) - accumulator
lower_mask = (row_tile > col_tile) | (
rows[:, None] >= cols[None, :]
)
tl.store(schur_ptr + schur_offsets, tl.where(lower_mask, values, 0.0))
@triton.jit
def _schur2048_split1024(
input_ptr,
l10_ptr,
schur_ptr,
matrix_stride: tl.constexpr,
l10_batch_stride: tl.constexpr,
l10_row_stride: tl.constexpr,
l10_col_stride: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // 136
tile_id = program_id % 136
row_tile = tl.cast(
(tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = tile_id - row_tile * (row_tile + 1) // 2
rows = row_tile * 64 + tl.arange(0, 64)
cols = col_tile * 64 + tl.arange(0, 64)
accumulator = tl.zeros((64, 64), tl.float32)
l10_base = matrix_id * l10_batch_stride
for k_start in range(0, 1024, 64):
k_ids = k_start + tl.arange(0, 64)
left = tl.load(
l10_ptr
+ l10_base
+ rows[:, None] * l10_row_stride
+ k_ids[None, :] * l10_col_stride
)
right = tl.load(
l10_ptr
+ l10_base
+ k_ids[:, None] * l10_col_stride
+ cols[None, :] * l10_row_stride
)
accumulator += tl.dot(
left,
right,
input_precision="tf32x3",
out_dtype=tl.float32,
)
input_offsets = (
matrix_id * matrix_stride
+ (rows[:, None] + 1024) * 2048
+ cols[None, :]
+ 1024
)
schur_offsets = (
matrix_id * 1024 * 1024 + rows[:, None] * 1024 + cols[None, :]
)
values = tl.load(input_ptr + input_offsets) - accumulator
lower_mask = (row_tile > col_tile) | (
rows[:, None] >= cols[None, :]
)
tl.store(schur_ptr + schur_offsets, tl.where(lower_mask, values, 0.0))
@triton.jit
def _schur4096_split2048(
input_ptr,
l10_ptr,
schur_ptr,
matrix_stride: tl.constexpr,
l10_batch_stride: tl.constexpr,
l10_row_stride: tl.constexpr,
l10_col_stride: tl.constexpr,
):
program_id = tl.program_id(0)
matrix_id = program_id // 528
tile_id = program_id % 528
row_tile = tl.cast(
(tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
tl.int32,
)
col_tile = tile_id - row_tile * (row_tile + 1) // 2
rows = row_tile * 64 + tl.arange(0, 64)
cols = col_tile * 64 + tl.arange(0, 64)
accumulator = tl.zeros((64, 64), tl.float32)
l10_base = matrix_id * l10_batch_stride
for k_start in range(0, 2048, 64):
k_ids = k_start + tl.arange(0, 64)
left = tl.load(
l10_ptr
+ l10_base
+ rows[:, None] * l10_row_stride
+ k_ids[None, :] * l10_col_stride
)
right = tl.load(
l10_ptr
+ l10_base
+ k_ids[:, None] * l10_col_stride
+ cols[None, :] * l10_row_stride
)
accumulator += tl.dot(
left,
right,
input_precision="tf32x3",
out_dtype=tl.float32,
)
input_offsets = (
matrix_id * matrix_stride
+ (rows[:, None] + 2048) * 4096
+ cols[None, :]
+ 2048
)
schur_offsets = (
matrix_id * 2048 * 2048 + rows[:, None] * 2048 + cols[None, :]
)
values = tl.load(input_ptr + input_offsets) - accumulator
lower_mask = (row_tile > col_tile) | (
rows[:, None] >= cols[None, :]
)
tl.store(schur_ptr + schur_offsets, tl.where(lower_mask, values, 0.0))
def _serial_cholesky(data):
output = torch.empty_like(data)
info = torch.empty(
(data.shape[0],),
dtype=torch.int32,
device=data.device,
)
for index in range(data.shape[0]):
torch.linalg.cholesky_ex(
data[index],
check_errors=False,
out=(output[index], info[index]),
)
return output
def _pure_deferred_large_blocks(data: torch.Tensor, block: int) -> torch.Tensor:
_, n, _ = data.shape
output = torch.zeros_like(data)
for tile_start in range(0, n, block):
tile_end = tile_start + block
if tile_start == 0:
source = data
else:
active_history = output[:, tile_start:, :tile_start]
pivot_history = output[:, tile_start:tile_end, :tile_start]
torch.baddbmm(
data[:, tile_start:, tile_start:tile_end],
active_history,
pivot_history.transpose(1, 2),
beta=1.0,
alpha=-1.0,
out=output[:, tile_start:, tile_start:tile_end],
)
source = output
diagonal = source[:, tile_start:tile_end, tile_start:tile_end]
diagonal_factor = torch.linalg.cholesky_ex(
diagonal,
check_errors=False,
).L
output[:, tile_start:tile_end, tile_start:tile_end].copy_(
diagonal_factor
)
if tile_end < n:
panel = source[:, tile_end:, tile_start:tile_end]
solved = torch.linalg.solve_triangular(
diagonal_factor,
panel.transpose(-1, -2),
upper=False,
)
output[:, tile_end:, tile_start:tile_end].copy_(
solved.transpose(-1, -2)
)
diagonal = torch.diagonal(output, dim1=-2, dim2=-1)
valid = torch.isfinite(diagonal).all() & (diagonal > 0.0).all()
if not bool(valid.item()):
return torch.linalg.cholesky_ex(data, check_errors=False).L
return output
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
if (n == 2048 and batch == 2) or (n == 4096 and batch == 2):
return _serial_cholesky(data)
if n == 2048 and batch >= 8:
output = torch.empty_like(data)
for tile_start in range(0, 2048, 128):
remaining_blocks = (1920 - tile_start) // 128
if tile_start == 0:
source = data
else:
active_history = output[:, tile_start:, :tile_start]
pivot_history = output[
:,
tile_start:tile_start + 128,
:tile_start,
]
torch.baddbmm(
data[:, tile_start:, tile_start:tile_start + 128],
active_history,
pivot_history.transpose(1, 2),
beta=1.0,
alpha=-1.0,
out=output[
:,
tile_start:,
tile_start:tile_start + 128,
],
)
source = output
_cuda_n32_module.factor64_strided128_cuda(
source,
output,
tile_start,
)
first_half_panels = 1 + remaining_blocks * 2
_cuda_n32_module.trsm2048_panel64_cuda(
source,
output,
tile_start,
first_half_panels,
)
_update2048_second_panel64[
(batch * first_half_panels,)
](
source,
output,
n * n,
tile_start=tile_start,
panel_count=first_half_panels,
dot_precision="tf32x3",
num_warps=4,
)
_cuda_n32_module.factor64_strided128_cuda(
output,
output,
tile_start + 64,
)
remaining_tiles = remaining_blocks * 2
if remaining_tiles:
_cuda_n32_module.trsm2048_panel64_cuda(
output,
output,
tile_start + 64,
remaining_tiles,
)
diagonal = torch.diagonal(output, dim1=-2, dim2=-1)
valid = torch.isfinite(diagonal).all() & (diagonal > 0.0).all()
if not bool(valid.item()):
return torch.linalg.cholesky_ex(data, check_errors=False).L
return output
if n == 1024 and batch >= 4:
output = torch.empty_like(data)
for tile_start in range(0, 1024, 128):
source = data if tile_start == 0 else output
remaining_blocks = (896 - tile_start) // 128
_cuda_n32_module.factor64_strided128_cuda(
source,
output,
tile_start,
)
first_half_panels = 1 + remaining_blocks * 2
_cuda_n32_module.trsm1024_panel64_cuda(
source,
output,
tile_start,
first_half_panels,
)
_update1024_second_panel64[
(batch * first_half_panels,)
](
source,
output,
n * n,
tile_start=tile_start,
panel_count=first_half_panels,
dot_precision="tf32x3",
num_warps=4,
)
_cuda_n32_module.factor64_strided128_cuda(
output,
output,
tile_start + 64,
)
remaining_tiles = remaining_blocks * 2
if remaining_tiles:
_cuda_n32_module.trsm1024_panel64_cuda(
output,
output,
tile_start + 64,
remaining_tiles,
)
pair_count = remaining_tiles * (remaining_tiles + 1) // 2
_update1024_tiles128[
(batch * pair_count,)
](
source,
output,
n * n,
tile_start=tile_start,
remaining_tiles=remaining_tiles,
dot_precision="tf32x3",
num_warps=4,
)
return output
if batch == 1 and n in (8192, 16384, 32768):
return _pure_deferred_large_blocks(data, 2048)
if batch == 1 and n == 4096:
return torch.linalg.cholesky_ex(data, check_errors=False).L
if n not in (32, 64, 128, 256, 512):
return torch.linalg.cholesky_ex(data, check_errors=False).L
output = torch.empty_like(data)
if n == 32:
_cuda_n32_module.cholesky32_cuda(data, output)
elif n == 64:
if batch % 4 == 0:
_cuda_n32_module.cholesky64_recursive_cuda(data, output)
else:
_cholesky64_fused[(batch,)](
data,
output,
n * n,
num_warps=1,
)
elif n == 128:
_cuda_n32_module.factor64_strided128_cuda(data, output, 0)
_cuda_n32_module.trsm128_panel64_cuda(data, output)
_syrk128_trailing64[(batch * 3,)](
data,
output,
n * n,
num_warps=2,
)
_cuda_n32_module.factor64_strided128_cuda(output, output, 64)
elif n == 256:
for tile_start, remaining_tiles in ((0, 3), (64, 2), (128, 1)):
source = data if tile_start == 0 else output
_cuda_n32_module.factor64_strided128_cuda(
source,
output,
tile_start,
)
_cuda_n32_module.trsm256_panel64_cuda(
source,
output,
tile_start,
remaining_tiles,
)
pair_count = remaining_tiles * (remaining_tiles + 1) // 2
_update256_tiles64[(batch * pair_count,)](
source,
output,
n * n,
tile_start=tile_start,
remaining_tiles=remaining_tiles,
num_warps=4,
)
_cuda_n32_module.factor64_strided128_cuda(output, output, 192)
elif batch == 60 or batch >= 128:
for tile_start, remaining_blocks in ((0, 3), (128, 2), (256, 1)):
source = data if tile_start == 0 else output
_cuda_n32_module.factor64_strided128_cuda(
source,
output,
tile_start,
)
first_half_panels = 1 + remaining_blocks * 2
_cuda_n32_module.trsm512_panel64_cuda(
source,
output,
tile_start,
first_half_panels,
)
_update512_second_panel64[(batch * first_half_panels,)](
source,
output,
n * n,
tile_start=tile_start,
panel_count=first_half_panels,
num_warps=4,
)
_cuda_n32_module.factor64_strided128_cuda(
output,
output,
tile_start + 64,
)
remaining_tiles = remaining_blocks * 2
_cuda_n32_module.trsm512_panel64_cuda(
output,
output,
tile_start + 64,
remaining_tiles,
)
pair_count = remaining_tiles * (remaining_tiles + 1) // 2
_update512_tiles128[(batch * pair_count,)](
source,
output,
n * n,
tile_start=tile_start,
remaining_tiles=remaining_tiles,
num_warps=4 if batch == 60 else 8,
)
_cuda_n32_module.factor64_strided128_cuda(
output,
output,
384,
)
_cuda_n32_module.trsm512_panel64_cuda(
output,
output,
384,
1,
)
_update512_second_panel64[(batch,)](
output,
output,
n * n,
tile_start=384,
panel_count=1,
num_warps=4,
)
_cuda_n32_module.factor64_strided128_cuda(
output,
output,
448,
)
else:
trsm_warps = 1 if batch >= 128 else 2
dot_precision = "tf32" if batch >= 128 else "tf32x3"
use_2d_grid = batch < 128
panel_steps = (
(0, 7),
(64, 6),
(128, 5),
(192, 4),
(256, 3),
(320, 2),
(384, 1),
)
for tile_start, remaining_tiles in panel_steps:
source = data if tile_start == 0 else output
_cuda_n32_module.factor64_strided128_cuda(
source,
output,
tile_start,
)
_cuda_n32_module.trsm512_panel64_cuda(
source,
output,
tile_start,
remaining_tiles,
)
pair_count = remaining_tiles * (remaining_tiles + 1) // 2
if batch >= 128:
_update512_tiles64[(batch * pair_count,)](
source,
output,
n * n,
tile_start=tile_start,
remaining_tiles=remaining_tiles,
dot_precision=dot_precision,
use_2d_grid=False,
num_warps=4,
)
else:
_update512_tiles32[(pair_count * 4, batch)](
source,
output,
n * n,
tile_start=tile_start,
remaining_tiles=remaining_tiles,
dot_precision=dot_precision,
use_2d_grid=True,
num_warps=2,
)
_cuda_n32_module.factor64_strided128_cuda(
output,
output,
448,
)
return output
scrolls · 2629 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