Skip to content
KernelIndex
Search⌘K

submission 826841

Light · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

codex_submission_v44.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-826841?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
3.88ms
#125 of 515
2026-06-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:65701d586601b17e05e6d72683b9f3ba10224886682c75e69f0307c54ce70070
license declaredunknown
license concludedunknown
authorsLight
imported2026-08-26

Techniques

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

num-warps = 8def _v2_triton_qr(A, block, num_warps=8, tf32_update=False):
shared-memorydef _s22n128_smem_row (row :Int32 )->Int32 :

Kernel source

codex_submission_v44.py1386 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



# Hybrid Triton path from triton_v0.py for sizes where it benchmarks faster than
# the CuteDSL v22 implementation above.
import triton
import triton.language as tl

@triton.jit
def _v2_triton_panel_kernel(
    P,
    TAU,
    T,
    VOUT,
    M,
    IB,
    spb,
    spr,
    spc,
    stb,
    sti,
    sTb,
    sTr,
    sTc,
    svb,
    svr,
    svc,
    BM: tl.constexpr,
    BNB: tl.constexpr,
):
    b = tl.program_id(0)
    r = tl.arange(0, BM)
    c = tl.arange(0, BNB)
    rm = r < M
    cm = c < IB
    p = P + b * spb + r[:, None] * spr + c[None, :] * spc
    tile = tl.load(p, mask=rm[:, None] & cm[None, :], other=0.0)
    tau_vec = tl.zeros((BNB,), dtype=tl.float32)
    for j in range(BNB):
        colj = tl.sum(tl.where(c[None, :] == j, tile, 0.0), axis=1)
        alpha = tl.sum(tl.where(r == j, colj, 0.0))
        xn2 = tl.sum(tl.where(r > j, colj * colj, 0.0))
        reflect = xn2 > 0.0
        sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = tl.where(reflect, -sgn * tl.sqrt(alpha * alpha + xn2), alpha)
        tau_j = tl.where(reflect, (beta - alpha) / tl.where(reflect, beta, 1.0), 0.0)
        denom = tl.where(reflect, alpha - beta, 1.0)
        vb = colj / denom
        v = tl.where(r == j, 1.0, tl.where(r > j, vb, 0.0))
        vmask = tl.where(r >= j, v, 0.0)
        w = tl.sum(tl.where(c[None, :] > j, vmask[:, None] * tile, 0.0), axis=0)
        tile = tile - tau_j * vmask[:, None] * w[None, :]
        newcol = tl.where(r < j, colj, tl.where(r == j, beta, vb))
        tile = tl.where(c[None, :] == j, newcol[:, None], tile)
        tau_vec = tl.where(c == j, tau_j, tau_vec)
    V = tl.where(
        r[:, None] == c[None, :], 1.0, tl.where(r[:, None] > c[None, :], tile, 0.0)
    )
    tl.store(
        VOUT + b * svb + r[:, None] * svr + c[None, :] * svc,
        V,
        mask=rm[:, None] & cm[None, :],
    )
    # in-kernel compact-WY T
    Tt = tl.zeros((BNB, BNB), dtype=tl.float32)
    tau0 = tl.sum(tl.where(c == 0, tau_vec, 0.0))
    Tt = tl.where((c[:, None] == 0) & (c[None, :] == 0), tau0, Tt)
    for i in range(1, BNB):
        tau_i = tl.sum(tl.where(c == i, tau_vec, 0.0))
        Vi = tl.sum(tl.where(c[None, :] == i, V, 0.0), axis=1)
        dots = tl.sum(V * Vi[:, None], axis=0)
        z = tl.where(c < i, -tau_i * dots, 0.0)
        Tz = tl.sum(tl.where(c[None, :] < i, Tt * z[None, :], 0.0), axis=1)
        newTcol = tl.where(c < i, Tz, tl.where(c == i, tau_i, 0.0))
        Tt = tl.where(c[None, :] == i, newTcol[:, None], Tt)
    tl.store(
        T + b * sTb + c[:, None] * sTr + c[None, :] * sTc,
        Tt,
        mask=cm[:, None] & cm[None, :],
    )
    tl.store(
        P + b * spb + r[:, None] * spr + c[None, :] * spc,
        tile,
        mask=rm[:, None] & cm[None, :],
    )
    tl.store(TAU + b * stb + c * sti, tau_vec, mask=cm)


def _v2_triton_qr(A, block, num_warps=8, tf32_update=False):
    B, m, n = A.shape
    bs = int(block)
    BNB = triton.next_power_of_2(bs)
    H = A.clone()
    tau = A.new_zeros(B, n)
    for k in range(0, n, bs):
        ib = min(bs, n - k)
        BM = triton.next_power_of_2(m - k)
        Hv = H[:, k:, k : k + ib]  # strided view, factored in place
        Tt = A.new_zeros(B, BNB, BNB)
        ts = A.new_zeros(B, BNB)
        Vb = A.new_zeros(B, m - k, ib)  # kernel writes unit-lower V here
        _v2_triton_panel_kernel[(B,)](
            Hv,
            ts,
            Tt,
            Vb,
            m - k,
            ib,
            Hv.stride(0),
            Hv.stride(1),
            Hv.stride(2),
            ts.stride(0),
            ts.stride(1),
            Tt.stride(0),
            Tt.stride(1),
            Tt.stride(2),
            Vb.stride(0),
            Vb.stride(1),
            Vb.stride(2),
            BM=BM,
            BNB=BNB,
            num_warps=num_warps,
        )
        tau[:, k : k + ib] = ts[:, :ib]
        hi = k + ib
        if hi < n:
            V = Vb
            T = Tt[:, :ib, :ib]
            C = H[:, k:, hi:]
            if tf32_update:
                old_allow_tf32 = torch.backends.cuda.matmul.allow_tf32
                old_precision = torch.get_float32_matmul_precision()
                torch.backends.cuda.matmul.allow_tf32 = True
                torch.set_float32_matmul_precision("high")
                try:
                    W = V.transpose(-1, -2) @ C
                    W = T.transpose(-1, -2) @ W
                    C.baddbmm_(V, W, beta=1, alpha=-1)
                finally:
                    torch.backends.cuda.matmul.allow_tf32 = old_allow_tf32
                    torch.set_float32_matmul_precision(old_precision)
            else:
                W = V.transpose(-1, -2) @ C
                W = T.transpose(-1, -2) @ W
                C.baddbmm_(V, W, beta=1, alpha=-1)  # your kept lever
    return H, tau



def _v2_triton_custom_kernel(data: input_t) -> output_t:
    A = data
    n = A.shape[-1]
    if n > 2048:
        return torch.geqrf(A.contiguous())
    if n == 1024:
        block, nw = 32, 8
    elif n >= 1024:
        block, nw = 16, 8
    elif n == 512:
        block, nw = 32, 4
    elif n >= 256:
        block, nw = 32, 4
    else:
        block, nw = 32, 4
    return _v2_triton_qr(A.contiguous(), block, nw, tf32_update=(n == 1024 or n == 2048))


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    if n == _N32 or n == _N128:
        return blocked_householder_qr(data)
    return _v2_triton_custom_kernel(data)
scrolls · 1386 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