submission 829206
steve · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1218 lines, June 9 Researcher Reciprocity License v1.0.
submission_v22.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-829206?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:96695714fb5173d106343de4f4e845567b24c43d821f09a89cdaac15c7b95323
license declaredunknown
license concludedunknown
authorssteve
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
def _s22n128_smem_row (row :Int32 )->Int32 :Kernel source
submission_v22.py1218 lines
import torch
from task import input_t, output_t
import cutlass
import cutlass.cute as cute
import cutlass.utils as utils
from cutlass import Float32, Int32
from cutlass.cute.arch.nvvm_wrappers import shuffle_sync, warp_reduction_sum
from cutlass.cute.runtime import make_ptr
# Standalone n128 v1 fused implementation.
import torch
from task import input_t ,output_t
import cutlass
import cutlass .cute as cute
import cutlass .utils as utils
from cutlass import Float32 ,Int32
from cutlass .cute .arch .nvvm_wrappers import shuffle_sync ,warp_reduction_sum
from cutlass .cute .runtime import make_ptr
_s22n128_QR32_THREADS =32
_s22n128_QR128_THREADS =1024
_s22n128_QR32_LOAD_VEC =4
_s22n128_N128 =128
_s22n128_N128_ELEMS =_s22n128_N128 *_s22n128_N128
_s22n128_N128_SMEM_STRIDE =129
_s22n128_N128_SMEM_ELEMS =_s22n128_N128 *_s22n128_N128_SMEM_STRIDE
_s22n128_compile_cache ={}
@cute .jit
def _s22n128_smem_row (row :Int32 )->Int32 :
return row ^(row >>Int32 (2 ))
@cute .kernel
def _s22n128_qr128_fused_kernel (
h :cute .Pointer ,
tau :cute .Pointer ,
batch_count :Int32 ,
):
tidx ,_ ,_ =cute .arch .thread_idx ()
bidx ,_ ,_ =cute .arch .block_idx ()
bdim ,_ ,_ =cute .arch .block_dim ()
smem =utils .SmemAllocator ()
scratch =smem .allocate_tensor (
Float32 ,
cute .make_layout ((_s22n128_QR128_THREADS ,),stride =(1 ,)),
byte_alignment =16 ,
)
v_smem =smem .allocate_tensor (
Float32 ,
cute .make_layout ((_s22n128_N128 ,),stride =(1 ,)),
byte_alignment =16 ,
)
tile =smem .allocate_tensor (
Float32 ,
cute .make_layout ((_s22n128_N128 ,_s22n128_N128_SMEM_STRIDE ),stride =(_s22n128_N128_SMEM_STRIDE ,1 )),
byte_alignment =16 ,
)
matrix_base =bidx *_s22n128_N128_ELEMS
tau_base =bidx *_s22n128_N128
g_h4 =cute .make_tensor (
h ,
cute .make_layout (
(batch_count ,_s22n128_N128 ,_s22n128_N128 //_s22n128_QR32_LOAD_VEC ,_s22n128_QR32_LOAD_VEC ),
stride =(_s22n128_N128_ELEMS ,_s22n128_N128 ,_s22n128_QR32_LOAD_VEC ,1 ),
),
)
g_h =cute .make_tensor (h ,cute .make_layout ((batch_count *_s22n128_N128_ELEMS ,),stride =(1 ,)))
g_tau =cute .make_tensor (tau ,cute .make_layout ((batch_count *_s22n128_N128 ,),stride =(1 ,)))
copy_atom_128 =cute .make_copy_atom (
cute .nvgpu .CopyUniversalOp (),Float32 ,num_bits_per_copy =128
)
load_vec =cute .make_rmem_tensor ((_s22n128_QR32_LOAD_VEC ,),Float32 )
for vec_base in cutlass .range (tidx *_s22n128_QR32_LOAD_VEC ,_s22n128_N128_ELEMS ,bdim *_s22n128_QR32_LOAD_VEC ):
row =vec_base //_s22n128_N128
col =vec_base -row *_s22n128_N128
group =col //_s22n128_QR32_LOAD_VEC
cute .copy (copy_atom_128 ,g_h4 [bidx ,row ,group ,None ],load_vec [None ])
for col_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
tile [col +col_off ,_s22n128_smem_row (row )]=load_vec [col_off ]
cute .arch .barrier ()
warp =tidx //_s22n128_QR32_THREADS
lane =tidx -warp *_s22n128_QR32_THREADS
col_base =warp *_s22n128_QR32_LOAD_VEC
row_base =lane *_s22n128_QR32_LOAD_VEC
vals =cute .make_rmem_tensor (
(_s22n128_QR32_LOAD_VEC ,_s22n128_QR32_LOAD_VEC ),
Float32 ,
)
for row_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
row =row_base +row_off
for col_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
vals [row_off ,col_off ]=tile [col_base +col_off ,_s22n128_smem_row (row )]
cute .arch .barrier ()
for col in cutlass .range (_s22n128_N128 ):
col_warp =col //_s22n128_QR32_LOAD_VEC
col_off =col -col_warp *_s22n128_QR32_LOAD_VEC
diag_lane =col //_s22n128_QR32_LOAD_VEC
diag_row_off =col -diag_lane *_s22n128_QR32_LOAD_VEC
owns_col =warp ==col_warp
tail_norm_sq_part =Float32 (0.0 )
if owns_col :
for row_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
row =row_base +row_off
if row >col :
x =vals [row_off ,col_off ]
tail_norm_sq_part =tail_norm_sq_part +x *x
tail_norm_sq =warp_reduction_sum (tail_norm_sq_part )
alpha_part =Float32 (0.0 )
if owns_col :
if lane ==diag_lane :
alpha_part =vals [diag_row_off ,col_off ]
alpha =shuffle_sync (alpha_part ,mask =-1 ,offset =diag_lane ,mask_and_clamp =31 )
beta =alpha
tau_col =Float32 (0.0 )
tail_scale =Float32 (0.0 )
if owns_col :
if tail_norm_sq !=Float32 (0.0 ):
x_norm =cute .math .sqrt (alpha *alpha +tail_norm_sq )
beta =-x_norm
if alpha <Float32 (0.0 ):
beta =x_norm
tau_col =(beta -alpha )/beta
tail_scale =Float32 (1.0 )/(alpha -beta )
if owns_col :
for row_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
row =row_base +row_off
if row >=col :
v_i =Float32 (1.0 )
if row !=col :
vals [row_off ,col_off ]=vals [row_off ,col_off ]*tail_scale
v_i =vals [row_off ,col_off ]
else :
vals [row_off ,col_off ]=beta
v_smem [row ]=v_i
if lane ==0 :
scratch [0 ]=tau_col
g_tau [tau_base +col ]=tau_col
cute .arch .barrier ()
tau_col =scratch [0 ]
if tau_col !=Float32 (0.0 ):
for target_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
target_col =col_base +target_off
dot_part =Float32 (0.0 )
if target_col >col :
for row_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
row =row_base +row_off
if row >=col :
a_val =vals [row_off ,target_off ]
v_val =v_smem [row ]
dot_part =dot_part +v_val *a_val
dot =warp_reduction_sum (dot_part )*tau_col
if target_col >col :
for row_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
row =row_base +row_off
if row >=col :
v_val =v_smem [row ]
vals [row_off ,target_off ]=vals [row_off ,target_off ]-v_val *dot
cute .arch .barrier ()
for row_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
row =row_base +row_off
for col_off in cutlass .range (_s22n128_QR32_LOAD_VEC ,unroll_full =True ):
tile [col_base +col_off ,_s22n128_smem_row (row )]=vals [row_off ,col_off ]
cute .arch .barrier ()
for idx in cutlass .range (tidx ,_s22n128_N128_ELEMS ,bdim ):
row =idx //_s22n128_N128
col =idx -row *_s22n128_N128
g_h [matrix_base +row *_s22n128_N128 +col ]=tile [col ,_s22n128_smem_row (row )]
@cute .jit
def _s22n128_launch_qr128_fused (
h :cute .Pointer ,
tau :cute .Pointer ,
batch_count :Int32 ,
):
_s22n128_qr128_fused_kernel (h ,tau ,batch_count ).launch (
grid =[batch_count ,1 ,1 ],
block =[_s22n128_QR128_THREADS ,1 ,1 ],
smem =(_s22n128_QR128_THREADS +_s22n128_N128 +_s22n128_N128_SMEM_ELEMS )*4 ,
)
def _s22n128_compiled_qr128_fused ():
if "qr128_fused_fp32"not in _s22n128_compile_cache :
ptr =make_ptr (cutlass .Float32 ,0 ,cute .AddressSpace .gmem ,assumed_align =16 )
_s22n128_compile_cache ["qr128_fused_fp32"]=cute .compile (
_s22n128_launch_qr128_fused ,
ptr ,
ptr ,
Int32 (1 ),
)
return _s22n128_compile_cache ["qr128_fused_fp32"]
def _s22n128_qr128_fused_rowmajor (h :torch .Tensor ,tau :torch .Tensor )->None :
batch_count =h .shape [0 ]
kernel =_s22n128_compiled_qr128_fused ()
kernel (
make_ptr (cutlass .Float32 ,h .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
make_ptr (cutlass .Float32 ,tau .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
Int32 (batch_count ),
)
def _s22n128_qr128_fused (data :torch .Tensor )->output_t :
h =data .contiguous ().clone ()
batch ,n ,_ =h .shape
if n !=_s22n128_N128 :
raise ValueError ("n_128_v1 expects input matrices with n=128")
tau =torch .zeros ((batch ,n ),device =h .device ,dtype =torch .float32 )
_s22n128_qr128_fused_rowmajor (h ,tau )
return h ,tau
def _s22n128_custom_kernel (data :input_t )->output_t :
return _s22n128_qr128_fused (data )
# Standalone n512 v16 implementation.
import torch
from task import input_t ,output_t
import cutlass
import cutlass .cute as cute
import cutlass .utils as utils
from cutlass import Float32 ,Int32
from cutlass .cute .arch .nvvm_wrappers import shuffle_sync ,warp_reduction_sum
from cutlass .cute .runtime import make_ptr
"""
n512 v16 is based on n512 v15, but fuses the trailing subtract into the
final batched GEMM with torch.baddbmm.
Requested panel mapping:
- Use 1024 threads, i.e. 32 warps.
- Two warps own the row work for the current panel column.
- Each thread in the warp pair handles four consecutive rows.
- A warp pair covers 2 * 32 * 4 = 256 row values, so two row chunks cover
a 512-row column.
- Stage only the active panel rows k_start:n in shared memory with the n128
v1 row swizzle: tile[panel_col, row ^ (row >> 2)].
Panel columns still run in order because Householder QR has column
dependencies inside the panel. After all 32-wide panels before the final
128x128 diagonal block, the last n128 block calls the n128 v1 fused kernel.
"""
_s22n512_PANEL_THREADS =1024
_s22n512_BUILD_VT_THREADS =256
_s22n512_MAX_PANEL =16
_s22n512_FINAL_PANEL =128
_s22n512_PANEL_SMEM_STRIDE =513
_s22n512_PANEL_SMEM_ELEMS =_s22n512_MAX_PANEL *_s22n512_PANEL_SMEM_STRIDE
_s22n512_QR_WARP_THREADS =32
_s22n512_WARPS_PER_COLUMN =2
_s22n512_ROWS_PER_THREAD =4
_s22n512_ROW_CHUNK =_s22n512_QR_WARP_THREADS *_s22n512_WARPS_PER_COLUMN *_s22n512_ROWS_PER_THREAD
_s22n512_N512 =512
_s22n512_compile_cache ={}
@cute .jit
def _s22n512_smem_row (row :Int32 )->Int32 :
return row ^(row >>2 )
@cute .jit
def _s22n512_block_sum_1024 (value :Float32 ,scratch :cute .Tensor ,tidx :Int32 )->Float32 :
scratch [tidx ]=value
cute .arch .barrier ()
if tidx <512 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +512 ]
cute .arch .barrier ()
if tidx <256 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +256 ]
cute .arch .barrier ()
if tidx <128 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +128 ]
cute .arch .barrier ()
if tidx <64 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +64 ]
cute .arch .barrier ()
if tidx <32 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +32 ]
cute .arch .barrier ()
if tidx <16 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +16 ]
cute .arch .barrier ()
if tidx <8 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +8 ]
cute .arch .barrier ()
if tidx <4 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +4 ]
cute .arch .barrier ()
if tidx <2 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +2 ]
cute .arch .barrier ()
if tidx <1 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +1 ]
cute .arch .barrier ()
return scratch [0 ]
@cute .jit
def _s22n512_block_sum_256 (value :Float32 ,scratch :cute .Tensor ,tidx :Int32 )->Float32 :
scratch [tidx ]=value
cute .arch .barrier ()
if tidx <128 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +128 ]
cute .arch .barrier ()
if tidx <64 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +64 ]
cute .arch .barrier ()
if tidx <32 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +32 ]
cute .arch .barrier ()
if tidx <16 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +16 ]
cute .arch .barrier ()
if tidx <8 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +8 ]
cute .arch .barrier ()
if tidx <4 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +4 ]
cute .arch .barrier ()
if tidx <2 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +2 ]
cute .arch .barrier ()
if tidx <1 :
scratch [tidx ]=scratch [tidx ]+scratch [tidx +1 ]
cute .arch .barrier ()
return scratch [0 ]
@cute .jit
def _s22n512_warp_pair_sum (
value :Float32 ,
scratch :cute .Tensor ,
warp :Int32 ,
lane :Int32 ,
pair_warp0 :Int32 ,
active_pair :cutlass .Boolean ,
)->Float32 :
warp_sum =warp_reduction_sum (value )
if active_pair :
if lane ==0 :
scratch [warp -pair_warp0 ]=warp_sum
cute .arch .barrier ()
result =scratch [0 ]+scratch [1 ]
cute .arch .barrier ()
return result
@cute .kernel
def _s22n512_factor_panel_kernel (
h :cute .Pointer ,
tau :cute .Pointer ,
batch_count :Int32 ,
n :Int32 ,
k_start :Int32 ,
panel_width :Int32 ,
):
tidx ,_ ,_ =cute .arch .thread_idx ()
bidx ,_ ,_ =cute .arch .block_idx ()
bdim ,_ ,_ =cute .arch .block_dim ()
smem =utils .SmemAllocator ()
scratch =smem .allocate_tensor (
Float32 ,
cute .make_layout ((_s22n512_PANEL_THREADS ,),stride =(1 ,)),
byte_alignment =16 ,
)
panel =smem .allocate_tensor (
Float32 ,
cute .make_layout ((_s22n512_MAX_PANEL ,_s22n512_PANEL_SMEM_STRIDE ),stride =(_s22n512_PANEL_SMEM_STRIDE ,1 )),
byte_alignment =16 ,
)
matrix_elems =n *n
matrix_base =bidx *matrix_elems
tau_base =bidx *n
panel_end =k_start +panel_width
active_rows =n -k_start
g_h =cute .make_tensor (h ,cute .make_layout ((batch_count *matrix_elems ,),stride =(1 ,)))
g_tau =cute .make_tensor (tau ,cute .make_layout ((batch_count *n ,),stride =(1 ,)))
for idx in cutlass .range (tidx ,active_rows *panel_width ,bdim ):
local_col =idx //active_rows
row =k_start +idx -local_col *active_rows
panel [local_col ,_s22n512_smem_row (row )]=g_h [matrix_base +row *n +k_start +local_col ]
cute .arch .barrier ()
for j in cutlass .range (panel_width ):
col =k_start +j
col_tail =n -col -1
tail_chunks =(col_tail +_s22n512_ROW_CHUNK -1 )//_s22n512_ROW_CHUNK
full_chunks =(n -col +_s22n512_ROW_CHUNK -1 )//_s22n512_ROW_CHUNK
warp =tidx //_s22n512_QR_WARP_THREADS
lane =tidx -warp *_s22n512_QR_WARP_THREADS
pair_slot =j -(j //16 )*16
pair_warp0 =pair_slot *_s22n512_WARPS_PER_COLUMN
active_pair =warp ==pair_warp0
if warp ==pair_warp0 +1 :
active_pair =True
tail_norm_sq_part =Float32 (0.0 )
if active_pair :
pair_lane =(warp -pair_warp0 )*_s22n512_QR_WARP_THREADS +lane
for chunk in cutlass .range (tail_chunks ):
row_base =col +1 +chunk *_s22n512_ROW_CHUNK +pair_lane *_s22n512_ROWS_PER_THREAD
for row_off in cutlass .range (_s22n512_ROWS_PER_THREAD ,unroll_full =True ):
row =row_base +row_off
if row <n :
x =panel [j ,_s22n512_smem_row (row )]
tail_norm_sq_part =tail_norm_sq_part +x *x
tail_norm_sq =_s22n512_warp_pair_sum (tail_norm_sq_part ,scratch ,warp ,lane ,pair_warp0 ,active_pair )
alpha =panel [j ,_s22n512_smem_row (col )]
beta =alpha
tau_col =Float32 (0.0 )
tail_scale =Float32 (0.0 )
if tail_norm_sq !=Float32 (0.0 ):
x_norm =cute .math .sqrt (alpha *alpha +tail_norm_sq )
beta =-x_norm
if alpha <Float32 (0.0 ):
beta =x_norm
tau_col =(beta -alpha )/beta
tail_scale =Float32 (1.0 )/(alpha -beta )
if active_pair :
pair_lane =(warp -pair_warp0 )*_s22n512_QR_WARP_THREADS +lane
for chunk in cutlass .range (tail_chunks ):
row_base =col +1 +chunk *_s22n512_ROW_CHUNK +pair_lane *_s22n512_ROWS_PER_THREAD
for row_off in cutlass .range (_s22n512_ROWS_PER_THREAD ,unroll_full =True ):
row =row_base +row_off
if row <n :
sw_row =_s22n512_smem_row (row )
panel [j ,sw_row ]=panel [j ,sw_row ]*tail_scale
if tidx ==0 :
panel [j ,_s22n512_smem_row (col )]=beta
g_tau [tau_base +col ]=tau_col
cute .arch .barrier ()
if tau_col !=Float32 (0.0 ):
for target_col in cutlass .range (col +1 ,panel_end ):
target_local_col =target_col -k_start
dot_part =Float32 (0.0 )
if active_pair :
pair_lane =(warp -pair_warp0 )*_s22n512_QR_WARP_THREADS +lane
for chunk in cutlass .range (full_chunks ):
row_base =col +chunk *_s22n512_ROW_CHUNK +pair_lane *_s22n512_ROWS_PER_THREAD
for row_off in cutlass .range (_s22n512_ROWS_PER_THREAD ,unroll_full =True ):
row =row_base +row_off
if row <n :
sw_row =_s22n512_smem_row (row )
a_val =panel [target_local_col ,sw_row ]
v_val =Float32 (1.0 )
if row !=col :
v_val =panel [j ,sw_row ]
dot_part =dot_part +v_val *a_val
dot =_s22n512_warp_pair_sum (dot_part ,scratch ,warp ,lane ,pair_warp0 ,active_pair )*tau_col
if active_pair :
pair_lane =(warp -pair_warp0 )*_s22n512_QR_WARP_THREADS +lane
for chunk in cutlass .range (full_chunks ):
row_base =col +chunk *_s22n512_ROW_CHUNK +pair_lane *_s22n512_ROWS_PER_THREAD
for row_off in cutlass .range (_s22n512_ROWS_PER_THREAD ,unroll_full =True ):
row =row_base +row_off
if row <n :
sw_row =_s22n512_smem_row (row )
a_val =panel [target_local_col ,sw_row ]
v_val =Float32 (1.0 )
if row !=col :
v_val =panel [j ,sw_row ]
panel [target_local_col ,sw_row ]=a_val -v_val *dot
cute .arch .barrier ()
for idx in cutlass .range (tidx ,active_rows *panel_width ,bdim ):
local_col =idx //active_rows
row =k_start +idx -local_col *active_rows
g_h [matrix_base +row *n +k_start +local_col ]=panel [local_col ,_s22n512_smem_row (row )]
@cute .kernel
def _s22n512_build_vt_kernel (
h :cute .Pointer ,
tau :cute .Pointer ,
v :cute .Pointer ,
t :cute .Pointer ,
batch_count :Int32 ,
n :Int32 ,
k_start :Int32 ,
panel_width :Int32 ,
v_rows :Int32 ,
):
tidx ,_ ,_ =cute .arch .thread_idx ()
bidx ,_ ,_ =cute .arch .block_idx ()
bdim ,_ ,_ =cute .arch .block_dim ()
smem =utils .SmemAllocator ()
scratch =smem .allocate_tensor (
Float32 ,
cute .make_layout ((_s22n512_BUILD_VT_THREADS +_s22n512_MAX_PANEL ,),stride =(1 ,)),
byte_alignment =16 ,
)
matrix_elems =n *n
matrix_base =bidx *matrix_elems
tau_base =bidx *n
v_base =bidx *v_rows *panel_width
t_base =bidx *panel_width *panel_width
g_h =cute .make_tensor (h ,cute .make_layout ((batch_count *matrix_elems ,),stride =(1 ,)))
g_tau =cute .make_tensor (tau ,cute .make_layout ((batch_count *n ,),stride =(1 ,)))
g_v =cute .make_tensor (v ,cute .make_layout ((batch_count *v_rows *panel_width ,),stride =(1 ,)))
g_t =cute .make_tensor (t ,cute .make_layout ((batch_count *panel_width *panel_width ,),stride =(1 ,)))
for idx in cutlass .range (tidx ,v_rows *panel_width ,bdim ):
row =idx //panel_width
col =idx -row *panel_width
out =Float32 (0.0 )
if row >col :
out =g_h [matrix_base +(k_start +row )*n +(k_start +col )]
if row ==col :
out =Float32 (1.0 )
g_v [v_base +row *panel_width +col ]=out
for idx in cutlass .range (tidx ,panel_width *panel_width ,bdim ):
g_t [t_base +idx ]=Float32 (0.0 )
cute .arch .barrier ()
for j in cutlass .range (panel_width ):
tau_j =g_tau [tau_base +k_start +j ]
if tidx ==0 :
g_t [t_base +j *panel_width +j ]=tau_j
cute .arch .barrier ()
if j !=0 :
for i in cutlass .range (j ):
dot_part =Float32 (0.0 )
for row in cutlass .range (tidx ,v_rows ,bdim ):
vi =g_v [v_base +row *panel_width +i ]
vj =g_v [v_base +row *panel_width +j ]
dot_part =dot_part +vi *vj
scratch [_s22n512_BUILD_VT_THREADS +i ]=_s22n512_block_sum_256 (dot_part ,scratch ,tidx )
cute .arch .barrier ()
for l in cutlass .range (tidx ,j ,bdim ):
accum =Float32 (0.0 )
for i in cutlass .range (j ):
accum =accum +scratch [_s22n512_BUILD_VT_THREADS +i ]*g_t [t_base +i *panel_width +l ]
g_t [t_base +j *panel_width +l ]=-tau_j *accum
cute .arch .barrier ()
@cute .jit
def _s22n512_launch_factor_panel (
h :cute .Pointer ,
tau :cute .Pointer ,
batch_count :Int32 ,
n :Int32 ,
k_start :Int32 ,
panel_width :Int32 ,
):
_s22n512_factor_panel_kernel (h ,tau ,batch_count ,n ,k_start ,panel_width ).launch (
grid =[batch_count ,1 ,1 ],
block =[_s22n512_PANEL_THREADS ,1 ,1 ],
smem =(_s22n512_PANEL_THREADS +_s22n512_PANEL_SMEM_ELEMS )*4 ,
)
@cute .jit
def _s22n512_launch_build_vt (
h :cute .Pointer ,
tau :cute .Pointer ,
v :cute .Pointer ,
t :cute .Pointer ,
batch_count :Int32 ,
n :Int32 ,
k_start :Int32 ,
panel_width :Int32 ,
v_rows :Int32 ,
):
_s22n512_build_vt_kernel (h ,tau ,v ,t ,batch_count ,n ,k_start ,panel_width ,v_rows ).launch (
grid =[batch_count ,1 ,1 ],
block =[_s22n512_BUILD_VT_THREADS ,1 ,1 ],
smem =(_s22n512_BUILD_VT_THREADS +_s22n512_MAX_PANEL )*4 ,
)
def _s22n512_compiled_panel ():
if "panel_fp32"not in _s22n512_compile_cache :
ptr =make_ptr (cutlass .Float32 ,0 ,cute .AddressSpace .gmem ,assumed_align =16 )
_s22n512_compile_cache ["panel_fp32"]=cute .compile (
_s22n512_launch_factor_panel ,
ptr ,
ptr ,
Int32 (1 ),
Int32 (1 ),
Int32 (1 ),
Int32 (1 ),
)
return _s22n512_compile_cache ["panel_fp32"]
def _s22n512_compiled_build_vt ():
if "build_vt_fp32"not in _s22n512_compile_cache :
ptr =make_ptr (cutlass .Float32 ,0 ,cute .AddressSpace .gmem ,assumed_align =16 )
_s22n512_compile_cache ["build_vt_fp32"]=cute .compile (
_s22n512_launch_build_vt ,
ptr ,
ptr ,
ptr ,
ptr ,
Int32 (1 ),
Int32 (1 ),
Int32 (1 ),
Int32 (1 ),
Int32 (1 ),
)
return _s22n512_compile_cache ["build_vt_fp32"]
def _s22n512_factor_panel (h :torch .Tensor ,tau :torch .Tensor ,k :int ,panel_width :int )->None :
batch_count ,n ,_ =h .shape
kernel =_s22n512_compiled_panel ()
kernel (
make_ptr (cutlass .Float32 ,h .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
make_ptr (cutlass .Float32 ,tau .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
Int32 (batch_count ),
Int32 (n ),
Int32 (k ),
Int32 (panel_width ),
)
def _s22n512_build_vt (h :torch .Tensor ,tau :torch .Tensor ,k :int ,panel_width :int )->tuple [torch .Tensor ,torch .Tensor ]:
batch_count ,n ,_ =h .shape
v_rows =n -k
v =torch .empty ((batch_count ,v_rows ,panel_width ),device =h .device ,dtype =torch .float32 )
t =torch .empty ((batch_count ,panel_width ,panel_width ),device =h .device ,dtype =torch .float32 )
kernel =_s22n512_compiled_build_vt ()
kernel (
make_ptr (cutlass .Float32 ,h .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
make_ptr (cutlass .Float32 ,tau .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
make_ptr (cutlass .Float32 ,v .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
make_ptr (cutlass .Float32 ,t .data_ptr (),cute .AddressSpace .gmem ,assumed_align =16 ),
Int32 (batch_count ),
Int32 (n ),
Int32 (k ),
Int32 (panel_width ),
Int32 (v_rows ),
)
return v ,t
def _s22n512_blocked_householder_qr (data :torch .Tensor )->output_t :
h =data .contiguous ().clone ()
batch ,n ,_ =h .shape
if n !=_s22n512_N512 :
raise ValueError ("n_512_v16 expects input matrices with n=512")
tau =torch .zeros ((batch ,n ),device =h .device ,dtype =torch .float32 )
for k in range (0 ,n -_s22n512_FINAL_PANEL ,_s22n512_MAX_PANEL ):
panel_width =min (_s22n512_MAX_PANEL ,n -_s22n512_FINAL_PANEL -k )
_s22n512_factor_panel (h ,tau ,k ,panel_width )
trailing_start =k +panel_width
if trailing_start >=n :
continue
v ,t =_s22n512_build_vt (h ,tau ,k ,panel_width )
trailing =h [:,k :,trailing_start :]
work =torch .bmm (v .transpose (1 ,2 ),trailing )
work =torch .bmm (t ,work )
torch .baddbmm (trailing ,v ,work ,beta =1.0 ,alpha =-1.0 ,out =trailing )
final_start =n -_s22n512_FINAL_PANEL
h128 =h [:,final_start :,final_start :].contiguous ()
h128 ,tau128 =_s22n128_custom_kernel (h128 )
h [:,final_start :,final_start :]=h128
tau [:,final_start :]=tau128
return h ,tau
def _s22n512_custom_kernel (data :input_t )->output_t :
return _s22n512_blocked_householder_qr (data )
_THREADS = 256
_QR32_THREADS = 32
_QR32_LOAD_VEC = 4
_MAX_PANEL = 32
_N32 = 32
_N128 = 128
_N512 = 512
_N32_ELEMS = _N32 * _N32
_compile_cache = {}
def _block_size_for_n(n: int) -> int:
if n <= 32:
return 8
if n <= 176:
return 16
return 32
@cute.jit
def _block_sum(value: Float32, scratch: cute.Tensor, tidx: Int32) -> Float32:
scratch[tidx] = value
cute.arch.barrier()
if tidx < 128:
scratch[tidx] = scratch[tidx] + scratch[tidx + 128]
cute.arch.barrier()
if tidx < 64:
scratch[tidx] = scratch[tidx] + scratch[tidx + 64]
cute.arch.barrier()
if tidx < 32:
scratch[tidx] = scratch[tidx] + scratch[tidx + 32]
cute.arch.barrier()
if tidx < 16:
scratch[tidx] = scratch[tidx] + scratch[tidx + 16]
cute.arch.barrier()
if tidx < 8:
scratch[tidx] = scratch[tidx] + scratch[tidx + 8]
cute.arch.barrier()
if tidx < 4:
scratch[tidx] = scratch[tidx] + scratch[tidx + 4]
cute.arch.barrier()
if tidx < 2:
scratch[tidx] = scratch[tidx] + scratch[tidx + 2]
cute.arch.barrier()
if tidx < 1:
scratch[tidx] = scratch[tidx] + scratch[tidx + 1]
cute.arch.barrier()
return scratch[0]
@cute.kernel
def _qr32_fused_kernel(
h: cute.Pointer,
tau: cute.Pointer,
batch_count: Int32,
):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
bdim, _, _ = cute.arch.block_dim()
smem = utils.SmemAllocator()
tile = smem.allocate_tensor(
Float32,
cute.make_layout((_N32_ELEMS,), stride=(1,)),
byte_alignment=16,
)
matrix_base = bidx * _N32_ELEMS
tau_base = bidx * _N32
g_h = cute.make_tensor(h, cute.make_layout((batch_count * _N32_ELEMS,), stride=(1,)))
g_tau = cute.make_tensor(tau, cute.make_layout((batch_count * _N32,), stride=(1,)))
for vec_base in cutlass.range(tidx * _QR32_LOAD_VEC, _N32_ELEMS, bdim * _QR32_LOAD_VEC):
tile[vec_base + 0] = g_h[matrix_base + vec_base + 0]
tile[vec_base + 1] = g_h[matrix_base + vec_base + 1]
tile[vec_base + 2] = g_h[matrix_base + vec_base + 2]
tile[vec_base + 3] = g_h[matrix_base + vec_base + 3]
cute.arch.sync_warp()
row = cute.make_rmem_tensor((_N32,), Float32)
row_base = tidx * _N32
for col in cutlass.range(_N32, unroll_full=True):
row[col] = tile[row_base + col]
for col in cutlass.range(_N32, unroll_full=True):
a_col = row[col]
tail_norm_sq_part = Float32(0.0)
if tidx > col:
tail_norm_sq_part = a_col * a_col
tail_norm_sq = warp_reduction_sum(tail_norm_sq_part)
alpha = shuffle_sync(a_col, mask=-1, offset=col, mask_and_clamp=31)
beta = alpha
tau_col = Float32(0.0)
tail_scale = Float32(0.0)
if tail_norm_sq != Float32(0.0):
x_norm = cute.math.sqrt(alpha * alpha + tail_norm_sq)
beta = -x_norm
if alpha < Float32(0.0):
beta = x_norm
tau_col = (beta - alpha) / beta
tail_scale = Float32(1.0) / (alpha - beta)
v_i = Float32(0.0)
if tidx == col:
row[col] = beta
g_tau[tau_base + col] = tau_col
v_i = Float32(1.0)
if tidx > col:
row[col] = a_col * tail_scale
v_i = row[col]
if tau_col != Float32(0.0):
for target_col in cutlass.range(col + 1, _N32, unroll_full=True):
dot_part = v_i * row[target_col]
if tidx < col:
dot_part = Float32(0.0)
dot = warp_reduction_sum(dot_part) * tau_col
if tidx >= col:
row[target_col] = row[target_col] - v_i * dot
for col in cutlass.range(_N32, unroll_full=True):
tile[row_base + col] = row[col]
cute.arch.sync_warp()
for vec_base in cutlass.range(tidx * _QR32_LOAD_VEC, _N32_ELEMS, bdim * _QR32_LOAD_VEC):
g_h[matrix_base + vec_base + 0] = tile[vec_base + 0]
g_h[matrix_base + vec_base + 1] = tile[vec_base + 1]
g_h[matrix_base + vec_base + 2] = tile[vec_base + 2]
g_h[matrix_base + vec_base + 3] = tile[vec_base + 3]
@cute.kernel
def _factor_panel_kernel(
h: cute.Pointer,
tau: cute.Pointer,
batch_count: Int32,
n: Int32,
k_start: Int32,
panel_width: Int32,
):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
bdim, _, _ = cute.arch.block_dim()
smem = utils.SmemAllocator()
scratch = smem.allocate_tensor(
Float32,
cute.make_layout((_THREADS,), stride=(1,)),
byte_alignment=16,
)
matrix_elems = n * n
matrix_base = bidx * matrix_elems
tau_base = bidx * n
panel_end = k_start + panel_width
g_h = cute.make_tensor(h, cute.make_layout((batch_count * matrix_elems,), stride=(1,)))
g_tau = cute.make_tensor(tau, cute.make_layout((batch_count * n,), stride=(1,)))
for j in cutlass.range(panel_width):
col = k_start + j
col_tail = n - col - 1
tail_norm_sq_part = Float32(0.0)
for r_tail in cutlass.range(tidx, col_tail, bdim):
row = col + 1 + r_tail
x = g_h[matrix_base + row * n + col]
tail_norm_sq_part = tail_norm_sq_part + x * x
tail_norm_sq = _block_sum(tail_norm_sq_part, scratch, tidx)
alpha = g_h[matrix_base + col * n + col]
beta = alpha
tau_col = Float32(0.0)
tail_scale = Float32(0.0)
if tail_norm_sq != Float32(0.0):
x_norm = cute.math.sqrt(alpha * alpha + tail_norm_sq)
beta = -x_norm
if alpha < Float32(0.0):
beta = x_norm
tau_col = (beta - alpha) / beta
tail_scale = Float32(1.0) / (alpha - beta)
for r_tail in cutlass.range(tidx, col_tail, bdim):
row = col + 1 + r_tail
idx = matrix_base + row * n + col
g_h[idx] = g_h[idx] * tail_scale
if tidx == 0:
g_h[matrix_base + col * n + col] = beta
g_tau[tau_base + col] = tau_col
cute.arch.barrier()
if tau_col != Float32(0.0):
for target_col in cutlass.range(col + 1, panel_end):
dot_part = Float32(0.0)
for r_local in cutlass.range(tidx, n - col, bdim):
row = col + r_local
a_val = g_h[matrix_base + row * n + target_col]
v_val = Float32(1.0)
if r_local != 0:
v_val = g_h[matrix_base + row * n + col]
dot_part = dot_part + v_val * a_val
dot = _block_sum(dot_part, scratch, tidx) * tau_col
for r_local in cutlass.range(tidx, n - col, bdim):
row = col + r_local
idx = matrix_base + row * n + target_col
a_val = g_h[idx]
v_val = Float32(1.0)
if r_local != 0:
v_val = g_h[matrix_base + row * n + col]
g_h[idx] = a_val - v_val * dot
cute.arch.barrier()
@cute.kernel
def _build_vt_kernel(
h: cute.Pointer,
tau: cute.Pointer,
v: cute.Pointer,
t: cute.Pointer,
batch_count: Int32,
n: Int32,
k_start: Int32,
panel_width: Int32,
v_rows: Int32,
):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
bdim, _, _ = cute.arch.block_dim()
smem = utils.SmemAllocator()
scratch = smem.allocate_tensor(
Float32,
cute.make_layout((_THREADS + _MAX_PANEL,), stride=(1,)),
byte_alignment=16,
)
matrix_elems = n * n
matrix_base = bidx * matrix_elems
tau_base = bidx * n
v_base = bidx * v_rows * panel_width
t_base = bidx * panel_width * panel_width
g_h = cute.make_tensor(h, cute.make_layout((batch_count * matrix_elems,), stride=(1,)))
g_tau = cute.make_tensor(tau, cute.make_layout((batch_count * n,), stride=(1,)))
g_v = cute.make_tensor(v, cute.make_layout((batch_count * v_rows * panel_width,), stride=(1,)))
g_t = cute.make_tensor(t, cute.make_layout((batch_count * panel_width * panel_width,), stride=(1,)))
for idx in cutlass.range(tidx, v_rows * panel_width, bdim):
row = idx // panel_width
col = idx - row * panel_width
out = Float32(0.0)
if row > col:
out = g_h[matrix_base + (k_start + row) * n + (k_start + col)]
if row == col:
out = Float32(1.0)
g_v[v_base + row * panel_width + col] = out
for idx in cutlass.range(tidx, panel_width * panel_width, bdim):
g_t[t_base + idx] = Float32(0.0)
cute.arch.barrier()
for j in cutlass.range(panel_width):
tau_j = g_tau[tau_base + k_start + j]
if tidx == 0:
g_t[t_base + j * panel_width + j] = tau_j
cute.arch.barrier()
if j != 0:
for i in cutlass.range(j):
dot_part = Float32(0.0)
for row in cutlass.range(tidx, v_rows, bdim):
vi = g_v[v_base + row * panel_width + i]
vj = g_v[v_base + row * panel_width + j]
dot_part = dot_part + vi * vj
scratch[_THREADS + i] = _block_sum(dot_part, scratch, tidx)
cute.arch.barrier()
for l in cutlass.range(tidx, j, bdim):
accum = Float32(0.0)
for i in cutlass.range(j):
accum = accum + scratch[_THREADS + i] * g_t[t_base + i * panel_width + l]
g_t[t_base + j * panel_width + l] = -tau_j * accum
cute.arch.barrier()
@cute.jit
def _launch_qr32_fused(
h: cute.Pointer,
tau: cute.Pointer,
batch_count: Int32,
):
_qr32_fused_kernel(h, tau, batch_count).launch(
grid=[batch_count, 1, 1],
block=[_QR32_THREADS, 1, 1],
smem=_N32_ELEMS * 4,
)
@cute.jit
def _launch_factor_panel(
h: cute.Pointer,
tau: cute.Pointer,
batch_count: Int32,
n: Int32,
k_start: Int32,
panel_width: Int32,
):
_factor_panel_kernel(h, tau, batch_count, n, k_start, panel_width).launch(
grid=[batch_count, 1, 1],
block=[_THREADS, 1, 1],
smem=_THREADS * 4,
)
@cute.jit
def _launch_build_vt(
h: cute.Pointer,
tau: cute.Pointer,
v: cute.Pointer,
t: cute.Pointer,
batch_count: Int32,
n: Int32,
k_start: Int32,
panel_width: Int32,
v_rows: Int32,
):
_build_vt_kernel(h, tau, v, t, batch_count, n, k_start, panel_width, v_rows).launch(
grid=[batch_count, 1, 1],
block=[_THREADS, 1, 1],
smem=(_THREADS + _MAX_PANEL) * 4,
)
def _compiled_qr32_fused():
if "qr32_fused_fp32" not in _compile_cache:
ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_compile_cache["qr32_fused_fp32"] = cute.compile(
_launch_qr32_fused,
ptr,
ptr,
Int32(1),
)
return _compile_cache["qr32_fused_fp32"]
def _compiled_panel():
if "panel_fp32" not in _compile_cache:
ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_compile_cache["panel_fp32"] = cute.compile(
_launch_factor_panel,
ptr,
ptr,
Int32(1),
Int32(1),
Int32(1),
Int32(1),
)
return _compile_cache["panel_fp32"]
def _compiled_build_vt():
if "build_vt_fp32" not in _compile_cache:
ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
_compile_cache["build_vt_fp32"] = cute.compile(
_launch_build_vt,
ptr,
ptr,
ptr,
ptr,
Int32(1),
Int32(1),
Int32(1),
Int32(1),
Int32(1),
)
return _compile_cache["build_vt_fp32"]
def _factor_panel(h: torch.Tensor, tau: torch.Tensor, k: int, panel_width: int) -> None:
batch_count, n, _ = h.shape
kernel = _compiled_panel()
kernel(
make_ptr(cutlass.Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
Int32(batch_count),
Int32(n),
Int32(k),
Int32(panel_width),
)
def _build_vt(h: torch.Tensor, tau: torch.Tensor, k: int, panel_width: int) -> tuple[torch.Tensor, torch.Tensor]:
batch_count, n, _ = h.shape
v_rows = n - k
v = torch.empty((batch_count, v_rows, panel_width), device=h.device, dtype=torch.float32)
t = torch.empty((batch_count, panel_width, panel_width), device=h.device, dtype=torch.float32)
kernel = _compiled_build_vt()
kernel(
make_ptr(cutlass.Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.Float32, v.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.Float32, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
Int32(batch_count),
Int32(n),
Int32(k),
Int32(panel_width),
Int32(v_rows),
)
return v, t
def blocked_householder_qr(data: torch.Tensor) -> output_t:
if data.shape[-1] > 1024:
return torch.geqrf(data)
n = data.shape[-1]
if n == _N128:
return _s22n128_custom_kernel(data)
if n == _N512:
return _s22n512_custom_kernel(data)
h = data.contiguous().clone()
batch, n, _ = h.shape
tau = torch.zeros((batch, n), device=h.device, dtype=torch.float32)
if n == _N32:
kernel = _compiled_qr32_fused()
kernel(
make_ptr(cutlass.Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
Int32(batch),
)
return h, tau
block_size = _block_size_for_n(n)
for k in range(0, n, block_size):
panel_width = min(block_size, n - k)
_factor_panel(h, tau, k, panel_width)
trailing_start = k + panel_width
if trailing_start >= n:
continue
v, t = _build_vt(h, tau, k, panel_width)
trailing = h[:, k:, trailing_start:]
work = torch.bmm(v.transpose(1, 2), trailing)
work = torch.bmm(t, work)
trailing -= torch.bmm(v, work)
return h, tau
def custom_kernel(data: input_t) -> output_t:
return blocked_householder_qr(data)
scrolls · 1218 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