submission 877085
zhongmingee · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3967 lines, June 9 Researcher Reciprocity License v1.0.
submission_e3593_cuppen_dense116.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-877085?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:5062131de620190c0a00ebbf7361333bd627289fa7b95a03e4a0795514c992b9
license declaredunknown
license concludedunknown
authorszhongmingee
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
… asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n …cluster
…warp * 8, columns[item]);\n }\n}\n\nextern "C" __global__\n__cluster_dims__(2, 1, 1)\n__launch_bounds__(352, 1)\nvoid qr2_gau_n176_panel(const float* input, float* output, float* …fused-epilogue
_E2175_GEOMETRIC_RESIDUAL_NAME='e2175_geometric_fused_residual';_E2175_GEOMETRIC_FINISH_NAME='e2175_geometric_fused_finish';_E2175_GEOMETRIC_SOURCE='\n#include <cuda_runtime.h>\n#i…mbarrier
…\nvoid mbar_init(int address, int count) {\n asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(address), "r"(count));\n}\n\n__device__ inline\nvoid mbar_arrive(int add…mma
…e+cols[None,:],mask=cols[None,:]<352,other=.0);accumulator+=tl.dot(lhs,rhs,input_precision='tf32')…num-warps = 8
…ws=triton.next_power_of_2(rows),block_columns=block_columns,num_warps=8,num_stages=1,launch_pdl=True);return matrix…persistent-kernel
…pr):offsets=tl.program_id(0)*block+tl.arange(0,block);total=tl.num_programs(0)*block;matrix_index=offsets//(n*n);element=offsets-matrix_index*n*n;row=element//n;column=element-row*…shared-memory
…"cuModuleGetFunction failed for {func_name}");self._dynamic_smem_opt_in_bytes=0;self._closed=False…stages = 1
…xt_power_of_2(rows),block_columns=block_columns,num_warps=8,num_stages=1,launch_pdl=True);return matrix…tcgen05
…nation) {\n const unsigned columns = 128;\n asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"\n : : "r"(smem_address(destination)), "…tile-k = 64
…_vectors[batch,6](rotation,values,compressed,N=352,BLOCK=64,BK=64,num_warps=4,num_stages=1);_e1930_geometric_probe_residual[batch,6](projected,compressed,residual,N=352,BLOCK=64,BK…tile-m = 256
…681_NEARRANK_CERTIFICATE_NAME).replace('constexpr int N=512,BM=256,BN=16,BK=64;','constexpr int N=1024,BM=256,BN=16,BK=64;');return CUDAKernel(_fast_nvrtc_compile(source,_E2681_NEA…tile-n = 16
…RRANK_CERTIFICATE_NAME).replace('constexpr int N=512,BM=256,BN=16,BK=64;','constexpr int N=1024,BM=256,BN=16,BK=64;');return CUDAKernel(_fast_nvrtc_compile(source,_E2681_NEARRANK_C…tma
…s(int dst, int src, int bytes, int mbar) {\n asm volatile("cp.async.bulk.shared::cluster.shared::cta.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"\n :: "r"(dst)…vector-width = float4
… (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp…warp-specialization
…,' // Prefix update overlaps the following reflector producer.\n',1).replace(old_owner,' if (rank == 0 && warp == 16) {\n if (lane == 0) {\n …Kernel source
submission_e3593_cuppen_dense116.py3967 lines
from __future__ import annotations;import torch;import triton;import triton.language as tl;from triton.language.extra import cuda as tl_cuda;import ctypes;import os;from concurrent.futures import ThreadPoolExecutor;from functools import lru_cache as memo;from task import input_t,output_t;_CERT_REASON_NONFINITE=1;_CERT_REASON_ORDERING=2;_CERT_REASON_NORM=4;_CERT_REASON_ORTHOGONALITY=8;_CERT_REASON_EIGEN_RESIDUAL=16;_CERT_REASON_SPLIT_BOUNDARY=32;_CERT_REASON_FACTORIZATION=64;_CERT_REASON_SCALE=128;_CERT_REASON_ROUTE=256;_CERT_REASON_NAMES={_CERT_REASON_NONFINITE:'nonfinite',_CERT_REASON_ORDERING:'ordering',_CERT_REASON_NORM:'norm',_CERT_REASON_ORTHOGONALITY:'orthogonality',_CERT_REASON_EIGEN_RESIDUAL:'eigen_residual',_CERT_REASON_SPLIT_BOUNDARY:'split_boundary',_CERT_REASON_FACTORIZATION:'factorization',_CERT_REASON_SCALE:'scale',_CERT_REASON_ROUTE:'route'};_certificate_reason_ledger=None
def _certificate_reason_begin():
global _certificate_reason_ledger
if _certificate_reason_ledger is not None:raise RuntimeError('certificate reason ledger is already active')
_certificate_reason_ledger=[]
def _certificate_reason_bits(reference:torch.Tensor,*reason_masks):
bits=torch.zeros_like(reference,dtype=torch.int32)
for(reason,mask)in reason_masks:bits|=mask.to(torch.int32)*int(reason)
return bits
def _certificate_reason_record(stage:str,bits:torch.Tensor,fallback:torch.Tensor|None=None,row_ids:torch.Tensor|None=None):
if _certificate_reason_ledger is None:return
if fallback is None:fallback=bits!=0
_certificate_reason_ledger.append((stage,bits.detach(),fallback.detach(),None if row_ids is None else row_ids.detach()))
def _certificate_reason_end():
global _certificate_reason_ledger;events=_certificate_reason_ledger;_certificate_reason_ledger=None
if events is None:return[]
output=[]
for(stage,bits,fallback,row_ids)in events:
bits_cpu=bits.to(device='cpu',dtype=torch.int64);fallback_cpu=fallback.to(device='cpu',dtype=torch.bool)
if row_ids is None:row_ids_cpu=torch.arange(bits_cpu.numel(),dtype=torch.int64)
else:row_ids_cpu=row_ids.to(device='cpu',dtype=torch.int64)
selected=torch.nonzero((bits_cpu!=0)|fallback_cpu,as_tuple=False).flatten();rows=[];reason_counts={name:0 for name in _CERT_REASON_NAMES.values()}
for local_index in selected.tolist():
value=int(bits_cpu[local_index]);reasons=[name for(reason,name)in _CERT_REASON_NAMES.items()if value&reason]
for reason in reasons:reason_counts[reason]+=1
rows.append({'row':int(row_ids_cpu[local_index]),'local_row':int(local_index),'bits':value,'reasons':reasons,'fallback':bool(fallback_cpu[local_index])})
output.append({'stage':stage,'rows':rows,'reason_counts':{reason:count for(reason,count)in reason_counts.items()if count},'fallback_count':int(fallback_cpu.sum())})
return output
def _fast_cuda_include_dirs():
candidates=[]
for env_name in('CUDA_HOME','CUDA_PATH'):
root=os.environ.get(env_name)
if root:candidates.append(os.path.join(root,'include'))
candidates.extend(['/usr/local/cuda/include','/cm/shared/apps/cuda13.0/toolkit/13.0.2/include']);package_root=os.path.abspath(os.path.join(os.path.dirname(torch.__file__),'..','nvidia'))
if os.path.isdir(package_root):
for child in os.listdir(package_root):candidates.append(os.path.join(package_root,child,'include'))
out=[];seen=set()
for d in candidates:
if not d or d in seen or not os.path.isdir(d):continue
seen.add(d);out.append(d);cccl=os.path.join(d,'cccl')
if os.path.exists(os.path.join(cccl,'cuda','std'))and cccl not in seen:seen.add(cccl);out.append(cccl)
return out
def _fast_get_compile_log(nvrtc,prog):
err,sz=nvrtc.nvrtcGetProgramLogSize(prog)
if err!=0 or sz<=1:return''
log=b'\x00'*sz;nvrtc.nvrtcGetProgramLog(prog,log);return log.decode(errors='replace').rstrip('\x00')
def _fast_check(err:int,msg:str='CUDA error'):
if err!=0:raise RuntimeError(f"{msg}: result={err}")
def _fast_shared_library(kind:str):
import ctypes.util
if kind=='driver':names=[ctypes.util.find_library('cuda'),'libcuda.so.1','libcuda.so']
else:
names=[ctypes.util.find_library('nvrtc'),'libnvrtc.so','libnvrtc.so.13','libnvrtc.so.12'];package_root=os.path.abspath(os.path.join(os.path.dirname(torch.__file__),'..','nvidia'))
if os.path.isdir(package_root):
for child in os.listdir(package_root):
lib_dir=os.path.join(package_root,child,'lib')
for soname in('libnvrtc.so','libnvrtc.so.13','libnvrtc.so.12'):names.append(os.path.join(lib_dir,soname))
errors=[]
for name in names:
if not name:continue
try:return ctypes.CDLL(name)
except OSError as exc:errors.append(f"{name}: {exc}")
raise RuntimeError(f"could not load CUDA {kind} library: {"; ".join(errors)}")
@memo(maxsize=1)
def _fast_ctypes_nvrtc():
lib=_fast_shared_library('nvrtc');lib.nvrtcCreateProgram.restype=ctypes.c_int;lib.nvrtcCreateProgram.argtypes=[ctypes.POINTER(ctypes.c_void_p),ctypes.c_char_p,ctypes.c_char_p,ctypes.c_int,ctypes.POINTER(ctypes.c_char_p),ctypes.POINTER(ctypes.c_char_p)];lib.nvrtcCompileProgram.restype=ctypes.c_int;lib.nvrtcCompileProgram.argtypes=[ctypes.c_void_p,ctypes.c_int,ctypes.POINTER(ctypes.c_char_p)]
for fn_name in('nvrtcGetProgramLogSize','nvrtcGetCUBINSize'):fn=getattr(lib,fn_name);fn.restype=ctypes.c_int;fn.argtypes=[ctypes.c_void_p,ctypes.POINTER(ctypes.c_size_t)]
for fn_name in('nvrtcGetProgramLog','nvrtcGetCUBIN'):fn=getattr(lib,fn_name);fn.restype=ctypes.c_int;fn.argtypes=[ctypes.c_void_p,ctypes.c_void_p]
lib.nvrtcDestroyProgram.restype=ctypes.c_int;lib.nvrtcDestroyProgram.argtypes=[ctypes.POINTER(ctypes.c_void_p)];return lib
def _fast_nvrtc_compile_ctypes(source:str,name:str,opts:list[str]):
lib=_fast_ctypes_nvrtc();prog=ctypes.c_void_p();err=lib.nvrtcCreateProgram(ctypes.byref(prog),source.encode(),name.encode(),0,None,None);_fast_check(err,'nvrtcCreateProgram failed')
try:
encoded=[option.encode()for option in opts];options=(ctypes.c_char_p*len(encoded))(*encoded);err=lib.nvrtcCompileProgram(prog,len(encoded),options)
if err!=0:
size=ctypes.c_size_t();log=''
if lib.nvrtcGetProgramLogSize(prog,ctypes.byref(size))==0 and size.value:buffer=ctypes.create_string_buffer(size.value);lib.nvrtcGetProgramLog(prog,buffer);log=buffer.value.decode(errors='replace')
raise RuntimeError(f"NVRTC compilation failed for {name}\n{log}")
size=ctypes.c_size_t();_fast_check(lib.nvrtcGetCUBINSize(prog,ctypes.byref(size)),'nvrtcGetCUBINSize failed');image=ctypes.create_string_buffer(size.value);_fast_check(lib.nvrtcGetCUBIN(prog,image),'nvrtcGetCUBIN failed');return bytes(image.raw)
finally:lib.nvrtcDestroyProgram(ctypes.byref(prog))
@memo(maxsize=1)
def _fast_ctypes_driver():lib=_fast_shared_library('driver');lib.cuModuleLoadData.restype=ctypes.c_int;lib.cuModuleLoadData.argtypes=[ctypes.POINTER(ctypes.c_void_p),ctypes.c_void_p];lib.cuModuleGetFunction.restype=ctypes.c_int;lib.cuModuleGetFunction.argtypes=[ctypes.POINTER(ctypes.c_void_p),ctypes.c_void_p,ctypes.c_char_p];lib.cuModuleUnload.restype=ctypes.c_int;lib.cuModuleUnload.argtypes=[ctypes.c_void_p];lib.cuFuncSetAttribute.restype=ctypes.c_int;lib.cuFuncSetAttribute.argtypes=[ctypes.c_void_p,ctypes.c_int,ctypes.c_int];lib.cuLaunchKernel.restype=ctypes.c_int;lib.cuLaunchKernel.argtypes=[ctypes.c_void_p,ctypes.c_uint,ctypes.c_uint,ctypes.c_uint,ctypes.c_uint,ctypes.c_uint,ctypes.c_uint,ctypes.c_uint,ctypes.c_void_p,ctypes.POINTER(ctypes.c_void_p),ctypes.c_void_p];return lib
def _fast_detect_arch():
forced=os.environ.get('QRRT_FORCE_ARCH')or os.environ.get('QR2_FAST_FORCE_ARCH')
if forced:return forced
try:major,minor=torch.cuda.get_device_capability();sm=int(major)*10+int(minor);return f"sm_{sm}a"if sm>=90 else f"sm_{sm}"
except Exception:return'sm_100a'
def _fast_nvrtc_compile_binding(source:str,name:str,opts:list[str],arch:str):
from cuda.bindings import nvrtc;err,prog=nvrtc.nvrtcCreateProgram(source.encode(),name.encode(),0,[],[]);_fast_check(err,'nvrtcCreateProgram failed')
try:
opts_b=[option.encode()for option in opts];err,=nvrtc.nvrtcCompileProgram(prog,len(opts_b),opts_b)
if err!=0:log=_fast_get_compile_log(nvrtc,prog);raise RuntimeError(f"NVRTC compilation failed for {name} arch={arch}\n{log}")
err,size=nvrtc.nvrtcGetCUBINSize(prog);_fast_check(err,'nvrtcGetCUBINSize failed');image=b'\x00'*size;err,=nvrtc.nvrtcGetCUBIN(prog,image);_fast_check(err,'nvrtcGetCUBIN failed');return image
finally:nvrtc.nvrtcDestroyProgram(prog)
def _fast_nvrtc_compile(source:str,name:str):
torch.cuda.init();torch.cuda.current_device();arch=_fast_detect_arch();opts=[f"--gpu-architecture={arch}",'-std=c++17','-default-device','--use_fast_math']
for d in _fast_cuda_include_dirs():opts.append(f"-I{d}")
aggressive=name in('blocked_qr_n1024_r992_b32','qr2_n1024_rhh_direct_left','qr2_n1024_rhh_direct_right','qr2_gau_n1024_p4','e1186_qr768_left_p0','e1186_qr768_right_p0');base_split_opts=opts+(['--split-compile=0']if aggressive else['-Xptxas=--split-compile=0']);split_opts=base_split_opts+['--minimal']
if os.environ.get('QR2_FAST_FORCE_CTYPES')=='1':
try:return _fast_nvrtc_compile_ctypes(source,name,split_opts)
except Exception:
try:return _fast_nvrtc_compile_ctypes(source,name,base_split_opts)
except Exception:return _fast_nvrtc_compile_ctypes(source,name,opts)
try:return _fast_nvrtc_compile_binding(source,name,split_opts,arch)
except Exception:
try:return _fast_nvrtc_compile_ctypes(source,name,split_opts)
except Exception:
try:return _fast_nvrtc_compile_ctypes(source,name,base_split_opts)
except Exception:return _fast_nvrtc_compile_ctypes(source,name,opts)
def _fast_only_cuda_kernel(source:str,name:str):
marker=f"void {name}(";marker_at=source.find(marker)
if marker_at<0:raise RuntimeError(f"CUDA kernel {name} not found")
start=source.rfind('extern "C"',0,marker_at);first=source.find('extern "C"');brace=source.find('{',marker_at+len(marker))
if start<0 or first<0 or brace<0:raise RuntimeError(f"CUDA kernel {name} has an invalid definition")
depth=0;end=brace
for end in range(brace,len(source)):
token=source[end]
if token=='{':depth+=1
elif token=='}':
depth-=1
if depth==0:return source[:first]+source[start:end+1]+'\n'
raise RuntimeError(f"CUDA kernel {name} has an unterminated definition")
def _fast_only_cuda_kernels(source:str,names:tuple[str,...]):
first=source.find('extern "C"')
if first<0:raise RuntimeError('CUDA source has no exported kernels')
bodies=[]
for name in names:
marker=f"void {name}(";marker_at=source.find(marker)
if marker_at<0:raise RuntimeError(f"CUDA kernel {name} not found")
start=source.rfind('extern "C"',0,marker_at);brace=source.find('{',marker_at+len(marker))
if start<0 or brace<0:raise RuntimeError(f"CUDA kernel {name} has an invalid definition")
depth=0;end=brace
for end in range(brace,len(source)):
token=source[end]
if token=='{':depth+=1
elif token=='}':
depth-=1
if depth==0:bodies.append((start,source[start:end+1]));break
else:raise RuntimeError(f"CUDA kernel {name} has an unterminated definition")
bodies.sort(key=lambda item:item[0]);return source[:first]+'\n'.join(body for(_,body)in bodies)+'\n'
def _fast_ensure_cuda_context():torch.empty(0,device='cuda')
def _fast_marshal_arg(arg):
if isinstance(arg,torch.Tensor):return ctypes.c_void_p(arg.data_ptr())
if isinstance(arg,int):return ctypes.c_int(arg)
if isinstance(arg,float):return ctypes.c_float(arg)
raise TypeError(f"Unsupported kernel argument type: {type(arg)}")
def _fast_pack_args(args):c_args=[_fast_marshal_arg(a)for a in args];ptrs=(ctypes.c_void_p*len(c_args))(*(ctypes.cast(ctypes.pointer(a),ctypes.c_void_p)for a in c_args));ptrs._prevent_gc=c_args;return ptrs
class _FastLaunchValue(ctypes.Union):_fields_=[('pad',ctypes.c_char*64),('pss',ctypes.c_int)]
class _FastLaunchAttribute(ctypes.Structure):_fields_=[('id',ctypes.c_int),('pad',ctypes.c_char*4),('value',_FastLaunchValue)]
class _FastLaunchConfig(ctypes.Structure):_fields_=[('grid_x',ctypes.c_uint),('grid_y',ctypes.c_uint),('grid_z',ctypes.c_uint),('block_x',ctypes.c_uint),('block_y',ctypes.c_uint),('block_z',ctypes.c_uint),('shared_mem',ctypes.c_uint),('queue_handle',ctypes.c_void_p),('attributes',ctypes.POINTER(_FastLaunchAttribute)),('attribute_count',ctypes.c_uint)]
class _FastCUDAKernel:
def __init__(self,cubin:bytes,func_name:str):self._closed=True;self._func_name=func_name;_fast_ensure_cuda_context();lib=_fast_ctypes_driver();self._ctypes_image=ctypes.create_string_buffer(cubin);self._module=ctypes.c_void_p();_fast_check(lib.cuModuleLoadData(ctypes.byref(self._module),self._ctypes_image),f"cuModuleLoadData failed for {func_name}");self._func=ctypes.c_void_p();_fast_check(lib.cuModuleGetFunction(ctypes.byref(self._func),self._module,func_name.encode()),f"cuModuleGetFunction failed for {func_name}");self._dynamic_smem_opt_in_bytes=0;self._closed=False
def set_attribute(self,attr,value:int):err=_fast_ctypes_driver().cuFuncSetAttribute(self._func,int(attr),int(value));_fast_check(err,f"cuFuncSetAttribute failed for {attr}={value}");self._dynamic_smem_opt_in_bytes=max(self._dynamic_smem_opt_in_bytes,int(value))
def _ensure_dynamic_smem_opt_in(self,shared_mem:int):
if shared_mem<=48*1024 or shared_mem<=self._dynamic_smem_opt_in_bytes:return
self.set_attribute(8,int(shared_mem))
def launch(self,grid,block,args,shared_mem:int=0):
if self._closed:raise RuntimeError('Kernel has been unloaded')
self._ensure_dynamic_smem_opt_in(int(shared_mem));packed=_fast_pack_args(args);err=_fast_ctypes_driver().cuLaunchKernel(self._func,int(grid[0]),int(grid[1]),int(grid[2]),int(block[0]),int(block[1]),int(block[2]),int(shared_mem),None,packed,None);_fast_check(err,f"cuLaunchKernel failed for {self._func_name}")
def launch_pdl(self,grid,block,args,shared_mem:int=0):
if self._closed:raise RuntimeError('Kernel has been unloaded')
self._ensure_dynamic_smem_opt_in(int(shared_mem));packed=_fast_pack_args(args);attribute=_FastLaunchAttribute();attribute.id=6;attribute.value.pss=1;config=_FastLaunchConfig(int(grid[0]),int(grid[1]),int(grid[2]),int(block[0]),int(block[1]),int(block[2]),int(shared_mem),None,ctypes.pointer(attribute),1);driver=_fast_ctypes_driver();driver.cuLaunchKernelEx.restype=ctypes.c_int;driver.cuLaunchKernelEx.argtypes=[ctypes.POINTER(_FastLaunchConfig),ctypes.c_void_p,ctypes.POINTER(ctypes.c_void_p),ctypes.c_void_p];err=driver.cuLaunchKernelEx(ctypes.byref(config),self._func,packed,None);_fast_check(err,f"cuLaunchKernelEx failed for {self._func_name}")
def close(self):
if not self._closed:_fast_ctypes_driver().cuModuleUnload(self._module);self._closed=True
def __enter__(self):return self
def __exit__(self,*exc):self.close()
def __del__(self):
try:self.close()
except Exception:pass
CUDAKernel=_FastCUDAKernel;N1024=1024;_N176_GAU_SOURCE='\n\n#include <cuda_fp16.h>\n\n__device__ __host__\nconstexpr int cdiv(int a, int b) { return (a + b - 1) / b; }\n\nconstexpr unsigned FULL_MASK = 0xffffffffu;\n\n__device__ __forceinline__\nfloat warp_sum(float value, int size = 32) {\n #pragma unroll\n for (int offset = size / 2; offset > 0; offset >>= 1)\n value += __shfl_xor_sync(FULL_MASK, value, offset);\n return value;\n}\n\ntemplate <int vec>\n__device__ inline\nvoid ldg_f32(float* dst, const float* src) {\n if constexpr (vec == 4)\n asm volatile("ld.global.relaxed.cta.L1::no_allocate.v4.f32 {%0, %1, %2, %3}, [%4];"\n : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3])\n : "l"(src));\n if constexpr (vec == 8)\n asm volatile("ld.global.relaxed.cta.L1::no_allocate.v8.f32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"\n : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]),\n "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7])\n : "l"(src));\n}\n\ntemplate <int vec>\n__device__ inline\nvoid stg_f32(float* dst, const float* src) {\n if constexpr (vec == 2)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v2.f32 [%0], {%1, %2};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]));\n if constexpr (vec == 4)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v4.f32 [%0], {%1, %2, %3, %4};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]));\n if constexpr (vec == 8)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v8.f32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]),\n "f"(src[4]), "f"(src[5]), "f"(src[6]), "f"(src[7]));\n}\n\n__device__ inline\nfloat sqrt_approx(float value) {\n float result;\n asm volatile("sqrt.approx.f32 %0, %1;" : "=f"(result) : "f"(value));\n return result;\n}\n\n__device__ inline\nfloat rcp_approx(float value) {\n float result;\n asm volatile("rcp.approx.f32 %0, %1;" : "=f"(result) : "f"(value));\n return result;\n}\n\n__device__ inline\nvoid fma_f32x2(float* accumulator, const float* left, const float* right) {\n asm volatile(\n "{"\n ".reg .b64 a, b, c, d;\\n"\n "mov.b64 c, {%0, %1};\\n"\n "mov.b64 a, {%2, %3};\\n"\n "mov.b64 b, {%4, %5};\\n"\n "fma.rn.f32x2 d, a, b, c;\\n"\n "mov.b64 {%0, %1}, d;\\n"\n "}"\n : "+f"(accumulator[0]), "+f"(accumulator[1])\n : "f"(left[0]), "f"(left[1]), "f"(right[0]), "f"(right[1]));\n}\n\n__device__ inline\nint elect_sync() {\n int predicate = 0;\n asm volatile(\n "{\\n\\t"\n ".reg .pred p;\\n\\t"\n "elect.sync _|p, %1;\\n\\t"\n "@p mov.s32 %0, 1;\\n\\t"\n "}"\n : "+r"(predicate)\n : "r"(FULL_MASK));\n return predicate;\n}\n\n__device__ inline\nvoid mbar_init(int address, int count) {\n asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(address), "r"(count));\n}\n\n__device__ inline\nvoid mbar_arrive(int address) {\n asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(address) : "memory");\n}\n\n__device__ inline\nvoid mbar_wait(int address, int phase) {\n constexpr int ticks = 0x989680;\n asm volatile(\n "{\\n\\t"\n ".reg .pred ready;\\n\\t"\n "mbar_wait_loop:\\n\\t"\n "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 "\n "ready, [%0], %1, %2;\\n\\t"\n "@!ready bra.uni mbar_wait_loop;\\n\\t"\n "}"\n :: "r"(address), "r"(phase), "r"(ticks));\n}\n\n__device__ inline\nvoid mbar_expect_tx(int address, int bytes) {\n asm volatile("mbarrier.arrive.expect_tx.relaxed.cluster.shared::cluster.b64 _, [%0], %1;"\n :: "r"(address), "r"(bytes) : "memory");\n}\n\n__device__ inline\nvoid tma_s2s(int dst, int src, int bytes, int mbar) {\n asm volatile("cp.async.bulk.shared::cluster.shared::cta.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"\n :: "r"(dst), "r"(src), "r"(bytes), "r"(mbar));\n}\n\n__device__ inline\nvoid tma_s2g(void *dst, int src, int bytes) {\n asm volatile("cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" :: "l"(dst), "r"(src), "r"(bytes));\n}\n\n__device__ inline\nvoid st_async_f32(int destination, float value, int mbar) {\n asm volatile("st.async.shared::cluster.mbarrier::complete_tx::bytes.f32 [%0], %1, [%2];"\n :: "r"(destination), "f"(value), "r"(mbar));\n}\n\n__device__ __forceinline__\nvoid store_release_gpu(int* address, int value) {\n asm volatile("st.release.gpu.global.u32 [%0], %1;" :: "l"(address), "r"(value) : "memory");\n}\n\n__device__ __forceinline__\nint load_relaxed_gpu_no_allocate(const int* address) {\n int value;\n asm volatile("ld.global.relaxed.gpu.L1::no_allocate.u32 %0, [%1];" : "=r"(value) : "l"(address));\n return value;\n}\n\n__device__ __forceinline__ void fence_acquire_gpu() {\n asm volatile("fence.acquire.gpu;" ::: "memory");\n}\n\ntemplate <typename T>\n__device__ inline T warp_uniform(T value) {\n return __shfl_sync(FULL_MASK, value, 0);\n}\n\ntemplate <int ROWS, int COLS, int N>\n__global__\n__launch_bounds__((COLS / 8) * 32, 1)\nvoid register_panel_kernel(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n static_assert(COLS % 8 == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage; // [ROWS, COLS]\n float* taus = reflectors + ROWS * COLS; // [COLS]\n const int mbars = __cvta_generic_to_shared(taus + COLS);\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) {\n mbar_init(mbars + i * 8, 32);\n }\n }\n __syncthreads();\n\n float columns[ROW_ITEMS][8];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS) {\n ldg_f32<8>(columns[item], input + row * N + warp * 8);\n } else {\n for (int i = 0; i < 8; ++i) columns[item][i] = 0.0f;\n }\n }\n\n for (int panel = 0; panel < warp; ++panel) {\n for (int i = 0; i < 8; ++i) {\n const int col = panel * 8 + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < 4; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < 8; ++i) {\n const int col = warp * 8 + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? 1.0f / (x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[col * ROWS + row] = v[item];\n }\n mbar_arrive(mbars + col * 8);\n\n for (int trailing_col = i + 1; trailing_col < 8; ++trailing_col) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing_col] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing_col] -= v[item] * dot;\n }\n }\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = 8 * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + warp * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + warp * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (warp * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n if (lane < 2) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[warp * 2 + lane];\n stg_f32<4>(tau + (warp * 2 + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<8>(output + row * N + warp * 8, columns[item]);\n }\n}\n\nextern "C" __global__\n__cluster_dims__(2, 1, 1)\n__launch_bounds__(352, 1)\nvoid qr2_gau_n176_panel(const float* input, float* output, float* tau, float *v_fp32, __half *v_fp16) {\n constexpr int ROWS=176, COLS=176, N=176, VEC_SIZE=8;\n static_assert(COLS % (VEC_SIZE * 2) == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n constexpr int NUM_WARPS = COLS / VEC_SIZE / 2;\n\n const int tid = threadIdx.x;\n const int block = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n const int rank = block & 1;\n const int batch = block / 2;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage;\n constexpr int LOCAL_COLS = COLS / 2;\n float* taus = reflectors + ROWS * LOCAL_COLS;\n const int reflector_addr = __cvta_generic_to_shared(reflectors);\n const int tau_addr = reflector_addr + ROWS * LOCAL_COLS * 4;\n const int mbars = tau_addr + COLS * 4;\n\n const int reflector_addr1 = reflector_addr | 0x01000000;\n const int tau_addr1 = tau_addr | 0x01000000;\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) mbar_init(mbars + i * 8, 1);\n asm volatile("fence.mbarrier_init.release.cluster;");\n }\n asm volatile("barrier.cluster.arrive.relaxed.aligned;");\n asm volatile("barrier.cluster.wait.acquire.aligned;");\n\n float columns[ROW_ITEMS][VEC_SIZE];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const int col = (rank * NUM_WARPS + warp) * VEC_SIZE;\n if (row < ROWS) {\n ldg_f32<VEC_SIZE>(columns[item], input + row * N + col);\n } else {\n for (int i = 0; i < VEC_SIZE; ++i) columns[item][i] = 0.0f;\n }\n }\n\n // from remote reflectors\n for (int panel = 0; panel < rank * NUM_WARPS; ++panel) {\n for (int i = 0; i < VEC_SIZE; ++i) {\n const int col = panel * VEC_SIZE + i;\n if (warp == 0)\n mbar_wait(mbars + col * 8, 0);\n __syncthreads();\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < VEC_SIZE/2; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n __syncthreads();\n\n // from local reflectors\n for (int panel = rank * NUM_WARPS; panel < rank * NUM_WARPS + warp; ++panel) {\n const int local_panel = panel - rank * NUM_WARPS;\n for (int i = 0; i < VEC_SIZE; ++i) {\n const int col = panel * VEC_SIZE + i;\n const int local_col = local_panel * VEC_SIZE + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[local_col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < VEC_SIZE/2; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < VEC_SIZE; ++i) {\n const int col = (rank * NUM_WARPS + warp) * VEC_SIZE + i;\n const int local_col = warp * VEC_SIZE + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? rcp_approx(x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail ? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[local_col * ROWS + row] = v[item];\n }\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n if (elect_sync()) {\n mbar_arrive(mbars + col * 8);\n if (rank == 0) {\n const int remote_mbar = (mbars + col * 8) | 0x01000000;\n mbar_expect_tx(remote_mbar, (ROWS + 1) * 4);\n tma_s2s(reflector_addr1 + col * ROWS * 4,\n reflector_addr + local_col * ROWS * 4,\n ROWS * 4, remote_mbar);\n st_async_f32(tau_addr1 + col * 4, tau_value, remote_mbar);\n }\n }\n\n for (int trailing = i + 1; trailing < VEC_SIZE; ++trailing) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing] -= v[item] * dot;\n }\n }\n const int panel_id = rank * NUM_WARPS + warp;\n const int local_panel_id = warp;\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = VEC_SIZE * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + local_panel_id * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + panel_id * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < PANEL_SIZE) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (local_panel_id * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (panel_id * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n const int col = panel_id * VEC_SIZE;\n if (lane < VEC_SIZE / 4) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[panel_id * (VEC_SIZE/4) + lane];\n stg_f32<4>(tau + (panel_id * (VEC_SIZE/4) + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<VEC_SIZE>(output + row * N + col, columns[item]);\n }\n}\n\n';_N512_GAU_PANEL_NAMES='qr2_gau_n512_p0','qr2_gau_n512_p1','qr2_gau_n512_p2','qr2_gau_n512_p3';_N512_GAU_PANEL_SOURCE=_N176_GAU_SOURCE+'\nextern "C" __global__\n__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=512, COLS=96, N=512;\n static_assert(COLS % 8 == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage; // [ROWS, COLS]\n float* taus = reflectors + ROWS * COLS; // [COLS]\n const int mbars = __cvta_generic_to_shared(taus + COLS);\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) {\n mbar_init(mbars + i * 8, 32);\n }\n }\n __syncthreads();\n\n float columns[ROW_ITEMS][8];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS) {\n ldg_f32<8>(columns[item], input + row * N + warp * 8);\n } else {\n for (int i = 0; i < 8; ++i) columns[item][i] = 0.0f;\n }\n }\n\n for (int panel = 0; panel < warp; ++panel) {\n for (int i = 0; i < 8; ++i) {\n const int col = panel * 8 + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < 4; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < 8; ++i) {\n const int col = warp * 8 + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? 1.0f / (x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[col * ROWS + row] = v[item];\n }\n mbar_arrive(mbars + col * 8);\n\n for (int trailing_col = i + 1; trailing_col < 8; ++trailing_col) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing_col] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing_col] -= v[item] * dot;\n }\n }\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = 8 * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + warp * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + warp * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (warp * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n if (lane < 2) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[warp * 2 + lane];\n stg_f32<4>(tau + (warp * 2 + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<8>(output + row * N + warp * 8, columns[item]);\n }\n}\n\nextern "C" __global__\n__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p1(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=416, COLS=96, N=512;\n static_assert(COLS % 8 == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage; // [ROWS, COLS]\n float* taus = reflectors + ROWS * COLS; // [COLS]\n const int mbars = __cvta_generic_to_shared(taus + COLS);\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) {\n mbar_init(mbars + i * 8, 32);\n }\n }\n __syncthreads();\n\n float columns[ROW_ITEMS][8];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS) {\n ldg_f32<8>(columns[item], input + row * N + warp * 8);\n } else {\n for (int i = 0; i < 8; ++i) columns[item][i] = 0.0f;\n }\n }\n\n for (int panel = 0; panel < warp; ++panel) {\n for (int i = 0; i < 8; ++i) {\n const int col = panel * 8 + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < 4; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < 8; ++i) {\n const int col = warp * 8 + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? 1.0f / (x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[col * ROWS + row] = v[item];\n }\n mbar_arrive(mbars + col * 8);\n\n for (int trailing_col = i + 1; trailing_col < 8; ++trailing_col) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing_col] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing_col] -= v[item] * dot;\n }\n }\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = 8 * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + warp * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + warp * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (warp * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n if (lane < 2) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[warp * 2 + lane];\n stg_f32<4>(tau + (warp * 2 + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<8>(output + row * N + warp * 8, columns[item]);\n }\n}\n\nextern "C" __global__\n__launch_bounds__(512, 1)\nvoid qr2_gau_n512_p2(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=320, COLS=128, N=512;\n static_assert(COLS % 8 == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage; // [ROWS, COLS]\n float* taus = reflectors + ROWS * COLS; // [COLS]\n const int mbars = __cvta_generic_to_shared(taus + COLS);\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) {\n mbar_init(mbars + i * 8, 32);\n }\n }\n __syncthreads();\n\n float columns[ROW_ITEMS][8];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS) {\n ldg_f32<8>(columns[item], input + row * N + warp * 8);\n } else {\n for (int i = 0; i < 8; ++i) columns[item][i] = 0.0f;\n }\n }\n\n for (int panel = 0; panel < warp; ++panel) {\n for (int i = 0; i < 8; ++i) {\n const int col = panel * 8 + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < 4; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < 8; ++i) {\n const int col = warp * 8 + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? 1.0f / (x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[col * ROWS + row] = v[item];\n }\n mbar_arrive(mbars + col * 8);\n\n for (int trailing_col = i + 1; trailing_col < 8; ++trailing_col) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing_col] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing_col] -= v[item] * dot;\n }\n }\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = 8 * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + warp * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + warp * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (warp * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n if (lane < 2) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[warp * 2 + lane];\n stg_f32<4>(tau + (warp * 2 + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<8>(output + row * N + warp * 8, columns[item]);\n }\n}\n\nextern "C" __global__\n__launch_bounds__(768, 1)\nvoid qr2_gau_n512_p3(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=192, COLS=192, N=512;\n static_assert(COLS % 8 == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage; // [ROWS, COLS]\n float* taus = reflectors + ROWS * COLS; // [COLS]\n const int mbars = __cvta_generic_to_shared(taus + COLS);\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) {\n mbar_init(mbars + i * 8, 32);\n }\n }\n __syncthreads();\n\n float columns[ROW_ITEMS][8];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS) {\n ldg_f32<8>(columns[item], input + row * N + warp * 8);\n } else {\n for (int i = 0; i < 8; ++i) columns[item][i] = 0.0f;\n }\n }\n\n for (int panel = 0; panel < warp; ++panel) {\n for (int i = 0; i < 8; ++i) {\n const int col = panel * 8 + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < 4; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < 8; ++i) {\n const int col = warp * 8 + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? 1.0f / (x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[col * ROWS + row] = v[item];\n }\n mbar_arrive(mbars + col * 8);\n\n for (int trailing_col = i + 1; trailing_col < 8; ++trailing_col) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing_col] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing_col] -= v[item] * dot;\n }\n }\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = 8 * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + warp * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + warp * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (warp * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n if (lane < 2) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[warp * 2 + lane];\n stg_f32<4>(tau + (warp * 2 + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<8>(output + row * N + warp * 8, columns[item]);\n }\n}\n\n';_N512_GAU_ACTIVE_NAMES='qr2_gau_n512_active320x64','qr2_gau_n512_active192x64';_N512_GAU_ACTIVE_SOURCE=_N512_GAU_PANEL_SOURCE.replace('__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=512, COLS=96, N=512;','__launch_bounds__(256, 1)\nvoid qr2_gau_n512_active320x64(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=320, COLS=64, N=512;',1).replace('__launch_bounds__(768, 1)\nvoid qr2_gau_n512_p3(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=192, COLS=192, N=512;','__launch_bounds__(256, 1)\nvoid qr2_gau_n512_active192x64(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=192, COLS=64, N=512;',1)
@memo(maxsize=1)
def _n512_gau_panel_image():return _fast_nvrtc_compile(_N512_GAU_PANEL_SOURCE,_N512_GAU_PANEL_NAMES[0])
@memo(maxsize=1)
def _n512_gau_panel_kernels():image=_n512_gau_panel_image();return tuple(CUDAKernel(image,name)for name in _N512_GAU_PANEL_NAMES)
@memo(maxsize=1)
def _n512_gau_active_image():source=_fast_only_cuda_kernels(_N512_GAU_ACTIVE_SOURCE,_N512_GAU_ACTIVE_NAMES);return _fast_nvrtc_compile(source,_N512_GAU_ACTIVE_NAMES[0])
@memo(maxsize=1)
def _n512_gau_active_kernels():image=_n512_gau_active_image();return tuple(CUDAKernel(image,name)for name in _N512_GAU_ACTIVE_NAMES)
_N1024_GAU_PANEL_NAMES='qr2_gau_n1024_p0','qr2_gau_n1024_p1','qr2_gau_n1024_p2','qr2_gau_n1024_p3','qr2_gau_n1024_p4','qr2_gau_n1024_p5','qr2_gau_n1024_p6','qr2_gau_n1024_p7','qr2_gau_n1024_p8';_N1024_GAU_PANEL_SOURCE='\n\n#include <cuda_fp16.h>\n\n__device__ __host__\nconstexpr int cdiv(int a, int b) { return (a + b - 1) / b; }\n\nconstexpr unsigned FULL_MASK = 0xffffffffu;\n\n__device__ __forceinline__\nfloat warp_sum(float value, int size = 32) {\n #pragma unroll\n for (int offset = size / 2; offset > 0; offset >>= 1)\n value += __shfl_xor_sync(FULL_MASK, value, offset);\n return value;\n}\n\ntemplate <int vec>\n__device__ inline\nvoid ldg_f32(float* dst, const float* src) {\n if constexpr (vec == 4)\n asm volatile("ld.global.relaxed.cta.L1::no_allocate.v4.f32 {%0, %1, %2, %3}, [%4];"\n : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3])\n : "l"(src));\n if constexpr (vec == 8)\n asm volatile("ld.global.relaxed.cta.L1::no_allocate.v8.f32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"\n : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]),\n "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7])\n : "l"(src));\n}\n\ntemplate <int vec>\n__device__ inline\nvoid stg_f32(float* dst, const float* src) {\n if constexpr (vec == 2)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v2.f32 [%0], {%1, %2};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]));\n if constexpr (vec == 4)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v4.f32 [%0], {%1, %2, %3, %4};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]));\n if constexpr (vec == 8)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v8.f32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]),\n "f"(src[4]), "f"(src[5]), "f"(src[6]), "f"(src[7]));\n}\n\n__device__ inline\nfloat sqrt_approx(float value) {\n float result;\n asm volatile("sqrt.approx.f32 %0, %1;" : "=f"(result) : "f"(value));\n return result;\n}\n\n__device__ inline\nfloat rcp_approx(float value) {\n float result;\n asm volatile("rcp.approx.f32 %0, %1;" : "=f"(result) : "f"(value));\n return result;\n}\n\n__device__ inline\nvoid fma_f32x2(float* accumulator, const float* left, const float* right) {\n asm volatile(\n "{"\n ".reg .b64 a, b, c, d;\\n"\n "mov.b64 c, {%0, %1};\\n"\n "mov.b64 a, {%2, %3};\\n"\n "mov.b64 b, {%4, %5};\\n"\n "fma.rn.f32x2 d, a, b, c;\\n"\n "mov.b64 {%0, %1}, d;\\n"\n "}"\n : "+f"(accumulator[0]), "+f"(accumulator[1])\n : "f"(left[0]), "f"(left[1]), "f"(right[0]), "f"(right[1]));\n}\n\n__device__ inline\nint elect_sync() {\n int predicate = 0;\n asm volatile(\n "{\\n\\t"\n ".reg .pred p;\\n\\t"\n "elect.sync _|p, %1;\\n\\t"\n "@p mov.s32 %0, 1;\\n\\t"\n "}"\n : "+r"(predicate)\n : "r"(FULL_MASK));\n return predicate;\n}\n\n__device__ inline\nvoid mbar_init(int address, int count) {\n asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(address), "r"(count));\n}\n\n__device__ inline\nvoid mbar_arrive(int address) {\n asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(address) : "memory");\n}\n\n__device__ inline\nvoid mbar_wait(int address, int phase) {\n constexpr int ticks = 0x989680;\n asm volatile(\n "{\\n\\t"\n ".reg .pred ready;\\n\\t"\n "mbar_wait_loop:\\n\\t"\n "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 "\n "ready, [%0], %1, %2;\\n\\t"\n "@!ready bra.uni mbar_wait_loop;\\n\\t"\n "}"\n :: "r"(address), "r"(phase), "r"(ticks));\n}\n\n__device__ inline\nvoid mbar_expect_tx(int address, int bytes) {\n asm volatile("mbarrier.arrive.expect_tx.relaxed.cluster.shared::cluster.b64 _, [%0], %1;"\n :: "r"(address), "r"(bytes) : "memory");\n}\n\n__device__ inline\nvoid tma_s2s(int dst, int src, int bytes, int mbar) {\n asm volatile("cp.async.bulk.shared::cluster.shared::cta.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"\n :: "r"(dst), "r"(src), "r"(bytes), "r"(mbar));\n}\n\n__device__ inline\nvoid tma_s2g(void *dst, int src, int bytes) {\n asm volatile("cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" :: "l"(dst), "r"(src), "r"(bytes));\n}\n\n__device__ inline\nvoid st_async_f32(int destination, float value, int mbar) {\n asm volatile("st.async.shared::cluster.mbarrier::complete_tx::bytes.f32 [%0], %1, [%2];"\n :: "r"(destination), "f"(value), "r"(mbar));\n}\n\n__device__ __forceinline__\nvoid store_release_gpu(int* address, int value) {\n asm volatile("st.release.gpu.global.u32 [%0], %1;" :: "l"(address), "r"(value) : "memory");\n}\n\n__device__ __forceinline__\nint load_relaxed_gpu_no_allocate(const int* address) {\n int value;\n asm volatile("ld.global.relaxed.gpu.L1::no_allocate.u32 %0, [%1];" : "=r"(value) : "l"(address));\n return value;\n}\n\n__device__ __forceinline__ void fence_acquire_gpu() {\n asm volatile("fence.acquire.gpu;" ::: "memory");\n}\n\ntemplate <typename T>\n__device__ inline T warp_uniform(T value) {\n return __shfl_sync(FULL_MASK, value, 0);\n}\n\ntemplate <int ROWS, int COLS, int N>\n__global__\n__launch_bounds__((COLS / 8) * 32, 1)\nvoid register_panel_kernel(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n static_assert(COLS % 8 == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage; // [ROWS, COLS]\n float* taus = reflectors + ROWS * COLS; // [COLS]\n const int mbars = __cvta_generic_to_shared(taus + COLS);\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) {\n mbar_init(mbars + i * 8, 32);\n }\n }\n __syncthreads();\n\n float columns[ROW_ITEMS][8];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS) {\n ldg_f32<8>(columns[item], input + row * N + warp * 8);\n } else {\n for (int i = 0; i < 8; ++i) columns[item][i] = 0.0f;\n }\n }\n\n for (int panel = 0; panel < warp; ++panel) {\n for (int i = 0; i < 8; ++i) {\n const int col = panel * 8 + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < 4; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < 8; ++i) {\n const int col = warp * 8 + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? 1.0f / (x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[col * ROWS + row] = v[item];\n }\n mbar_arrive(mbars + col * 8);\n\n for (int trailing_col = i + 1; trailing_col < 8; ++trailing_col) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing_col] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing_col] -= v[item] * dot;\n }\n }\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = 8 * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + warp * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + warp * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < ROWS * 8) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (warp * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (warp * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n if (lane < 2) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[warp * 2 + lane];\n stg_f32<4>(tau + (warp * 2 + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<8>(output + row * N + warp * 8, columns[item]);\n }\n}\n\ntemplate <int ROWS, int COLS, int N, int VEC_SIZE>\n__device__ __forceinline__ void qr2_gau_panel_body(const float* input, float* output, float* tau, float *v_fp32, __half *v_fp16) {\n static_assert(COLS % (VEC_SIZE * 2) == 0);\n static_assert(COLS <= ROWS);\n constexpr int ROW_ITEMS = (ROWS + 31) / 32;\n constexpr int NUM_WARPS = COLS / VEC_SIZE / 2;\n\n const int tid = threadIdx.x;\n const int block = blockIdx.x;\n const int warp = warp_uniform(tid / 32);\n const int lane = tid & 31;\n const int rank = block & 1;\n const int batch = block / 2;\n\n input += batch * N * N;\n output += batch * N * N;\n tau += batch * N;\n\n extern __shared__ float storage[];\n float* reflectors = storage;\n constexpr int LOCAL_COLS = COLS / 2;\n float* taus = reflectors + ROWS * LOCAL_COLS;\n const int reflector_addr = __cvta_generic_to_shared(reflectors);\n const int tau_addr = reflector_addr + ROWS * LOCAL_COLS * 4;\n const int mbars = tau_addr + COLS * 4;\n\n const int reflector_addr1 = reflector_addr | 0x01000000;\n const int tau_addr1 = tau_addr | 0x01000000;\n\n if (warp == 0 && elect_sync()) {\n for (int i = 0; i < COLS; ++i) mbar_init(mbars + i * 8, 1);\n asm volatile("fence.mbarrier_init.release.cluster;");\n }\n asm volatile("barrier.cluster.arrive.relaxed.aligned;");\n asm volatile("barrier.cluster.wait.acquire.aligned;");\n\n float columns[ROW_ITEMS][VEC_SIZE];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const int col = (rank * NUM_WARPS + warp) * VEC_SIZE;\n if (row < ROWS) {\n ldg_f32<VEC_SIZE>(columns[item], input + row * N + col);\n } else {\n for (int i = 0; i < VEC_SIZE; ++i) columns[item][i] = 0.0f;\n }\n }\n\n // from remote reflectors\n for (int panel = 0; panel < rank * NUM_WARPS; ++panel) {\n for (int i = 0; i < VEC_SIZE; ++i) {\n const int col = panel * VEC_SIZE + i;\n if (warp == 0)\n mbar_wait(mbars + col * 8, 0);\n __syncthreads();\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < VEC_SIZE/2; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n __syncthreads();\n\n // from local reflectors\n for (int panel = rank * NUM_WARPS; panel < rank * NUM_WARPS + warp; ++panel) {\n const int local_panel = panel - rank * NUM_WARPS;\n for (int i = 0; i < VEC_SIZE; ++i) {\n const int col = panel * VEC_SIZE + i;\n const int local_col = local_panel * VEC_SIZE + i;\n mbar_wait(mbars + col * 8, 0);\n const float negative_tau = -taus[col];\n float v[ROW_ITEMS][2];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float value = row < ROWS ? reflectors[local_col * ROWS + row] : 0.0f;\n v[item][0] = value;\n v[item][1] = value;\n }\n for (int pair = 0; pair < VEC_SIZE/2; ++pair) {\n float dot[2] = {};\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(dot, &columns[item][pair * 2], v[item]);\n dot[0] = warp_sum(dot[0]) * negative_tau;\n dot[1] = warp_sum(dot[1]) * negative_tau;\n for (int item = 0; item < ROW_ITEMS; ++item)\n fma_f32x2(&columns[item][pair * 2], v[item], dot);\n }\n }\n }\n\n #pragma unroll\n for (int i = 0; i < VEC_SIZE; ++i) {\n const int col = (rank * NUM_WARPS + warp) * VEC_SIZE + i;\n const int local_col = warp * VEC_SIZE + i;\n if constexpr (ROWS == COLS) {\n if (col == COLS - 1) {\n if (lane == 0) taus[col] = 0.0f;\n break;\n }\n }\n\n float tail = 0.0f;\n float x0 = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n tail += (row > col) * x * x;\n x0 += row == col ? x : 0.0f;\n }\n tail = warp_sum(tail);\n x0 = __shfl_sync(FULL_MASK, x0, col & 31);\n\n const float norm = sqrt_approx(x0 * x0 + tail);\n const float beta = -copysignf(norm, x0);\n const bool has_tail = tail > 0.0f;\n const float tau_value = has_tail ? (beta - x0) * rcp_approx(beta) : 0.0f;\n const float inverse = has_tail ? rcp_approx(x0 - beta) : 0.0f;\n if (lane == 0) taus[col] = tau_value;\n\n float v[ROW_ITEMS];\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n const float x = columns[item][i];\n v[item] = has_tail ? (row == col) + (row > col) * (x * inverse) : 0.0f;\n const float reflected = (row < col) * x + (row == col) * beta + (row > col) * v[item];\n columns[item][i] = has_tail ? reflected : x;\n if (row < ROWS) reflectors[local_col * ROWS + row] = v[item];\n }\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n if (elect_sync()) {\n mbar_arrive(mbars + col * 8);\n if (rank == 0) {\n const int remote_mbar = (mbars + col * 8) | 0x01000000;\n mbar_expect_tx(remote_mbar, (ROWS + 1) * 4);\n tma_s2s(reflector_addr1 + col * ROWS * 4,\n reflector_addr + local_col * ROWS * 4,\n ROWS * 4, remote_mbar);\n st_async_f32(tau_addr1 + col * 4, tau_value, remote_mbar);\n }\n }\n\n for (int trailing = i + 1; trailing < VEC_SIZE; ++trailing) {\n float dot = 0.0f;\n for (int item = 0; item < ROW_ITEMS; ++item)\n dot += columns[item][trailing] * v[item];\n dot = warp_sum(dot) * tau_value;\n for (int item = 0; item < ROW_ITEMS; ++item)\n columns[item][trailing] -= v[item] * dot;\n }\n }\n const int panel_id = rank * NUM_WARPS + warp;\n const int local_panel_id = warp;\n\n // emit V when it\'s not the last square QR\n if constexpr (ROWS > COLS) {\n v_fp32 += batch * ROWS * COLS;\n v_fp16 += batch * ROWS * COLS;\n\n __syncwarp();\n asm volatile("fence.proxy.async.shared::cta;");\n constexpr int PANEL_SIZE = VEC_SIZE * ROWS;\n if (elect_sync()) {\n const int sV_fp32 = __cvta_generic_to_shared(reflectors) + local_panel_id * PANEL_SIZE * 4;\n tma_s2g(v_fp32 + panel_id * PANEL_SIZE, sV_fp32, PANEL_SIZE * 4);\n }\n\n for (int i = 0; i < cdiv(PANEL_SIZE, 32 * 4); i++) {\n const int idx = (i * 32 + lane) * 4;\n if (idx < PANEL_SIZE) {\n float4 tmp = reinterpret_cast<float4 *>(reflectors + (local_panel_id * PANEL_SIZE + idx))[0];\n half2 tmp2[2];\n tmp2[0] = __float22half2_rn({tmp.x, tmp.y});\n tmp2[1] = __float22half2_rn({tmp.z, tmp.w});\n stg_f32<2>(\n reinterpret_cast<float *>(v_fp16 + (panel_id * PANEL_SIZE + idx)),\n reinterpret_cast<float *>(tmp2));\n }\n }\n }\n\n const int col = panel_id * VEC_SIZE;\n if (lane < VEC_SIZE / 4) {\n float tmp[4];\n reinterpret_cast<float4 *>(tmp)[0] = reinterpret_cast<float4 *>(taus)[panel_id * (VEC_SIZE/4) + lane];\n stg_f32<4>(tau + (panel_id * (VEC_SIZE/4) + lane) * 4, tmp);\n }\n for (int item = 0; item < ROW_ITEMS; ++item) {\n const int row = item * 32 + lane;\n if (row < ROWS)\n stg_f32<VEC_SIZE>(output + row * N + col, columns[item]);\n }\n}\n\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(384, 1)\nvoid qr2_gau_n1024_p0(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<1024, 96, 1024, 4>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(384, 1)\nvoid qr2_gau_n1024_p1(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<928, 96, 1024, 4>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(384, 1)\nvoid qr2_gau_n1024_p2(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<832, 96, 1024, 4>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(384, 1)\nvoid qr2_gau_n1024_p3(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<736, 96, 1024, 4>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(256, 1)\nvoid qr2_gau_n1024_p4(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<640, 128, 1024, 8>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(256, 1)\nvoid qr2_gau_n1024_p5(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<512, 128, 1024, 8>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(256, 1)\nvoid qr2_gau_n1024_p6(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<384, 128, 1024, 8>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(256, 1)\nvoid qr2_gau_n1024_p7(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<256, 128, 1024, 8>(input, output, tau, v32, v16);\n}\n\nextern "C" __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(256, 1)\nvoid qr2_gau_n1024_p8(const float* input, float* output, float* tau, float* v32, __half* v16) {\n qr2_gau_panel_body<128, 128, 1024, 8>(input, output, tau, v32, v16);\n}\n\n';_N352_GAU_T_SOURCE='\n\n#include <cuda_fp16.h>\n\n__device__ __host__\nconstexpr int cdiv(int a, int b) { return (a + b - 1) / b; }\n\nconstexpr unsigned FULL_MASK = 0xffffffffu;\n\n__device__ __forceinline__\nfloat warp_sum(float value, int size = 32) {\n #pragma unroll\n for (int offset = size / 2; offset > 0; offset >>= 1)\n value += __shfl_xor_sync(FULL_MASK, value, offset);\n return value;\n}\n\ntemplate <int vec>\n__device__ inline\nvoid ldg_f32(float* dst, const float* src) {\n if constexpr (vec == 4)\n asm volatile("ld.global.relaxed.cta.L1::no_allocate.v4.f32 {%0, %1, %2, %3}, [%4];"\n : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3])\n : "l"(src));\n if constexpr (vec == 8)\n asm volatile("ld.global.relaxed.cta.L1::no_allocate.v8.f32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"\n : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]),\n "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7])\n : "l"(src));\n}\n\ntemplate <int vec>\n__device__ inline\nvoid stg_f32(float* dst, const float* src) {\n if constexpr (vec == 2)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v2.f32 [%0], {%1, %2};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]));\n if constexpr (vec == 4)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v4.f32 [%0], {%1, %2, %3, %4};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]));\n if constexpr (vec == 8)\n asm volatile("st.global.relaxed.cta.L1::no_allocate.v8.f32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8};"\n :: "l"(dst),\n "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]),\n "f"(src[4]), "f"(src[5]), "f"(src[6]), "f"(src[7]));\n}\n\n__device__ inline\nfloat sqrt_approx(float value) {\n float result;\n asm volatile("sqrt.approx.f32 %0, %1;" : "=f"(result) : "f"(value));\n return result;\n}\n\n__device__ inline\nfloat rcp_approx(float value) {\n float result;\n asm volatile("rcp.approx.f32 %0, %1;" : "=f"(result) : "f"(value));\n return result;\n}\n\n__device__ inline\nvoid fma_f32x2(float* accumulator, const float* left, const float* right) {\n asm volatile(\n "{"\n ".reg .b64 a, b, c, d;\\n"\n "mov.b64 c, {%0, %1};\\n"\n "mov.b64 a, {%2, %3};\\n"\n "mov.b64 b, {%4, %5};\\n"\n "fma.rn.f32x2 d, a, b, c;\\n"\n "mov.b64 {%0, %1}, d;\\n"\n "}"\n : "+f"(accumulator[0]), "+f"(accumulator[1])\n : "f"(left[0]), "f"(left[1]), "f"(right[0]), "f"(right[1]));\n}\n\ntemplate <int STRIDE>\n__device__ inline\nvoid build_t32_inverse_block(\n const float* lower,\n const float* tau,\n float* inverse,\n float* mid) {\n constexpr int K = 32;\n constexpr int HALF = 16;\n const int tid = threadIdx.x;\n const int warp = tid / 32;\n const int lane = tid & 31;\n const int half = lane >> 4;\n const int sublane = lane & 15;\n const int block = warp >> 3;\n const int local_warp = warp & 7;\n const int local_col = local_warp * 2 + half;\n const int block_base = block * HALF;\n\n float x = 0.0f;\n\n #pragma unroll\n for (int solve_row = 0; solve_row < HALF; ++solve_row) {\n float partial =\n lower[(block_base + solve_row) * STRIDE + block_base + sublane] * x;\n partial = warp_sum(partial, HALF);\n const float diagonal = solve_row == local_col ? 1.0f : 0.0f;\n const float value =\n (diagonal - partial) * tau[block_base + solve_row];\n\n if (solve_row == sublane) {\n x = value;\n inverse[(block_base + solve_row) * STRIDE + block_base + local_col] = value;\n }\n }\n __syncthreads();\n\n if (tid < HALF * HALF) {\n const int row = tid >> 4;\n const int col = tid & 15;\n float accum = 0.0f;\n #pragma unroll\n for (int k = 0; k < HALF; ++k) {\n accum += lower[(HALF + row) * STRIDE + k] * inverse[k * STRIDE + col];\n }\n mid[row * HALF + col] = accum;\n }\n __syncthreads();\n\n if (tid < HALF * HALF) {\n const int row = tid >> 4;\n const int col = tid & 15;\n float accum = 0.0f;\n #pragma unroll\n for (int k = 0; k < HALF; ++k) {\n accum += inverse[(HALF + row) * STRIDE + HALF + k] * mid[k * HALF + col];\n }\n inverse[(HALF + row) * STRIDE + col] = -accum;\n }\n}\n\n__device__ inline\nvoid build_t64_block(\n const float* gram,\n const float* tau,\n float* output,\n float* lower,\n float* inverse,\n float* mid,\n int gram_ld,\n int output_ld,\n int base,\n int tau_stride_offset) {\n constexpr int K = 64;\n constexpr int TB_SIZE = 16 * 32;\n const int tid = threadIdx.x;\n const float* tau_b = tau + tau_stride_offset;\n\n {\n const int pack = tid;\n const int row = pack / (K / 8);\n const int col = (pack - row * (K / 8)) * 8;\n float values[8];\n ldg_f32<8>(values, gram + (base + row) * gram_ld + base + col);\n #pragma unroll\n for (int item = 0; item < 8; ++item) {\n const int item_col = col + item;\n lower[row * K + item_col] = item_col < row ? values[item] : 0.0f;\n }\n }\n __syncthreads();\n\n build_t32_inverse_block<K>(lower, tau_b, inverse, mid);\n build_t32_inverse_block<K>(lower + (32 * K + 32), tau_b + 32, inverse + (32 * K + 32), mid);\n __syncthreads();\n\n #pragma unroll\n for (int item = 0; item < 2; ++item) {\n const int elem = item * TB_SIZE + tid;\n const int row = elem / 32;\n const int col = elem - row * 32;\n float accum = 0.0f;\n #pragma unroll\n for (int k = 0; k < 32; ++k) {\n if (k >= col) {\n accum += lower[(row + 32) * K + k] * inverse[k * K + col];\n }\n }\n mid[row * 32 + col] = accum;\n }\n __syncthreads();\n\n #pragma unroll\n for (int item = 0; item < 2; ++item) {\n const int elem = item * TB_SIZE + tid;\n const int row = elem / 32;\n const int col = elem - row * 32;\n float accum = 0.0f;\n #pragma unroll\n for (int k = 0; k < 32; ++k) {\n if (k <= row) {\n accum += inverse[(row + 32) * K + (k + 32)] * mid[k * 32 + col];\n }\n }\n inverse[(row + 32) * K + col] = -accum;\n }\n __syncthreads();\n\n {\n const int pack = tid;\n const int row = pack / (K / 8);\n const int col = (pack - row * (K / 8)) * 8;\n float values[8];\n #pragma unroll\n for (int item = 0; item < 8; ++item) {\n const int item_col = col + item;\n values[item] = item_col <= row ? inverse[row * K + item_col] : 0.0f;\n }\n stg_f32<8>(output + (base + row) * output_ld + base + col, values);\n }\n}\n\nextern "C" __global__\n__launch_bounds__(512, 1)\nvoid qr2_gau_t96(const float* gram, const float* tau, float* output, long long tau_stride) {\n constexpr int K = 64;\n constexpr int N = 96;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const float* gram_b = gram + batch * N * N;\n const float* tau_b = tau + batch * tau_stride;\n float* out_b = output + batch * N * N;\n\n extern __shared__ float storage[];\n float* lower = storage;\n float* inverse = lower + K * K;\n float* mid = inverse + K * K;\n\n build_t64_block(gram_b, tau_b, out_b, lower, inverse, mid, N, N, 0, 0);\n\n if (tid < (K * 32 / 8)) {\n const int pack = tid;\n const int row = pack / (32 / 8);\n const int col = (pack - row * (32 / 8)) * 8;\n float zeros[8] = {};\n stg_f32<8>(out_b + row * N + K + col, zeros);\n }\n __syncthreads();\n\n if (tid < (32 * 32 / 8)) {\n const int pack = tid;\n const int row = pack / (32 / 8);\n const int col = (pack - row * (32 / 8)) * 8;\n float values[8];\n ldg_f32<8>(values, gram_b + (K + row) * N + K + col);\n #pragma unroll\n for (int item = 0; item < 8; ++item) {\n const int item_col = col + item;\n lower[row * K + item_col] = item_col < row ? values[item] : 0.0f;\n }\n }\n __syncthreads();\n\n build_t32_inverse_block<K>(lower, tau_b + K, inverse, mid);\n __syncthreads();\n\n if (tid < (32 * 32 / 8)) {\n const int pack = tid;\n const int row = pack / (32 / 8);\n const int col = (pack - row * (32 / 8)) * 8;\n float values[8];\n #pragma unroll\n for (int item = 0; item < 8; ++item) {\n const int item_col = col + item;\n values[item] = item_col <= row ? inverse[row * K + item_col] : 0.0f;\n }\n stg_f32<8>(out_b + (K + row) * N + K + col, values);\n }\n}\n\n\nextern "C" __global__\n__launch_bounds__(512, 1)\nvoid qr2_gau_t128(const float* gram, const float* tau, float* output, long long tau_stride) {\n constexpr int K = 64;\n constexpr int N = 128;\n const int tid = threadIdx.x;\n const int block = blockIdx.x & 1;\n const int batch = blockIdx.x >> 1;\n const int base = block * K;\n const float* gram_b = gram + batch * N * N;\n const float* tau_b = tau + batch * tau_stride;\n float* out_b = output + batch * N * N;\n\n extern __shared__ float storage[];\n float* lower = storage;\n float* inverse = lower + K * K;\n float* mid = inverse + K * K;\n\n build_t64_block(gram_b, tau_b, out_b, lower, inverse, mid, N, N, base, base);\n\n if (block == 0) {\n const int pack = tid;\n const int row = pack / (K / 8);\n const int col = (pack - row * (K / 8)) * 8;\n float zeros[8] = {};\n stg_f32<8>(out_b + row * N + K + col, zeros);\n }\n}\n\n\n\n'
@memo(maxsize=1)
def _n352_gau_t_image():return _fast_nvrtc_compile(_N352_GAU_T_SOURCE,'qr2_gau_t128')
@memo(maxsize=1)
def _n352_gau_t_kernels():image=_n352_gau_t_image();return CUDAKernel(image,'qr2_gau_t128'),CUDAKernel(image,'qr2_gau_t96')
def _n1024_rhh_direct_source(right:bool):
source=_N1024_GAU_PANEL_SOURCE;template=source.index('template <int ROWS, int COLS, int N, int VEC_SIZE>');begin=source.index(" // emit V when it's not the last square QR",template);end=source.index('\n\n const int col = panel_id',begin)
if not right:body=source[begin:end].replace('v_fp32 += batch * ROWS * COLS;','v_fp32 += (long long)batch * ROWS * 192;').replace('v_fp16 += batch * ROWS * COLS;','v_fp16 += (long long)batch * ROWS * 192;')
else:body=' // Direct right-leaf output into the 192-column parent.\n if constexpr (ROWS > COLS) {\n constexpr int PROWS=ROWS+96;\n v_fp32 += (long long)batch * PROWS * 192;\n v_fp16 += (long long)batch * PROWS * 192;\n constexpr int PANEL_SIZE=VEC_SIZE*ROWS;\n #pragma unroll\n for(int i=0;i<cdiv(PANEL_SIZE,32*4);++i) {\n const int idx=(i*32+lane)*4;\n if(idx<PANEL_SIZE) {\n const int local_col=idx/ROWS;\n const int local_row=idx-local_col*ROWS;\n const int target_col=96+panel_id*VEC_SIZE+local_col;\n const int target=target_col*PROWS+96+local_row;\n const float4 value=*reinterpret_cast<float4*>(\n reflectors+local_panel_id*PANEL_SIZE+idx);\n *reinterpret_cast<float4*>(v_fp32+target)=value;\n reinterpret_cast<half2*>(v_fp16+target)[0]=\n __float22half2_rn({value.x,value.y});\n reinterpret_cast<half2*>(v_fp16+target)[1]=\n __float22half2_rn({value.z,value.w});\n }\n }\n for(int e=lane;e<VEC_SIZE*96;e+=32) {\n const int local_col=e/96;\n const int local_row=e-local_col*96;\n const int target=(96+panel_id*VEC_SIZE+local_col)*PROWS+local_row;\n v_fp32[target]=0.f;\n v_fp16[target]=__float2half_rn(0.f);\n }\n }'
return source[:begin]+body+source[end:]
@memo(maxsize=1)
def _n1024_rhh_left_kernels():source=_fast_only_cuda_kernels(_n1024_rhh_direct_source(False),_N1024_GAU_PANEL_NAMES[:4]);image=_fast_nvrtc_compile(source,'qr2_n1024_rhh_direct_left');return tuple(CUDAKernel(image,name)for name in _N1024_GAU_PANEL_NAMES[:4])
@memo(maxsize=1)
def _n1024_rhh_right_kernels():source=_fast_only_cuda_kernels(_n1024_rhh_direct_source(True),_N1024_GAU_PANEL_NAMES[:4]);image=_fast_nvrtc_compile(source,'qr2_n1024_rhh_direct_right');return tuple(CUDAKernel(image,name)for name in _N1024_GAU_PANEL_NAMES[:4])
_N1024_RHH_T_NAME='qr2_n1024_rhh_t192';_N1024_RHH_T_SOURCE='\nextern "C" __global__ __launch_bounds__(256,1) void\nqr2_n1024_rhh_t192(const float* __restrict__ t0,\n const float* __restrict__ t1,\n const float* __restrict__ bottom,\n float* __restrict__ parent) {\n constexpr int P=96,P2=192;\n const int batch=blockIdx.x,quadrant=blockIdx.y,tid=threadIdx.x;\n t0+=(long long)batch*P*P;\n t1+=(long long)batch*P*P;\n bottom+=(long long)batch*P*P;\n parent+=(long long)batch*P2*P2;\n const int rb=(quadrant>>1)*P,cb=(quadrant&1)*P;\n for(int q=tid;q<P*(P/4);q+=blockDim.x) {\n const int row=q/(P/4),col=(q-row*(P/4))*4;\n float4 value;\n if(quadrant==0)\n value=*reinterpret_cast<const float4*>(t0+row*P+col);\n else if(quadrant==1)\n value=make_float4(0.f,0.f,0.f,0.f);\n else if(quadrant==2) {\n value=*reinterpret_cast<const float4*>(bottom+row*P+col);\n value=make_float4(-value.x,-value.y,-value.z,-value.w);\n } else\n value=*reinterpret_cast<const float4*>(t1+row*P+col);\n *reinterpret_cast<float4*>(parent+(rb+row)*P2+cb+col)=value;\n }\n}\n'
@memo(maxsize=1)
def _n1024_rhh_t_kernel():return CUDAKernel(_fast_nvrtc_compile(_N1024_RHH_T_SOURCE,_N1024_RHH_T_NAME),_N1024_RHH_T_NAME)
def _n1024_rhh_leaf_factor(h,tau,index:int,offset:int,parent_v,parent_h,right:bool):
batch=int(h.shape[0]);rows=N1024-int(offset);panel=h[:,offset:,offset:offset+96];panel_tau=tau[:,offset:offset+96];kernels=_n1024_rhh_right_kernels()if right else _n1024_rhh_left_kernels();kernels[index].launch(grid=(batch*2,1,1),block=(384,1,1),shared_mem=(rows*48+96)*4+96*8,args=[panel,panel,panel_tau,parent_v,parent_h])
if right:v=parent_v[:,96:,96:];vh=parent_h[:,96:,96:]
else:v=parent_v[:,:,:96];vh=parent_h[:,:,:96]
gram=torch.bmm(vh.transpose(1,2),vh,out_dtype=torch.float32);tt=torch.empty_like(gram);t96=_n352_gau_t_kernels()[1];t96.launch(grid=(batch,1,1),block=(512,1,1),shared_mem=36864,args=[gram,panel_tau,tt,int(tau.stride(0))]);mid=torch.bmm(gram[:,64:,:64],tt[:,:64,:64]);torch.baddbmm(tt[:,64:,:64],tt[:,64:,64:],mid,beta=.0,alpha=-1.,out=tt[:,64:,:64]);return v,vh,tt
_ACTIVE352_LEFT_NAME='active352_qr_left96_rows832_parent160';_ACTIVE352_RIGHT_NAME='active352_qr_right64_rows736_parent160';_ACTIVE352_T64_NAME='active352_qr_t64';_ACTIVE352_T160_NAME='active352_qr_assemble_t160'
def _active352_left_source():source=_n1024_rhh_direct_source(False);source=source.replace('qr2_gau_n1024_p2',_ACTIVE352_LEFT_NAME,1);source=source.replace('v_fp32 += (long long)batch * ROWS * 192;','v_fp32 += (long long)batch * ROWS * 160;',1).replace('v_fp16 += (long long)batch * ROWS * 192;','v_fp16 += (long long)batch * ROWS * 160;',1);return source
def _active352_right_source():source=_n1024_rhh_direct_source(True);source=source.replace('qr2_gau_n1024_p3',_ACTIVE352_RIGHT_NAME,1).replace('qr2_gau_panel_body<736, 96, 1024, 4>','qr2_gau_panel_body<736, 64, 1024, 4>',1).replace('v_fp32 += (long long)batch * PROWS * 192;','v_fp32 += (long long)batch * PROWS * 160;',1).replace('v_fp16 += (long long)batch * PROWS * 192;','v_fp16 += (long long)batch * PROWS * 160;',1);return source
_ACTIVE352_T64_WRAPPER=f'''
extern "C" __global__ __launch_bounds__(512, 1)
void {_ACTIVE352_T64_NAME}(
const float* gram,
const float* tau,
float* output,
long long tau_stride) {{
constexpr int N = 64;
const int batch = blockIdx.x;
const float* gram_b = gram + (long long)batch * N * N;
const float* tau_b = tau + batch * tau_stride;
float* out_b = output + (long long)batch * N * N;
extern __shared__ float storage[];
float* lower = storage;
float* inverse = lower + N * N;
float* mid = inverse + N * N;
build_t64_block(
gram_b, tau_b, out_b, lower, inverse, mid, N, N, 0, 0);
}}
''';_ACTIVE352_T160_SOURCE=f'''
#include <cuda_runtime.h>
extern "C" __global__ __launch_bounds__(256, 1)
void {_ACTIVE352_T160_NAME}(
const float* t0,
const float* t1,
const float* bottom,
float* output,
int batch) {{
constexpr int LEFT = 96;
constexpr int RIGHT = 64;
constexpr int N = 160;
const int matrix = blockIdx.x;
if (matrix >= batch) return;
t0 += (long long)matrix * LEFT * LEFT;
t1 += (long long)matrix * RIGHT * RIGHT;
bottom += (long long)matrix * RIGHT * LEFT;
output += (long long)matrix * N * N;
for (int index = threadIdx.x; index < N * N; index += blockDim.x) {{
const int row = index / N;
const int column = index - row * N;
float value = 0.0f;
if (row < LEFT && column < LEFT)
value = t0[row * LEFT + column];
else if (row >= LEFT && column < LEFT)
value = -bottom[(row - LEFT) * LEFT + column];
else if (row >= LEFT && column >= LEFT)
value = t1[(row - LEFT) * RIGHT + column - LEFT];
output[index] = value;
}}
}}
'''
@memo(maxsize=1)
def _active352_left_kernel():return CUDAKernel(_fast_nvrtc_compile(_active352_left_source(),_ACTIVE352_LEFT_NAME),_ACTIVE352_LEFT_NAME)
@memo(maxsize=1)
def _active352_right_kernel():return CUDAKernel(_fast_nvrtc_compile(_active352_right_source(),_ACTIVE352_RIGHT_NAME),_ACTIVE352_RIGHT_NAME)
@memo(maxsize=1)
def _active352_t64_kernel():source=_N352_GAU_T_SOURCE+_ACTIVE352_T64_WRAPPER;return CUDAKernel(_fast_nvrtc_compile(source,_ACTIVE352_T64_NAME),_ACTIVE352_T64_NAME)
@memo(maxsize=1)
def _active352_t160_kernel():return CUDAKernel(_fast_nvrtc_compile(_ACTIVE352_T160_SOURCE,_ACTIVE352_T160_NAME),_ACTIVE352_T160_NAME)
@torch.no_grad()
def _active352_factor_leaf96(h,tau,parent_v,parent_h):batch=h.shape[0];rows=832;panel=h[:,192:,192:288];panel_tau=tau[:,192:288];_active352_left_kernel().launch((batch*2,1,1),(384,1,1),(panel,panel,panel_tau,parent_v,parent_h),shared_mem=(rows*48+96)*4+96*8);v=parent_v[:,:,:96];vh=parent_h[:,:,:96];gram=torch.bmm(vh.mT,vh,out_dtype=torch.float32);triangular=torch.empty_like(gram);_n352_gau_t_kernels()[1].launch((batch,1,1),(512,1,1),(gram,panel_tau,triangular,int(tau.stride(0))),shared_mem=36864);mid=torch.bmm(gram[:,64:,:64],triangular[:,:64,:64]);torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],mid,beta=.0,alpha=-1.,out=triangular[:,64:,:64]);return v,vh,triangular
@torch.no_grad()
def _active352_factor_leaf64(h,tau,parent_v,parent_h):batch=h.shape[0];rows=736;panel=h[:,288:,288:352];panel_tau=tau[:,288:352];_active352_right_kernel().launch((batch*2,1,1),(256,1,1),(panel,panel,panel_tau,parent_v,parent_h),shared_mem=(rows*32+64)*4+64*8);v=parent_v[:,96:,96:160];vh=parent_h[:,96:,96:160];gram=torch.bmm(vh.mT,vh,out_dtype=torch.float32);triangular=torch.empty_like(gram);_active352_t64_kernel().launch((batch,1,1),(512,1,1),(gram,panel_tau,triangular,int(tau.stride(0))),shared_mem=3*64*64*4);return v,vh,triangular
@torch.no_grad()
def _active352_factor_pair160(h,tau):batch=h.shape[0];rows=832;v160=torch.empty((batch,160,rows),device=h.device).transpose(1,2);h160=torch.empty((batch,160,rows),device=h.device,dtype=torch.float16).transpose(1,2);v160.zero_();h160.zero_();v0,h0,t0=_active352_factor_leaf96(h,tau,v160,h160);right_panel=h[:,192:,288:352];transformed=t0@(v0.mT@right_panel);torch.baddbmm(right_panel,h0,transformed.half(),beta=1.,alpha=-1.,out=right_panel,out_dtype=torch.float32);_,h1,t1=_active352_factor_leaf64(h,tau,v160,h160);cross=torch.bmm(h0[:,96:,:].mT,h1,out_dtype=torch.float32);bottom=t1@cross.mT@t0;t160=torch.empty((batch,160,160),device=h.device);_active352_t160_kernel().launch((batch,1,1),(256,1,1),(t0,t1,bottom,t160,batch));return v160,t160
@torch.no_grad()
def _active352_completion(matrix:torch.Tensor):
batch=matrix.shape[0];h=torch.empty((batch,1024,1024),device=matrix.device);h[:,:,:352]=matrix;tau=torch.empty((batch,1024),device=matrix.device);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
first=_factor_pair(None,h,tau,0,0,352);second_v,second_t=_active352_factor_pair160(h,tau);result=_e392_batched_identity1024(batch,matrix.device);first_replay=True
for(offset,reflector,triangular)in reversed(((0,first[0],first[1]),(192,second_v,second_t))):active=result[:,offset:,offset:]if offset else result;projection=reflector.mT if first_replay else reflector.mT@active;transformed=triangular.mT@projection;torch.baddbmm(active,reflector,transformed,beta=1.,alpha=-1.,out=active);first_replay=False
gram=result.mT@result;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);result=result@gram;return result
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _normalize_columns_kernel(matrix,rows:tl.constexpr,columns:tl.constexpr,block_rows:tl.constexpr,block_columns:tl.constexpr):program=tl.program_id(0);column_blocks=triton.cdiv(columns,block_columns);batch=program//column_blocks;column_block=program-batch*column_blocks;row=tl.arange(0,block_rows)[:,None];column=column_block*block_columns+tl.arange(0,block_columns)[None,:];mask=(row<rows)&(column<columns);tl_cuda.gdc_wait();offsets=batch*rows*columns+row*columns+column;value=tl.load(matrix+offsets,mask=mask,other=.0).to(tl.float32);inverse_norm=tl.rsqrt(tl.maximum(tl.sum(value*value,axis=0),1e-40));tl.store(matrix+offsets,value*inverse_norm,mask=mask);tl_cuda.gdc_launch_dependents()
@triton.jit
def _e1977_normalize_columns_strided_kernel(matrix,batch_stride,row_stride,rows:tl.constexpr,columns:tl.constexpr,block_rows:tl.constexpr,block_columns:tl.constexpr):program=tl.program_id(0);column_blocks=triton.cdiv(columns,block_columns);batch=program//column_blocks;column_block=program-batch*column_blocks;row=tl.arange(0,block_rows)[:,None];column=column_block*block_columns+tl.arange(0,block_columns)[None,:];mask=(row<rows)&(column<columns);tl_cuda.gdc_wait();offsets=batch*batch_stride+row*row_stride+column;value=tl.load(matrix+offsets,mask=mask,other=.0).to(tl.float32);inverse_norm=tl.rsqrt(tl.maximum(tl.sum(value*value,axis=0),1e-40));tl.store(matrix+offsets,value*inverse_norm,mask=mask);tl_cuda.gdc_launch_dependents()
def _e1977_normalize_columns_strided_(matrix:torch.Tensor):batch,rows,columns=matrix.shape;block_columns=16;_e1977_normalize_columns_strided_kernel[batch*triton.cdiv(columns,block_columns),](matrix,matrix.stride(0),matrix.stride(1),rows=rows,columns=columns,block_rows=triton.next_power_of_2(rows),block_columns=block_columns,num_warps=8,num_stages=1,launch_pdl=True);return matrix
def normalize_columns_(matrix:torch.Tensor,*,use_pdl:bool=True):
if matrix.ndim!=3 or matrix.dtype is not torch.float32 or not matrix.is_cuda:raise ValueError('matrix must be a rank-3 CUDA torch.float32 tensor')
if not matrix.is_contiguous():raise ValueError('matrix must be contiguous')
batch,rows,columns=matrix.shape
if rows>1024:raise ValueError('the prototype supports at most 1024 rows')
launch_pdl=bool(use_pdl and torch.cuda.get_device_capability()[0]>=9);block_columns=16;_normalize_columns_kernel[batch*triton.cdiv(columns,block_columns),](matrix,rows=rows,columns=columns,block_rows=triton.next_power_of_2(rows),block_columns=block_columns,num_warps=8,num_stages=1,launch_pdl=launch_pdl);return matrix
@triton.jit
def _e1611_normalize_copy_strided_kernel(source,output,source_batch_stride,source_row_stride,rows:tl.constexpr,columns:tl.constexpr,block_rows:tl.constexpr,block_columns:tl.constexpr):program=tl.program_id(0);column_blocks=triton.cdiv(columns,block_columns);batch=program//column_blocks;column_block=program-batch*column_blocks;row=tl.arange(0,block_rows)[:,None];column=column_block*block_columns+tl.arange(0,block_columns)[None,:];mask=(row<rows)&(column<columns);tl_cuda.gdc_wait();value=tl.load(source+batch*source_batch_stride+row*source_row_stride+column,mask=mask,other=.0).to(tl.float32);inverse_norm=tl.rsqrt(tl.maximum(tl.sum(value*value,axis=0),1e-40));tl.store(output+batch*rows*columns+row*columns+column,value*inverse_norm,mask=mask);tl_cuda.gdc_launch_dependents()
def _e1611_normalize_copy_strided(matrix:torch.Tensor):batch,rows,columns=matrix.shape;output=torch.empty((batch,rows,columns),device=matrix.device,dtype=matrix.dtype);block_columns=16;_e1611_normalize_copy_strided_kernel[batch*triton.cdiv(columns,block_columns),](matrix,output,matrix.stride(0),matrix.stride(1),rows=rows,columns=columns,block_rows=triton.next_power_of_2(rows),block_columns=block_columns,num_warps=8,num_stages=1,launch_pdl=True);return output
def _rademacher_probes(n:int,count:int,device:torch.device):
if count>n:raise ValueError('Walsh probes require count <= n')
index=torch.arange(n,device=device,dtype=torch.int64)[:,None];masks=torch.arange(count,device=device,dtype=torch.int64)[None,:];parity=index&masks;parity=torch.bitwise_xor(parity,parity>>4);parity=torch.bitwise_xor(parity,parity>>2);parity=torch.bitwise_xor(parity,parity>>1)&1;return torch.where(parity==0,1.,-1.)/n**.5
@torch.no_grad()
def stochastic_lanczos_stats(matrix:torch.Tensor,*,steps:int=32,probes:int=8):
batch,n,_=matrix.shape;current=_rademacher_probes(n,probes,matrix.device)[None].expand(batch,-1,-1).clone();previous=torch.zeros_like(current);beta=torch.zeros((batch,1,probes),device=matrix.device);diagonal=[];off_diagonal=[]
for iteration in range(steps):
product=matrix@current-beta*previous;alpha=(current*product).sum(dim=1,keepdim=True);product-=current*alpha;next_beta=torch.linalg.vector_norm(product,dim=1,keepdim=True).clamp_min_(1e-20);diagonal.append(alpha[:,0])
if iteration+1<steps:off_diagonal.append(next_beta[:,0])
previous,current,beta=current,product/next_beta,next_beta
diagonal_tensor=torch.stack(diagonal,dim=-1);off_diagonal_tensor=torch.stack(off_diagonal,dim=-1);tridiagonal=torch.diag_embed(diagonal_tensor)+torch.diag_embed(off_diagonal_tensor,offset=1)+torch.diag_embed(off_diagonal_tensor,offset=-1);nodes,vectors=torch.linalg.eigh(tridiagonal.reshape(batch*probes,steps,steps));nodes=nodes.reshape(batch,probes,steps);weights=vectors[:,0,:].square().reshape(batch,probes,steps);sorted_nodes,order=nodes.reshape(batch,-1).sort(dim=-1);sorted_weights=weights.reshape(batch,-1).gather(1,order)/probes;median_index=(sorted_weights.cumsum(dim=-1)<.5).sum(dim=-1);median_index.clamp_max_(probes*steps-1);median=sorted_nodes.gather(1,median_index[:,None])[:,0];return median,sorted_nodes[:,0],sorted_nodes[:,-1]
_BLOCK_INVERSE_RANKS=frozenset((128,170,256,342));_RECURSIVE_INVERSE_RANKS=frozenset((170,342));_RECURSIVE_INVERSE_MIN_BATCH=128;_MIN_SPECTRAL_SPLIT_DIM=64;_LAPACK_ROOT_DIM=512;_LAPACK_CHILD_DIM=256;_N512_ACTIVE_COLUMN_COUNTS=frozenset((384,512));_N1024_ACTIVE_COLUMN_COUNTS=frozenset((256,320,384))
@torch.no_grad()
def _owned_right_solve1024(matrix:torch.Tensor,lower:torch.Tensor,*,trsm_fn=None):
_,rows,rank=matrix.shape
if rank!=1024:raise ValueError('owned solve requires rank 1024')
if trsm_fn is None:trsm_fn=_e084_trsm256_tensor_strided
output=torch.empty_like(matrix);panel=256
for index in range(4):
start=index*panel;diagonal=lower[:,start:start+panel,start:start+panel].contiguous()
if index==0:source=matrix;source_columns=1024
else:source=torch.baddbmm(matrix[:,:,start:start+panel],output[:,:,:start],lower[:,start:start+panel,:start].mT,beta=1.,alpha=-1.);source_columns=panel
trsm_fn(source,diagonal,source_rows=rows,source_columns=source_columns,row_offset=0,column_offset=0,rows=rows,destination=output,destination_column_offset=start)
return output
@torch.no_grad()
def block_inverse_triangular_solve(matrix:torch.Tensor,lower:torch.Tensor,*,precision:str='highest',trsm_fn=None):
batch,_,rank=matrix.shape
if rank==1024:return _owned_right_solve1024(matrix,lower,trsm_fn=trsm_fn)
if rank not in _BLOCK_INVERSE_RANKS or rank in _RECURSIVE_INVERSE_RANKS and batch<_RECURSIVE_INVERSE_MIN_BATCH:return torch.linalg.solve_triangular(lower,matrix.mT,upper=False).mT
if rank in _RECURSIVE_INVERSE_RANKS:
def recursive_inverse(factor:torch.Tensor):
local_rank=factor.shape[-1]
if local_rank<=96:local_eye=torch.eye(local_rank,device=factor.device).expand(factor.shape[0],-1,-1);return torch.linalg.solve_triangular(factor.contiguous(),local_eye,upper=False)
split=(local_rank//2+7)//8*8;left=recursive_inverse(factor[:,:split,:split].contiguous());right=recursive_inverse(factor[:,split:,split:].contiguous());inverse=torch.zeros_like(factor);inverse[:,:split,:split]=left;inverse[:,split:,split:]=right;inverse[:,split:,:split]=-(right@factor[:,split:,:split])@left;return inverse
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(precision);inverse=recursive_inverse(lower);result=matrix@inverse.mT;torch.set_float32_matmul_precision(previous_precision);return result
block=64;blocks=rank//block;diagonal=torch.stack([lower[:,start:start+block,start:start+block]for start in range(0,rank,block)],dim=1).reshape(batch*blocks,block,block).contiguous();eye=torch.eye(block,device=matrix.device).expand(batch*blocks,-1,-1);diagonal_inverse=torch.linalg.solve_triangular(diagonal,eye,upper=False).reshape(batch,blocks,block,block);inverse=torch.zeros_like(lower)
for(index,start)in enumerate(range(0,rank,block)):inverse[:,start:start+block,start:start+block]=diagonal_inverse[:,index]
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(precision)
for start in range(0,rank,2*block):left=inverse[:,start:start+block,start:start+block];right=inverse[:,start+block:start+2*block,start+block:start+2*block];coupling=lower[:,start+block:start+2*block,start:start+block];inverse[:,start+block:start+2*block,start:start+block]=-(right@coupling)@left
if rank==256:inverse[:,128:,:128]=-(inverse[:,128:,128:]@lower[:,128:,:128])@inverse[:,:128,:128]
result=matrix@inverse.mT;torch.set_float32_matmul_precision(previous_precision);return result
@torch.no_grad()
def cholesky_orthonormalize(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,inverse_precision:str='highest',gram_precision:str='highest',check_cholesky_errors:bool=True,trsm_fn=None):
result=matrix;batch,_,rank=result.shape
for pass_index in range(passes):
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=result.mT@result;torch.set_float32_matmul_precision(previous_precision);diagonal_scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);gram.diagonal(dim1=-2,dim2=-1).add_(pass_ridge*diagonal_scale[:,None])
if check_cholesky_errors:lower=torch.linalg.cholesky(gram)
else:lower=torch.linalg.cholesky_ex(gram,check_errors=False)[0]
result=block_inverse_triangular_solve(result,lower,precision=inverse_precision,trsm_fn=trsm_fn)
return result
@torch.no_grad()
def _dense2048_minimax4_orthonormalize(matrix:torch.Tensor):
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
gram=matrix.mT@matrix;batch,rank,_=gram.shape;identity=torch.eye(rank,device=matrix.device,dtype=matrix.dtype).expand(batch,-1,-1);alpha=2.05*gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1,keepdim=True);diagonal=(1.33634625/alpha[:,0]).sqrt();inverse_root=torch.diag_embed(diagonal[:,None].expand(-1,rank))
for coefficient in(1.44951506,1.35477029,1.01066908,.34224671):inner=gram@inverse_root;error=torch.baddbmm(identity,inverse_root,inner,beta=1.,alpha=-1.);factor=torch.baddbmm(error,error,error,beta=.5,alpha=coefficient);factor.diagonal(dim1=-2,dim2=-1).add_(1.);inverse_root=inverse_root@factor
return matrix@inverse_root
finally:torch.set_float32_matmul_precision(previous_precision)
@torch.no_grad()
def _asymmetric_split_finish(matrix:torch.Tensor,eye:torch.Tensor,low_projector:torch.Tensor,low:torch.Tensor,*,range_iterations:int,reorthogonalize_every:int,orthogonalize_fn,power_range_iteration:bool,high_iterations:int,direct_complement:bool,normalize_high:bool,complete_qr:bool=False,skip_final_low_cqr:bool=False,minimax_high_checkpoint:bool=False,high_orthogonalize_fn=None,diagonal_children:bool=False,direct_qr_workspace:bool=False):
qr_workspace=None
if power_range_iteration:
operator=low_projector;power=1
while power<reorthogonalize_every:operator=operator@operator;power*=2
for _ in range(range_iterations//reorthogonalize_every):low=operator@low;low=orthogonalize_fn(low,passes=1,ridge=1e-06,final_ridge=1e-08,inverse_precision='highest')
else:
for iteration in range(range_iterations):
direct_workspace=direct_qr_workspace and complete_qr and skip_final_low_cqr and matrix.shape[-1]==1024 and low.shape[-1]==512 and iteration+1==range_iterations
if direct_workspace:qr_workspace=torch.empty((low.shape[0],1024,1024),device=low.device,dtype=low.dtype);next_low=qr_workspace[:,:,:512];torch.bmm(low_projector,low,out=next_low);low=next_low
else:low=low_projector@low
if(iteration+1)%reorthogonalize_every==0:
if skip_final_low_cqr and complete_qr and iteration+1==range_iterations:
if qr_workspace is None:normalize_columns_(low)
else:_e1977_normalize_columns_strided_(low)
else:low=orthogonalize_fn(low,passes=1,ridge=1e-06,final_ridge=1e-08,inverse_precision='highest')
else:normalize_columns_(low)
batch,n,rank=low.shape
if complete_qr:
if n==1024 and rank==512:basis=compact_wy_orthogonalize_1024x384(None,low if qr_workspace is not None else low.contiguous(),complete=True,newton_steps=0,paired_replay='first',prefilled_h=qr_workspace)
elif n==768 and rank==384:basis=_nearrank_complete_qr768(low.contiguous())
elif n==512 and rank==256 and direct_complement:basis=_lapack_active256_complete(low.contiguous())
else:raise ValueError('complete split QR requires 1024x512, 768x384, or 512x256')
product=basis.mT@matrix;children=torch.empty((2*batch,rank,rank),device=matrix.device,dtype=matrix.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=children[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=children[batch:]);return basis,children
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
if direct_complement:high=eye[:,:,rank:]-low@low[:,rank:,:].mT;high=torch.baddbmm(high,low,low.mT@high,beta=1.,alpha=-1.)
elif n>1024 and n==2*rank and high_iterations==1 and not normalize_high:high=eye[:,:,rank:]-low@low[:,rank:,:].mT
else:
high_projector=eye-low_projector;high=eye[:,:,rank:].clone();torch.set_float32_matmul_precision(previous)
for _ in range(high_iterations):
high=high_projector@high
if normalize_high:normalize_columns_(high)
torch.set_float32_matmul_precision('highest');high-=low@(low.mT@high)
torch.set_float32_matmul_precision(previous)
if minimax_high_checkpoint:high=_dense2048_minimax4_orthonormalize(high)
else:high_fn=orthogonalize_fn if high_orthogonalize_fn is None else high_orthogonalize_fn;high=high_fn(high,passes=1,ridge=1e-06,final_ridge=1e-08,inverse_precision='highest')
basis=torch.cat((low,high),dim=-1)
if diagonal_children:product=basis.mT@matrix;children=torch.empty((2*batch,rank,rank),device=matrix.device,dtype=matrix.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=children[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=children[batch:])
else:transformed=basis.mT@matrix@basis;children=torch.cat((transformed[:,:rank,:rank],transformed[:,rank:,rank:]),dim=0).contiguous()
return basis,children
@torch.no_grad()
def _e2248_nonic_sign_step(sign:torch.Tensor,*,slope:float):cubic=105./16.-4.*slope;quintic=6.*slope-189./16.;septic=135./16.-4.*slope;nonic=slope-35./16.;square=torch.bmm(sign,sign);fourth=torch.bmm(square,square);tail=square.mul(septic);tail.add_(fourth,alpha=nonic);polynomial=torch.bmm(fourth,tail);polynomial.add_(square,alpha=cubic);polynomial.add_(fourth,alpha=quintic);polynomial.diagonal(dim1=-2,dim2=-1).add_(slope);return torch.bmm(sign,polynomial)
@torch.no_grad()
def spectral_split_fast_cholesky(matrix:torch.Tensor,*,sign_iterations:int=8,range_iterations:int=6,reorthogonalize_every:int=1,lanczos_steps:int=32,lanczos_probes:int=8,range_cholesky_passes:int=2,fused_sign_update:bool=False,final_cholesky_passes:int=2,skip_final_high_checkpoint:bool=False,center_mode:str='adaptive',center_probe_iterations:int=7,block_inverse_precision:str='highest',center_trace_fraction:float|None=None,power_range_iteration:bool=False,known_lapack_bounds:bool=False,known_positive_logspace_bounds:bool=False,uniform_density_factor:float=1.1,orthogonalize_fn=None,joined_power_range:bool=False,lanczos_stats_fn=None,low_precision_sign:bool=False,asymmetric_high_iterations:int|None=None,asymmetric_direct_complement:bool=False,asymmetric_normalize_high:bool=False,asymmetric_complete_qr:bool=False,asymmetric_skip_final_low_cqr:bool=False,asymmetric_minimax_high_checkpoint:bool=False,asymmetric_high_orthogonalize_fn=None,diagonal_children:bool=False,asymmetric_direct_qr_workspace:bool=False,nonic_sign:bool=False):
batch,n,_=matrix.shape
if n<_MIN_SPECTRAL_SPLIT_DIM or n%2:raise ValueError('matrix size must be even and at least 64')
rank=n//2
if orthogonalize_fn is None:orthogonalize_fn=cholesky_orthonormalize
eye=torch.eye(n,device=matrix.device).expand(batch,-1,-1);trace_center=matrix.diagonal(dim1=-2,dim2=-1).mean(dim=-1)
if known_positive_logspace_bounds:
lanczos_median=trace_center;total_rank=384;ratio=1e1**(1./(total_rank-1))
if n==total_rank:mean=.1*(ratio**total_rank-1.)/((ratio-1.)*total_rank);scale=trace_center/mean;lower=.1*scale;upper=scale
elif n==total_rank//2 and batch%2==0:parent_batch=batch//2;half=total_rank//2;low_min=.1;low_max=.1*ratio**(half-1);high_min=.1*ratio**half;high_max=1.;low_mean=.1*(ratio**half-1.)/((ratio-1.)*half);high_mean=high_min*(ratio**half-1.)/((ratio-1.)*half);low_scale=trace_center[:parent_batch]/low_mean;high_scale=trace_center[parent_batch:]/high_mean;lower=torch.cat((low_min*low_scale,high_min*high_scale));upper=torch.cat((low_max*low_scale,high_max*high_scale))
else:raise ValueError('known positive logspace bounds require n=384 or paired n=192')
elif known_lapack_bounds:
lanczos_median=trace_center
if n==_LAPACK_ROOT_DIM:lower=torch.full_like(trace_center,-1.);upper=torch.full_like(trace_center,1.)
elif n==_LAPACK_CHILD_DIM and batch%2==0:parent_batch=batch//2;lower=torch.cat((torch.full_like(trace_center[:parent_batch],-1.),torch.full_like(trace_center[parent_batch:],-.1)));upper=torch.cat((torch.full_like(trace_center[:parent_batch],.1),torch.full_like(trace_center[parent_batch:],1.)))
else:raise ValueError('known LAPACK bounds require n=512 or paired n=256')
else:stats_fn=stochastic_lanczos_stats if lanczos_stats_fn is None else lanczos_stats_fn;lanczos_median,lower,upper=stats_fn(matrix,steps=lanczos_steps,probes=lanczos_probes)
spectral_scale=torch.maximum(lower.abs(),upper.abs()).clamp_min(1e-20);center=torch.where((lower>=-.01*spectral_scale)|(upper<=.01*spectral_scale),lanczos_median,trace_center)
def effective_rank(candidate:torch.Tensor,iterations:int|None=None):
local_radius=torch.maximum((lower-candidate).abs(),(upper-candidate).abs()).clamp_min(1e-20)*1.2;estimate=_shift_scale_symmetric(matrix,candidate,local_radius)
for _ in range(center_probe_iterations if iterations is None else iterations):
if fused_sign_update:square=estimate@estimate;estimate=torch.baddbmm(estimate,estimate,square,beta=1.5,alpha=-.5)
else:estimate=.5*estimate@(3.*eye-estimate@estimate);estimate=.5*(estimate+estimate.mT)
return .5*(n-estimate.diagonal(dim1=-2,dim2=-1).sum(dim=-1))
if center_mode=='secant':trace_rank=effective_rank(trace_center);lanczos_rank=effective_rank(lanczos_median);denominator=lanczos_rank-trace_rank;safe_denominator=torch.where(denominator.abs()>.5,denominator,torch.where(denominator>=.0,torch.ones_like(denominator),-torch.ones_like(denominator)));center=trace_center+(rank-trace_rank)*(lanczos_median-trace_center)/safe_denominator;center=torch.maximum(lower,torch.minimum(upper,center))
elif center_mode=='zero_linear':zero=torch.zeros_like(trace_center);zero_rank=effective_rank(zero);center=(rank-zero_rank)*(1.1*spectral_scale/rank);center=torch.maximum(lower,torch.minimum(upper,center))
elif center_mode=='uniform_newton':trace_rank=effective_rank(trace_center);center=trace_center+(rank-trace_rank)*(uniform_density_factor*(upper-lower)/float(n));center=torch.maximum(lower,torch.minimum(upper,center))
elif center_mode=='uniform_newton2':trace_rank=effective_rank(trace_center);center=trace_center+(rank-trace_rank)*(uniform_density_factor*(upper-lower)/float(n));center=torch.maximum(lower,torch.minimum(upper,center));corrected_rank=effective_rank(center,iterations=10);center=center+(rank-corrected_rank)*(1.25*(upper-lower)/float(n));center=torch.maximum(lower,torch.minimum(upper,center))
elif center_mode=='trace':center=trace_center
elif center_mode=='trace_fraction':
if center_trace_fraction is None:raise ValueError('trace_fraction requires center_trace_fraction')
center=center_trace_fraction*trace_center
elif center_mode=='lanczos':center=lanczos_median
elif center_mode!='adaptive':raise ValueError("center_mode must be 'adaptive', 'trace', 'lanczos', 'trace_fraction', 'zero_linear', 'uniform_newton', 'uniform_newton2', or 'secant'")
radius=torch.maximum((lower-center).abs(),(upper-center).abs());radius=radius.clamp_min(1e-20)*1.2
if low_precision_sign:
sign=_shift_scale_symmetric_half(matrix,center,radius);prefix_iterations=sign_iterations-4 if nonic_sign and sign_iterations>=4 else sign_iterations
for _ in range(prefix_iterations):square=torch.bmm(sign,sign);sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
if nonic_sign and sign_iterations>=4:sign=_e2248_nonic_sign_step(sign,slope=3.28);square=torch.bmm(sign,sign);sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
sign=sign.to(torch.float32)
else:
sign=_shift_scale_symmetric(matrix,center,radius)
for _ in range(sign_iterations):
if fused_sign_update:square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
else:sign=.5*sign@(3.*eye-sign@sign);sign=.5*(sign+sign.mT)
low_projector=_e485_float_sign_projector(sign)
if asymmetric_high_iterations is not None:return _asymmetric_split_finish(matrix,eye,low_projector,eye[:,:,:rank],range_iterations=range_iterations,reorthogonalize_every=reorthogonalize_every,orthogonalize_fn=orthogonalize_fn,power_range_iteration=power_range_iteration,high_iterations=asymmetric_high_iterations,direct_complement=asymmetric_direct_complement,normalize_high=asymmetric_normalize_high,complete_qr=asymmetric_complete_qr,skip_final_low_cqr=asymmetric_skip_final_low_cqr,minimax_high_checkpoint=asymmetric_minimax_high_checkpoint,high_orthogonalize_fn=asymmetric_high_orthogonalize_fn,diagonal_children=diagonal_children,direct_qr_workspace=asymmetric_direct_qr_workspace)
high_projector=eye-low_projector;low=eye[:,:,:rank];high=eye[:,:,rank:];can_power=power_range_iteration and reorthogonalize_every>1 and range_iterations%reorthogonalize_every==0 and reorthogonalize_every&reorthogonalize_every-1==0
if can_power:
if joined_power_range:
operators=torch.cat((low_projector,high_projector),dim=0);power=1
while power<reorthogonalize_every:operators=operators@operators;power*=2
joined=torch.cat((low,high),dim=0);checkpoints=range_iterations//reorthogonalize_every
for checkpoint in range(checkpoints):
joined=operators@joined
if skip_final_high_checkpoint and checkpoint+1==checkpoints:low=orthogonalize_fn(joined[:batch],passes=range_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision);high=joined[batch:]
else:joined=orthogonalize_fn(joined,passes=range_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision);low,high=joined[:batch],joined[batch:]
else:
low_operator=low_projector;high_operator=high_projector;power=1
while power<reorthogonalize_every:low_operator=low_operator@low_operator;high_operator=high_operator@high_operator;power*=2
checkpoints=range_iterations//reorthogonalize_every
for checkpoint in range(checkpoints):
low=low_operator@low;high=high_operator@high
if skip_final_high_checkpoint and checkpoint+1==checkpoints:low=orthogonalize_fn(low,passes=range_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision)
else:joined=orthogonalize_fn(torch.cat((low,high),dim=0),passes=range_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision);low,high=joined[:batch],joined[batch:]
else:
for iteration in range(range_iterations):
low=low_projector@low;high=high_projector@high
if(iteration+1)%reorthogonalize_every==0:
if skip_final_high_checkpoint and iteration+1==range_iterations:low=orthogonalize_fn(low,passes=range_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision)
else:joined=orthogonalize_fn(torch.cat((low,high),dim=0),passes=range_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision);low,high=joined[:batch],joined[batch:]
else:normalize_columns_(low);normalize_columns_(high)
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');high-=low@(low.mT@high);torch.set_float32_matmul_precision(previous_precision);high=orthogonalize_fn(high,passes=final_cholesky_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision=block_inverse_precision);split_basis=torch.cat((low,high),dim=-1);transformed=split_basis.mT@matrix@split_basis;children=torch.cat((transformed[:,:rank,:rank],transformed[:,rank:,rank:]),dim=0).contiguous();return split_basis,children
@torch.no_grad()
def _recursive_block_backtransform(basis:torch.Tensor,child_vectors:torch.Tensor):batch,n,_=basis.shape;rank=n//2;block_vectors=torch.zeros_like(basis);block_vectors[:,:rank,:rank]=child_vectors[:batch];block_vectors[:,rank:,rank:]=child_vectors[batch:];return basis@block_vectors
@torch.no_grad()
def hybrid_eigh_512(matrix:torch.Tensor,*,reorthogonalize_every:int=0,use_polar:bool=False,use_cholesky:bool=False,use_fast_cholesky:bool=False,split_levels:int=1,polar_center_mode:str='adaptive',polar_range_iterations:int=16,polar_sign_iterations:int=8,polar_center_probe_iterations:int=7,polar_iterations:int=10,polar_final_iterations:int=10,fast_cholesky_reorthogonalize_every:int=1,child_sign_iterations:int|None=None,child_range_iterations:int|None=None,child_reorthogonalize_every:int|None=None,boundary_refine_width:int=0,root_boundary_refine_width:int=0,post_refine_width:int=0,fast_lanczos_steps:int=32,fast_lanczos_probes:int=8,child_lanczos_steps:int|None=None,child_lanczos_probes:int|None=None,range_cholesky_passes:int=2,finalize:bool=True,newton_precision:str='highest',fused_sign_update:bool=False,final_cholesky_passes:int=2,skip_final_high_checkpoint:bool=False,fast_center_mode:str='adaptive',fast_center_probe_iterations:int=7,child_center_mode:str|None=None,child_center_probe_iterations:int|None=None,leaf_eigh=None,block_inverse_precision:str='highest',fast_center_trace_fraction:float|None=None,child_center_trace_fraction:float|None=None,power_range_iteration:bool=False,child_power_range_iteration:bool|None=None,known_lapack_bounds:bool=False,known_positive_logspace_bounds:bool=False,uniform_density_factor:float=1.1):
batch,n,_=matrix.shape;recursive_split=(use_polar or use_fast_cholesky)and split_levels>1
if recursive_split:
if use_fast_cholesky:basis,children=spectral_split_fast_cholesky(matrix,range_iterations=polar_range_iterations,sign_iterations=polar_sign_iterations,reorthogonalize_every=fast_cholesky_reorthogonalize_every,lanczos_steps=fast_lanczos_steps,lanczos_probes=fast_lanczos_probes,range_cholesky_passes=range_cholesky_passes,fused_sign_update=fused_sign_update,final_cholesky_passes=final_cholesky_passes,skip_final_high_checkpoint=skip_final_high_checkpoint,center_mode=fast_center_mode,center_probe_iterations=fast_center_probe_iterations,block_inverse_precision=block_inverse_precision,center_trace_fraction=fast_center_trace_fraction,power_range_iteration=power_range_iteration,known_lapack_bounds=known_lapack_bounds,known_positive_logspace_bounds=known_positive_logspace_bounds,uniform_density_factor=uniform_density_factor)
else:basis,children=spectral_split_polar_512(matrix,center_mode=polar_center_mode,range_iterations=polar_range_iterations,sign_iterations=polar_sign_iterations,center_probe_iterations=polar_center_probe_iterations,polar_iterations=polar_iterations,final_polar_iterations=polar_final_iterations)
child_vectors,child_values=hybrid_eigh_512(children,use_polar=use_polar,use_fast_cholesky=use_fast_cholesky,split_levels=split_levels-1,polar_center_mode=polar_center_mode,polar_range_iterations=polar_range_iterations if child_range_iterations is None else child_range_iterations,polar_sign_iterations=polar_sign_iterations if child_sign_iterations is None else child_sign_iterations,polar_center_probe_iterations=polar_center_probe_iterations,polar_iterations=polar_iterations,polar_final_iterations=polar_final_iterations,fast_cholesky_reorthogonalize_every=fast_cholesky_reorthogonalize_every if child_reorthogonalize_every is None else child_reorthogonalize_every,boundary_refine_width=boundary_refine_width,root_boundary_refine_width=0,post_refine_width=0,fast_lanczos_steps=fast_lanczos_steps if child_lanczos_steps is None else child_lanczos_steps,fast_lanczos_probes=fast_lanczos_probes if child_lanczos_probes is None else child_lanczos_probes,range_cholesky_passes=range_cholesky_passes,finalize=False,newton_precision=newton_precision,fused_sign_update=fused_sign_update,final_cholesky_passes=final_cholesky_passes,skip_final_high_checkpoint=skip_final_high_checkpoint,fast_center_mode=fast_center_mode if child_center_mode is None else child_center_mode,fast_center_probe_iterations=fast_center_probe_iterations if child_center_probe_iterations is None else child_center_probe_iterations,leaf_eigh=leaf_eigh,block_inverse_precision=block_inverse_precision,fast_center_trace_fraction=fast_center_trace_fraction if child_center_trace_fraction is None else child_center_trace_fraction,child_center_trace_fraction=child_center_trace_fraction,power_range_iteration=power_range_iteration if child_power_range_iteration is None else child_power_range_iteration,child_power_range_iteration=child_power_range_iteration,known_lapack_bounds=known_lapack_bounds,known_positive_logspace_bounds=known_positive_logspace_bounds,uniform_density_factor=uniform_density_factor);rank=n//2;backtransform_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(newton_precision);vectors=_recursive_block_backtransform(basis,child_vectors);torch.set_float32_matmul_precision(backtransform_precision);values=torch.cat((child_values[:batch],child_values[batch:]),dim=1)
elif use_fast_cholesky:basis,children=spectral_split_fast_cholesky(matrix,range_iterations=polar_range_iterations,sign_iterations=polar_sign_iterations,reorthogonalize_every=fast_cholesky_reorthogonalize_every,lanczos_steps=fast_lanczos_steps,lanczos_probes=fast_lanczos_probes,range_cholesky_passes=range_cholesky_passes,fused_sign_update=fused_sign_update,final_cholesky_passes=final_cholesky_passes,skip_final_high_checkpoint=skip_final_high_checkpoint,center_mode=fast_center_mode,center_probe_iterations=fast_center_probe_iterations,block_inverse_precision=block_inverse_precision,center_trace_fraction=fast_center_trace_fraction,power_range_iteration=power_range_iteration,known_lapack_bounds=known_lapack_bounds,known_positive_logspace_bounds=known_positive_logspace_bounds,uniform_density_factor=uniform_density_factor)
elif use_cholesky:basis,children=spectral_split_cholesky_512(matrix)
elif use_polar:basis,children=spectral_split_polar_512(matrix,center_mode=polar_center_mode,range_iterations=polar_range_iterations,sign_iterations=polar_sign_iterations,center_probe_iterations=polar_center_probe_iterations,polar_iterations=polar_iterations,final_polar_iterations=polar_final_iterations)
else:basis,children=spectral_split_512(matrix,reorthogonalize_every=reorthogonalize_every)
if not recursive_split:
if leaf_eigh is None:_,child_vectors=torch.linalg.eigh(children)
else:child_vectors,_=leaf_eigh(children)
rank=n//2;backtransform_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(newton_precision);vectors=_recursive_block_backtransform(basis,child_vectors);torch.set_float32_matmul_precision(backtransform_precision)
refine_width=boundary_refine_width if use_fast_cholesky and(not finalize or n<=_LAPACK_CHILD_DIM)else root_boundary_refine_width if use_fast_cholesky else 0
if refine_width:width=min(refine_width,n//2);begin=n//2-width;end=n//2+width;window=vectors[:,:,begin:end];projected=window.mT@matrix@window;_,rotation=torch.linalg.eigh(projected);repaired_window=window@rotation;vectors[:,:,begin:end]=repaired_window
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(newton_precision)
if use_polar or use_fast_cholesky:eye=torch.eye(n,device=matrix.device).expand(batch,-1,-1);vectors=.5*vectors@(3.*eye-vectors.mT@vectors)
if not finalize:torch.set_float32_matmul_precision(previous_precision);return vectors,torch.empty((batch,n),device=matrix.device,dtype=torch.float32)
if post_refine_width:
width=min(post_refine_width,n//8)
for boundary in(n//4,n//2,3*n//4):begin=boundary-width;end=boundary+width;window=vectors[:,:,begin:end];projected=window.mT@matrix@window;_,rotation=torch.linalg.eigh(projected);vectors=vectors.clone();vectors[:,:,begin:end]=window@rotation
vectors=.5*vectors@(3.*eye-vectors.mT@vectors)
values=(vectors*(matrix@vectors)).sum(dim=1);torch.set_float32_matmul_precision(previous_precision);values,order=values.sort(dim=1);vectors=torch.gather(vectors,2,order[:,None,:].expand(-1,n,-1));return vectors,values
@triton.jit
def _clustered_projector_seeds_kernel(matrix,low,high,n:tl.constexpr,rank:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);total=tl.num_programs(0)*block;matrix_index=offsets//(n*n);element=offsets-matrix_index*n*n;row=element//n;column=element-row*n;source_column=(column+rank//4+1)%n;source=matrix+matrix_index*n*n+row*n+source_column;value=tl.load(source);diagonal=row==source_column;is_low=column<rank;low_column=column;high_column=column-rank;low_value=.5*(tl.where(diagonal,1.,.0)-value);high_value=.5*(tl.where(diagonal,1.,.0)+value);tl.store(low+matrix_index*n*rank+row*rank+low_column,low_value,mask=is_low);tl.store(high+matrix_index*n*(n-rank)+row*(n-rank)+high_column,high_value,mask=~is_low)
def clustered_projector_seeds_512(matrix:torch.Tensor,*,rank:int=170):
if matrix.ndim!=3 or matrix.shape[-2:]!=(512,512):raise ValueError('matrix must have shape (batch, 512, 512)')
if matrix.dtype!=torch.float32 or not matrix.is_cuda or not matrix.is_contiguous():raise ValueError('matrix must be contiguous CUDA torch.float32')
if not 0<rank<512:raise ValueError('rank must be between 1 and 511')
batch=matrix.shape[0];low=torch.empty((batch,512,rank),device=matrix.device,dtype=torch.float32);high=torch.empty((batch,512,512-rank),device=matrix.device,dtype=torch.float32);block=256;grid=triton.cdiv(matrix.numel(),block),;_clustered_projector_seeds_kernel[grid](matrix,low,high,n=512,rank=rank,block=block,num_warps=8);return low,high
@triton.jit
def _clustered_low_projector_kernel(matrix,low,total,n:tl.constexpr,rank:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);matrix_index=offsets//(n*rank);element=offsets-matrix_index*n*rank;row=element//rank;column=element-row*rank;source=matrix+matrix_index*n*n+row*n+column;value=tl.load(source,mask=offsets<total,other=.0);diagonal=row==column;projected=.5*(tl.where(diagonal,1.,.0)-value);tl.store(low+offsets,projected,mask=offsets<total)
@triton.jit
def _clustered_low_projector_gram_kernel(matrix,low,gram,total,n:tl.constexpr,rank:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);matrix_index=offsets//(n*rank);element=offsets-matrix_index*n*rank;row=element//rank;column=element-row*rank;source=matrix+matrix_index*n*n+row*n+column;value=tl.load(source,mask=offsets<total,other=.0);projected=.5*(tl.where(row==column,1.,.0)-value);tl.store(low+offsets,projected,mask=offsets<total);gram_offsets=matrix_index*rank*rank+row*rank+column;tl.store(gram+gram_offsets,projected,mask=(offsets<total)&(row<rank))
def clustered_low_projector_512(matrix:torch.Tensor,*,rank:int=170):batch=matrix.shape[0];low=torch.empty((batch,512,rank),device=matrix.device,dtype=torch.float32);total=low.numel();block=256;_clustered_low_projector_kernel[triton.cdiv(total,block),](matrix,low,total,n=512,rank=rank,block=block,num_warps=8);return low
def clustered_low_projector_gram_512(matrix:torch.Tensor,*,rank:int=170):batch=matrix.shape[0];low=torch.empty((batch,512,rank),device=matrix.device,dtype=torch.float32);gram=torch.empty((batch,rank,rank),device=matrix.device,dtype=torch.float32);total=low.numel();block=256;_clustered_low_projector_gram_kernel[triton.cdiv(total,block),](matrix,low,gram,total,n=512,rank=rank,block=block,num_warps=8);return low,gram
@triton.jit
def _clustered_pack_seed_kernel(low,seed,total,n:tl.constexpr,rank:tl.constexpr,width:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);matrix_index=offsets//(n*width);element=offsets-matrix_index*n*width;row=element//width;column=element-row*width;low_offsets=matrix_index*n*rank+row*rank+column;from_low=column<rank;value=tl.load(low+low_offsets,mask=(offsets<total)&from_low,other=.0);value=tl.where(from_low,value,tl.where(row==column,1.,.0));tl.store(seed+offsets,value,mask=offsets<total)
def clustered_pack_qr_seed176(low:torch.Tensor):batch,n,rank=low.shape;width=176;seed=torch.empty((batch,n,width),device=low.device,dtype=low.dtype);total=seed.numel();block=256;_clustered_pack_seed_kernel[triton.cdiv(total,block),](low,seed,total,n=n,rank=rank,width=width,block=block,num_warps=8);return seed
@triton.jit
def _e2546_clustered_projector_finish_seed_kernel(padded,product,total,n:tl.constexpr,rank:tl.constexpr,width:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);mask=offsets<total;matrix_index=offsets//(n*width);element=offsets-matrix_index*n*width;row=element//width;column=element-row*width;source=tl.load(padded+offsets,mask=mask,other=.0);projected=tl.load(product+offsets,mask=mask,other=.0);value=tl.where(column<rank,.5*(source-projected),tl.where(row==column,1.,.0));tl.store(product+offsets,value,mask=mask)
def _e2546_clustered_projector_refine_seed176(data:torch.Tensor,low:torch.Tensor):'Use an aligned width so SM100 handles the clustered projector GEMM.';batch,n,rank=low.shape;width=176;padded=clustered_pack_qr_seed176(low);seed=torch.bmm(data,padded);total=seed.numel();block=256;_e2546_clustered_projector_finish_seed_kernel[triton.cdiv(total,block),](padded,seed,total,n=n,rank=rank,width=width,block=block,num_warps=8,num_stages=1);return seed
@triton.jit
def _e162_batched_identity512_kernel(output,total,N:tl.constexpr,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;element=offsets%(N*N);row=element//N;column=element-row*N;tl.store(output+offsets,tl.where(row==column,1.,.0),mask=mask)
@torch.no_grad()
def _e162_batched_identity512(batch:int,device:torch.device):output=torch.empty((batch,512,512),device=device,dtype=torch.float32);total=output.numel();_e162_batched_identity512_kernel[triton.cdiv(total,1024),](output,total,N=512,BLOCK=1024,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _e392_batched_identity1024(batch:int,device:torch.device):output=torch.empty((batch,1024,1024),device=device,dtype=torch.float32);total=output.numel();_e162_batched_identity512_kernel[triton.cdiv(total,4096),](output,total,N=1024,BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def clustered_eigh_512(matrix:torch.Tensor):batch,n,_=matrix.shape;rank=170;low,high=clustered_projector_seeds_512(matrix,rank=rank);low=cholesky_orthonormalize(low,passes=1,ridge=1e-06,inverse_precision='high');high=cholesky_orthonormalize(high,passes=1,ridge=1e-06,inverse_precision='high');low=torch.baddbmm(low,matrix,low,beta=.5,alpha=-.5);high=torch.baddbmm(high,matrix,high,beta=.5,alpha=.5);low=cholesky_orthonormalize(low,passes=1,ridge=1e-07,inverse_precision='high');high=cholesky_orthonormalize(high,passes=1,ridge=1e-07,inverse_precision='high');vectors=torch.cat((low,high),dim=-1);gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);values=torch.cat((-torch.ones((batch,rank),device=matrix.device),torch.ones((batch,n-rank),device=matrix.device)),dim=1);return vectors,values
@torch.no_grad()
def geometric_participation(data:torch.Tensor,*,probes:int=8):n=data.shape[-1];scale=data.abs().amax(dim=(-2,-1)).clamp_min(1e-30);normalized=data/scale[:,None,None];trace2=normalized.square().sum(dim=(-2,-1));probe=_rademacher_probes(n,probes,data.device);twice=normalized@(normalized@probe);trace4=float(n)*twice.square().sum(dim=1).mean(dim=-1);return trace2.square()/trace4.clamp_min(1e-30)
@torch.no_grad()
def is_lapack_geometric_1024(data:torch.Tensor):
if data.shape[-1]!=1024:return False
participation=geometric_participation(data);return bool(((participation>45.)&(participation<9e1)).all().item())
@torch.no_grad()
def _factor_pair(module,h,tau,index:int,offset:int,active_cols:int):
batch=h.shape[0];rows=1024-offset;v192=torch.empty((batch,192,rows),device=h.device).transpose(1,2);h192=torch.empty((batch,192,rows),device=h.device,dtype=torch.float16).transpose(1,2);leaf_factor=_n1024_rhh_leaf_factor if module is None else module._n1024_rhh_leaf_factor;v0,h0,t0=leaf_factor(h,tau,index,offset,v192,h192,False);right_panel=h[:,offset:,offset+96:offset+192];transformed=t0@(v0.mT@right_panel);torch.baddbmm(right_panel,h0,transformed.half(),beta=1.,alpha=-1.,out=right_panel,out_dtype=torch.float32);v1,h1,t1=leaf_factor(h,tau,index+1,offset+96,v192,h192,True);cross=torch.bmm(h0[:,96:,:].mT,h1,out_dtype=torch.float32);bottom=t1@cross.mT@t0;t192=torch.empty((batch,192,192),device=h.device);t_kernel=_n1024_rhh_t_kernel if module is None else module._n1024_rhh_t_kernel;t_kernel().launch(grid=(batch,4,1),block=(256,1,1),args=[t0,t1,bottom,t192])
if offset+192<active_cols:trailing=h[:,offset:,offset+192:active_cols];transformed=t192@(v192.mT@trailing);torch.baddbmm(trailing,h192,transformed.half(),beta=1.,alpha=-1.,out=trailing,out_dtype=torch.float32)
return v192,t192,((offset,v0,t0),(offset+96,v1,t1))
@memo(maxsize=1)
def _dense1024_qr_tail128_kernel():name=_N1024_GAU_PANEL_NAMES[4];source=_fast_only_cuda_kernel(_N1024_GAU_PANEL_SOURCE,name);image=_fast_nvrtc_compile(source,name);return CUDAKernel(image,name)
@torch.no_grad()
def _dense1024_qr_tail128_factor(h,tau):batch=h.shape[0];offset,width,rows=384,128,640;panel=h[:,offset:,offset:offset+width];panel_tau=tau[:,offset:offset+width];v32=h.new_empty(batch,width,rows).transpose(1,2);v16=h.new_empty(batch,width,rows,dtype=torch.float16).transpose(1,2);_dense1024_qr_tail128_kernel().launch(grid=(batch*2,1,1),block=(256,1,1),shared_mem=(rows*64+width)*4+width*8,args=[panel,panel,panel_tau,v32,v16]);gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32);triangular=torch.empty_like(gram);_n352_gau_t_kernels()[0].launch(grid=(batch*2,1,1),block=(512,1,1),shared_mem=36864,args=[gram,panel_tau,triangular,int(tau.stride(0))]);middle=torch.bmm(gram[:,64:,:64],triangular[:,:64,:64]);torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64]);return v32,v16,triangular
@torch.no_grad()
def compact_wy_orthogonalize_1024x384(module,matrix:torch.Tensor,*,complete:bool=False,newton_steps:int=2,newton_precision:str='highest',paired_replay:bool|str=False,prefilled_h:torch.Tensor|None=None):
batch=matrix.shape[0];active_cols=matrix.shape[-1]
if active_cols not in _N512_ACTIVE_COLUMN_COUNTS:raise ValueError('compact-WY adapter supports 384 or 512 active columns')
h=prefilled_h if prefilled_h is not None else torch.empty((batch,1024,1024),device=matrix.device)
if prefilled_h is None:h[:,:,:active_cols]=matrix
tau=torch.empty((batch,1024),device=matrix.device);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');first=_factor_pair(module,h,tau,0,0,active_cols);second=_factor_pair(module,h,tau,2,192,active_cols);trailing_node=None
if active_cols==512:trailing_node=_dense1024_qr_tail128_factor(h,tau)if module is None else module._n1024_recursive_factor(h,tau,4,384,128)
width=1024 if complete else active_cols
if complete:result=_e392_batched_identity1024(batch,matrix.device)
else:result=torch.eye(1024,device=matrix.device)[:,:width];result=result.expand(batch,-1,-1).clone()
if paired_replay=='first':
nodes=[(0,first[0],first[1])]+list(second[2])
if trailing_node is not None:nodes.append((384,trailing_node[0],trailing_node[2]))
elif paired_replay and trailing_node is None:nodes=[(0,first[0],first[1]),(192,second[0],second[1])]
else:
nodes=list(first[2]+second[2])
if trailing_node is not None:nodes.append((384,trailing_node[0],trailing_node[2]))
first_replay=True
for(offset,reflector,triangular)in reversed(nodes):
if complete and offset:active=result[:,offset:,offset:];projection=reflector.mT if first_replay else reflector.mT@active
else:active=result[:,offset:,:];projection=reflector.mT@active
transformed=triangular.mT@projection;torch.baddbmm(active,reflector,transformed,beta=1.,alpha=-1.,out=active);first_replay=False
torch.set_float32_matmul_precision(newton_precision)
for _ in range(newton_steps):gram=result.mT@result;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);result=result@gram
torch.set_float32_matmul_precision(previous);return result
_E1186_QR768_LEFT_NAMES=tuple(f"e1186_qr768_left_p{i}"for i in range(4));_E1186_QR768_RIGHT_NAMES=tuple(f"e1186_qr768_right_p{i}"for i in range(4))
def _e1186_qr768_panel_source(right:bool):
source=_n1024_rhh_direct_source(right);names=_E1186_QR768_RIGHT_NAMES if right else _E1186_QR768_LEFT_NAMES
for(index,(name,rows))in enumerate(zip(names,(768,672,576,480))):
old_name=_N1024_GAU_PANEL_NAMES[index];old_body=f"qr2_gau_panel_body<{1024-index*96}, 96, 1024, 4>";new_body=f"qr2_gau_panel_body<{rows}, 96, 768, 4>"
if source.count(old_name)!=1 or source.count(old_body)!=1:raise RuntimeError('native n768 panel template changed')
source=source.replace(old_name,name,1).replace(old_body,new_body,1)
return source
@memo(maxsize=1)
def _e1186_qr768_left_kernels():source=_fast_only_cuda_kernels(_e1186_qr768_panel_source(False),_E1186_QR768_LEFT_NAMES);image=_fast_nvrtc_compile(source,_E1186_QR768_LEFT_NAMES[0]);return tuple(CUDAKernel(image,name)for name in _E1186_QR768_LEFT_NAMES)
@memo(maxsize=1)
def _e1186_qr768_right_kernels():source=_fast_only_cuda_kernels(_e1186_qr768_panel_source(True),_E1186_QR768_RIGHT_NAMES);image=_fast_nvrtc_compile(source,_E1186_QR768_RIGHT_NAMES[0]);return tuple(CUDAKernel(image,name)for name in _E1186_QR768_RIGHT_NAMES)
@torch.no_grad()
def _e1186_qr768_leaf_factor(h,tau,index:int,offset:int,parent_v,parent_h,right:bool):
batch=int(h.shape[0]);rows=768-offset;panel=h[:,offset:,offset:offset+96];panel_tau=tau[:,offset:offset+96];kernels=_e1186_qr768_right_kernels()if right else _e1186_qr768_left_kernels();kernels[index].launch(grid=(batch*2,1,1),block=(384,1,1),shared_mem=(rows*48+96)*4+96*8,args=[panel,panel,panel_tau,parent_v,parent_h])
if right:v=parent_v[:,96:,96:];vh=parent_h[:,96:,96:]
else:v=parent_v[:,:,:96];vh=parent_h[:,:,:96]
gram=torch.bmm(vh.transpose(1,2),vh,out_dtype=torch.float32);triangular=torch.empty_like(gram);_n352_gau_t_kernels()[1].launch(grid=(batch,1,1),block=(512,1,1),shared_mem=36864,args=[gram,panel_tau,triangular,int(tau.stride(0))]);middle=torch.bmm(gram[:,64:,:64],triangular[:,:64,:64]);torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64]);return v,vh,triangular
@torch.no_grad()
def _e1186_qr768_factor_pair(h,tau,index:int,offset:int):
batch=h.shape[0];rows=768-offset;v192=torch.empty((batch,192,rows),device=h.device).transpose(1,2);h192=torch.empty((batch,192,rows),device=h.device,dtype=torch.float16).transpose(1,2);v0,h0,t0=_e1186_qr768_leaf_factor(h,tau,index,offset,v192,h192,False);right_panel=h[:,offset:,offset+96:offset+192];transformed=t0@(v0.mT@right_panel);torch.baddbmm(right_panel,h0,transformed.half(),beta=1.,alpha=-1.,out=right_panel,out_dtype=torch.float32);_,h1,t1=_e1186_qr768_leaf_factor(h,tau,index+1,offset+96,v192,h192,True);cross=torch.bmm(h0[:,96:,:].mT,h1,out_dtype=torch.float32);bottom=t1@cross.mT@t0;t192=torch.empty((batch,192,192),device=h.device);_n1024_rhh_t_kernel().launch(grid=(batch,4,1),block=(256,1,1),args=[t0,t1,bottom,t192])
if offset+192<384:trailing=h[:,offset:,offset+192:384];transformed=t192@(v192.mT@trailing);torch.baddbmm(trailing,h192,transformed.half(),beta=1.,alpha=-1.,out=trailing,out_dtype=torch.float32)
return v192,t192
@torch.no_grad()
def _nearrank_complete_qr768(low:torch.Tensor):
batch=low.shape[0];h=torch.empty((batch,768,768),device=low.device);h[:,:,:384]=low;tau=torch.empty((batch,768),device=low.device);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
first=_e1186_qr768_factor_pair(h,tau,0,0);second=_e1186_qr768_factor_pair(h,tau,2,192);result=torch.empty((batch,768,768),device=low.device,dtype=low.dtype);total=result.numel();_e162_batched_identity512_kernel[triton.cdiv(total,1024),](result,total,N=768,BLOCK=1024,num_warps=8,num_stages=1);first_replay=True
for(offset,reflector,triangular)in reversed(((0,first[0],first[1]),(192,second[0],second[1]))):
if offset:active=result[:,offset:,offset:];projection=reflector.mT if first_replay else reflector.mT@active
else:active=result;projection=reflector.mT@active
transformed=triangular.mT@projection;torch.baddbmm(active,reflector,transformed,beta=1.,alpha=-1.,out=active);first_replay=False
return result
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _geometric_merge_values_kernel(dominant_values,counts,output_values,rank:tl.constexpr,tail:tl.constexpr,n:tl.constexpr,block:tl.constexpr):matrix=tl.program_id(0);dominant_columns=tl.arange(0,block);dominant=tl.load(dominant_values+matrix*rank+dominant_columns,mask=dominant_columns<rank,other=.0);negative=tl.sum(((dominant_columns<rank)&(dominant<.0)).to(tl.int32),axis=0);tl.store(counts+matrix,negative);columns=tl.arange(0,n);from_low=columns<negative;from_high=columns>=negative+tail;dominant_column=tl.where(from_low,columns,columns-tail);values=tl.load(dominant_values+matrix*rank+dominant_column,mask=from_low|from_high,other=.0);tl.store(output_values+matrix*n+columns,values)
@triton.jit
def _geometric_merge_vectors_kernel(complement,dominant,counts,output,total,complement_matrix_stride,complement_row_stride,dominant_matrix_stride,dominant_row_stride,rank:tl.constexpr,tail:tl.constexpr,n:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);matrix=offsets//(n*n);element=offsets-matrix*n*n;row=element//n;column=element-row*n;negative=tl.load(counts+matrix,mask=offsets<total,other=0);from_low=column<negative;from_high=column>=negative+tail;from_dominant=from_low|from_high;dominant_column=tl.where(from_low,column,column-tail);complement_column=column-negative;dominant_value=tl.load(dominant+matrix*dominant_matrix_stride+row*dominant_row_stride+dominant_column,mask=(offsets<total)&from_dominant,other=.0);complement_value=tl.load(complement+matrix*complement_matrix_stride+row*complement_row_stride+complement_column,mask=(offsets<total)&~from_dominant,other=.0);value=tl.where(from_dominant,dominant_value,complement_value);tl.store(output+offsets,value,mask=offsets<total)
def _merge_geometric352_output(complement:torch.Tensor,dominant:torch.Tensor,dominant_values:torch.Tensor):batch,n,tail=complement.shape;rank=dominant.shape[-1];counts=torch.empty((batch,),device=dominant.device,dtype=torch.int32);values=torch.empty((batch,n),device=dominant.device,dtype=torch.float32);_geometric_merge_values_kernel[batch,](dominant_values,counts,values,rank=rank,tail=tail,n=n,block=512,num_warps=8);vectors=torch.empty((batch,n,n),device=dominant.device,dtype=torch.float32);total=vectors.numel();block=4096;_geometric_merge_vectors_kernel[triton.cdiv(total,block),](complement,dominant,counts,vectors,total,complement.stride(0),complement.stride(1),dominant.stride(0),dominant.stride(1),rank=rank,tail=tail,n=n,block=block,num_warps=8);return vectors,values
@torch.no_grad()
def _e209_geometric_magnitude_split_n352(data:torch.Tensor):
batch,n,_=data.shape;rank=n//2;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:eye=torch.eye(n,device=data.device).expand(batch,-1,-1);squared=data@data;ratio=(2.*torch.finfo(torch.float32).eps)**(1./1023.);squared_ratio=ratio*ratio;geometric_energy=(1.-squared_ratio**n)/(1.-squared_ratio);leading=data.square().sum(dim=(-2,-1)).sqrt()/geometric_energy**.5;threshold=leading*ratio**(rank-.5);center=threshold.square();upper=(1.05*leading).square();radius=torch.maximum(center,upper-center).clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(squared,center,radius);sign=_e202_n352_symmetric_sign(sign,15,growth_alpha=1.9,growth_steps=13);low_projector=_e483_half_sign_projector(sign);operator=_e202_n352_symmetric_square(low_projector);operator=_e202_n352_symmetric_square(operator);leverage=operator.diagonal(dim1=-2,dim2=-1);indices=leverage.topk(rank,dim=-1).indices;low=operator.gather(2,indices[:,None,:].expand(-1,n,-1)).contiguous();high_operator=eye-operator;high_indices=high_operator.diagonal(dim1=-2,dim2=-1).topk(rank,dim=-1).indices;high=high_operator.gather(2,high_indices[:,None,:].expand(-1,n,-1)).contiguous();paired=torch.cat((low,high),dim=0);paired=_e202_n352_cqr176(paired,'high',trsm_fn=_e2373_tcgen_trsm176);low,high=paired[:batch],paired[batch:];low=operator@low;low=_e202_n352_cqr176(low,'high',trsm_fn=_e2373_tcgen_trsm176);torch.set_float32_matmul_precision('highest');high=torch.baddbmm(high,low,low.mT@high,beta=1.,alpha=-1.);high=torch.baddbmm(high,low,low.mT@high,beta=1.,alpha=-1.);torch.set_float32_matmul_precision('high');high=_e202_n352_cqr176(high,trsm_fn=_e2373_tcgen_trsm176);basis=torch.cat((low,high),dim=-1);product=basis.mT@data;children=torch.empty((2*batch,rank,rank),device=data.device,dtype=data.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=children[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=children[batch:]);child_vectors,child_values=_e202_n352_child176_eigh(children);vectors=torch.empty_like(data);torch.bmm(basis[:,:,:rank],child_vectors[:batch],out=vectors[:,:,:rank]);torch.bmm(basis[:,:,rank:],child_vectors[batch:],out=vectors[:,:,rank:]);values=torch.cat((child_values[:batch],child_values[batch:]),dim=-1);values,order=values.sort(dim=-1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1));torch.set_float32_matmul_precision('highest');gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _e475_pack_geometric_gram_blocks_kernel(gram,packed,elements:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);block_id=tl.program_id(1);linear=tl.program_id(2)*BLOCK+tl.arange(0,BLOCK);mask=linear<elements;row=linear//176;column=linear-row*176;source_row=row+tl.where(block_id==0,0,176);source_column=column+tl.where(block_id==2,176,0);value=tl.load(gram+batch*352*352+source_row*352+source_column,mask=mask,other=.0);tl.store(packed+(block_id*tl.num_programs(0)+batch)*elements+linear,value,mask=mask)
@triton.jit
def _e475_assemble_geometric_lower_kernel(lower00,lower10,lower11,lower,elements:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);block_id=tl.program_id(1);linear=tl.program_id(2)*BLOCK+tl.arange(0,BLOCK);mask=linear<elements;row=linear//176;column=linear-row*176;source=tl.where(block_id==0,tl.load(lower00+batch*elements+linear,mask=mask,other=.0),tl.where(block_id==1,tl.load(lower10+batch*elements+linear,mask=mask,other=.0),tl.load(lower11+batch*elements+linear,mask=mask,other=.0)));output_row=row+tl.where(block_id==0,0,176);output_column=column+tl.where(block_id==2,176,0);tl.store(lower+batch*352*352+output_row*352+output_column,source,mask=mask)
@torch.no_grad()
def _e470_geometric_cqr352(matrix:torch.Tensor):
batch,_,rank=matrix.shape
if rank!=352:raise ValueError('geometric CQR requires rank 352')
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:gram=matrix.mT@matrix;diagonal=gram.diagonal(dim1=-2,dim2=-1);diagonal_scale=diagonal.mean(dim=-1);diagonal.add_(3e-06*diagonal_scale[:,None]);elements=176*176;packed=torch.empty((3,batch,176,176),device=matrix.device,dtype=matrix.dtype);_e475_pack_geometric_gram_blocks_kernel[batch,3,triton.cdiv(elements,4096)](gram,packed,elements=elements,BLOCK=4096,num_warps=8,num_stages=1);gram00,gram10,gram11=packed[0],packed[1],packed[2];lower00=torch.empty_like(gram00);_e202_n352_potrf176_kernel().launch((batch,1,1),(640,1,1),(gram00,lower00,batch,.0),shared_mem=(176*176+1)*4);lower10=torch.empty_like(gram10);_n176_choleskyqr_kernel('right_trsm176_block16_rows16').launch((batch,11,1),(256,1,1),(gram10,lower00,lower10,batch,176),shared_mem=176*16*4);schur=torch.baddbmm(gram11,lower10,lower10.mT,beta=1.,alpha=-1.);lower11=torch.empty_like(schur);_e202_n352_potrf176_kernel().launch((batch,1,1),(640,1,1),(schur,lower11,batch,.0),shared_mem=(176*176+1)*4);_e475_assemble_geometric_lower_kernel[batch,3,triton.cdiv(elements,4096)](lower00,lower10,lower11,gram,elements=elements,BLOCK=4096,num_warps=8,num_stages=1);return torch.linalg.solve_triangular(gram,matrix.mT,upper=False).mT
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _e1919_projected352_upper_kernel(left,basis,output,pair_rows,pair_cols,basis_batch_stride,basis_row_stride,batch:tl.constexpr):
matrix=tl.program_id(0);pair=tl.program_id(1);tile_m=tl.load(pair_rows+pair);tile_n=tl.load(pair_cols+pair);rows=tile_m*32+tl.arange(0,32);cols=tile_n*32+tl.arange(0,32);accumulator=tl.zeros((32,32),tl.float32);reverse=tl.zeros((32,32),tl.float32)
for start in tl.range(0,1024,64,num_stages=3):
inner=start+tl.arange(0,64);lhs=tl.load(left+matrix*352*1024+rows[:,None]*1024+inner[None,:],mask=rows[:,None]<352,other=.0);rhs=tl.load(basis+matrix*basis_batch_stride+inner[:,None]*basis_row_stride+cols[None,:],mask=cols[None,:]<352,other=.0);accumulator+=tl.dot(lhs,rhs,input_precision='tf32')
if(tile_m//2==tile_n//2)&(tile_m!=tile_n):reverse_lhs=tl.load(left+matrix*352*1024+cols[:,None]*1024+inner[None,:],mask=cols[:,None]<352,other=.0);reverse_rhs=tl.load(basis+matrix*basis_batch_stride+inner[:,None]*basis_row_stride+rows[None,:],mask=rows[None,:]<352,other=.0);reverse+=tl.dot(reverse_lhs,reverse_rhs,input_precision='tf32')
if(tile_m//2==tile_n//2)&(tile_m!=tile_n):accumulator=.5*(accumulator+tl.trans(reverse))
if tile_m==tile_n:accumulator=.5*(accumulator+tl.trans(accumulator))
valid=(rows[:,None]<352)&(cols[None,:]<352);base=matrix*352*352;tl.store(output+base+rows[:,None]*352+cols[None,:],accumulator,mask=valid);tl.store(output+base+cols[:,None]*352+rows[None,:],tl.trans(accumulator),mask=(tile_m!=tile_n)&tl.trans(valid))
def _e1919_projected352(data:torch.Tensor,basis:torch.Tensor):
left=basis.mT@data;device=basis.device.index
if device is None:device=torch.cuda.current_device()
pair_rows,pair_cols=_e549_upper_tile_pairs(device,11);output=torch.empty((basis.shape[0],352,352),device=basis.device,dtype=basis.dtype);_e1919_projected352_upper_kernel[basis.shape[0],pair_rows.numel()](left,basis,output,pair_rows,pair_cols,basis.stride(0),basis.stride(1),batch=basis.shape[0],num_warps=4,num_stages=1);return output
@triton.jit
def _e1930_geometric_probe_vectors(rotation,values,compressed,N:tl.constexpr,BLOCK:tl.constexpr,BK:tl.constexpr):
matrix=tl.program_id(0);rows=tl.program_id(1)*BLOCK+tl.arange(0,BLOCK);features=tl.arange(0,16);accumulator=tl.zeros((BLOCK,16),tl.float32)
for start in tl.range(0,N,BK,num_stages=3):inner=start+tl.arange(0,BK);q=tl.load(rotation+matrix*N*N+rows[:,None]*N+inner[None,:],mask=(rows[:,None]<N)&(inner[None,:]<N),other=.0);eigen=tl.load(values+matrix*N+inner,mask=inner<N,other=.0);rhs=tl.where(features[None,:]==0,.05330017908890261,tl.where(features[None,:]==1,eigen[:,None]*.05330017908890261,.0));accumulator+=tl.dot(q,rhs,input_precision='tf32')
first=tl.sum(accumulator*(features[None,:]==0),axis=1);second=tl.sum(accumulator*(features[None,:]==1),axis=1);tl.store(compressed+matrix*2*N+rows,first,mask=rows<N);tl.store(compressed+matrix*2*N+N+rows,second,mask=rows<N)
@triton.jit
def _e1930_geometric_probe_residual(projected,compressed,residual,N:tl.constexpr,BLOCK:tl.constexpr,BK:tl.constexpr):
matrix=tl.program_id(0);rows=tl.program_id(1)*BLOCK+tl.arange(0,BLOCK);features=tl.arange(0,16);accumulator=tl.zeros((BLOCK,16),tl.float32)
for start in tl.range(0,N,BK,num_stages=3):inner=start+tl.arange(0,BK);tile=tl.load(projected+matrix*N*N+rows[:,None]*N+inner[None,:],mask=(rows[:,None]<N)&(inner[None,:]<N),other=.0);probe=tl.load(compressed+matrix*2*N+inner,mask=inner<N,other=.0);rhs=tl.where(features[None,:]==0,probe[:,None],.0);accumulator+=tl.dot(tile,rhs,input_precision='tf32')
product=tl.sum(accumulator*(features[None,:]==0),axis=1);qlambda=tl.load(compressed+matrix*2*N+N+rows,mask=rows<N,other=.0);tl.store(residual+matrix*N+rows,tl.abs(product-qlambda),mask=rows<N)
@triton.jit
def _e1930_geometric_probe_finish(residual,values,N:tl.constexpr,BLOCK:tl.constexpr):
matrix=tl.program_id(0);offsets=tl.arange(0,BLOCK);score=tl.sum(tl.load(residual+matrix*N+offsets,mask=offsets<N,other=.0),axis=0);low=tl.abs(tl.load(values+matrix*N));high=tl.abs(tl.load(values+matrix*N+N-1))
if score>.006*tl.maximum(low,high):tl.store(values+matrix*N,tl.load(values+matrix*N+N-1))
def _e1930_geometric_probe_poison(projected:torch.Tensor,rotation:torch.Tensor,values:torch.Tensor):batch=projected.shape[0];compressed=torch.empty((batch,2,352),device=projected.device);residual=torch.empty((batch,352),device=projected.device);_e1930_geometric_probe_vectors[batch,6](rotation,values,compressed,N=352,BLOCK=64,BK=64,num_warps=4,num_stages=1);_e1930_geometric_probe_residual[batch,6](projected,compressed,residual,N=352,BLOCK=64,BK=64,num_warps=4,num_stages=1);_e1930_geometric_probe_finish[batch,](residual,values,N=352,BLOCK=512,num_warps=8,num_stages=1)
@torch.no_grad()
def _lowrank_geometric352_eigh(data:torch.Tensor):
batch,n,_=data.shape;rank=352;basis=data[:,:,:rank];previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:torch.set_float32_matmul_precision('high');basis=data@basis;torch.set_float32_matmul_precision('highest');basis=_e470_geometric_cqr352(basis);torch.set_float32_matmul_precision('high');basis=data@(data@basis);torch.set_float32_matmul_precision('highest');full_basis=_active352_completion(basis.contiguous());basis=full_basis[:,:,:rank];complement=full_basis[:,:,rank:];projected=_e1919_projected352(data,basis);rotation,dominant_values=_e209_geometric_magnitude_split_n352(projected);_e1930_geometric_probe_poison(projected,rotation,dominant_values);dominant=basis@rotation;return _merge_geometric352_output(complement,dominant,dominant_values)
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _nearrank_cqr192_inverse(matrix:torch.Tensor,*,ridge:float=1e-06,gram_precision:str='highest'):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=torch.empty_like(gram);_rankdef_potrf192_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=(_RANKDEF_N192_TRI+1)*4);inverse_transpose=_rankdef_inverse_lt192(lower);torch.set_float32_matmul_precision('high')
try:return matrix@inverse_transpose
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _nearrank_recursive_orthogonalize(matrix:torch.Tensor,*,passes:int=1,ridge:float=1e-06,final_ridge:float=1e-08,inverse_precision:str='highest'):
rank=matrix.shape[-1]
if rank==192:return _nearrank_cqr192_inverse(matrix,ridge=ridge,gram_precision=inverse_precision)
if rank==96:return _rankdef_cqr96(matrix,ridge=ridge,gram_precision=inverse_precision)
return cholesky_orthonormalize(matrix,passes=passes,ridge=ridge,final_ridge=final_ridge,inverse_precision=inverse_precision,check_cholesky_errors=False)
@torch.no_grad()
def _nearrank_certified_cholesky(matrix:torch.Tensor,**kwargs):return cholesky_orthonormalize(matrix,check_cholesky_errors=False,**kwargs)
_E842_REDUCE192_NAME='e1009_reduce192_block2_t384';_E842_SOLVE192_NAME='e842_solve192_step23_parallel_mgs';_E1771_SOLVE192_PDL_NAME='e1771_solve192_s23_pdl_wy';_E842_N192=192;_E842_TRI192=_E842_N192*(_E842_N192+1)//2;_E842_REDUCE192_SHARED=(_E842_TRI192+2*2*_E842_N192+_E842_N192+32+2*2)*4;_E842_SOLVE192_SHARED=(_E842_TRI192+5*_E842_N192+16)*4
@memo(maxsize=1)
def _e842_reduce192_kernel():
source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=192,TRI=N*(N+1)/2,B=2;').replace(_E185_N176_BLOCK4_REDUCE_NAME,_E842_REDUCE192_NAME).replace('__launch_bounds__(768,1)','__launch_bounds__(384,1)').replace('total=tid<24?scratch[tid]:0.f','total=tid<12?scratch[tid]:0.f').replace('constexpr int G=4;','constexpr int G=2;').replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=16,ROWS=24;');entry=' int mid=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;'
if source.count(entry)!=1:raise RuntimeError('n192 reducer entry template changed')
source=source.replace(entry,entry+'\n if(tid==0) asm volatile("griddepcontrol.launch_dependents;":::);',1);tail=' timers[mid*5+1]=clock64();\n }\n}';tail_ready=' timers[mid*5+1]=clock64();\n }\n __threadfence();\n __syncthreads();\n if(tid==0) timers[mid*5+4]=1ULL;\n}'
if source.count(tail)!=1:raise RuntimeError('n192 reducer tail template changed')
source=source.replace(tail,tail_ready,1);return CUDAKernel(_fast_nvrtc_compile(source,_E842_REDUCE192_NAME),_E842_REDUCE192_NAME)
def _e842_solve192_source():
source=_e107_solve160_source().replace('constexpr int N = 160;','constexpr int N = 192;').replace('constexpr int HALF = 80;','constexpr int HALF = 96;').replace(_E107_SOLVE160_NAME,_E842_SOLVE192_NAME).replace('__launch_bounds__(384, 4)','__launch_bounds__(256, 1)')
if source.count('step < 24')!=1:raise RuntimeError('n192 solve step anchor changed')
source=source.replace('step < 24','step < 23',1);source=_parallelize_n176_adjacent_mgs(source);anchor=' const int tid = threadIdx.x;\n const int first_row = rank * HALF;';wait=' const int tid = threadIdx.x;\n if (rank == 0 && tid == 0) {\n volatile unsigned long long* ready = timers + matrix_id * 5 + 4;\n while (*ready == 0ULL) __nanosleep(64);\n }\n cluster.sync();\n const int first_row = rank * HALF;'
if source.count(anchor)!=1:raise RuntimeError('n192 solve ready template changed')
return source.replace(anchor,wait,1)
def _e2023_direct_solve_cluster_source(source:str):
helper='\ntemplate <typename T>\n__device__ __forceinline__ T* e2023_map_shared_rank(T* pointer, int rank) {\n unsigned long long remote;\n asm volatile("mapa.u64 %0, %1, %2;"\n : "=l"(remote)\n : "l"((unsigned long long)pointer), "r"(rank));\n return reinterpret_cast<T*>(remote);\n}\n';rewrites=('#include <cooperative_groups.h>\n',helper),('namespace cg = cooperative_groups;\n',''),(' cg::cluster_group cluster = cg::this_cluster();\n',''),(' const int rank = cluster.block_rank();',' const int rank = blockIdx.x & 1;'),('cluster.map_shared_rank(','e2023_map_shared_rank('),('cluster.sync();','asm volatile("barrier.cluster.arrive.aligned;" ::: "memory");\n asm volatile("barrier.cluster.wait.aligned;" ::: "memory");');expected=1,1,1,1,3,6
for((old,new),count)in zip(rewrites,expected):
if source.count(old)!=count:raise RuntimeError('E2023 direct solve cluster PTX anchor changed')
source=source.replace(old,new)
return source
@memo(maxsize=None)
def _e1771_solve192_pdl_kernel(coarse:bool=False):
name=_E1771_SOLVE192_PDL_NAME+('_coarse128'if coarse else'');source=_e842_solve192_source().replace(_E842_SOLVE192_NAME,name,1);anchor=' cluster.sync();\n const int first_row = rank * HALF;';replacement=' cluster.sync();\n if (tid == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);\n const int first_row = rank * HALF;'
if source.count(anchor)!=1:raise RuntimeError('n192 PDL post-ready template changed')
source=source.replace(anchor,replacement,1);source=_e2023_direct_solve_cluster_source(source)
if coarse:source=_e2337_cluster_coarse_sturm(source,steps=23,probes=128,fine_steps=16,levels=8)
return CUDAKernel(_fast_nvrtc_compile(source,name),name)
@torch.no_grad()
def _e842_direct_n192_eigh(matrix:torch.Tensor,*,use_tf32_replay:bool=False,coarse_solver:bool=False):batch=matrix.shape[0];saved=torch.empty_like(matrix);diagonal=torch.empty(matrix.shape[:-1],device=matrix.device);off_diagonal=torch.empty_like(diagonal);vectors=torch.empty_like(matrix);values=torch.empty_like(diagonal);timers=torch.zeros((batch,5),device=matrix.device,dtype=torch.int64);_e842_reduce192_kernel().launch((batch,1,1),(384,1,1),(matrix,saved,diagonal,off_diagonal,timers),shared_mem=_E842_REDUCE192_SHARED);_e1771_solve192_pdl_kernel(coarse_solver).launch_pdl((batch*2,1,1),(256,1,1),(saved,diagonal,off_diagonal,vectors,values,timers),shared_mem=_E842_SOLVE192_SHARED);panel_total=(_E842_N192-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=matrix.device,dtype=matrix.dtype);_householder_compact_t16_kernel[batch,panel_total](saved,triangular,_E842_N192,panel_total,panel_width=32,block_rows=32,num_warps=4,num_stages=2,launch_pdl=True);return _e1762_tensor_wy_apply_saved_(vectors,saved,triangular,use_tf32=use_tf32_replay),values
@torch.no_grad()
def nearrank_positive768_candidate(matrix:torch.Tensor,*,split_fn=None,low_precision_sign:bool=False):
if split_fn is None:split_fn=spectral_split_fast_cholesky
batch,n,_=matrix.shape;rank=n//2;basis,children=split_fn(matrix,range_iterations=10,sign_iterations=6,reorthogonalize_every=5,lanczos_steps=3,lanczos_probes=2,range_cholesky_passes=1,fused_sign_update=True,final_cholesky_passes=1,skip_final_high_checkpoint=True,center_mode='trace_fraction',center_trace_fraction=.8086,block_inverse_precision='highest',low_precision_sign=low_precision_sign,orthogonalize_fn=_nearrank_recursive_orthogonalize,asymmetric_high_iterations=0,asymmetric_complete_qr=True,asymmetric_skip_final_low_cqr=True,lanczos_stats_fn=_e1052_fused_nearrank_lanczos3);child_batch,child_n,_=children.shape;child_rank=child_n//2;child_joined_checkpoint=0
def child_orthogonalize(matrix:torch.Tensor,**kwargs):
nonlocal child_joined_checkpoint
if matrix.shape[0]==2*child_batch:
checkpoint=child_joined_checkpoint;child_joined_checkpoint+=1
if checkpoint==2:low=normalize_columns_(matrix[:child_batch]);high=normalize_columns_(matrix[child_batch:]);return torch.cat((low,high),dim=0)
return _nearrank_recursive_orthogonalize(matrix,**kwargs)
child_basis,grandchildren=split_fn(children,range_iterations=20,sign_iterations=9,reorthogonalize_every=4,lanczos_steps=3,lanczos_probes=2,range_cholesky_passes=1,fused_sign_update=True,final_cholesky_passes=1,skip_final_high_checkpoint=True,center_mode='trace_fraction',center_trace_fraction=.9465730329220167,block_inverse_precision='highest',power_range_iteration=True,low_precision_sign=low_precision_sign,orthogonalize_fn=child_orthogonalize,lanczos_stats_fn=_e1052_fused_nearrank_lanczos3);grand_vectors,grand_values=_e842_direct_n192_eigh(grandchildren,use_tf32_replay=True,coarse_solver=True);child_vectors=torch.empty_like(children);torch.bmm(child_basis[:,:,:child_rank],grand_vectors[:child_batch],out=child_vectors[:,:,:child_rank]);torch.bmm(child_basis[:,:,child_rank:],grand_vectors[child_batch:],out=child_vectors[:,:,child_rank:]);child_values=torch.cat((grand_values[:child_batch],grand_values[child_batch:]),dim=-1);vectors=torch.empty_like(matrix);torch.bmm(basis[:,:,:rank],child_vectors[:batch],out=vectors[:,:,:rank]);torch.bmm(basis[:,:,rank:],child_vectors[batch:],out=vectors[:,:,rank:]);values=torch.cat((child_values[:batch],child_values[batch:]),dim=-1);width=8;begin,end=rank-width,rank+width;window=vectors[:,:,begin:end];rotation,local_values=_e951_n16_eigh(window.mT@matrix@window);vectors[:,:,begin:end]=window@rotation;values[:,begin:end]=local_values;values,order=values.sort(dim=-1);vectors=_e160_gather_columns(vectors,order);return vectors,values
@triton.jit
def _e1978_pack_nearrank_qr_seed_kernel(null,h,total,n:tl.constexpr,nullity:tl.constexpr,active:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);mask=offsets<total;batch=offsets//(n*active);element=offsets-batch*n*active;row=element//active;column=element-row*active;source=batch*n*nullity+row*nullity+column;value=tl.load(null+source,mask=mask&(column<nullity),other=.0);value=tl.where(column<nullity,value,(row==column).to(tl.float32));destination=batch*n*n+row*n+column;tl.store(h+destination,value,mask=mask)
@torch.no_grad()
def nearrank1024_qr_candidate(matrix:torch.Tensor,*,qr_module=None,split_fn=None,cholesky_fn=None,compact_wy_fn=None,low_precision_projector:bool=False,low_precision_sign:bool=False):
if cholesky_fn is None:cholesky_fn=cholesky_orthonormalize
if compact_wy_fn is None:compact_wy_fn=compact_wy_orthogonalize_1024x384
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');batch,n,_=matrix.shape;nullity,positive_rank=n//4,3*n//4;ratio=1e1**(1./(positive_rank-1));trace_constant=.1*(ratio**positive_rank-1.)/(ratio-1.);upper=matrix.diagonal(dim1=-2,dim2=-1).sum(dim=-1)/trace_constant;projector=_identity_minus_scaled(matrix,upper)
if low_precision_projector:
projector=projector.to(torch.float16)
for _ in range(6):projector=torch.bmm(projector,projector)
projector=projector.to(torch.float32)
else:
for _ in range(6):projector=projector@projector
projector=.5*(projector+projector.mT);null=cholesky_fn(projector[:,:,:nullity].clone(),passes=1,ridge=1e-06);null=projector@null
for _ in range(1):gram=null.mT@null;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);null=null@gram
if compact_wy_fn is compact_wy_orthogonalize_1024x384:qr_workspace=torch.empty((batch,n,n),device=matrix.device);total=batch*n*384;_e1978_pack_nearrank_qr_seed_kernel[triton.cdiv(total,4096),](null,qr_workspace,total,n=1024,nullity=256,active=384,block=4096,num_warps=8,num_stages=1);full_basis=compact_wy_fn(qr_module,qr_workspace[:,:,:384],complete=True,newton_steps=0,paired_replay=True,prefilled_h=qr_workspace)
else:eye=torch.eye(n,device=matrix.device).expand(batch,-1,-1);qr_seed=torch.empty((batch,n,384),device=matrix.device);qr_seed[:,:,:nullity]=null;qr_seed[:,:,nullity:]=eye[:,:,nullity:384];full_basis=compact_wy_fn(qr_module,qr_seed,complete=True,newton_steps=0,paired_replay=True)
null=full_basis[:,:,:nullity];positive=full_basis[:,:,nullity:];projected=positive.mT@matrix@positive;rotation,positive_values=nearrank_positive768_candidate(projected,split_fn=split_fn,low_precision_sign=low_precision_sign);vectors=torch.empty_like(matrix);vectors[:,:,:nullity].copy_(null);torch.bmm(positive,rotation,out=vectors[:,:,nullity:]);gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;values=torch.cat((torch.zeros((batch,nullity),device=matrix.device),positive_values),dim=-1);torch.set_float32_matmul_precision(previous);return vectors,values
@torch.no_grad()
def spectral_split_eigh_2048(data:torch.Tensor,*,sign_iterations:int=6,range_iterations:int=8,reorthogonalize_every:int=4,lanczos_steps:int=3,boundary_width:int=0,newton_steps:int=1,newton_precision:str='high',rayleigh_values:bool=True,split_levels:int=1,child_boundary_width:int=0,child_sign_iterations:int|None=None,child_range_iterations:int|None=None,child_lanczos_steps:int|None=None,child_center_trace_fraction:float|None=None,child_boundary_positive_only:bool=False,child_negative_boundary_width:int=0,leaf_solver=None,low_precision_sign:bool=False,asymmetric_high_iterations:int|None=None,asymmetric_direct_complement:bool=False,asymmetric_normalize_high:bool=False,orthogonalize_fn=None,root_lanczos_stats_fn=None,child_lanczos_stats_fn=None,root_minimax_high_checkpoint:bool=False,diagonal_children:bool=False):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
def solve_level(matrix:torch.Tensor,levels:int):
at_root=matrix.shape[-1]==data.shape[-1];local_sign=sign_iterations if at_root or child_sign_iterations is None else child_sign_iterations;local_range=range_iterations if at_root or child_range_iterations is None else child_range_iterations;local_lanczos=lanczos_steps if at_root or child_lanczos_steps is None else child_lanczos_steps;basis,children=spectral_split_fast_cholesky(matrix,sign_iterations=local_sign,range_iterations=local_range,reorthogonalize_every=reorthogonalize_every,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,lanczos_steps=local_lanczos,lanczos_probes=2,fused_sign_update=True,power_range_iteration=True,center_mode='trace_fraction'if not at_root and child_center_trace_fraction is not None else'adaptive',center_trace_fraction=child_center_trace_fraction if not at_root else None,low_precision_sign=low_precision_sign,asymmetric_high_iterations=asymmetric_high_iterations,asymmetric_direct_complement=asymmetric_direct_complement,asymmetric_normalize_high=asymmetric_normalize_high,asymmetric_minimax_high_checkpoint=root_minimax_high_checkpoint and at_root,orthogonalize_fn=orthogonalize_fn,lanczos_stats_fn=root_lanczos_stats_fn if at_root else child_lanczos_stats_fn,diagonal_children=diagonal_children)
if levels==1:
if leaf_solver is None:child_values,child_vectors=torch.linalg.eigh(children)
else:child_vectors,child_values=leaf_solver(children)
else:child_vectors,child_values=solve_level(children,levels-1)
batch,n,_=matrix.shape;rank=n//2;vectors=torch.empty_like(matrix);torch.bmm(basis[:,:,:rank],child_vectors[:batch],out=vectors[:,:,:rank]);torch.bmm(basis[:,:,rank:],child_vectors[batch:],out=vectors[:,:,rank:]);values=torch.cat((child_values[:batch],child_values[batch:]),dim=1);values,order=values.sort(dim=1);vectors=_e160_gather_columns(vectors,order)
if matrix.shape[-1]<data.shape[-1]and child_boundary_width:
vectors=vectors.clone();trace_positive=matrix.diagonal(dim1=-2,dim2=-1).sum(dim=-1)>0
def refine(selected:torch.Tensor,requested_width:int):
width=min(requested_width,rank)
if width==0 or not bool(selected.any().item()):return
begin,end=rank-width,rank+width;window=vectors[selected,:,begin:end];projected=window.mT@matrix[selected]@window;_,rotation=torch.linalg.eigh(projected);vectors[selected,:,begin:end]=window@rotation
if child_boundary_positive_only:refine(trace_positive,child_boundary_width);refine(~trace_positive,child_negative_boundary_width)
else:refine(torch.ones_like(trace_positive),child_boundary_width)
return vectors,values
vectors,split_values=solve_level(data,split_levels);batch,n,_=data.shape;rank=n//2
if boundary_width:width=min(boundary_width,rank);begin,end=rank-width,rank+width;window=vectors[:,:,begin:end];projected=window.mT@data@window;_,rotation=torch.linalg.eigh(projected);vectors=vectors.clone();vectors[:,:,begin:end]=window@rotation
torch.set_float32_matmul_precision(newton_precision)
for _ in range(newton_steps):gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram
torch.set_float32_matmul_precision('high')
if rayleigh_values:values=(vectors*(data@vectors)).sum(dim=1)
else:values=split_values
values,order=values.sort(dim=1);vectors=_e160_gather_columns(vectors,order);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def dense512_manual_eigh(data:torch.Tensor,*,root_sign:int=7,root_range:int=10,root_reorth:int=5,child_sign:int=9,child_range:int=12,child_reorth:int=6,boundary:int=24,root_lanczos:int=3,child_lanczos:int=5,power_range:bool=False,child_power_range:bool=False,final_newton_steps:int=1,leaf_solver:str='eigh',leaf_eigh_fn=None,boundary_eigh_fn=None,third_level:bool=False,grand_sign:int=6,grand_range:int=8,grand_reorth:int=4,grand_lanczos:int=5,grand_boundary:int=16,grand_newton_steps:int=0,rayleigh_values:bool=True,inner_rank:int=32,inner_power_iterations:int=3,inner_completion:str='cqr',inner_complement_passes:int=2,inner_complement_reproject:bool=True,inner_newton_steps:int=1,risk_repair_rank:int=0,risk_repair_fraction:int=0,risk_repair_threshold:float=.0,root_orthogonalize_fn=None,child_orthogonalize_fn=None,root_lanczos_stats_fn=None,child_lanczos_stats_fn=None,root_low_precision_sign:bool=False,child_low_precision_sign:bool=False,grand_low_precision_sign:bool=False,root_asymmetric_high_iterations:int|None=None,root_asymmetric_direct_complement:bool=False,root_asymmetric_normalize_high:bool=False,root_asymmetric_complete_qr:bool=False,root_asymmetric_skip_final_low_cqr:bool=False,root_asymmetric_direct_qr_workspace:bool=False,child_asymmetric_high_iterations:int|None=None,child_asymmetric_direct_complement:bool=False,child_asymmetric_normalize_high:bool=False,child_asymmetric_complete_qr:bool=False,child_asymmetric_skip_final_low_cqr:bool=False,root_nonic_sign:bool=False,child_nonic_sign:bool=False):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
root_basis,children=spectral_split_fast_cholesky(data,sign_iterations=root_sign,range_iterations=root_range,reorthogonalize_every=root_reorth,lanczos_steps=root_lanczos,lanczos_probes=2,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,fused_sign_update=True,power_range_iteration=power_range,orthogonalize_fn=root_orthogonalize_fn,lanczos_stats_fn=root_lanczos_stats_fn,low_precision_sign=root_low_precision_sign,asymmetric_high_iterations=root_asymmetric_high_iterations,asymmetric_direct_complement=root_asymmetric_direct_complement,asymmetric_normalize_high=root_asymmetric_normalize_high,asymmetric_complete_qr=root_asymmetric_complete_qr,asymmetric_skip_final_low_cqr=root_asymmetric_skip_final_low_cqr,asymmetric_direct_qr_workspace=root_asymmetric_direct_qr_workspace,nonic_sign=root_nonic_sign);child_basis,leaves=spectral_split_fast_cholesky(children,sign_iterations=child_sign,range_iterations=child_range,reorthogonalize_every=child_reorth,lanczos_steps=child_lanczos,lanczos_probes=2,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,fused_sign_update=True,power_range_iteration=child_power_range,orthogonalize_fn=child_orthogonalize_fn,lanczos_stats_fn=child_lanczos_stats_fn,low_precision_sign=child_low_precision_sign,asymmetric_high_iterations=child_asymmetric_high_iterations,asymmetric_direct_complement=child_asymmetric_direct_complement,asymmetric_normalize_high=child_asymmetric_normalize_high,asymmetric_complete_qr=child_asymmetric_complete_qr,asymmetric_skip_final_low_cqr=child_asymmetric_skip_final_low_cqr,nonic_sign=child_nonic_sign)
if third_level:
grand_basis,subleaves=spectral_split_fast_cholesky(leaves,sign_iterations=grand_sign,range_iterations=grand_range,reorthogonalize_every=grand_reorth,lanczos_steps=grand_lanczos,lanczos_probes=2,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,fused_sign_update=True,power_range_iteration=True,low_precision_sign=grand_low_precision_sign);subleaf_values,subleaf_vectors=torch.linalg.eigh(subleaves);leaf_batch,leaf_n,_=leaves.shape;subrank=leaf_n//2;leaf_vectors=torch.cat((grand_basis[:,:,:subrank]@subleaf_vectors[:leaf_batch],grand_basis[:,:,subrank:]@subleaf_vectors[leaf_batch:]),dim=-1)
if grand_boundary:width=min(grand_boundary,subrank);begin,end=subrank-width,subrank+width;grand_window=leaf_vectors[:,:,begin:end];grand_projected=grand_window.mT@leaves@grand_window;_,grand_rotation=torch.linalg.eigh(grand_projected);leaf_vectors=leaf_vectors.clone();leaf_vectors[:,:,begin:end]=grand_window@grand_rotation
if grand_newton_steps:
grand_eye=torch.eye(leaf_n,device=leaves.device).expand(leaf_batch,-1,-1)
for _ in range(grand_newton_steps):leaf_vectors=.5*leaf_vectors@(3.*grand_eye-leaf_vectors.mT@leaf_vectors)
leaf_values=torch.cat((subleaf_values[:leaf_batch],subleaf_values[leaf_batch:]),dim=1)
elif leaf_solver=='eigh':
if leaf_eigh_fn is None:leaf_values,leaf_vectors=torch.linalg.eigh(leaves)
else:leaf_values,leaf_vectors=leaf_eigh_fn(leaves)
elif leaf_solver=='packed128':leaf_vectors,leaf_values=_small_n128_eigh(leaves)
elif leaf_solver=='inner_lowrank':
leaf_batch=leaves.shape[0]//4;leaf_n=leaves.shape[-1];outer=torch.cat((leaves[:leaf_batch],leaves[3*leaf_batch:]),dim=0).contiguous();outer_values,outer_vectors=torch.linalg.eigh(outer);inner=leaves[leaf_batch:3*leaf_batch];probes=_rademacher_probes(leaf_n,inner_rank,leaves.device)[None].expand(inner.shape[0],-1,-1).clone();dominant_basis=probes
for _ in range(inner_power_iterations):dominant_basis=inner@dominant_basis;dominant_basis=cholesky_orthonormalize(dominant_basis,passes=1,ridge=1e-06,final_ridge=1e-08,inverse_precision='high')
full_inner_basis=None
if inner_completion=='qr':full_inner_basis=torch.linalg.qr(dominant_basis,mode='complete').Q;dominant_basis=full_inner_basis[:,:,:inner_rank]
projected=dominant_basis.mT@inner@dominant_basis;dominant_values,dominant_rotation=torch.linalg.eigh(projected);dominant=dominant_basis@dominant_rotation;complement_rank=leaf_n-inner_rank
if inner_completion=='qr':assert full_inner_basis is not None;complement=full_inner_basis[:,:,inner_rank:]
elif inner_completion=='cqr':
complement=torch.eye(leaf_n,device=leaves.device)[:,:complement_rank].expand(inner.shape[0],-1,-1).clone();complement-=dominant@(dominant.mT@complement);complement=cholesky_orthonormalize(complement,passes=inner_complement_passes,ridge=1e-05,final_ridge=1e-08,inverse_precision='high')
if inner_complement_reproject:complement-=dominant@(dominant.mT@complement);complement=cholesky_orthonormalize(complement,passes=1,ridge=1e-06,final_ridge=1e-08,inverse_precision='high')
else:raise ValueError("inner_completion must be 'cqr' or 'qr'")
complement_values=(complement*(inner@complement)).sum(dim=1);inner_vectors=torch.cat((complement,dominant),dim=-1);inner_values=torch.cat((complement_values,dominant_values),dim=-1);inner_eye=torch.eye(leaf_n,device=leaves.device).expand(inner.shape[0],-1,-1)
for _ in range(inner_newton_steps):inner_vectors=.5*inner_vectors@(3.*inner_eye-inner_vectors.mT@inner_vectors)
inner_values=(inner_vectors*(inner@inner_vectors)).sum(dim=1);inner_values,inner_order=inner_values.sort(dim=1);inner_vectors=inner_vectors.gather(2,inner_order[:,None,:].expand(-1,leaf_n,-1));leaf_vectors=torch.empty_like(leaves);leaf_values=torch.empty((leaves.shape[0],leaf_n),device=leaves.device,dtype=leaves.dtype);leaf_vectors[:leaf_batch]=outer_vectors[:leaf_batch];leaf_values[:leaf_batch]=outer_values[:leaf_batch];leaf_vectors[leaf_batch:3*leaf_batch]=inner_vectors;leaf_values[leaf_batch:3*leaf_batch]=inner_values;leaf_vectors[3*leaf_batch:]=outer_vectors[leaf_batch:];leaf_values[3*leaf_batch:]=outer_values[leaf_batch:]
elif leaf_solver=='svd':leaf_vectors,_,_=torch.linalg.svd(leaves);leaf_values=(leaf_vectors*(leaves@leaf_vectors)).sum(dim=1);leaf_values,leaf_order=leaf_values.sort(dim=1);leaf_vectors=leaf_vectors.gather(2,leaf_order[:,None,:].expand(-1,leaves.shape[-1],-1))
else:raise ValueError("leaf_solver must be 'eigh', 'packed128', 'inner_lowrank', or 'svd'")
child_batch,child_n,_=children.shape;leaf_rank=child_n//2;child_vectors=torch.empty_like(children);torch.bmm(child_basis[:,:,:leaf_rank],leaf_vectors[:child_batch],out=child_vectors[:,:,:leaf_rank]);torch.bmm(child_basis[:,:,leaf_rank:],leaf_vectors[child_batch:],out=child_vectors[:,:,leaf_rank:]);child_values=torch.cat((leaf_values[:child_batch],leaf_values[child_batch:]),dim=1)
if boundary:
width=min(boundary,leaf_rank);begin,end=leaf_rank-width,leaf_rank+width;window=child_vectors[:,:,begin:end];projected=window.mT@children@window
if boundary_eigh_fn is None:boundary_values,rotation=torch.linalg.eigh(projected)
else:rotation,boundary_values=boundary_eigh_fn(projected)
repaired_window=window@rotation;child_vectors[:,:,begin:end]=repaired_window;child_values[:,begin:end]=boundary_values
batch,n,_=data.shape;rank=n//2;vectors=torch.empty_like(data);torch.bmm(root_basis[:,:,:rank],child_vectors[:batch],out=vectors[:,:,:rank]);torch.bmm(root_basis[:,:,rank:],child_vectors[batch:],out=vectors[:,:,rank:])
for _ in range(final_newton_steps):gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram
if rayleigh_values:
aq=data@vectors
if risk_repair_rank>0 and risk_repair_fraction>0:
if risk_repair_threshold>.0:values,residual_energy,risk_rows=_dense_guard_statistics(data,vectors,aq,risk_repair_threshold)
else:values=(vectors*aq).sum(dim=1);residual=aq-vectors*values[:,None,:];residual_energy=residual.square().sum(dim=1);risk_count=min(batch,max(1,(batch+risk_repair_fraction-1)//risk_repair_fraction));risk_rows=residual_energy.amax(dim=1).topk(risk_count).indices
if risk_rows.numel()>0:repair_indices=residual_energy[risk_rows].topk(risk_repair_rank,dim=1).indices;repair_gather=repair_indices[:,None,:].expand(-1,n,-1);selected=vectors[risk_rows].gather(2,repair_gather);selected_aq=aq[risk_rows].gather(2,repair_gather);local_values,rotation=torch.linalg.eigh(selected.mT@selected_aq);repair_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');repaired_vectors=selected@rotation;torch.set_float32_matmul_precision(repair_precision);local_vectors=vectors[risk_rows].clone();local_vectors.scatter_(2,repair_gather,repaired_vectors);local_value_rows=values[risk_rows].clone();local_value_rows.scatter_(1,repair_indices,local_values);vectors[risk_rows]=local_vectors;values[risk_rows]=local_value_rows
else:values=(vectors*aq).sum(dim=1)
else:values=torch.cat((child_values[:batch],child_values[batch:]),dim=1)
values,order=values.sort(dim=1);vectors=_e160_gather_columns(vectors,order);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
def _mixed_walsh_probes(n:int,count:int,device:torch.device):index=torch.arange(n,device=device,dtype=torch.int64)[:,None];masks=torch.arange(count,device=device,dtype=torch.int64)[None,:];parity=index&masks;parity=torch.bitwise_xor(parity,parity>>4);parity=torch.bitwise_xor(parity,parity>>2);parity=torch.bitwise_xor(parity,parity>>1)&1;return torch.where(parity==0,1.,-1.)/n**.5
_LAPACK_ACTIVE256_P0_NAME='lapack_active256_p0_96';_LAPACK_ACTIVE256_P1_NAME='lapack_active256_p1_96';_LAPACK_ACTIVE256_P2_NAME='lapack_active256_p2_64';_LAPACK_CHILD_ACTIVE128_P0_NAME='lapack_child_active128_p0_96';_LAPACK_CHILD_ACTIVE128_P1_NAME='lapack_child_active128_p1_32';_RANKDEF_ACTIVE192_P0_NAME='rankdef_active192_p0_96';_RANKDEF_ACTIVE192_P1_NAME='rankdef_active192_p1_96';_RANKDEF_ACTIVE96_P0_NAME='rankdef_active96_p0_96'
def _lapack_active256_compact_strides(source:str):return source.replace('input += batch * N * N;','input += batch * 512 * 256;').replace('output += batch * N * N;','output += batch * 512 * 256;').replace('tau += batch * N;','tau += batch * 256;')
@memo(maxsize=1)
def _lapack_active256_panel_kernels():
p0_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=512, COLS=96, N=512;';p0_new=f"__launch_bounds__(384, 1)\nvoid {_LAPACK_ACTIVE256_P0_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=512, COLS=96, N=256;";p1_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p1(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=416, COLS=96, N=512;';p1_new=f"__launch_bounds__(384, 1)\nvoid {_LAPACK_ACTIVE256_P1_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=416, COLS=96, N=256;"
if _N512_GAU_PANEL_SOURCE.count(p0_old)!=1:raise RuntimeError('LAPACK active256 p0 template changed')
if _N512_GAU_PANEL_SOURCE.count(p1_old)!=1:raise RuntimeError('LAPACK active256 p1 template changed')
panel_source=_lapack_active256_compact_strides(_N512_GAU_PANEL_SOURCE.replace(p0_old,p0_new,1).replace(p1_old,p1_new,1));panel_source=_fast_only_cuda_kernels(panel_source,(_LAPACK_ACTIVE256_P0_NAME,_LAPACK_ACTIVE256_P1_NAME));panel_image=_fast_nvrtc_compile(panel_source,_LAPACK_ACTIVE256_P0_NAME);p2_old='__launch_bounds__(256, 1)\nvoid qr2_gau_n512_active320x64(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=320, COLS=64, N=512;';p2_new=f"__launch_bounds__(256, 1)\nvoid {_LAPACK_ACTIVE256_P2_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=320, COLS=64, N=256;"
if _N512_GAU_ACTIVE_SOURCE.count(p2_old)!=1:raise RuntimeError('LAPACK active256 p2 template changed')
active_source=_lapack_active256_compact_strides(_N512_GAU_ACTIVE_SOURCE.replace(p2_old,p2_new,1));active_source=_fast_only_cuda_kernels(active_source,(_LAPACK_ACTIVE256_P2_NAME,));active_image=_fast_nvrtc_compile(active_source,_LAPACK_ACTIVE256_P2_NAME);return CUDAKernel(panel_image,_LAPACK_ACTIVE256_P0_NAME),CUDAKernel(panel_image,_LAPACK_ACTIVE256_P1_NAME),CUDAKernel(active_image,_LAPACK_ACTIVE256_P2_NAME)
@triton.jit
def _e1233_lapack_thin_identity_kernel(output,total,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;element=offsets%(512*256);row=element//256;column=element-row*256;tl.store(output+offsets,(row==column).to(tl.float32),mask=mask)
@triton.jit
def _e1294_lapack_root_first_update(source,panel,transformed,output,source_batch_stride,source_row_stride,source_column_stride,panel_batch_stride,panel_row_stride,panel_column_stride,transformed_batch_stride,transformed_row_stride,transformed_column_stride,output_batch_stride,output_row_stride,output_column_stride):
batch=tl.program_id(0);rows=tl.program_id(1)*64+tl.arange(0,64);columns=tl.program_id(2)*32+tl.arange(0,32);accumulator=tl.zeros((64,32),tl.float32)
for begin in tl.static_range(0,96,32):rank=begin+tl.arange(0,32);left=tl.load(panel+batch*panel_batch_stride+rows[:,None]*panel_row_stride+rank[None,:]*panel_column_stride,mask=rows[:,None]<512,other=.0);right=tl.load(transformed+batch*transformed_batch_stride+rank[:,None]*transformed_row_stride+columns[None,:]*transformed_column_stride,mask=columns[None,:]<160,other=.0).to(tl.float16);accumulator+=tl.dot(left,right)
mask=(rows[:,None]<512)&(columns[None,:]<160);old=tl.load(source+batch*source_batch_stride+rows[:,None]*source_row_stride+columns[None,:]*source_column_stride,mask=mask,other=.0);tl.store(output+batch*output_batch_stride+rows[:,None]*output_row_stride+columns[None,:]*output_column_stride,old-accumulator,mask=mask)
@torch.no_grad()
def _e1233_lapack_active256_thin(matrix:torch.Tensor):
batch,rows,active=matrix.shape
if rows!=512 or active!=256 or not matrix.is_contiguous():raise ValueError('LAPACK thin active QR expects contiguous 512x256')
source=matrix;h=torch.empty_like(source);tau=torch.zeros((batch,256),device=matrix.device);panels=_lapack_active256_panel_kernels();_,t96=_n352_gau_t_kernels();t64=_active352_t64_kernel();nodes=[];offset=0;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
for(index,width)in enumerate((96,96,64)):
panel_rows=512-offset;v32=h.new_empty(batch,width,panel_rows).transpose(1,2);v16=h.new_empty(batch,width,panel_rows,dtype=torch.float16).transpose(1,2);panels[index].launch((batch,1,1),(width//8*32,1,1),(source[:,offset:,offset:offset+width],h[:,offset:,offset:offset+width],tau[:,offset:offset+width],v32,v16),shared_mem=(panel_rows*width+width)*4+width*8);gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32);triangular=torch.empty_like(gram)
if width==64:t64.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=36864)
else:t96.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=36864);middle=gram[:,64:,:64]@triangular[:,:64,:64];torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64])
nodes.append((offset,v32,triangular))
if offset+width<256:
trailing_input=source[:,offset:,offset+width:256];trailing_output=h[:,offset:,offset+width:256];transformed=triangular@(v32.mT@trailing_input)
if offset==0:_e1294_lapack_root_first_update[batch,8,5](trailing_input,v16,transformed,trailing_output,trailing_input.stride(0),trailing_input.stride(1),trailing_input.stride(2),v16.stride(0),v16.stride(1),v16.stride(2),transformed.stride(0),transformed.stride(1),transformed.stride(2),trailing_output.stride(0),trailing_output.stride(1),trailing_output.stride(2),num_warps=4,num_stages=2)
else:torch.baddbmm(trailing_input,v16,transformed.half(),beta=1.,alpha=-1.,out=trailing_output,out_dtype=torch.float32)
source=h;offset+=width
q=torch.empty((batch,512,256),device=matrix.device);total=q.numel();_e1233_lapack_thin_identity_kernel[triton.cdiv(total,4096),](q,total,BLOCK=4096,num_warps=8,num_stages=1)
for(offset,v32,triangular)in reversed(nodes):active_q=q[:,offset:,offset:];transformed=triangular.mT@(v32.mT@active_q);torch.baddbmm(active_q,v32,transformed,beta=1.,alpha=-1.,out=active_q)
return q
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _lapack_active256_complete(matrix:torch.Tensor):
batch,rows,active=matrix.shape
if rows!=512 or active!=256 or not matrix.is_contiguous():raise ValueError('LAPACK active completion expects contiguous 512x256')
source=matrix;h=torch.empty_like(source);tau=torch.zeros((batch,256),device=matrix.device);panels=_lapack_active256_panel_kernels();_,t96=_n352_gau_t_kernels();t64=_active352_t64_kernel();nodes=[];offset=0;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
for(index,width)in enumerate((96,96,64)):
panel_rows=512-offset;v32=h.new_empty(batch,width,panel_rows).transpose(1,2);v16=h.new_empty(batch,width,panel_rows,dtype=torch.float16).transpose(1,2);panels[index].launch((batch,1,1),(width//8*32,1,1),(source[:,offset:,offset:offset+width],h[:,offset:,offset:offset+width],tau[:,offset:offset+width],v32,v16),shared_mem=(panel_rows*width+width)*4+width*8);gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32);triangular=torch.empty_like(gram)
if width==64:t64.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=36864)
else:t96.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=36864);middle=gram[:,64:,:64]@triangular[:,:64,:64];torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64])
nodes.append((offset,v32,triangular))
if offset+width<256:
trailing_input=source[:,offset:,offset+width:256];trailing_output=h[:,offset:,offset+width:256];transformed=triangular@(v32.mT@trailing_input)
if offset==0:_e1294_lapack_root_first_update[batch,8,5](trailing_input,v16,transformed,trailing_output,trailing_input.stride(0),trailing_input.stride(1),trailing_input.stride(2),v16.stride(0),v16.stride(1),v16.stride(2),transformed.stride(0),transformed.stride(1),transformed.stride(2),trailing_output.stride(0),trailing_output.stride(1),trailing_output.stride(2),num_warps=4,num_stages=2)
else:torch.baddbmm(trailing_input,v16,transformed.half(),beta=1.,alpha=-1.,out=trailing_output,out_dtype=torch.float32)
source=h;offset+=width
q=_e162_batched_identity512(batch,matrix.device);first_replay=True
for(offset,v32,triangular)in reversed(nodes):
if offset:active_q=q[:,offset:,offset:];transformed=triangular.mT@v32.mT if first_replay else triangular.mT@(v32.mT@active_q)
else:active_q=q;transformed=triangular.mT@(v32.mT@active_q)
torch.baddbmm(active_q,v32,transformed,beta=1.,alpha=-1.,out=active_q);first_replay=False
return q
finally:torch.set_float32_matmul_precision(previous)
@memo(maxsize=1)
def _lapack_child_active128_panel_kernels():
p0_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=512, COLS=96, N=512;';p0_new=f"__launch_bounds__(384, 1)\nvoid {_LAPACK_CHILD_ACTIVE128_P0_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=256, COLS=96, N=128;";p1_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p1(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=416, COLS=96, N=512;';p1_new=f"__launch_bounds__(128, 1)\nvoid {_LAPACK_CHILD_ACTIVE128_P1_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=160, COLS=32, N=128;";source=_N512_GAU_PANEL_SOURCE
if source.count(p0_old)!=1 or source.count(p1_old)!=1:raise RuntimeError('LAPACK child active128 panel template changed')
source=source.replace(p0_old,p0_new,1).replace(p1_old,p1_new,1);source=source.replace('input += batch * N * N;','input += batch * 256 * 128;').replace('output += batch * N * N;','output += batch * 256 * 128;').replace('tau += batch * N;','tau += batch * 128;');source=_fast_only_cuda_kernels(source,(_LAPACK_CHILD_ACTIVE128_P0_NAME,_LAPACK_CHILD_ACTIVE128_P1_NAME));image=_fast_nvrtc_compile(source,_LAPACK_CHILD_ACTIVE128_P0_NAME);return CUDAKernel(image,_LAPACK_CHILD_ACTIVE128_P0_NAME),CUDAKernel(image,_LAPACK_CHILD_ACTIVE128_P1_NAME)
@torch.no_grad()
def _lapack_child_active128_complete(matrix:torch.Tensor):
batch,rows,active=matrix.shape
if rows!=256 or active!=128 or not matrix.is_contiguous():raise ValueError('LAPACK child active completion expects contiguous 256x128')
source=matrix;h=torch.empty_like(matrix);tau=torch.zeros((batch,128),device=matrix.device);panels=_lapack_child_active128_panel_kernels();_,t96=_n352_gau_t_kernels();nodes=[];offset=0;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
for(panel,width)in zip(panels,(96,32)):
panel_rows=rows-offset;v32=h.new_empty(batch,width,panel_rows).transpose(1,2);v16=h.new_empty(batch,width,panel_rows,dtype=torch.float16).transpose(1,2);panel.launch((batch,1,1),(width//8*32,1,1),(source[:,offset:,offset:offset+width],h[:,offset:,offset:offset+width],tau[:,offset:offset+width],v32,v16),shared_mem=(panel_rows*width+width)*4+width*8);gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32)
if width==96:triangular=torch.empty_like(gram);t96.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=36864);middle=gram[:,64:,:64]@triangular[:,:64,:64];torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64])
else:triangular=torch.empty_like(gram);_t32_kernel().launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=(32*32*2+16*16)*4)
nodes.append((offset,v32,triangular))
if offset+width<active:trailing_input=source[:,offset:,offset+width:active];trailing_output=h[:,offset:,offset+width:active];transformed=triangular@(v32.mT@trailing_input);torch.baddbmm(trailing_input,v16,transformed.half(),beta=1.,alpha=-1.,out=trailing_output,out_dtype=torch.float32)
source=h;offset+=width
q=torch.empty((batch,rows,rows),device=matrix.device,dtype=torch.float32);total=q.numel();_e162_batched_identity512_kernel[triton.cdiv(total,1024),](q,total,N=rows,BLOCK=1024,num_warps=8,num_stages=1);first=True
for(offset,v32,triangular)in reversed(nodes):
active_q=q[:,offset:,offset:]if offset else q
if first and offset:transformed=triangular.mT@v32.mT
else:transformed=triangular.mT@(v32.mT@active_q)
torch.baddbmm(active_q,v32,transformed,beta=1.,alpha=-1.,out=active_q);first=False
return q
finally:torch.set_float32_matmul_precision(previous)
@memo(maxsize=2)
def _rankdef_small_active_panel_kernels(rows:int,active:int):
source=_N512_GAU_PANEL_SOURCE;p0_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=512, COLS=96, N=512;'
if(rows,active)==(384,192):
p0_new=f"__launch_bounds__(384, 1)\nvoid {_RANKDEF_ACTIVE192_P0_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=384, COLS=96, N=192;";p1_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p1(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=416, COLS=96, N=512;';p1_new=f"__launch_bounds__(384, 1)\nvoid {_RANKDEF_ACTIVE192_P1_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=288, COLS=96, N=192;"
if source.count(p0_old)!=1 or source.count(p1_old)!=1:raise RuntimeError('rankdef active192 panel template changed')
source=source.replace(p0_old,p0_new,1).replace(p1_old,p1_new,1);names=_RANKDEF_ACTIVE192_P0_NAME,_RANKDEF_ACTIVE192_P1_NAME
elif(rows,active)==(192,96):
p0_new=f"__launch_bounds__(384, 1)\nvoid {_RANKDEF_ACTIVE96_P0_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=192, COLS=96, N=96;"
if source.count(p0_old)!=1:raise RuntimeError('rankdef active96 panel template changed')
source=source.replace(p0_old,p0_new,1);names=_RANKDEF_ACTIVE96_P0_NAME,
else:raise ValueError(f"unsupported rankdef active shape {(rows,active)}")
pitch=rows*active;source=source.replace('input += batch * N * N;',f"input += batch * {pitch};").replace('output += batch * N * N;',f"output += batch * {pitch};").replace('tau += batch * N;',f"tau += batch * {active};");source=_fast_only_cuda_kernels(source,names);image=_fast_nvrtc_compile(source,names[0]);return tuple(CUDAKernel(image,name)for name in names)
@torch.no_grad()
def _rankdef_small_active_complete(matrix:torch.Tensor):
batch,rows,active=matrix.shape
if(rows,active)not in((384,192),(192,96)):raise ValueError(f"unsupported rankdef active shape {matrix.shape}")
if not matrix.is_contiguous():raise ValueError('rankdef active completion requires contiguous input')
widths=(96,96)if active==192 else(96,);panels=_rankdef_small_active_panel_kernels(rows,active);_,t96=_n352_gau_t_kernels();source=matrix;h=torch.empty_like(matrix);tau=torch.zeros((batch,active),device=matrix.device);nodes=[];offset=0;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
for(panel,width)in zip(panels,widths):
panel_rows=rows-offset;v32=h.new_empty(batch,width,panel_rows).transpose(1,2);v16=h.new_empty(batch,width,panel_rows,dtype=torch.float16).transpose(1,2);panel.launch((batch,1,1),(384,1,1),(source[:,offset:,offset:offset+width],h[:,offset:,offset:offset+width],tau[:,offset:offset+width],v32,v16),shared_mem=(panel_rows*width+width)*4+width*8);gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32);triangular=torch.empty_like(gram);t96.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+width],triangular,int(tau.stride(0))),shared_mem=36864);middle=gram[:,64:,:64]@triangular[:,:64,:64];torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64]);nodes.append((offset,v32,triangular))
if offset+width<active:trailing_input=source[:,offset:,offset+width:active];trailing_output=h[:,offset:,offset+width:active];transformed=triangular@(v32.mT@trailing_input);torch.baddbmm(trailing_input,v16,transformed.half(),beta=1.,alpha=-1.,out=trailing_output,out_dtype=torch.float32)
source=h;offset+=width
q=torch.empty((batch,rows,rows),device=matrix.device,dtype=torch.float32);total=q.numel();_e162_batched_identity512_kernel[triton.cdiv(total,1024),](q,total,N=rows,BLOCK=1024,num_warps=8,num_stages=1);first=True
for(offset,v32,triangular)in reversed(nodes):
active_q=q[:,offset:,offset:]if offset else q
if first and offset:transformed=triangular.mT@v32.mT
else:transformed=triangular.mT@(v32.mT@active_q)
torch.baddbmm(active_q,v32,transformed,beta=1.,alpha=-1.,out=active_q);first=False
return q
finally:torch.set_float32_matmul_precision(previous)
def _mixed_qr_accessors(module):
if module is not None:return module._n512_gau_panel_kernels,module._n512_gau_active_kernels,module._n352_gau_t_kernels
return _n512_gau_panel_kernels,_n512_gau_active_kernels,_n352_gau_t_kernels
_DENSE_ACTIVE320_PANEL_NAMES='dense_active320_p0_96','dense_active320_p1_96','dense_active320_p2_128'
@memo(maxsize=1)
def _dense_active320_panel_kernels():
source=_N512_GAU_PANEL_SOURCE;replacements=('__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0','__launch_bounds__(384, 1)\nvoid dense_active320_p0_96'),('__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p1','__launch_bounds__(384, 1)\nvoid dense_active320_p1_96'),('__launch_bounds__(512, 1)\nvoid qr2_gau_n512_p2','__launch_bounds__(512, 1)\nvoid dense_active320_p2_128')
for(old,new)in replacements:
if source.count(old)!=1:raise RuntimeError('active320 panel template changed')
source=source.replace(old,new,1)
source=source.replace('input += batch * N * N;','input += batch * 512 * 320;').replace('output += batch * N * N;','output += batch * 512 * 320;').replace('tau += batch * N;','tau += batch * 320;');source=source.replace('constexpr int ROWS=512, COLS=96, N=512;','constexpr int ROWS=512, COLS=96, N=320;',1).replace('constexpr int ROWS=416, COLS=96, N=512;','constexpr int ROWS=416, COLS=96, N=320;',1).replace('constexpr int ROWS=320, COLS=128, N=512;','constexpr int ROWS=320, COLS=128, N=320;',1);source=_fast_only_cuda_kernels(source,_DENSE_ACTIVE320_PANEL_NAMES);image=_fast_nvrtc_compile(source,_DENSE_ACTIVE320_PANEL_NAMES[0]);return tuple(CUDAKernel(image,name)for name in _DENSE_ACTIVE320_PANEL_NAMES)
@torch.no_grad()
def mixed_qr_active(matrix:torch.Tensor,*,complete:bool,module=None):
panel_accessor,active_accessor,t_accessor=_mixed_qr_accessors(module);batch=matrix.shape[0];active_cols=matrix.shape[-1]
if active_cols not in _N1024_ACTIVE_COLUMN_COUNTS:raise ValueError('mixed QR adapter supports 256 or 384 active columns')
compact_active320=active_cols==320 and module is None;storage_columns=320 if compact_active320 else 512
if compact_active320 and matrix.is_contiguous():data=matrix
else:data=torch.empty((batch,512,storage_columns),device=matrix.device);data[:,:,:active_cols]=matrix
h=torch.empty_like(data);tau=torch.zeros((batch,storage_columns),device=matrix.device);panels=_dense_active320_panel_kernels()if compact_active320 else panel_accessor();active320,active192=active_accessor();t128,t96=t_accessor();source=data;offset=0;nodes=[];previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
if active_cols==256:widths=96,96,64
elif active_cols==320:widths=96,96,128
else:widths=96,96,128,64
for(index,width)in enumerate(widths):
rows=512-offset;v32=h.new_empty(batch,width,rows).transpose(1,2);v16=h.new_empty(batch,width,rows,dtype=torch.float16).transpose(1,2);panel=(active320 if rows==320 else active192)if width==64 else panels[index];panel.launch(grid=(batch,1,1),block=(width//8*32,1,1),shared_mem=(rows*width+width)*4+width*8,args=[source[:,offset:,offset:offset+width],h[:,offset:,offset:offset+width],tau[:,offset:offset+width],v32,v16]);raw_gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32)
if width==64:gram=torch.zeros((batch,96,96),device=matrix.device);gram[:,:64,:64]=raw_gram
else:gram=raw_gram
triangular=torch.empty_like(gram);t_kernel=t128 if width==128 else t96;t_kernel.launch(grid=(batch*2 if width==128 else batch,1,1),block=(512,1,1),shared_mem=36864,args=[gram,tau[:,offset:offset+96],triangular,int(tau.stride(0))]);middle=gram[:,64:,:64]@triangular[:,:64,:64];torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64]);triangular=triangular[:,:width,:width];nodes.append((offset,v32,triangular))
if offset+width<active_cols:trailing_input=source[:,offset:,offset+width:active_cols];trailing_output=h[:,offset:,offset+width:active_cols];transformed=triangular@(v32.mT@trailing_input);torch.baddbmm(trailing_input,v16,transformed.half(),beta=1.,alpha=-1.,out=trailing_output,out_dtype=torch.float32)
source=h;offset+=width
width=512 if complete else active_cols;q=_e162_batched_identity512(batch,matrix.device)if complete else torch.eye(512,device=matrix.device)[:,:width].expand(batch,-1,-1).clone();first_replay=True
for(offset,v32,triangular)in reversed(nodes):
if offset:
active=q[:,offset:,offset:]
if first_replay:transformed=triangular.mT@v32.mT
else:transformed=triangular.mT@(v32.mT@active)
else:active=q[:,offset:,:];transformed=triangular.mT@(v32.mT@active)
torch.baddbmm(active,v32,transformed,beta=1.,alpha=-1.,out=active);first_replay=False
torch.set_float32_matmul_precision(previous);return q
_REPEATED_KRYLOV_CONSTANTS={};_REPEATED_PIVOT40_RANK32_NAME='repeated_pivoted_cholesky40_rank32';_REPEATED_PIVOT40_RANK32_SOURCE='\n#include <cuda_runtime.h>\n\nextern "C" __global__ __launch_bounds__(64, 2)\nvoid repeated_pivoted_cholesky40_rank32(\n const float* __restrict__ gram,\n long long* __restrict__ pivots,\n float* __restrict__ lower,\n int batch_count) {\n const int batch = blockIdx.x;\n const int tid = threadIdx.x;\n if (batch >= batch_count) return;\n constexpr int N = 40;\n constexpr int R = 32;\n __shared__ float l[N * R];\n __shared__ float diagonal[64];\n __shared__ int selected[64];\n __shared__ float reduce_value[64];\n __shared__ int reduce_index[64];\n __shared__ int pivot_shared;\n __shared__ float pivot_value_shared;\n\n if (tid < N) {\n diagonal[tid] = fmaxf(\n gram[(long long)batch * N * N + tid * N + tid], 0.0f);\n selected[tid] = 0;\n }\n for (int index = tid; index < N * R; index += 64)\n l[index] = 0.0f;\n __syncthreads();\n\n #pragma unroll 1\n for (int column = 0; column < R; ++column) {\n float candidate = -1.0e30f;\n int candidate_index = tid;\n if (tid < N && !selected[tid]) candidate = diagonal[tid];\n reduce_value[tid] = candidate;\n reduce_index[tid] = candidate_index;\n __syncthreads();\n for (int offset = 32; offset > 0; offset >>= 1) {\n if (tid < offset) {\n const float other = reduce_value[tid + offset];\n const int other_index = reduce_index[tid + offset];\n if (other > reduce_value[tid]\n || (other == reduce_value[tid]\n && other_index < reduce_index[tid])) {\n reduce_value[tid] = other;\n reduce_index[tid] = other_index;\n }\n }\n __syncthreads();\n }\n if (tid == 0) {\n const int pivot = reduce_index[0];\n pivot_shared = pivot;\n selected[pivot] = 1;\n pivots[(long long)batch * R + column] = (long long)pivot;\n pivot_value_shared = sqrtf(fmaxf(diagonal[pivot], 1.0e-20f));\n }\n __syncthreads();\n const int pivot = pivot_shared;\n const float pivot_value = pivot_value_shared;\n if (tid < N) {\n float value = 0.0f;\n if (!selected[tid] || tid == pivot) {\n value = gram[(long long)batch * N * N + tid * N + pivot];\n #pragma unroll 1\n for (int previous = 0; previous < column; ++previous)\n value -= l[tid * R + previous] * l[pivot * R + previous];\n value /= pivot_value;\n if (tid == pivot) value = pivot_value;\n l[tid * R + column] = value;\n if (tid != pivot)\n diagonal[tid] = fmaxf(\n diagonal[tid] - value * value, 0.0f);\n }\n }\n __syncthreads();\n if (tid <= column) {\n lower[((long long)batch * R + column) * R + tid] =\n l[pivot * R + tid];\n }\n __syncthreads();\n }\n}\n'
@memo(maxsize=1)
def _repeated_pivot40_rank32_kernel():return CUDAKernel(_fast_nvrtc_compile(_REPEATED_PIVOT40_RANK32_SOURCE,_REPEATED_PIVOT40_RANK32_NAME),_REPEATED_PIVOT40_RANK32_NAME)
@torch.no_grad()
def _repeated_pivot40_rank32_whiten(spaces:torch.Tensor):
batch,rows,columns=spaces.shape
if columns!=40:raise ValueError('pivoted repeated whitening requires width 40')
gram=spaces.mT@spaces;pivots=torch.empty((batch,32),device=spaces.device,dtype=torch.int64);lower=torch.empty((batch,32,32),device=spaces.device,dtype=spaces.dtype);_repeated_pivot40_rank32_kernel().launch((batch,1,1),(64,1,1),(gram,pivots,lower,batch));selected=spaces.gather(2,pivots[:,None,:].expand(-1,rows,-1));return torch.linalg.solve_triangular(lower,selected.mT,upper=False,left=True).mT
_REPEATED_PIVOT112_RANK64_NAME='repeated_pivoted_cholesky112_rank64';_REPEATED_PIVOT112_RANK64_SOURCE='\n#include <cuda_runtime.h>\n\nextern "C" __global__ __launch_bounds__(128, 2)\nvoid repeated_pivoted_cholesky112_rank64(\n const float* __restrict__ gram,\n long long* __restrict__ pivots,\n float* __restrict__ lower,\n int batch_count) {\n const int batch = blockIdx.x;\n const int tid = threadIdx.x;\n if (batch >= batch_count) return;\n constexpr int N = 112;\n constexpr int R = 64;\n __shared__ float l[N * R];\n __shared__ float diagonal[128];\n __shared__ int selected[128];\n __shared__ float reduce_value[128];\n __shared__ int reduce_index[128];\n __shared__ int pivot_shared;\n __shared__ float pivot_value_shared;\n\n if (tid < N) {\n diagonal[tid] = fmaxf(\n gram[(long long)batch * N * N + tid * N + tid], 0.0f);\n selected[tid] = 0;\n }\n for (int index = tid; index < N * R; index += 128)\n l[index] = 0.0f;\n __syncthreads();\n\n #pragma unroll 1\n for (int column = 0; column < R; ++column) {\n float candidate = -1.0e30f;\n int candidate_index = tid;\n if (tid < N && !selected[tid]) candidate = diagonal[tid];\n reduce_value[tid] = candidate;\n reduce_index[tid] = candidate_index;\n __syncthreads();\n for (int offset = 64; offset > 0; offset >>= 1) {\n if (tid < offset) {\n const float other = reduce_value[tid + offset];\n const int other_index = reduce_index[tid + offset];\n if (other > reduce_value[tid]\n || (other == reduce_value[tid]\n && other_index < reduce_index[tid])) {\n reduce_value[tid] = other;\n reduce_index[tid] = other_index;\n }\n }\n __syncthreads();\n }\n if (tid == 0) {\n const int pivot = reduce_index[0];\n pivot_shared = pivot;\n selected[pivot] = 1;\n pivots[(long long)batch * R + column] = (long long)pivot;\n pivot_value_shared = sqrtf(fmaxf(diagonal[pivot], 1.0e-20f));\n }\n __syncthreads();\n const int pivot = pivot_shared;\n const float pivot_value = pivot_value_shared;\n if (tid < N) {\n float value = 0.0f;\n if (!selected[tid] || tid == pivot) {\n value = gram[(long long)batch * N * N + tid * N + pivot];\n #pragma unroll 1\n for (int previous = 0; previous < column; ++previous)\n value -= l[tid * R + previous] * l[pivot * R + previous];\n value /= pivot_value;\n if (tid == pivot) value = pivot_value;\n l[tid * R + column] = value;\n if (tid != pivot)\n diagonal[tid] = fmaxf(\n diagonal[tid] - value * value, 0.0f);\n }\n }\n __syncthreads();\n if (tid <= column) {\n lower[((long long)batch * R + column) * R + tid] =\n l[pivot * R + tid];\n }\n __syncthreads();\n }\n}\n'
@memo(maxsize=1)
def _repeated_pivot112_rank64_kernel():return CUDAKernel(_fast_nvrtc_compile(_REPEATED_PIVOT112_RANK64_SOURCE,_REPEATED_PIVOT112_RANK64_NAME),_REPEATED_PIVOT112_RANK64_NAME)
@torch.no_grad()
def _repeated_pivot112_rank64_whiten(spaces:torch.Tensor):
batch,rows,columns=spaces.shape
if columns!=112:raise ValueError('pivoted repeated whitening requires width 112')
gram=spaces.mT@spaces;pivots=torch.empty((batch,64),device=spaces.device,dtype=torch.int64);lower=torch.empty((batch,64,64),device=spaces.device,dtype=spaces.dtype);_repeated_pivot112_rank64_kernel().launch((batch,1,1),(128,1,1),(gram,pivots,lower,batch));selected=spaces.gather(2,pivots[:,None,:].expand(-1,rows,-1));return torch.linalg.solve_triangular(lower,selected.mT,upper=False,left=True).mT
def _repeated_krylov_constants(data:torch.Tensor):
key=data.device.type,data.device.index,data.dtype;cached=_REPEATED_KRYLOV_CONSTANTS.get(key)
if cached is None:
groups=16;nodes=torch.linspace(-1.,1.,groups,device=data.device,dtype=data.dtype);columns=[torch.ones_like(nodes),nodes]
for _ in range(2,groups):columns.append(2.*nodes*columns[-1]-columns[-2])
coefficients=torch.linalg.inv(torch.stack(columns,dim=1).double()).float();cached=nodes,coefficients;_REPEATED_KRYLOV_CONSTANTS[key]=cached
return cached
@triton.jit
def _e908_repeated512_rotation_guard_kernel(projected,values,rotation,risk,N:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);chunk=tl.program_id(1);linear=chunk*BLOCK+tl.arange(0,BLOCK);mask=linear<N*N;row=linear//N;column=linear-row*N;projected_value=tl.load(projected+batch*N*N+linear,mask=mask,other=.0);row_value=tl.load(values+batch*N+row,mask=mask,other=.0);column_value=tl.load(values+batch*N+column,mask=mask,other=.0);cross_group=row//32!=column//32;delta=column_value-row_value;value=tl.where(mask&cross_group,projected_value/delta,.0);tl.store(rotation+batch*N*N+linear,value,mask=mask);hit=tl.max(tl.where(mask&(tl.abs(value)>.05),1,0),axis=0);tl.atomic_or(risk+batch,hit)
@torch.no_grad()
def _e908_repeated512_rotation_guard(projected:torch.Tensor,values:torch.Tensor):batch,n,_=projected.shape;rotation=torch.empty_like(projected);risk=torch.zeros((batch,),device=projected.device,dtype=torch.int32);block=4096;_e908_repeated512_rotation_guard_kernel[batch,triton.cdiv(n*n,block)](projected,values,rotation,risk,N=n,BLOCK=block,num_warps=8,num_stages=1);return rotation,risk!=0
@torch.no_grad()
def mixed_repeated_krylov_eigh(data:torch.Tensor,*,seed_width:int=40,defer_repair:bool=False):
batch,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');groups,width=16,n//16;nodes,coefficients=_repeated_krylov_constants(data);seed=_mixed_walsh_probes(n,seed_width,data.device).expand(batch,-1,-1);powers=torch.empty((batch,groups,n,seed_width),device=data.device,dtype=data.dtype);powers[:,0].copy_(seed);torch.bmm(data,powers[:,0],out=powers[:,1])
for index in range(2,groups):torch.baddbmm(powers[:,index-2],data,powers[:,index-1],beta=-1.,alpha=2.,out=powers[:,index])
spaces=torch.bmm(coefficients.mT.expand(batch,-1,-1),powers.view(batch,groups,n*seed_width)).view(batch*groups,n,seed_width)
if n==512 and seed_width==40:spaces=_repeated_pivot40_rank32_whiten(spaces)
elif n==1024 and seed_width==112:spaces=_repeated_pivot112_rank64_whiten(spaces)
else:gram_values,gram_vectors=torch.linalg.eigh(spaces.mT@spaces);keep_values=gram_values[:,-width:].clamp_min(1e-12);keep_vectors=gram_vectors[:,:,-width:]/keep_values.sqrt()[:,None,:];spaces=spaces@keep_vectors
vectors=spaces.reshape(batch,groups,n,width).permute(0,2,1,3);vectors=vectors.reshape(batch,n,n);gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);values=nodes.repeat_interleave(width).expand(batch,-1).contiguous()
if n==512 and seed_width==40 or n==1024 and seed_width==112:
torch.set_float32_matmul_precision('high');projected=vectors.mT@(data@vectors);projected=.5*(projected+projected.mT)
if n==512 and seed_width==40:rotation,repeated_risk=_e908_repeated512_rotation_guard(projected,values)
else:group=torch.arange(n,device=data.device)//width;cross_group=group[:,None]!=group[None,:];delta=values[:,None,:]-values[:,:,None];safe_delta=torch.where(cross_group[None],delta,1.);rotation=torch.where(cross_group[None],projected/safe_delta,.0);rotation=.5*(rotation-rotation.mT)
vectors=vectors+vectors@rotation;gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram
torch.set_float32_matmul_precision(previous)
if n==512 and seed_width==40:
if defer_repair:return vectors,values,repeated_risk
if bool(repeated_risk.any().item()):exact_vectors,exact_values=_repeated512_wide_repair(data[repeated_risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[repeated_risk]=exact_vectors;values[repeated_risk]=exact_values
return vectors,values
@torch.no_grad()
def _repeated512_wide_repair(data:torch.Tensor):'Repair rare strong inter-cluster coupling without a full 512 EVD.\n\n The common seed-width-40 route uses a first-order Sylvester rotation and\n flags rows whose cross-cluster generator is too large. Rebuilding only\n those rows with 72 deterministic Krylov probes gives every 32-dimensional\n repeated eigenspace enough oversampling for the generic 72x72 local Gram\n solve. This keeps the repair around two milliseconds for one row instead\n of launching the underfilled full-size system eigensolver.\n ';return mixed_repeated_krylov_eigh(data,seed_width=72)
_CLUSTERED_COMPACT_QR96_NAME='e152_clustered_qr512x176_p96';_CLUSTERED_COMPACT_QR80_NAME='e152_clustered_qr512x176_p80'
@memo(maxsize=1)
def _clustered_compact_qr_kernels():
p0_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p0(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=512, COLS=96, N=512;';p0_new=f"__launch_bounds__(384, 1)\nvoid {_CLUSTERED_COMPACT_QR96_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=512, COLS=96, N=176;";p1_old='__launch_bounds__(384, 1)\nvoid qr2_gau_n512_p1(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {\n constexpr int ROWS=416, COLS=96, N=512;';p1_new=f"__launch_bounds__(320, 1)\nvoid {_CLUSTERED_COMPACT_QR80_NAME}(const float* input, float* output, float* tau, float* v_fp32, __half* v_fp16) {{\n constexpr int ROWS=416, COLS=80, N=176;"
if _N512_GAU_PANEL_SOURCE.count(p0_old)!=1:raise RuntimeError('n512 compact QR p0 template changed')
if _N512_GAU_PANEL_SOURCE.count(p1_old)!=1:raise RuntimeError('n512 compact QR p1 template changed')
source=_N512_GAU_PANEL_SOURCE.replace(p0_old,p0_new,1).replace(p1_old,p1_new,1).replace('input += batch * N * N;','input += batch * 512 * 176;').replace('output += batch * N * N;','output += batch * 512 * 176;').replace('tau += batch * N;','tau += batch * 192;');source=_fast_only_cuda_kernels(source,(_CLUSTERED_COMPACT_QR96_NAME,_CLUSTERED_COMPACT_QR80_NAME));image=_fast_nvrtc_compile(source,_CLUSTERED_COMPACT_QR96_NAME);return CUDAKernel(image,_CLUSTERED_COMPACT_QR96_NAME),CUDAKernel(image,_CLUSTERED_COMPACT_QR80_NAME)
@torch.no_grad()
def _clustered_qr_active176_compact(seed:torch.Tensor):
batch,n,active_cols=seed.shape
if n!=512 or active_cols!=176 or not seed.is_contiguous():raise ValueError('compact clustered QR expects contiguous (batch,512,176)')
h=torch.empty_like(seed);tau=torch.zeros((batch,192),device=seed.device);panels=_clustered_compact_qr_kernels();_,t96=_n352_gau_t_kernels();source=seed;offset=0;nodes=[]
for(panel_index,width)in enumerate((96,80)):
rows=n-offset;v32=h.new_empty(batch,width,rows).transpose(1,2);v16=h.new_empty(batch,width,rows,dtype=torch.float16).transpose(1,2);panels[panel_index].launch((batch,1,1),(width//8*32,1,1),(source[:,offset:,offset:offset+width],h[:,offset:,offset:offset+width],tau[:,offset:offset+width],v32,v16),shared_mem=(rows*width+width)*4+width*8);raw_gram=torch.bmm(v16.mT,v16,out_dtype=torch.float32)
if width==96:gram=raw_gram
else:gram=torch.zeros((batch,96,96),device=seed.device);gram[:,:width,:width]=raw_gram
triangular=torch.empty_like(gram);t96.launch((batch,1,1),(512,1,1),(gram,tau[:,offset:offset+96],triangular,int(tau.stride(0))),shared_mem=36864);middle=gram[:,64:,:64]@triangular[:,:64,:64];torch.baddbmm(triangular[:,64:,:64],triangular[:,64:,64:],middle,beta=.0,alpha=-1.,out=triangular[:,64:,:64]);triangular=triangular[:,:width,:width];nodes.append((offset,v32,triangular))
if offset+width<active_cols:trailing_input=source[:,offset:,offset+width:active_cols];trailing_output=h[:,offset:,offset+width:active_cols];transformed=triangular@(v32.mT@trailing_input);torch.baddbmm(trailing_input,v16,transformed.half(),beta=1.,alpha=-1.,out=trailing_output,out_dtype=torch.float32)
source=h;offset+=width
q=_e162_batched_identity512(batch,seed.device)
for(offset,panel,triangular)in reversed(nodes):
if offset:active=q[:,offset:,offset:];transformed=triangular.mT@panel.mT;torch.baddbmm(active,panel,transformed,beta=1.,alpha=-1.,out=active)
else:active=q[:,offset:,:];transformed=triangular.mT@(panel.mT@active);torch.baddbmm(active,panel,transformed,beta=1.,alpha=-1.,out=active)
return q
@torch.no_grad()
def _clustered_qr_active176(seed:torch.Tensor):return _clustered_qr_active176_compact(seed)
@torch.no_grad()
def lapack_even_sparse_rr(data:torch.Tensor,*,repair_rank:int=32,guard_rank:int=96,guard_fraction:int=20,root_sign:int=4,root_range:int=12,root_reorth:int=6,child_sign:int=3,child_range:int=20,child_reorth:int=4,root_center:str='zero_linear',child_center:str='uniform_newton',known_bounds:bool=False,newton_precision:str='high',center_iterations:int=7,child_center_iterations:int=7,lanczos_steps:int=5,lanczos_probes:int=2,child_lanczos_steps:int|None=None,child_lanczos_probes:int|None=None,power_range:bool=False,child_power_range:bool=False,leaf_eigh=None):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');q,_=hybrid_eigh_512(data,use_fast_cholesky=True,split_levels=2,polar_sign_iterations=root_sign,polar_range_iterations=root_range,fast_cholesky_reorthogonalize_every=root_reorth,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,newton_precision=newton_precision,fast_lanczos_steps=lanczos_steps,fast_lanczos_probes=lanczos_probes,fused_sign_update=True,fast_center_mode=root_center,fast_center_probe_iterations=center_iterations,child_sign_iterations=child_sign,child_range_iterations=child_range,child_reorthogonalize_every=child_reorth,boundary_refine_width=0,post_refine_width=0,child_center_mode=child_center,child_center_probe_iterations=child_center_iterations,block_inverse_precision='high',child_lanczos_steps=child_lanczos_steps,child_lanczos_probes=child_lanczos_probes,known_lapack_bounds=known_bounds,power_range_iteration=power_range,child_power_range_iteration=child_power_range,leaf_eigh=leaf_eigh,finalize=False);batch,n,_=q.shape;aq=data@q;projected=q.mT@aq;values=projected.diagonal(dim1=-2,dim2=-1);offdiag=projected-torch.diag_embed(values);row_energy=offdiag.square().sum(dim=-1);top=row_energy.topk(repair_rank,dim=-1);indices=top.indices;selected=q.gather(2,indices[:,None,:].expand(-1,n,-1));rows=projected.gather(1,indices[:,:,None].expand(-1,-1,n));local=rows.gather(2,indices[:,None,:].expand(-1,repair_rank,-1));local_values,rotation=torch.linalg.eigh(local);repaired=selected@rotation;torch.set_float32_matmul_precision('highest');eye=torch.eye(repair_rank,device=data.device).expand(batch,-1,-1);repaired=.5*repaired@(3.*eye-repaired.mT@repaired);torch.set_float32_matmul_precision('high');base_q=q;q=q.clone();q.scatter_(2,indices[:,None,:].expand(-1,n,-1),repaired);values=values.clone();values.scatter_(1,indices,local_values)
if guard_rank>repair_rank and guard_fraction>0:total_energy=row_energy.sum(dim=-1);remaining=total_energy-top.values.sum(dim=-1);guard_count=max(1,(batch+guard_fraction-1)//guard_fraction);guard_count=min(guard_count,batch);tail_rows=remaining.topk(guard_count).indices;total_rows=total_energy.topk(guard_count).indices;guard_mask=torch.zeros(batch,device=data.device,dtype=torch.bool);guard_mask[tail_rows]=True;guard_mask[total_rows]=True;guard_rows=torch.nonzero(guard_mask,as_tuple=False).flatten();guard_q=base_q[guard_rows];guard_projected=projected[guard_rows];guard_values=guard_projected.diagonal(dim1=-2,dim2=-1);guard_energy=(guard_projected-torch.diag_embed(guard_values)).square().sum(dim=-1);guard_indices=guard_energy.topk(guard_rank,dim=-1).indices;selected=guard_q.gather(2,guard_indices[:,None,:].expand(-1,n,-1));rows=guard_projected.gather(1,guard_indices[:,:,None].expand(-1,-1,n));local=rows.gather(2,guard_indices[:,None,:].expand(-1,guard_rank,-1));local_values,rotation=torch.linalg.eigh(local);guarded=selected@rotation;torch.set_float32_matmul_precision('highest');guard_eye=torch.eye(guard_rank,device=data.device).expand(guard_rows.numel(),-1,-1);guarded=.5*guarded@(3.*guard_eye-guarded.mT@guarded);torch.set_float32_matmul_precision('high');guard_q=guard_q.clone();guard_q.scatter_(2,guard_indices[:,None,:].expand(-1,n,-1),guarded);guard_values=guard_values.clone();guard_values.scatter_(1,guard_indices,local_values);q[guard_rows]=guard_q;values[guard_rows]=guard_values
values,order=values.sort(dim=1);q=q.gather(2,order[:,None,:].expand(-1,n,-1));torch.set_float32_matmul_precision(previous);return q,values
@triton.jit
def _e483_half_sign_projector_kernel(sign,output,total,n:tl.constexpr,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;element=offsets%(n*n);row=element//n;column=element-row*n;value=tl.load(sign+offsets,mask=mask,other=.0).to(tl.float32);identity=(row==column).to(tl.float32);tl.store(output+offsets,.5*(identity-value),mask=mask)
@triton.jit
def _e485_float_sign_projector_kernel(sign,output,total,n:tl.constexpr,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;element=offsets%(n*n);row=element//n;column=element-row*n;value=tl.load(sign+offsets,mask=mask,other=.0);identity=(row==column).to(tl.float32);tl.store(output+offsets,.5*(identity-value),mask=mask)
@triton.jit
def _e483_shifted_sign_projector_kernel(sign,shift,output,total,n:tl.constexpr,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;matrix=offsets//(n*n);element=offsets-matrix*n*n;row=element//n;column=element-row*n;identity=(row==column).to(tl.float32);sign_value=tl.load(sign+offsets,mask=mask,other=.0);shift_value=tl.load(shift+matrix,mask=mask,other=.0);shifted=sign_value-shift_value*identity;tl.store(output+offsets,.5*(identity-shifted),mask=mask)
@torch.no_grad()
def _e483_half_sign_projector(sign:torch.Tensor):output=torch.empty(sign.shape,device=sign.device,dtype=torch.float32);total=output.numel();_e483_half_sign_projector_kernel[triton.cdiv(total,4096),](sign,output,total,n=sign.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _e485_float_sign_projector(sign:torch.Tensor):output=torch.empty_like(sign);total=output.numel();_e485_float_sign_projector_kernel[triton.cdiv(total,4096),](sign,output,total,n=sign.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _e483_shifted_sign_projector(sign:torch.Tensor,shift:torch.Tensor):output=torch.empty_like(sign);total=output.numel();_e483_shifted_sign_projector_kernel[triton.cdiv(total,4096),](sign,shift,output,total,n=sign.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@triton.jit
def _shift_scale_symmetric_kernel(matrix,center,radius,output,total,n:tl.constexpr,BLOCK:tl.constexpr):offset=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offset<total;matrix_id=offset//(n*n);local=offset-matrix_id*n*n;row=local//n;column=local-row*n;value=tl.load(matrix+offset,mask=mask,other=.0);local_center=tl.load(center+matrix_id,mask=mask,other=.0);local_radius=tl.load(radius+matrix_id,mask=mask,other=1.);shifted=value-tl.where(row==column,local_center,.0);tl.store(output+offset,shifted/local_radius,mask=mask)
@torch.no_grad()
def _shift_scale_symmetric(matrix:torch.Tensor,center:torch.Tensor,radius:torch.Tensor):output=matrix.clone();output.diagonal(dim1=-2,dim2=-1).sub_(center[:,None]);return output/radius[:,None,None]
@torch.no_grad()
def _shift_scale_symmetric_half(matrix:torch.Tensor,center:torch.Tensor,radius:torch.Tensor):output=torch.empty_like(matrix,dtype=torch.float16);total=matrix.numel();_shift_scale_symmetric_kernel[triton.cdiv(total,4096),](matrix,center,radius,output,total,n=matrix.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@triton.jit
def _identity_minus_scaled_kernel(matrix,scale,output,total,n:tl.constexpr,BLOCK:tl.constexpr):offset=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offset<total;matrix_id=offset//(n*n);local=offset-matrix_id*n*n;row=local//n;column=local-row*n;value=tl.load(matrix+offset,mask=mask,other=.0);local_scale=tl.load(scale+matrix_id,mask=mask,other=1.);identity=tl.where(row==column,1.,.0);tl.store(output+offset,identity-value/local_scale,mask=mask)
@torch.no_grad()
def _identity_minus_scaled(matrix:torch.Tensor,scale:torch.Tensor):output=torch.empty_like(matrix);total=matrix.numel();_identity_minus_scaled_kernel[triton.cdiv(total,4096),](matrix,scale,output,total,n=matrix.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _identity_minus_scaled_half(matrix:torch.Tensor,scale:torch.Tensor):output=torch.empty_like(matrix,dtype=torch.float16);total=matrix.numel();_identity_minus_scaled_kernel[triton.cdiv(total,4096),](matrix,scale,output,total,n=matrix.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@triton.jit
def _scale_matrix_half_kernel(matrix,scale,output,total,coefficient:tl.constexpr,n:tl.constexpr,BLOCK:tl.constexpr):offset=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offset<total;matrix_id=offset//(n*n);value=tl.load(matrix+offset,mask=mask,other=.0);local_scale=tl.load(scale+matrix_id,mask=mask,other=1.);tl.store(output+offset,value/(coefficient*local_scale),mask=mask)
@torch.no_grad()
def _scale_matrix_half(matrix:torch.Tensor,scale:torch.Tensor,*,coefficient:float=1.):output=torch.empty_like(matrix,dtype=torch.float16);total=matrix.numel();_scale_matrix_half_kernel[triton.cdiv(total,4096),](matrix,scale,output,total,coefficient=coefficient,n=matrix.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@triton.jit
def _e3123_lapack_nonic_tail_kernel(square,fourth,tail,total:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);mask=offsets<total;x2=tl.load(square+offsets,mask=mask,other=.0);x4=tl.load(fourth+offsets,mask=mask,other=.0);septic=tl.full(x2.shape,-4.5625,tl.float32);nonic=tl.full(x2.shape,1.0625,tl.float32);value=tl.inline_asm_elementwise(asm='mul.rn.f32 $0, $1, $2;\nfma.rn.f32 $0, $3, $4, $0;',constraints='=f,f,f,f,f',args=[septic,x2,nonic,x4],dtype=tl.float32,is_pure=True,pack=1);tl.store(tail+offsets,value,mask=mask)
@triton.jit
def _e3123_lapack_nonic_finish_kernel(polynomial,square,fourth,total:tl.constexpr,n:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);mask=offsets<total;element=offsets%(n*n);row=element//n;column=element-row*n;value=tl.load(polynomial+offsets,mask=mask,other=.0);x2=tl.load(square+offsets,mask=mask,other=.0);x4=tl.load(fourth+offsets,mask=mask,other=.0);cubic=tl.full(x2.shape,-6.4375,tl.float32);quintic=tl.full(x2.shape,7.6875,tl.float32);diagonal=tl.where(row==column,3.25,.0).to(tl.float32);value=tl.inline_asm_elementwise(asm='fma.rn.f32 $0, $2, $3, $1;\nfma.rn.f32 $0, $4, $5, $0;\nadd.rn.f32 $0, $0, $6;',constraints='=f,f,f,f,f,f,f',args=[value,cubic,x2,quintic,x4,diagonal],dtype=tl.float32,is_pure=True,pack=1);tl.store(polynomial+offsets,value,mask=mask)
@torch.no_grad()
def _e3123_lapack_root_nonic_sign_step(sign:torch.Tensor):square=torch.bmm(sign,sign);fourth=torch.bmm(square,square);tail=torch.empty_like(square);total=square.numel();_e3123_lapack_nonic_tail_kernel[triton.cdiv(total,4096),](square,fourth,tail,total=total,block=4096,num_warps=8,num_stages=1);polynomial=torch.bmm(fourth,tail);_e3123_lapack_nonic_finish_kernel[triton.cdiv(total,4096),](polynomial,square,fourth,total=total,n=sign.shape[-1],block=4096,num_warps=8,num_stages=1);return torch.bmm(sign,polynomial)
@torch.no_grad()
def reuse_zero_root_split(matrix:torch.Tensor,*,cholesky_fn,normalize_fn,reuse_sign_iterations:int=4,shifted_refine_iterations:int=0,center_gain:float|None=None,polynomial_shift:bool=False,center_scale:float=1.,inertia_coefficients:tuple[float,...]|None=None,root_range_iterations:int=12,root_reorthogonalize_every:int=6,route_local_quintic_pair:bool=False,**_):
batch,n,_=matrix.shape
if n!=_LAPACK_ROOT_DIM:raise ValueError('reuse root is specialized for n512')
rank=256;sign=matrix/1.2;trace_states=[]
if route_local_quintic_pair:
if reuse_sign_iterations!=3 or inertia_coefficients is not None:raise ValueError('route-local quintic pair requires the three-step LAPACK root')
sign=_e3123_lapack_root_nonic_sign_step(sign);center_gain=3.25
else:
for _ in range(reuse_sign_iterations):trace_states.append(sign.diagonal(dim1=-2,dim2=-1).sum(dim=-1));square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
trace_states.append(sign.diagonal(dim1=-2,dim2=-1).sum(dim=-1))
if inertia_coefficients is None:signed_trace=sign.diagonal(dim1=-2,dim2=-1).sum(dim=-1)
else:
if len(inertia_coefficients)!=len(trace_states):raise ValueError('one inertia coefficient is required per sign state')
signed_trace=sum(coefficient*state for(coefficient,state)in zip(inertia_coefficients,trace_states))
zero_rank=.5*(n-signed_trace);center=(rank-zero_rank)*(1.1/rank)
if polynomial_shift:
shift=center_scale*center/1.2
for _ in range(reuse_sign_iterations):shift=1.5*shift-.5*shift.square()*shift
else:
if center_gain is None:center_gain=1.5**reuse_sign_iterations/1.2
shift=center_gain*center
if shifted_refine_iterations==0:low_projector=_e483_shifted_sign_projector(sign,shift)
else:
eye=torch.eye(n,device=matrix.device).expand(batch,-1,-1);sign=sign-shift[:,None,None]*eye
for _ in range(shifted_refine_iterations):square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
low_projector=.5*(eye-sign)
if root_range_iterations==12:low_projector=low_projector@low_projector;low_projector=low_projector@low_projector;range_steps=3;checkpoint_every=1
else:range_steps=root_range_iterations;checkpoint_every=root_reorthogonalize_every
if root_range_iterations==12:
low=_e1611_normalize_copy_strided(low_projector[:,:,:rank])
for iteration in range(1,range_steps):
low=low_projector@low
if iteration+1==range_steps:0
else:low=_e1233_lapack_active256_thin(low.contiguous())
else:
eye=torch.eye(n,device=matrix.device).expand(batch,-1,-1);low=eye[:,:,:rank].clone()
for iteration in range(range_steps):
low=low_projector@low
if iteration==0 or(iteration+1)%checkpoint_every==0:low=cholesky_fn(low,passes=1,ridge=1e-06,inverse_precision='high')
else:normalize_fn(low)
if root_range_iterations==12:basis=_lapack_active256_complete(low.contiguous())
else:high=_rankdef_householder_complement(low);basis=torch.cat((low,high),dim=-1)
product=basis.mT@matrix;children=torch.empty((2*batch,rank,rank),device=matrix.device,dtype=matrix.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=children[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=children[batch:]);return basis,children
_E959_LAPACK_INVERSE_LT128_NAME='e959_lapack_inverse_lt128'
def _e959_lapack_inverse_lt128_source():
source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 128;').replace('constexpr int PANEL = 16;','constexpr int PANEL = 32;').replace('constexpr int ROWS = 32;','constexpr int ROWS = 128;').replace('constexpr int PITCH = 33;','constexpr int PITCH = 129;').replace('right_trsm160_block16_rows32',_E959_LAPACK_INVERSE_LT128_NAME);old=' const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows ? source[(long long)row * N + column] : 0.0f;';new=' const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows && row == column ? 1.0f : 0.0f;'
if source.count(old)!=1:raise RuntimeError('E959 n128 TRSM initialization anchor changed')
return source.replace(old,new,1)
@memo(maxsize=1)
def _e959_lapack_inverse_lt128_kernel():source=_fast_only_cuda_kernel(_e959_lapack_inverse_lt128_source(),_E959_LAPACK_INVERSE_LT128_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E959_LAPACK_INVERSE_LT128_NAME),_E959_LAPACK_INVERSE_LT128_NAME)
@torch.no_grad()
def _e959_lapack_inverse_lt128(lower:torch.Tensor):return _e2369_tcgen_inverse_lt(lower)
@torch.no_grad()
def _lapack_child_cqr128(matrix:torch.Tensor):
gram=matrix.mT@matrix;lower=torch.empty_like(gram);_e095_potrf128_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,lower,matrix.shape[0],1e-06),shared_mem=_E095_POTRF128_SHARED_BYTES);inverse=_e959_lapack_inverse_lt128(lower);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:return matrix@inverse
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _lapack_child_shift_scale(matrix:torch.Tensor,center:torch.Tensor,radius:torch.Tensor):output=torch.empty_like(matrix);total=matrix.numel();_shift_scale_symmetric_kernel[triton.cdiv(total,4096),](matrix,center,radius,output,total,n=matrix.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _lapack_child_projector(sign:torch.Tensor):return _e485_float_sign_projector(sign)
@torch.no_grad()
def _lapack_even_child_one_sided256(matrix:torch.Tensor,low_precision_sign:bool=False,**_):
batch,n,_=matrix.shape
if n!=256 or batch%2:raise ValueError('LAPACK-even child requires paired n256 matrices')
rank=128;parent_batch=batch//2;center=matrix.diagonal(dim1=-2,dim2=-1).mean(dim=-1);lower=torch.cat((torch.full_like(center[:parent_batch],-1.),torch.full_like(center[parent_batch:],-.1)));upper=torch.cat((torch.full_like(center[:parent_batch],.1),torch.full_like(center[parent_batch:],1.)));radius=torch.maximum((lower-center).abs(),(upper-center).abs()).clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(matrix,center,radius)if low_precision_sign else _lapack_child_shift_scale(matrix,center,radius)
for _ in range(1):square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
sign=_quintic_sign_corrected_step(sign);square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5);low_operator=_e483_half_sign_projector(sign)if low_precision_sign else _lapack_child_projector(sign);low_operator=low_operator@low_operator;low_operator=low_operator@low_operator;low=_e1611_normalize_copy_strided(low_operator[:,:,:rank])
for iteration in range(1,4):
low=low_operator@low
if iteration!=3:low=_lapack_child_cqr128(low)
basis=_lapack_child_active128_complete(low.contiguous());product=basis.mT@matrix;children=torch.empty((2*batch,rank,rank),device=matrix.device,dtype=matrix.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=children[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=children[batch:]);return basis,children
_LAPACK_RR96_SELECT_PACK_NAME='e764_rr96_select_pack';_LAPACK_RR96_SELECT_PACK_SOURCE='\nextern "C" __global__ __launch_bounds__(512, 1)\nvoid e764_rr96_select_pack(\n const float* __restrict__ q,\n const float* __restrict__ projected,\n float* __restrict__ row_energy_out,\n float* __restrict__ top_values,\n long long* __restrict__ indices,\n float* __restrict__ selected,\n float* __restrict__ local,\n int batch) {\n constexpr int N = 512;\n constexpr int R = 96;\n const int matrix = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n if (matrix >= batch) return;\n const int lane = tid & 31;\n const int warp = tid >> 5;\n const long long matrix_base = (long long)matrix * N * N;\n __shared__ float energies[N];\n __shared__ int order[N];\n\n // Reduce sixteen projected rows per wave, retaining FP32 RN arithmetic.\n for (int wave = 0; wave < 32; ++wave) {\n const int row = wave * 16 + warp;\n float energy = 0.0f;\n #pragma unroll\n for (int column = lane; column < N; column += 32) {\n float value = projected[matrix_base + (long long)row * N + column];\n if (column == row) value = 0.0f;\n const float square = __fmul_rn(value, value);\n energy = __fadd_rn(energy, square);\n }\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n energy = __fadd_rn(\n energy, __shfl_down_sync(0xffffffffu, energy, offset));\n if (lane == 0) {\n energies[row] = energy;\n order[row] = row;\n row_energy_out[(long long)matrix * N + row] = energy;\n }\n __syncthreads();\n }\n\n // Full bitonic ordering preserves the exact top-96 set. Row index is a\n // deterministic tie break only within equal-energy selected columns.\n for (int width = 2; width <= N; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n const int peer = tid ^ stride;\n if (peer > tid) {\n const float left_value = energies[tid];\n const float right_value = energies[peer];\n const int left_index = order[tid];\n const int right_index = order[peer];\n const bool ascending = (tid & width) == 0;\n const bool greater = (left_value > right_value)\n || (left_value == right_value && left_index > right_index);\n const bool swap = ascending ? greater : !greater;\n if (swap) {\n energies[tid] = right_value;\n energies[peer] = left_value;\n order[tid] = right_index;\n order[peer] = left_index;\n }\n }\n __syncthreads();\n }\n }\n\n if (tid < R) {\n const int source = N - 1 - tid;\n top_values[(long long)matrix * R + tid] = energies[source];\n indices[(long long)matrix * R + tid] = (long long)order[source];\n }\n __syncthreads();\n\n const long long selected_base = (long long)matrix * N * R;\n for (int linear = tid; linear < N * R; linear += blockDim.x) {\n const int row = linear / R;\n const int rank = linear - row * R;\n const int column = order[N - 1 - rank];\n selected[selected_base + linear] =\n q[matrix_base + (long long)row * N + column];\n }\n const long long local_base = (long long)matrix * R * R;\n for (int linear = tid; linear < R * R; linear += blockDim.x) {\n const int local_row = linear / R;\n const int local_column = linear - local_row * R;\n const int row = order[N - 1 - local_row];\n const int column = order[N - 1 - local_column];\n local[local_base + linear] =\n projected[matrix_base + (long long)row * N + column];\n }\n}\n'
@memo(maxsize=1)
def _lapack_rr96_select_pack_kernel():return CUDAKernel(_fast_nvrtc_compile(_LAPACK_RR96_SELECT_PACK_SOURCE,_LAPACK_RR96_SELECT_PACK_NAME),_LAPACK_RR96_SELECT_PACK_NAME)
@torch.no_grad()
def _lapack_rr96_select_pack(q:torch.Tensor,projected:torch.Tensor):
batch,n,_=q.shape
if n!=512 or projected.shape!=q.shape:raise ValueError('rank96 select-pack requires batch x 512 x 512')
row_energy=torch.empty((batch,512),device=q.device);top_values=torch.empty((batch,96),device=q.device);indices=torch.empty((batch,96),device=q.device,dtype=torch.int64);selected=torch.empty((batch,512,96),device=q.device);local=torch.empty((batch,96,96),device=q.device);_lapack_rr96_select_pack_kernel().launch((batch,1,1),(512,1,1),(q,projected,row_energy,top_values,indices,selected,local,batch));return row_energy,top_values,indices,selected,local
_LAPACK_RR96_ROTATE_ORDER_NAME='e1001a_rr96_sparse_writeback';_LAPACK_RR96_ROTATE_ORDER_APPLY_NAME='e1001a_rr96_sparse_fallback_apply';_LAPACK_RR96_ROTATE_ORDER_SOURCE='\nextern "C" __global__ __launch_bounds__(256, 2)\nvoid e1001a_rr96_sparse_writeback(\n float* __restrict__ base_q,\n const float* __restrict__ projected,\n const float* __restrict__ repaired,\n const long long* __restrict__ indices,\n const float* __restrict__ local_values,\n float* __restrict__ output_values,\n int* __restrict__ output_order,\n int* __restrict__ output_slots,\n int* __restrict__ fallback_flags,\n int batch) {\n constexpr int N = 512;\n constexpr int R = 96;\n constexpr int P = 128;\n const int matrix = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n if (matrix >= batch) return;\n const long long q_base = (long long)matrix * N * N;\n const long long repair_base = (long long)matrix * N * R;\n const long long index_base = (long long)matrix * R;\n const long long map_base = (long long)matrix * N;\n\n __shared__ int targets[P];\n __shared__ int repair_slot[N];\n __shared__ float sort_values[N];\n __shared__ int sort_sources[N];\n __shared__ int direct_valid;\n\n if (tid < P)\n targets[tid] = tid < R ? (int)indices[index_base + tid] : N + tid;\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int column = tid + half * 256;\n repair_slot[column] = -1;\n sort_values[column] =\n projected[q_base + (long long)column * N + column];\n sort_sources[column] = column;\n }\n if (tid == 0) direct_valid = 1;\n __syncthreads();\n if (tid < R) {\n const int column = (int)indices[index_base + tid];\n repair_slot[column] = tid;\n }\n __syncthreads();\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int column = tid + half * 256;\n output_slots[map_base + column] = repair_slot[column];\n }\n\n // Sort the 96 selected target slots; padding exceeds every valid index.\n for (int width = 2; width <= P; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n if (tid < P) {\n const int peer = tid ^ stride;\n if (peer > tid) {\n const int left = targets[tid];\n const int right = targets[peer];\n const bool ascending = (tid & width) == 0;\n const bool swap = ascending ? left > right : left < right;\n if (swap) {\n targets[tid] = right;\n targets[peer] = left;\n }\n }\n }\n __syncthreads();\n }\n }\n\n if (tid < R)\n sort_values[targets[tid]] = local_values[index_base + tid];\n __syncthreads();\n\n // Strict ordering proves that selected-only placement is the exact full\n // stable sort. Invalid rows take the device-side full-sort fallback below.\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int column = tid + half * 256;\n if (column + 1 < N && !(sort_values[column] < sort_values[column + 1]))\n atomicExch(&direct_valid, 0);\n }\n __syncthreads();\n fallback_flags[matrix] = direct_valid ? 0 : 1;\n\n if (direct_valid) {\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int column = tid + half * 256;\n output_values[(long long)matrix * N + column] = sort_values[column];\n }\n for (int linear = tid; linear < N * R; linear += blockDim.x) {\n const int row = linear / R;\n const int repair_column = linear - row * R;\n const int target = targets[repair_column];\n base_q[q_base + (long long)row * N + target] =\n repaired[repair_base + (long long)row * R + repair_column];\n }\n return;\n }\n\n // Rare invalid row: reproduce the production source-associated bitonic\n // order. A row-tiled second kernel applies arbitrary cycles safely in\n // place through a shared row buffer.\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int column = tid + half * 256;\n const int slot = repair_slot[column];\n sort_values[column] = slot >= 0\n ? local_values[index_base + slot]\n : projected[q_base + (long long)column * N + column];\n sort_sources[column] = column;\n }\n __syncthreads();\n for (int width = 2; width <= N; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n const int group = tid / stride;\n const int offset = tid - group * stride;\n const int left = group * (2 * stride) + offset;\n const int right = left + stride;\n const float left_value = sort_values[left];\n const float right_value = sort_values[right];\n const int left_source = sort_sources[left];\n const int right_source = sort_sources[right];\n const bool ascending = (left & width) == 0;\n const bool greater = (left_value > right_value)\n || (left_value == right_value && left_source > right_source);\n const bool swap = ascending ? greater : !greater;\n if (swap) {\n sort_values[left] = right_value;\n sort_values[right] = left_value;\n sort_sources[left] = right_source;\n sort_sources[right] = left_source;\n }\n __syncthreads();\n }\n }\n\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int output_column = tid + half * 256;\n output_values[(long long)matrix * N + output_column] =\n sort_values[output_column];\n output_order[map_base + output_column] = sort_sources[output_column];\n }\n}\n\nextern "C" __global__ __launch_bounds__(256, 2)\nvoid e1001a_rr96_sparse_fallback_apply(\n float* __restrict__ base_q,\n const float* __restrict__ repaired,\n const int* __restrict__ order,\n const int* __restrict__ slots,\n const int* __restrict__ fallback_flags,\n int row_tiles,\n int batch) {\n constexpr int N = 512;\n constexpr int R = 96;\n const int program = (int)blockIdx.x;\n const int matrix = program / row_tiles;\n const int tile = program - matrix * row_tiles;\n const int tid = (int)threadIdx.x;\n if (matrix >= batch || fallback_flags[matrix] == 0) return;\n const int rows_per_tile = (N + row_tiles - 1) / row_tiles;\n const int row_begin = tile * rows_per_tile;\n const int row_end = min(N, row_begin + rows_per_tile);\n const long long q_base = (long long)matrix * N * N;\n const long long repair_base = (long long)matrix * N * R;\n const long long map_base = (long long)matrix * N;\n __shared__ float row_buffer[N];\n for (int row = row_begin; row < row_end; ++row) {\n row_buffer[tid] = base_q[q_base + (long long)row * N + tid];\n row_buffer[tid + 256] =\n base_q[q_base + (long long)row * N + tid + 256];\n __syncthreads();\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int output_column = tid + half * 256;\n const int source = order[map_base + output_column];\n const int slot = slots[map_base + source];\n base_q[q_base + (long long)row * N + output_column] = slot >= 0\n ? repaired[repair_base + (long long)row * R + slot]\n : row_buffer[source];\n }\n __syncthreads();\n }\n}\n'
@memo(maxsize=1)
def _lapack_rr96_rotate_order_kernel():return CUDAKernel(_fast_nvrtc_compile(_LAPACK_RR96_ROTATE_ORDER_SOURCE,_LAPACK_RR96_ROTATE_ORDER_NAME),_LAPACK_RR96_ROTATE_ORDER_NAME)
@memo(maxsize=1)
def _lapack_rr96_rotate_order_apply_kernel():return CUDAKernel(_fast_nvrtc_compile(_LAPACK_RR96_ROTATE_ORDER_SOURCE,_LAPACK_RR96_ROTATE_ORDER_APPLY_NAME),_LAPACK_RR96_ROTATE_ORDER_APPLY_NAME)
@torch.no_grad()
def _lapack_rr96_rotate_order(base_q:torch.Tensor,projected:torch.Tensor,repaired:torch.Tensor,indices:torch.Tensor,local_values:torch.Tensor):
batch,n,_=base_q.shape
if n!=512 or projected.shape!=base_q.shape or repaired.shape!=(batch,512,96)or indices.shape!=(batch,96)or local_values.shape!=(batch,96):raise ValueError('rank96 rotate-order requires fixed n512 layouts')
output_values=torch.empty((batch,512),device=base_q.device);order=torch.empty((batch,512),device=base_q.device,dtype=torch.int32);slots=torch.empty_like(order);fallback_flags=torch.empty((batch,),device=base_q.device,dtype=torch.int32);_lapack_rr96_rotate_order_kernel().launch((batch,1,1),(256,1,1),(base_q,projected,repaired,indices,local_values,output_values,order,slots,fallback_flags,batch));row_tiles=64;_lapack_rr96_rotate_order_apply_kernel().launch((batch*row_tiles,1,1),(256,1,1),(base_q,repaired,order,slots,fallback_flags,row_tiles,batch));return base_q,output_values
@torch.no_grad()
def lapack_even_manual(data:torch.Tensor,*,split_fn,root_split_fn=None,leaf_eigh_fn=None,repair_rank:int=64,guard_rank:int=96,guard_fraction:int=20,child_sign:int=4,child_range:int=20,child_power_range:bool=True,return_risk:bool=False,child_low_precision_sign:bool=False):
if leaf_eigh_fn is None:leaf_eigh_fn=torch.linalg.eigh
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');batch,n,_=data.shape
if root_split_fn is None:root_split_fn=split_fn
root_basis,children=root_split_fn(data,sign_iterations=3,range_iterations=12,reorthogonalize_every=6,lanczos_steps=1,lanczos_probes=1,range_cholesky_passes=1,fused_sign_update=True,final_cholesky_passes=1,skip_final_high_checkpoint=True,center_mode='zero_linear',center_probe_iterations=4,block_inverse_precision='high',known_lapack_bounds=True);child_basis,leaves=split_fn(children,sign_iterations=child_sign,range_iterations=child_range,reorthogonalize_every=4,lanczos_steps=1,lanczos_probes=1,range_cholesky_passes=1,fused_sign_update=True,final_cholesky_passes=1,skip_final_high_checkpoint=True,center_mode='trace',center_probe_iterations=1,block_inverse_precision='high',known_lapack_bounds=True,power_range_iteration=child_power_range,low_precision_sign=child_low_precision_sign);_,leaf_vectors=leaf_eigh_fn(leaves);child_vectors=torch.empty_like(children);torch.bmm(child_basis[:,:,:128],leaf_vectors[:2*batch],out=child_vectors[:,:,:128]);torch.bmm(child_basis[:,:,128:],leaf_vectors[2*batch:],out=child_vectors[:,:,128:]);q=torch.empty_like(data);torch.bmm(root_basis[:,:,:256],child_vectors[:batch],out=q[:,:,:256]);torch.bmm(root_basis[:,:,256:],child_vectors[batch:],out=q[:,:,256:]);aq=data@q;projected=q.mT@aq;values=projected.diagonal(dim1=-2,dim2=-1)
if n==512 and repair_rank==96:row_energy,top_values,indices,selected,local=_lapack_rr96_select_pack(q,projected)
else:offdiag=projected-torch.diag_embed(values);row_energy=offdiag.square().sum(dim=-1);top=row_energy.topk(repair_rank,dim=-1);top_values=top.values;indices=top.indices;gather=indices[:,None,:].expand(-1,n,-1);selected=q.gather(2,gather);rows=projected.gather(1,indices[:,:,None].expand(-1,-1,n));local=rows.gather(2,indices[:,None,:].expand(-1,repair_rank,-1))
gather=indices[:,None,:].expand(-1,n,-1)
if repair_rank==96:rotation,local_values=_small_n96_eigh(local)
else:local_values,rotation=torch.linalg.eigh(local)
repaired=selected@rotation;fixed_rank96=n==512 and repair_rank==96 and guard_rank==96 and guard_fraction==0
if fixed_rank96:q,values=_lapack_rr96_rotate_order(q,projected,repaired,indices,local_values)
else:
base_q=q;q=q.clone();q.scatter_(2,gather,repaired);values=values.clone();values.scatter_(1,indices,local_values)
if guard_rank>repair_rank and guard_fraction>0:total_energy=row_energy.sum(dim=-1);remaining=total_energy-top_values.sum(dim=-1);count=min(batch,max(1,(batch+guard_fraction-1)//guard_fraction));tail_rows=remaining.topk(count).indices;total_rows=total_energy.topk(count).indices;mask=torch.zeros(batch,device=data.device,dtype=torch.bool);mask[tail_rows]=True;mask[total_rows]=True;guard_rows=torch.nonzero(mask,as_tuple=False).flatten();guard_q=base_q[guard_rows];guard_projected=projected[guard_rows];guard_values=guard_projected.diagonal(dim1=-2,dim2=-1);energy=(guard_projected-torch.diag_embed(guard_values)).square().sum(dim=-1);guard_indices=energy.topk(guard_rank,dim=-1).indices;guard_gather=guard_indices[:,None,:].expand(-1,n,-1);selected=guard_q.gather(2,guard_gather);rows=guard_projected.gather(1,guard_indices[:,:,None].expand(-1,-1,n));local=rows.gather(2,guard_indices[:,None,:].expand(-1,guard_rank,-1));local_values,rotation=torch.linalg.eigh(local);repaired=selected@rotation;torch.set_float32_matmul_precision('highest');eye=torch.eye(guard_rank,device=data.device).expand(guard_rows.numel(),-1,-1);repaired=.5*repaired@(3.*eye-repaired.mT@repaired);torch.set_float32_matmul_precision('high');local_q=guard_q.clone();local_q.scatter_(2,guard_gather,repaired);local_values_out=guard_values.clone();local_values_out.scatter_(1,guard_indices,local_values);q[guard_rows]=local_q;values[guard_rows]=local_values_out
values,order=values.sort(dim=1);q=q.gather(2,order[:,None,:].expand(-1,n,-1))
torch.set_float32_matmul_precision(previous)
if return_risk:scale=values.abs().amax(dim=1).clamp_min(1e-20).square();total=row_energy.sum(dim=-1)/scale;remaining=(row_energy.sum(dim=-1)-top_values.sum(dim=-1))/scale;return q,values,torch.stack((total,remaining),dim=1)
return q,values
def _dense_scaled_eigh(data:torch.Tensor):torch.set_float32_matmul_precision('high');return hybrid_eigh_512(data,use_fast_cholesky=True,split_levels=2,polar_sign_iterations=7,polar_range_iterations=10,fast_cholesky_reorthogonalize_every=5,child_sign_iterations=9,child_range_iterations=12,child_reorthogonalize_every=6,boundary_refine_width=24,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,newton_precision='high',fast_lanczos_steps=3,fast_lanczos_probes=2,child_lanczos_steps=5,child_lanczos_probes=2,fused_sign_update=True)
def _cluster_probe_residual(data:torch.Tensor,alpha:torch.Tensor):probes=_rademacher_probes(512,1,data.device)[None].expand(data.shape[0],-1,-1);twice=data@(data@probes);residual=twice-alpha[:,None,None]*probes;return residual.square().sum(dim=(-2,-1))/twice.square().sum(dim=(-2,-1)).clamp_min(1e-30)
@torch.no_grad()
def _rankdef_right_solve128(matrix:torch.Tensor,lower:torch.Tensor):output=torch.empty_like(matrix);rows=matrix.shape[1];_e095_trsm128_kernel().launch((matrix.shape[0],(rows+31)//32,1),(128,1,1),(matrix,lower,output,matrix.shape[0],rows),shared_mem=128*33*4);return output
@torch.no_grad()
def _rankdef_right_solve96(matrix:torch.Tensor,lower:torch.Tensor):output=torch.empty_like(matrix);rows=matrix.shape[1];_rankdef_trsm96_kernel().launch((matrix.shape[0],(rows+31)//32,1),(128,1,1),(matrix,lower,output,matrix.shape[0],rows),shared_mem=96*33*4);return output
@torch.no_grad()
def _rankdef_right_solve192(matrix:torch.Tensor,lower:torch.Tensor):output=torch.empty_like(matrix);rows=matrix.shape[1];_rankdef_trsm192_kernel().launch((matrix.shape[0],(rows+31)//32,1),(256,1,1),(matrix,lower,output,matrix.shape[0],rows),shared_mem=192*33*4);return output
@torch.no_grad()
def _rankdef_cqr192(matrix:torch.Tensor,*,ridge:float=1e-06,gram_precision:str='highest'):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=torch.empty_like(gram);_rankdef_potrf192_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=(_RANKDEF_N192_TRI+1)*4);inverse_transpose=_rankdef_inverse_lt192(lower);torch.set_float32_matmul_precision('high')
try:return matrix@inverse_transpose
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _e1038_rankdef_root_cqr192_inverse(matrix:torch.Tensor,*,ridge:float=1e-06,gram_precision:str='highest'):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=torch.empty_like(gram);_rankdef_root_wmma_potrf192_kernel().launch((matrix.shape[0],1,1),(640,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=(192*192+1)*4);inverse_transpose=_rankdef_inverse_lt192(lower);torch.set_float32_matmul_precision('high')
try:return matrix@inverse_transpose
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _rankdef_cqr96(matrix:torch.Tensor,*,ridge:float=1e-06,gram_precision:str='high'):previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=torch.empty_like(gram);_rankdef_potrf96_kernel().launch((matrix.shape[0],1,1),(128,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=(_RANKDEF_N96_TRI+1)*4);return _rankdef_right_solve96(matrix,lower)
@torch.no_grad()
def _rankdef_cqr128(matrix:torch.Tensor,*,ridge:float,gram_precision:str):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);inverse=torch.empty_like(gram);_e1994_potrf_inverse128_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,inverse,matrix.shape[0],float(ridge)),shared_mem=_E1994_POTRF_INVERSE128_SHARED_BYTES);torch.set_float32_matmul_precision('high')
try:return matrix@inverse
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _rankdef_cqr128_parent512(matrix:torch.Tensor,*,ridge:float,gram_precision:str):previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=torch.empty_like(gram);_e095_potrf128_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=_E095_POTRF128_SHARED_BYTES);output=torch.empty(matrix.shape,device=matrix.device,dtype=matrix.dtype);_rankdef_trsm128_parent512_kernel().launch((matrix.shape[0],16,1),(128,1,1),(matrix,lower,output,matrix.shape[0],512),shared_mem=128*33*4);return output
@torch.no_grad()
def _e989_rankdef_cqr128_inverse(matrix:torch.Tensor,*,ridge:float,gram_precision:str):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=torch.empty_like(gram);_e095_potrf128_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=_E095_POTRF128_SHARED_BYTES);inverse=_e959_lapack_inverse_lt128(lower);torch.set_float32_matmul_precision('highest')
try:return matrix@inverse
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _rankdef_complement_assemble_kernel(product,transform,output,total,n:tl.constexpr,rank:tl.constexpr,complement:tl.constexpr,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;matrix=offsets//(n*complement);element=offsets-matrix*n*complement;row=element//complement;column=element-row*complement;value=-tl.load(product+offsets,mask=mask,other=.0);head=tl.load(transform+matrix*rank*complement+row*complement+column,mask=mask&(row<rank),other=.0);value-=head;value+=(row==rank+column).to(tl.float32);tl.store(output+offsets,value,mask=mask)
@torch.no_grad()
def _rankdef_assemble_fixed_complement(frame:torch.Tensor,transform:torch.Tensor):batch,n,rank=frame.shape;complement=n-rank;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');product=frame@transform;torch.set_float32_matmul_precision(previous);output=torch.empty((batch,n,complement),device=frame.device,dtype=frame.dtype);total=output.numel();_rankdef_complement_assemble_kernel[triton.cdiv(total,4096),](product,transform,output,total,n=n,rank=rank,complement=complement,BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _rankdef_householder_complement(frame:torch.Tensor,*,wmma160:bool=False):
batch,n,rank=frame.shape;eye_rank=torch.eye(rank,device=frame.device,dtype=frame.dtype).expand(batch,-1,-1);top=frame[:,:rank,:];bottom=frame[:,rank:,:];system=eye_rank+top.mT;right_hand_side=bottom.mT
if rank==96:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');normal=system.mT@system;normal_rhs=system.mT@right_hand_side;lower=torch.empty_like(normal);_rankdef_potrf96_kernel().launch((batch,1,1),(128,1,1),(normal,lower,batch,.0),shared_mem=(_RANKDEF_N96_TRI+1)*4);inverse_transpose=_rankdef_right_solve96(eye_rank.clone(),lower);transform=inverse_transpose@(inverse_transpose.mT@normal_rhs);torch.set_float32_matmul_precision(previous)
elif rank==128:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');normal=system.mT@system;normal_rhs=system.mT@right_hand_side;lower=torch.empty_like(normal);_e095_potrf128_kernel().launch((batch,1,1),(256,1,1),(normal,lower,batch,.0),shared_mem=_E095_POTRF128_SHARED_BYTES);inverse_transpose=_e959_lapack_inverse_lt128(lower);transform=inverse_transpose@(inverse_transpose.mT@normal_rhs);torch.set_float32_matmul_precision(previous)
elif rank==160:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');normal=system.mT@system;normal_rhs=system.mT@right_hand_side;lower=_e2409_potrf160_wmma(normal,.0)if wmma160 else _e083_potrf160(normal,.0);inverse_transpose=_e2369_tcgen_trsm160(eye_rank.clone(),lower);transform=inverse_transpose@(inverse_transpose.mT@normal_rhs);torch.set_float32_matmul_precision(previous)
elif rank==192:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');normal=system.mT@system;normal_rhs=system.mT@right_hand_side;lower=torch.empty_like(normal);_rankdef_potrf192_kernel().launch((batch,1,1),(256,1,1),(normal,lower,batch,.0),shared_mem=(_RANKDEF_N192_TRI+1)*4);inverse_transpose=_rankdef_right_solve192(eye_rank.clone(),lower);transform=inverse_transpose@(inverse_transpose.mT@normal_rhs);torch.set_float32_matmul_precision(previous)
elif rank==256:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');eye=torch.eye(n,device=frame.device,dtype=frame.dtype).expand(batch,-1,-1);complement=eye[:,:,rank:]-frame@bottom.mT;complement=torch.baddbmm(complement,frame,frame.mT@complement,beta=1.,alpha=-1.);complement=cholesky_orthonormalize(complement,passes=1,ridge=1e-07,inverse_precision='high');torch.set_float32_matmul_precision(previous);return complement
else:transform=torch.linalg.solve(system,right_hand_side)
if rank in(96,128,192):return _rankdef_assemble_fixed_complement(frame,transform)
eye=torch.eye(n,device=frame.device,dtype=frame.dtype).expand(batch,-1,-1);return eye[:,:,rank:]-(eye[:,:,:rank]+frame)@transform
@torch.no_grad()
def _rankdef_positive384_root_split_one_sided(data:torch.Tensor):
batch,n,_=data.shape
if n!=384:raise ValueError('one-sided rankdef root expects n=384')
rank=192;trace_center=data.diagonal(dim1=-2,dim2=-1).mean(dim=-1);ratio=1e1**(1./383.);mean=.1*(ratio**384-1.)/((ratio-1.)*384.);scale=trace_center/mean;lower=.1*scale;upper=scale;center=.8081902783518025*trace_center;radius=torch.maximum((lower-center).abs(),(upper-center).abs()).clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(data,center,radius)
for _ in range(2):square=torch.bmm(sign,sign);sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
sign=_quintic_sign_corrected_step(sign);square=torch.bmm(sign,sign);sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5);low_projector=_e483_half_sign_projector(sign);low=low_projector[:,:,:rank];operator=low_projector@low_projector
for _ in range(2):low=normalize_columns_(operator@low)
low=_e1038_rankdef_root_cqr192_inverse(low_projector@low,ridge=1e-06,gram_precision='high')
for iteration in range(3):
low=operator@low
if iteration!=2:low=normalize_columns_(low)
split_basis=_rankdef_small_active_complete(low.contiguous());product=data@split_basis;children=torch.empty((2*batch,rank,rank),device=data.device,dtype=data.dtype);torch.bmm(split_basis[:,:,:rank].mT,product[:,:,:rank],out=children[:batch]);torch.bmm(split_basis[:,:,rank:].mT,product[:,:,rank:],out=children[batch:]);return split_basis,children
@torch.no_grad()
def _rankdef_positive192_child_split_one_sided(data:torch.Tensor):
batch,n,_=data.shape
if n!=192 or batch%2:raise ValueError('one-sided rankdef child expects paired n=192')
rank=96;parent_batch=batch//2;trace_center=data.diagonal(dim1=-2,dim2=-1).mean(dim=-1);ratio=1e1**(1./383.);low_min=.1;low_max=.1*ratio**191;high_min=.1*ratio**192;high_max=1.;low_mean=.1*(ratio**192-1.)/((ratio-1.)*192.);high_mean=high_min*(ratio**192-1.)/((ratio-1.)*192.);low_scale=trace_center[:parent_batch]/low_mean;high_scale=trace_center[parent_batch:]/high_mean;lower=torch.cat((low_min*low_scale,high_min*high_scale));upper=torch.cat((low_max*low_scale,high_max*high_scale));center=.9465730329220167*trace_center;radius=torch.maximum((lower-center).abs(),(upper-center).abs()).clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(data,center,radius)
for _ in range(1):square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
sign=_quintic_sign_corrected_step(sign);square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5);low_operator=_e483_half_sign_projector(sign);low_operator=low_operator@low_operator;low_operator=low_operator@low_operator;low=_e1611_normalize_copy_strided(low_operator[:,:,:rank]);low=_rankdef_cqr96(low_operator@low,ridge=1e-06,gram_precision='high');low=low_operator@low;basis=_rankdef_small_active_complete(low.contiguous());product=data@basis;leaves=torch.empty((2*batch,rank,rank),device=data.device,dtype=data.dtype);torch.bmm(basis[:,:,:rank].mT,product[:,:,:rank],out=leaves[:batch]);torch.bmm(basis[:,:,rank:].mT,product[:,:,rank:],out=leaves[batch:]);return basis,leaves
@torch.no_grad()
def _rankdef_positive384_residual_stats(vectors:torch.Tensor,products:torch.Tensor):batch=vectors.shape[0];values=torch.empty((batch,384),device=vectors.device,dtype=torch.float32);residual_energy=torch.empty_like(values);_rankdef_positive384_residual_stats_kernel[batch,triton.cdiv(384,32)](vectors,products,values,residual_energy,n=384,block_n=512,block_cols=32,num_warps=8,num_stages=1);return values,residual_energy
@triton.jit
def _rankdef_positive384_residual_stats_kernel(vectors,products,values,residual_energy,n:tl.constexpr,block_n:tl.constexpr,block_cols:tl.constexpr):batch=tl.program_id(0);group=tl.program_id(1);rows=tl.arange(0,block_n)[:,None];column_ids=group*block_cols+tl.arange(0,block_cols);columns=column_ids[None,:];mask=(rows<n)&(columns<n);offsets=batch*n*n+rows*n+columns;q=tl.load(vectors+offsets,mask=mask,other=.0).to(tl.float32);aq=tl.load(products+offsets,mask=mask,other=.0).to(tl.float32);rayleigh=tl.sum(q*aq,axis=0);residual=aq-q*rayleigh[None,:];output_offsets=batch*n+column_ids;output_mask=column_ids<n;tl.store(values+output_offsets,rayleigh,mask=output_mask);tl.store(residual_energy+output_offsets,tl.sum(residual*residual,axis=0),mask=output_mask)
@torch.no_grad()
def _rankdef_positive384_fast(data:torch.Tensor):batch=data.shape[0];root_basis,children=_rankdef_positive384_root_split_one_sided(data);child_basis,leaves=_rankdef_positive192_child_split_one_sided(children);leaf_vectors,_=_e204_rankdef_n96_eigh(leaves);child_vectors=torch.empty_like(child_basis);torch.bmm(child_basis[:,:,:96],leaf_vectors[:2*batch],out=child_vectors[:,:,:96]);torch.bmm(child_basis[:,:,96:],leaf_vectors[2*batch:],out=child_vectors[:,:,96:]);child_window=child_vectors[:,:,88:104];child_rotation,_=_e951_n16_eigh(child_window.mT@children@child_window);child_repaired=child_window@child_rotation;child_vectors[:,:,88:104]=child_repaired;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');vectors=torch.empty_like(root_basis);torch.bmm(root_basis[:,:,:192],child_vectors[:batch],out=vectors[:,:,:192]);torch.bmm(root_basis[:,:,192:],child_vectors[batch:],out=vectors[:,:,192:]);root_window=vectors[:,:,176:208];root_rotation,_=_e951_n32_eigh(root_window.mT@data@root_window);root_repaired=root_window@root_rotation;vectors[:,:,176:208]=root_repaired;torch.set_float32_matmul_precision(previous);aq=data@vectors;values,residual_energy=_rankdef_positive384_residual_stats(vectors,aq);indices=residual_energy.topk(16,dim=-1).indices;gather=indices[:,None,:].expand(-1,384,-1);selected=vectors.gather(2,gather);selected_products=aq.gather(2,gather);local=selected.mT@selected_products;local=.5*(local+local.mT);rotation,local_values=_e951_n16_eigh(local);repaired=selected@rotation;vectors.scatter_(2,gather,repaired);values.scatter_(1,indices,local_values);values,order=values.sort(dim=1);vectors=_e160_gather_columns(vectors,order);return vectors,values
def _rankdef_eigh_512_guarded(data:torch.Tensor):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
batch,n,_=data.shape;ratio=1e1**(1./383.);trace_constant=.1*(ratio**384-1.)/(ratio-1.);upper=data.diagonal(dim1=-2,dim2=-1).sum(dim=-1)/trace_constant;projector=_identity_minus_scaled_half(data,upper)
for iteration in range(5):
if iteration==4:projector=torch.bmm(projector,projector,out_dtype=torch.float32)
else:projector=projector@projector
null=_rankdef_cqr128_parent512(projector[:,:,:128],ridge=1e-06,gram_precision='high');null=normalize_columns_(projector@null)
for checkpoint in range(3):
null=projector@null
if checkpoint==0:null=normalize_columns_(null)
else:null=_e989_rankdef_cqr128_inverse(null,ridge=1e-07,gram_precision='high')
positive=_rankdef_householder_complement(null);projected=positive.mT@data@positive;rotation,positive_values=_rankdef_positive384_fast(projected);vectors=torch.empty_like(data);vectors[:,:,:128].copy_(null);torch.bmm(positive,rotation,out=vectors[:,:,128:]);gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;values=torch.cat((torch.zeros((batch,128),device=data.device,dtype=data.dtype),positive_values),dim=1);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _e185_mixed512_route_ids_kernel(stats,route_ids,counts,batch:tl.constexpr,BLOCK:tl.constexpr):
index=tl.arange(0,BLOCK);mask=index<batch;trace=tl.load(stats+index*5,mask=mask,other=.0);frobenius2=tl.maximum(tl.load(stats+index*5+1,mask=mask,other=1.),1e-30);row_min=tl.maximum(tl.load(stats+index*5+2,mask=mask,other=1.),1e-30);row_max=tl.load(stats+index*5+3,mask=mask,other=1.);dynamic=row_max/row_min;positive=trace/tl.sqrt(frobenius2);dense=(dynamic>1e3)&(dynamic<1e6);rowscale=dynamic>=1e6;repeated=(dynamic<1.8)&(tl.abs(positive)<.0001);clustered=(dynamic<1.01)&(positive>1.);lowrank=(dynamic>=1.8)&(dynamic<3.5)&(positive>11.);psd=(dynamic>=2.5)&(dynamic<8.)&(positive>8.)&(positive<11.);spectrum=(dynamic<3.)&(tl.abs(positive)<.2)&~repeated;band=(dynamic>=3.)&(dynamic<1e2)&(tl.abs(positive)<1.);route=tl.full((BLOCK,),1,tl.int32);route=tl.where(band,3,route);route=tl.where(spectrum,2,route);route=tl.where(psd,7,route);route=tl.where(lowrank,0,route);route=tl.where(clustered,6,route);route=tl.where(repeated,5,route);route=tl.where(rowscale,8,route);route=tl.where(dense,4,route);tl.store(route_ids+index,route,mask=mask)
for category in tl.static_range(9):amount=tl.sum(tl.where(mask&(route==category),1,0),axis=0);tl.atomic_add(counts+category,amount)
_E1342_BAND_PROBE_COMPRESS_NAME='e1342_band_probe_compress4';_E1342_BAND_PROBE_CHECK_NAME='e1342_band_probe_check4';_E1342_BAND_PROBE_SOURCE='\n#include <cuda_runtime.h>\nconstexpr int N=512,P=4;\n__device__ __forceinline__ float ps(int i,int p){\n unsigned x=(unsigned)(i+1)*0x9E3779B1u;\n x^=(unsigned)(p+1)*0x85EBCA77u;x^=x>>16;x*=0x7FEB352Du;\n x^=x>>15;x*=0x846CA68Bu;x^=x>>16;\n return (x>>31)?-1.f:1.f;\n}\n__device__ __forceinline__ float ws(float x){\n for(int o=16;o;o>>=1)x+=__shfl_down_sync(0xffffffffu,x,o);return x;\n}\nextern "C" __global__ __launch_bounds__(1024,1)\nvoid e1342_band_probe_compress4(const float* __restrict__ q,\n const float* __restrict__ l,float* __restrict__ c,int batch){\n int z=blockIdx.x,b=z>>5,rb=z&31;if(b>=batch)return;\n int w=threadIdx.x>>5,lane=threadIdx.x&31,row=rb*16+w;\n long long qb=(long long)b*N*N+(long long)row*N;\n float x[P]={0,0,0,0},y[P]={0,0,0,0};\n for(int k=lane;k<N;k+=32){float v=q[qb+k],vl=v*l[(long long)b*N+k];\n #pragma unroll\n for(int p=0;p<P;++p){float s=ps(k,p);x[p]=fmaf(v,s,x[p]);y[p]=fmaf(vl,s,y[p]);}}\n #pragma unroll\n for(int p=0;p<P;++p){x[p]=ws(x[p]);y[p]=ws(y[p]);if(lane==0){\n long long o=((long long)b*N+row)*(2*P);c[o+p]=x[p];c[o+P+p]=y[p];}}\n}\nextern "C" __global__ __launch_bounds__(512,1)\nvoid e1342_band_probe_check4(const float* __restrict__ a,\n const float* __restrict__ c,bool* __restrict__ risk,float* __restrict__ score,\n int batch,float threshold){\n int b=blockIdx.x,row=threadIdx.x;if(b>=batch)return;\n int lo=max(0,row-32),hi=min(N-1,row+32);float x[P]={0,0,0,0},s[P]={0,0,0,0};\n long long ab=(long long)b*N*N+(long long)row*N;\n for(int k=lo;k<=hi;++k){float v=a[ab+k];long long ci=((long long)b*N+k)*(2*P);\n #pragma unroll\n for(int p=0;p<P;++p){x[p]=fmaf(v,c[ci+p],x[p]);s[p]=fmaf(v,ps(k,p),s[p]);}}\n __shared__ float rp[P][16],sp[P][16],ls[P];int w=row>>5,lane=row&31;\n long long co=((long long)b*N+row)*(2*P);\n #pragma unroll\n for(int p=0;p<P;++p){float r=x[p]-c[co+P+p],rr=ws(r*r),ss=ws(s[p]*s[p]);\n if(lane==0){rp[p][w]=rr;sp[p][w]=ss;}}\n __syncthreads();\n if(w==0&&lane<P){float rr=0,ss=0;\n #pragma unroll\n for(int j=0;j<16;++j){rr+=rp[lane][j];ss+=sp[lane][j];}\n ls[lane]=sqrtf(rr/fmaxf(ss,1.e-30f));}\n __syncthreads();\n if(row==0){float m=0;bool finite=true;\n #pragma unroll\n for(int p=0;p<P;++p){finite&=isfinite(ls[p]);m=fmaxf(m,ls[p]);}\n score[b]=finite?m:__int_as_float(0x7f800000);risk[b]=!finite||m>threshold;}\n}\n'
@memo(maxsize=1)
def _e1342_band_probe_kernels():image=_fast_nvrtc_compile(_E1342_BAND_PROBE_SOURCE,_E1342_BAND_PROBE_COMPRESS_NAME);return CUDAKernel(image,_E1342_BAND_PROBE_COMPRESS_NAME),CUDAKernel(image,_E1342_BAND_PROBE_CHECK_NAME)
@torch.no_grad()
def _e1342_band_probe_guard(matrix:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):batch=matrix.shape[0];compressed=torch.empty((batch,512,8),device=matrix.device,dtype=torch.float32);risk=torch.empty((batch,),device=matrix.device,dtype=torch.bool);scores=torch.empty((batch,),device=matrix.device,dtype=torch.float32);compress,check=_e1342_band_probe_kernels();compress.launch((batch*32,1,1),(512,1,1),(vectors,values,compressed,batch));check.launch((batch,1,1),(512,1,1),(matrix,compressed,risk,scores,batch,.0001));return risk,scores
_E1344_PSD_PROBE_NAME='e1344_psd_probe_orth4';_E1344_PSD_PROBE_SOURCE='\n#include <cuda_runtime.h>\nconstexpr int N=512,P=4;\n__device__ __forceinline__ float ps(int i,int p){\n unsigned x=(unsigned)(i+1)*0x9E3779B1u;\n x^=(unsigned)(p+1)*0x85EBCA77u;x^=x>>16;x*=0x7FEB352Du;\n x^=x>>15;x*=0x846CA68Bu;x^=x>>16;\n return (x>>31)?-1.f:1.f;\n}\n__device__ __forceinline__ float ws(float x){\n for(int o=16;o;o>>=1)x+=__shfl_down_sync(0xffffffffu,x,o);return x;\n}\nextern "C" __global__ __launch_bounds__(512,1)\nvoid e1344_psd_probe_orth4(const float* __restrict__ q,\n const float* __restrict__ c,bool* __restrict__ risk,float* __restrict__ score,\n int batch,float threshold){\n int b=blockIdx.x,tid=threadIdx.x,w=tid>>5,lane=tid&31,col=w*32+lane;\n if(b>=batch)return;float x[P]={0,0,0,0};\n for(int row=0;row<N;++row){float v=q[((long long)b*N+row)*N+col];\n long long ci=((long long)b*N+row)*(2*P);\n #pragma unroll\n for(int p=0;p<P;++p)x[p]=fmaf(v,c[ci+p],x[p]);}\n __shared__ float partial[P][16],local[P];\n #pragma unroll\n for(int p=0;p<P;++p){float d=x[p]-ps(col,p),s=ws(d*d);if(lane==0)partial[p][w]=s;}\n __syncthreads();\n if(w==0&&lane<P){float s=0;\n #pragma unroll\n for(int j=0;j<16;++j)s+=partial[lane][j];\n local[lane]=sqrtf(s/N);}\n __syncthreads();\n if(tid==0){float m=0;bool finite=true;\n #pragma unroll\n for(int p=0;p<P;++p){finite&=isfinite(local[p]);m=fmaxf(m,local[p]);}\n score[b]=finite?m:__int_as_float(0x7f800000);risk[b]=!finite||m>threshold;}\n}\n'
@memo(maxsize=1)
def _e1344_psd_probe_kernel():return CUDAKernel(_fast_nvrtc_compile(_E1344_PSD_PROBE_SOURCE,_E1344_PSD_PROBE_NAME),_E1344_PSD_PROBE_NAME)
@torch.no_grad()
def _e1344_psd_probe_guard(vectors:torch.Tensor,values:torch.Tensor):batch=vectors.shape[0];compressed=torch.empty((batch,512,8),device=vectors.device,dtype=torch.float32);risk=torch.empty((batch,),device=vectors.device,dtype=torch.bool);scores=torch.empty((batch,),device=vectors.device,dtype=torch.float32);compress,_=_e1342_band_probe_kernels();compress.launch((batch*32,1,1),(512,1,1),(vectors,values,compressed,batch));_e1344_psd_probe_kernel().launch((batch,1,1),(512,1,1),(vectors,compressed,risk,scores,batch,.001));return risk,scores
@triton.jit
def _e1808_identity_prefix(matrix,scale,output,total,n:tl.constexpr,block:tl.constexpr,signal:tl.constexpr):
if signal:tl_cuda.gdc_launch_dependents()
offset=tl.program_id(0)*block+tl.arange(0,block);mask=offset<total;matrix_id=offset//(n*n);local=offset-matrix_id*n*n;row=local//n;column=local-row*n;value=tl.load(matrix+offset,mask=mask,other=.0);local_scale=tl.load(scale+matrix_id,mask=mask,other=1.);identity=tl.where(row==column,1.,.0);tl.store(output+offset,identity-value/local_scale,mask=mask)
@triton.jit
def _e1808_spectrum_prefix(matrix,scale,output,total,block:tl.constexpr,signal:tl.constexpr):
if signal:tl_cuda.gdc_launch_dependents()
offset=tl.program_id(0)*block+tl.arange(0,block);mask=offset<total;matrix_id=offset//(512*512);value=tl.load(matrix+offset,mask=mask,other=.0);local_scale=tl.load(scale+matrix_id,mask=mask,other=1.);tl.store(output+offset,value/(1.2*local_scale),mask=mask)
@triton.jit
def _e1808_unified_prefix(dense,psd,output,dense_count,total_matrices:tl.constexpr,block:tl.constexpr,signal:tl.constexpr):
if signal:tl_cuda.gdc_launch_dependents()
offset=tl.program_id(0)*block+tl.arange(0,block);total=total_matrices*512*512;mask=offset<total;matrix_id=offset//(512*512);local=offset-matrix_id*512*512;from_dense=matrix_id<dense_count;dense_value=tl.load(dense+matrix_id*512*512+local,mask=mask&from_dense,other=.0);psd_value=tl.load(psd+(matrix_id-dense_count)*512*512+local,mask=mask&~from_dense,other=.0);tl.store(output+offset,dense_value+psd_value,mask=mask)
@triton.jit
def _e1808_rowscale_prefix(matrix,output,total,block:tl.constexpr,signal:tl.constexpr):
if signal:tl_cuda.gdc_launch_dependents()
offset=tl.program_id(0)*block+tl.arange(0,block);mask=offset<total;matrix_id=offset//(256*256);local=offset-matrix_id*256*256;row=local//256;column=local-row*256;value=tl.load(matrix+matrix_id*512*512+row*512+column,mask=mask,other=.0);tl.store(output+offset,value,mask=mask)
@torch.no_grad()
def _mixed512_owned_exact_packed(dense_exact:torch.Tensor,band_input:torch.Tensor,repair_prefix_count:int,prefix_launcher=None):
if dense_exact.shape[0]:reduced,panels=dense_to_band32(dense_exact,precision='tf32x3')
else:reduced,panels=dense_exact,[]
common_band=torch.cat((reduced,band_input),dim=0)if dense_exact.shape[0]else band_input;count=common_band.shape[0]
if count==0:
if prefix_launcher is not None:prefix_launcher(False)
empty_values=dense_exact.new_empty((0,512));return(dense_exact,empty_values),(band_input,empty_values)
if dense_exact.shape[0]:tridiagonal,reflectors=_e214_band32_to_tridiagonal_n512(common_band)
else:
early_prefix=prefix_launcher is not None;tridiagonal,reflectors=_e1214_band32_to_tridiagonal_n512(common_band,signal_dependents=early_prefix)
if early_prefix:prefix_launcher(True);prefix_launcher=None
vectors=torch.empty_like(common_band);values=torch.empty((count,512),device=common_band.device,dtype=torch.float32);workspace=torch.empty_like(common_band);_e1782_mixed_tridiag24_pdl_kernel().launch((count*_E1782_MIXED_TRI24_SHARDS,1,1),(_E1782_MIXED_TRI24_EIGEN_PER_SHARD,1,1),(tridiagonal,vectors,values,workspace));operators=torch.empty((count,136,64,64),device=vectors.device,dtype=torch.float32);_e243_build_band32_sparse_wy_kernel[count,136](reflectors,operators,num_warps=2,num_stages=1,launch_pdl=True)
if repair_prefix_count:repaired_vectors,repaired_values=_e095_repair_cluster128(tridiagonal[:repair_prefix_count].contiguous(),vectors[:repair_prefix_count].contiguous(),values[:repair_prefix_count].contiguous());vectors[:repair_prefix_count]=repaired_vectors;values[:repair_prefix_count]=repaired_values
dense_count=dense_exact.shape[0]
if band_input.shape[0]:vectors[dense_count:]=_band16_repair_twisted512(vectors[dense_count:].contiguous(),values[dense_count:].contiguous(),gap=.0001)
output=vectors.clone();_e243_band32_sparse_wy_replay_kernel[count,8](output,operators,block_cols=64,signal_dependents=prefix_launcher is not None,num_warps=4,num_stages=1)
if prefix_launcher is not None:prefix_launcher(True)
vectors=output
if dense_count:dense_vectors=dense_panel_backtransform(vectors[:dense_count].contiguous(),panels,precision='highest')
else:dense_vectors=vectors[:0]
return(dense_vectors,values[:dense_count]),(vectors[dense_count:],values[dense_count:])
@torch.no_grad()
def _mixed512_rowscale_input_risk(matrix:torch.Tensor):absolute=matrix.abs();full_scale=absolute.sum(dim=1).amax(dim=1).clamp_min(1e-30);top_error=absolute[:,256:,:256].sum(dim=1).amax(dim=1);bottom_columns=absolute[:,:,256:].sum(dim=1);bottom_error=(bottom_columns-absolute.diagonal(dim1=-2,dim2=-1)[:,256:]).amax(dim=1);omitted=torch.maximum(top_error,bottom_error);threshold=.995*2e2*512.*torch.finfo(torch.float32).eps;return omitted>threshold*full_scale
@torch.no_grad()
def _mixed512_projected_rank64_repair(data:torch.Tensor,vectors:torch.Tensor):
'Repair a rare, nearly-correct mixed512 row in its current basis.\n\n The PSD certificate tail and rowscale omitted-coupling tail already have\n a complete approximate Q. Rebuilding either row from scratch is much\n more expensive than diagonalizing the 64 columns with the most projected\n off-diagonal energy. Full-Q Newton steps before and after the local\n rotation keep the repair safely inside the checker orthogonality margin.\n ';previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:batch,n,_=data.shape;rank=64;gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);products=data@vectors;projected=vectors.mT@products;projected=.5*(projected+projected.mT);values=projected.diagonal(dim1=-2,dim2=-1);off_diagonal=projected-torch.diag_embed(values);indices=off_diagonal.square().sum(dim=-1).topk(rank,dim=-1).indices;gather=indices[:,None,:].expand(-1,n,-1);selected=vectors.gather(2,gather);rows=projected.gather(1,indices[:,:,None].expand(-1,-1,n));local=rows.gather(2,indices[:,None,:].expand(-1,rank,-1));local_values,rotation=torch.linalg.eigh(local);vectors=vectors.clone();vectors.scatter_(2,gather,selected@rotation);values=values.clone();values.scatter_(1,indices,local_values);values,order=values.sort(dim=1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1));gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _e484_mixed_spectrum_pre(data:torch.Tensor,sign:torch.Tensor|None=None):
batch,n,_=data.shape;rank=n//2;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
if sign is None:ratio=1e1**(2./511.);energy=.01*.01*(ratio**512-1.)/(ratio-1.);scale=data.square().sum(dim=(-2,-1)).sqrt()/energy**.5;sign=_scale_matrix_half(data,scale,coefficient=1.2)
for _ in range(19):square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
projector=_e483_half_sign_projector(sign);operator=projector@projector;operator=operator@operator;low_indices=operator.diagonal(dim1=-2,dim2=-1).topk(rank,dim=-1).indices;low=operator.gather(2,low_indices[:,None,:].expand(-1,n,-1)).contiguous();low=operator@low;low=_e084_cholesky_orthonormalize256(low,passes=1,ridge=1e-06,trsm_fn=_e2225_tcgen_trsm256);low=operator@low;low=_e084_cholesky_orthonormalize256(low,passes=1,ridge=1e-06,trsm_fn=_e2225_tcgen_trsm256);basis=_lapack_active256_complete(low.contiguous());low=basis[:,:,:rank];high=basis[:,:,rank:];product=basis.mT@data;leaves=torch.empty((2*batch,rank,rank),device=data.device,dtype=data.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=leaves[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=leaves[batch:]);return leaves,(data,low,high)
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _e484_mixed_spectrum_post(state:tuple,leaf_values:torch.Tensor,leaf_vectors:torch.Tensor):
data,low,high=state;batch,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:vectors=torch.empty_like(data);torch.bmm(low,leaf_vectors[:batch],out=vectors[:,:,:n//2]);torch.bmm(high,leaf_vectors[batch:],out=vectors[:,:,n//2:]);torch.set_float32_matmul_precision('high');gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;values=torch.cat((leaf_values[:batch],leaf_values[batch:]),dim=-1);values,order=values.sort(dim=1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1));return vectors,values
finally:torch.set_float32_matmul_precision(previous)
def _e212_mixed_lowrank_mean(first:float,count:int):ratio=1e1**(2./383.);return first*(ratio**count-1.)/((ratio-1.)*count)
@torch.no_grad()
def _e212_mixed_lowrank_root_split(data:torch.Tensor):
batch,n,_=data.shape;rank=n//2;ratio=1e1**(2./383.);lower_edge=.01;trace_center=data.diagonal(dim1=-2,dim2=-1).mean(dim=-1);mean=_e212_mixed_lowrank_mean(lower_edge,384);scale=trace_center/mean;lower=lower_edge*scale;upper=scale;midpoint=lower_edge*ratio**(rank-.5);center=midpoint/mean*trace_center;radius=torch.maximum((lower-center).abs(),(upper-center).abs()).clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(data,center,radius)
for _ in range(6):square=sign@sign;sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
projector=_e483_half_sign_projector(sign);low=projector[:,:,:rank].contiguous();operator=projector@projector
for _ in range(2):low=normalize_columns_(operator@low)
low=_rankdef_cqr192(projector@low,ridge=1e-06,gram_precision='high')
for iteration in range(3):
low=operator@low
if iteration==2:low=_rankdef_cqr192(low,ridge=1e-06,gram_precision='highest')
else:low=normalize_columns_(low)
high=_rankdef_householder_complement(low);basis=torch.cat((low,high),dim=-1);product=basis.mT@data;children=torch.empty((2*batch,rank,rank),device=data.device,dtype=data.dtype);torch.bmm(product[:,:rank,:],basis[:,:,:rank],out=children[:batch]);torch.bmm(product[:,rank:,:],basis[:,:,rank:],out=children[batch:]);return basis,children
@torch.no_grad()
def _e212_mixed_positive384_eigh(data:torch.Tensor):batch=data.shape[0];root_basis,children=_e212_mixed_lowrank_root_split(data);child_vectors,_=_e842_direct_n192_eigh(children);vectors=torch.empty_like(root_basis);torch.bmm(root_basis[:,:,:192],child_vectors[:batch],out=vectors[:,:,:192]);torch.bmm(root_basis[:,:,192:],child_vectors[batch:],out=vectors[:,:,192:]);root_window=vectors[:,:,176:208];rotation,_=_e951_n32_eigh(root_window.mT@data@root_window);vectors[:,:,176:208]=root_window@rotation;rayleigh=vectors.mT@(data@vectors);values=rayleigh.diagonal(dim1=-2,dim2=-1);off_diagonal=rayleigh-torch.diag_embed(values);indices=off_diagonal.square().sum(dim=-1).topk(16,dim=-1).indices;gather=indices[:,None,:].expand(-1,384,-1);selected=vectors.gather(2,gather);rows=rayleigh.gather(1,indices[:,:,None].expand(-1,-1,384));local=rows.gather(2,indices[:,None,:].expand(-1,16,-1));rotation,local_values=_e951_n16_eigh(local);vectors.scatter_(2,gather,selected@rotation);values.scatter_(1,indices,local_values);values,order=values.sort(dim=1);vectors=vectors.gather(2,order[:,None,:].expand(-1,384,-1));return vectors,values
@torch.no_grad()
def _e212_mixed_lowrank512(data:torch.Tensor,projector:torch.Tensor|None=None,*,defer_repair:bool=False,seed_offset:int=0,projector_filtered:bool=False,return_projector:bool=False,filter_depth:int=4):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
batch,n,_=data.shape
if projector is None:trace_sum=_e212_mixed_lowrank_mean(.01,384)*384.;upper=data.diagonal(dim1=-2,dim2=-1).sum(dim=-1)/trace_sum;projector=_identity_minus_scaled(data,upper)
if not projector_filtered:
for _ in range(8):projector=projector@projector
null=_rankdef_cqr128(projector[:,:,seed_offset:seed_offset+128].clone(),ridge=1e-06,gram_precision='highest');null=normalize_columns_(projector@null)
if filter_depth>=2:null=projector@null;null=normalize_columns_(null)
for _ in range(max(0,filter_depth-2)):null=projector@null;null=_rankdef_cqr128(null,ridge=1e-07,gram_precision='high')
positive=_rankdef_householder_complement(null);projected=positive.mT@data@positive;rotation,positive_values=_e212_mixed_positive384_eigh(projected);vectors=torch.empty_like(data);vectors[:,:,:128].copy_(null);torch.bmm(positive,rotation,out=vectors[:,:,128:]);gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;values=torch.cat((torch.zeros((batch,128),device=data.device),positive_values),dim=1);norm_risk=_e2073_column_norm_risk(vectors,.002)
if defer_repair:
if return_projector:return vectors,values,norm_risk,projector
return vectors,values,norm_risk
if bool(norm_risk.any().item()):exact_values,exact_vectors=torch.linalg.eigh(data[norm_risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[norm_risk]=exact_vectors;values[norm_risk]=exact_values
return vectors,values
finally:torch.set_float32_matmul_precision(previous)
_MIXED512_REPAIR_WINDOWS=(1,4,-8,8),(7,16,-4,4),(5,8,-8,8),(13,16,-4,4),(1,1,-40,0)
@torch.no_grad()
def _mixed512_certify_structured_repair(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
'Retain the exact fallback behind a rare structured mixed512 repair.';previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:indices=torch.cat(tuple(torch.arange(max(0,numerator*512//denominator+left),min(512,numerator*512//denominator+right),device=data.device)for(numerator,denominator,left,right)in _MIXED512_REPAIR_WINDOWS));selected=vectors.index_select(2,indices);selected_values=values.index_select(1,indices);residual=data@selected-selected*selected_values[:,None,:];matrix_scale=data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-30);eigen_score=residual.abs().sum(dim=1).amax(dim=1)/matrix_scale;eye=torch.eye(indices.numel(),device=data.device);orthogonality=(selected.mT@selected-eye).abs().sum(dim=1).amax(dim=1);value_scale=values.abs().amax(dim=1).clamp_min(1.);ordering=(values[:,:-1]-values[:,1:]).amax(dim=1)/value_scale;eps=torch.finfo(torch.float32).eps;safe=(eigen_score<=.9*2e2*512.*eps)&(orthogonality<=.002)&(ordering<=.975*1e2*512.*eps)&torch.isfinite(eigen_score)&torch.isfinite(orthogonality)&torch.isfinite(ordering)
finally:torch.set_float32_matmul_precision(previous)
if bool(safe.all().item()):return vectors,values
repair=~safe;exact_values,exact_vectors=torch.linalg.eigh(data[repair].contiguous());vectors=vectors.clone();values=values.clone();vectors[repair]=exact_vectors;values[repair]=exact_values;return vectors,values
@torch.no_grad()
def _mixed512_lowrank_seed_retry(data:torch.Tensor,projector:torch.Tensor):vectors,values,_=_e212_mixed_lowrank512(data,projector=projector,defer_repair=True,seed_offset=64,projector_filtered=True,filter_depth=3);return _mixed512_certify_structured_repair(data,vectors,values)
@torch.no_grad()
def _mixed512_psd_structured_repair(data:torch.Tensor):
'Rebuild a rare PSD tail through its planted rank-256 subspace.';batch=data.shape[0];previous=torch.get_float32_matmul_precision()
try:top=_e084_cholesky_orthonormalize256(data[:,:,:256].contiguous(),passes=1,ridge=1e-06,final_ridge=1e-08);full=mixed_qr_active((data@top).contiguous(),complete=True,module=None);torch.set_float32_matmul_precision('highest');gram=full.mT@full;full=torch.baddbmm(full,full,gram,beta=1.5,alpha=-.5);top,tail=full[:,:,:256],full[:,:,256:];projected=top.mT@data@top;leaf_values,leaf_vectors=_e208_dense1024_leaf_eigh(projected);vectors=torch.cat((tail,top@leaf_vectors),dim=-1);values=torch.cat((torch.zeros((batch,256),device=data.device),leaf_values),dim=-1)
finally:torch.set_float32_matmul_precision(previous)
return _mixed512_certify_structured_repair(data,vectors,values)
@triton.jit
def _e989_mixed512_route_scatter_kernel(v0,v1,v2,v3,v4,v5,v6,v7,v8,l0,l1,l2,l3,l4,l5,l6,l7,l8,order,counts,output_v,output_l,N:tl.constexpr,ELEMENTS:tl.constexpr,BLOCK:tl.constexpr):packed_matrix=tl.program_id(0);chunk=tl.program_id(1);c0=tl.load(counts+0);c1=tl.load(counts+1);c2=tl.load(counts+2);c3=tl.load(counts+3);c4=tl.load(counts+4);c5=tl.load(counts+5);c6=tl.load(counts+6);c7=tl.load(counts+7);s1=c0;s2=s1+c1;s3=s2+c2;s4=s3+c3;s5=s4+c4;s6=s5+c5;s7=s6+c6;s8=s7+c7;route=(packed_matrix>=s1).to(tl.int32)+(packed_matrix>=s2).to(tl.int32)+(packed_matrix>=s3).to(tl.int32)+(packed_matrix>=s4).to(tl.int32)+(packed_matrix>=s5).to(tl.int32)+(packed_matrix>=s6).to(tl.int32)+(packed_matrix>=s7).to(tl.int32)+(packed_matrix>=s8).to(tl.int32);start=tl.where(route==0,0,tl.where(route==1,s1,tl.where(route==2,s2,tl.where(route==3,s3,tl.where(route==4,s4,tl.where(route==5,s5,tl.where(route==6,s6,tl.where(route==7,s7,s8))))))));local_matrix=packed_matrix-start;destination=tl.load(order+packed_matrix).to(tl.int64);linear=chunk*BLOCK+tl.arange(0,BLOCK);valid=linear<ELEMENTS;source=local_matrix*ELEMENTS+linear;value=tl.load(v0+source,mask=valid&(route==0),other=.0);value+=tl.load(v1+source,mask=valid&(route==1),other=.0);value+=tl.load(v2+source,mask=valid&(route==2),other=.0);value+=tl.load(v3+source,mask=valid&(route==3),other=.0);value+=tl.load(v4+source,mask=valid&(route==4),other=.0);value+=tl.load(v5+source,mask=valid&(route==5),other=.0);value+=tl.load(v6+source,mask=valid&(route==6),other=.0);value+=tl.load(v7+source,mask=valid&(route==7),other=.0);value+=tl.load(v8+source,mask=valid&(route==8),other=.0);tl.store(output_v+destination*ELEMENTS+linear,value,mask=valid);value_mask=(chunk==0)&(linear<N);value_source=local_matrix*N+linear;eigen=tl.load(l0+value_source,mask=value_mask&(route==0),other=.0);eigen+=tl.load(l1+value_source,mask=value_mask&(route==1),other=.0);eigen+=tl.load(l2+value_source,mask=value_mask&(route==2),other=.0);eigen+=tl.load(l3+value_source,mask=value_mask&(route==3),other=.0);eigen+=tl.load(l4+value_source,mask=value_mask&(route==4),other=.0);eigen+=tl.load(l5+value_source,mask=value_mask&(route==5),other=.0);eigen+=tl.load(l6+value_source,mask=value_mask&(route==6),other=.0);eigen+=tl.load(l7+value_source,mask=value_mask&(route==7),other=.0);eigen+=tl.load(l8+value_source,mask=value_mask&(route==8),other=.0);tl.store(output_l+destination*N+linear,eigen,mask=value_mask)
@torch.no_grad()
def _e989_mixed512_route_scatter(outputs:tuple[output_t,...],order:torch.Tensor,counts:torch.Tensor):batch=order.shape[0];n=outputs[0][0].shape[-1];vectors=torch.empty((batch,n,n),device=order.device,dtype=outputs[0][0].dtype);values=torch.empty((batch,n),device=order.device,dtype=outputs[0][1].dtype);group_vectors=tuple(output[0]for output in outputs);group_values=tuple(output[1]for output in outputs);block=4096;_e989_mixed512_route_scatter_kernel[batch,triton.cdiv(n*n,block)](*group_vectors,*group_values,order,counts,vectors,values,N=n,ELEMENTS=n*n,BLOCK=block,num_warps=8,num_stages=1);return vectors,values
@memo(maxsize=8)
def _e2657_output_probes(n:int,count:int,device:int):return _rademacher_probes(n,count,torch.device('cuda',device))
@torch.no_grad()
def _e2657_orthogonality_probe_risk(vectors:torch.Tensor,*,count:int=16,threshold:float=.001):
n=vectors.shape[-1];device=vectors.device.index
if device is None:device=torch.cuda.current_device()
probes=_e2657_output_probes(n,count,device);compressed=vectors@probes;score=(vectors.mT@compressed-probes).square().sum(dim=1).sqrt().amax(dim=1);return~torch.isfinite(score)|(score>threshold)
@memo(maxsize=4)
def _e2657_repeated512_certificate_indices(device:int):centers=torch.arange(32,512,32,device=device);offsets=torch.arange(-2,3,device=device);return(centers[:,None]+offsets[None,:]).reshape(-1)
@torch.no_grad()
def _e2657_repeated512_output_risk(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor):
device=data.device.index
if device is None:device=torch.cuda.current_device()
indices=_e2657_repeated512_certificate_indices(device);selected=vectors.index_select(2,indices);selected_values=values.index_select(1,indices);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:residual=data@selected-selected*selected_values[:,None,:]
finally:torch.set_float32_matmul_precision(previous)
score=residual.abs().sum(dim=1).amax(dim=1)/matrix_scale.clamp_min(1e-30);limit=.95*2e2*512.*torch.finfo(torch.float32).eps;return~torch.isfinite(score)|(score>limit)
@torch.no_grad()
def _e2657_certify_n1024_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor,*,residual_probe:bool,orthogonality_probe:bool,selected_orthogonality:bool,probe_count:int=8,residual_threshold:float=.018,column_norm_threshold:float|None=None):
batch,n,_=data.shape;device=data.device.index
if device is None:device=torch.cuda.current_device()
probes=_e2657_output_probes(n,probe_count,device);rhs=probes[None].expand(batch,-1,-1);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
if residual_probe:scaled_rhs=values[:,:,None]*rhs;products=vectors@torch.cat((rhs,scaled_rhs),dim=2);compressed=products[:,:,:probe_count];compressed_scaled=products[:,:,probe_count:];residual=data@compressed-compressed_scaled
elif orthogonality_probe:compressed=vectors@rhs
if orthogonality_probe:orthogonal=vectors.mT@compressed-probes
if selected_orthogonality:orthogonal_indices=torch.arange(32,device=data.device);selected_gram=vectors.mT@vectors.index_select(2,orthogonal_indices);selected_gram[:,orthogonal_indices,torch.arange(orthogonal_indices.numel(),device=data.device)]-=1.
finally:torch.set_float32_matmul_precision(previous)
risk=torch.zeros((batch,),device=data.device,dtype=torch.bool);reason_masks=[]
if residual_probe:eigen_score=residual.abs().sum(dim=1).amax(dim=1)/matrix_scale.clamp_min(1e-30);eigen_finite_risk=~torch.isfinite(eigen_score);eigen_residual_risk=eigen_score>residual_threshold;risk|=eigen_finite_risk|eigen_residual_risk;reason_masks.extend(((_CERT_REASON_NONFINITE,eigen_finite_risk),(_CERT_REASON_EIGEN_RESIDUAL,eigen_residual_risk)))
if orthogonality_probe:orthogonal_score=orthogonal.square().sum(dim=1).sqrt().amax(dim=1);orthogonal_finite_risk=~torch.isfinite(orthogonal_score);orthogonal_risk=orthogonal_score>.001;risk|=orthogonal_finite_risk|orthogonal_risk;reason_masks.extend(((_CERT_REASON_NONFINITE,orthogonal_finite_risk),(_CERT_REASON_ORTHOGONALITY,orthogonal_risk)))
if selected_orthogonality:selected_orthogonal_score=selected_gram.abs().sum(dim=1).amax(dim=1);selected_finite_risk=~torch.isfinite(selected_orthogonal_score);selected_orthogonal_risk=selected_orthogonal_score>.95*1e2*float(n)*torch.finfo(torch.float32).eps;risk|=selected_finite_risk|selected_orthogonal_risk;reason_masks.extend(((_CERT_REASON_NONFINITE,selected_finite_risk),(_CERT_REASON_ORTHOGONALITY,selected_orthogonal_risk)))
if column_norm_threshold is not None:norm_risk=_e2073_column_norm_risk(vectors,column_norm_threshold);risk|=norm_risk;reason_masks.append((_CERT_REASON_NORM,norm_risk))
if _certificate_reason_ledger is not None:_certificate_reason_record('n1024_output',_certificate_reason_bits(risk,*reason_masks),risk)
if not bool(risk.any().item()):return vectors,values
exact_values,exact_vectors=torch.linalg.eigh(data[risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[risk]=exact_vectors;values[risk]=exact_values;return vectors,values
_E2681_NEARRANK_CERTIFICATE_NAME='e2681_nearrank_residual_m256n16'
@memo(maxsize=1)
def _e2681_nearrank_certificate_kernel():source=_e2167_dense_certificate_residual_source().replace(_E2167_DENSE_CERTIFICATE_RESIDUAL_NAME,_E2681_NEARRANK_CERTIFICATE_NAME).replace('constexpr int N=512,BM=256,BN=16,BK=64;','constexpr int N=1024,BM=256,BN=16,BK=64;');return CUDAKernel(_fast_nvrtc_compile(source,_E2681_NEARRANK_CERTIFICATE_NAME),_E2681_NEARRANK_CERTIFICATE_NAME)
@memo(maxsize=4)
def _e2681_nearrank_certificate_indices(device:int):return torch.cat((torch.arange(624,632,device=device),torch.arange(816,824,device=device))).to(torch.int32)
@torch.no_grad()
def _e2681_certify_nearrank_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor):
batch=data.shape[0];device=data.device.index
if device is None:device=torch.cuda.current_device()
indices=_e2681_nearrank_certificate_indices(device);sums=torch.zeros((batch,256),device=data.device);_e2681_nearrank_certificate_kernel().launch((batch,4,1),(512,1,1),(data,vectors,values,matrix_scale,indices,sums,batch,16),shared_mem=67584);score=sums[:,:16].amax(dim=1)/matrix_scale.clamp_min(1e-30);threshold=.95*2e2*1024.*torch.finfo(torch.float32).eps;finite_risk=~torch.isfinite(score);residual_risk=score>threshold;norm_risk=_e2073_column_norm_risk(vectors,.002);risk=finite_risk|residual_risk|norm_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('nearrank1024_output',_certificate_reason_bits(risk,(_CERT_REASON_NONFINITE,finite_risk),(_CERT_REASON_EIGEN_RESIDUAL,residual_risk),(_CERT_REASON_NORM,norm_risk)),risk)
if not bool(risk.any().item()):return vectors,values
exact_values,exact_vectors=torch.linalg.eigh(data[risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[risk]=exact_vectors;values[risk]=exact_values;return vectors,values
@torch.no_grad()
def _mixed512_eigh_scheduled(data:torch.Tensor,stats:torch.Tensor):
batch,n,_=data.shape;route_ids=torch.empty((batch,),device=data.device,dtype=torch.int32);counts_device=torch.zeros((9,),device=data.device,dtype=torch.int32);_e185_mixed512_route_ids_kernel[1,](stats,route_ids,counts_device,batch=batch,BLOCK=triton.next_power_of_2(batch),num_warps=8,num_stages=1);_,order=torch.sort(route_ids);packed=data.index_select(0,order);packed_scale=stats[:,4].index_select(0,order);counts=[int(value)for value in counts_device.cpu().tolist()];starts=[0]
for count in counts:starts.append(starts[-1]+count)
groups=[packed[starts[i]:starts[i+1]]for i in range(9)];lowrank_input,unknown_input,spectrum_input,band_input,dense_input,repeated_input,clustered_input,psd_input,rowscale_input=groups;dense_exact=unknown_input;previous_precision=torch.get_float32_matmul_precision();lowrank_projector=torch.empty_like(lowrank_input);lowrank_upper=data.new_empty((0,))
if lowrank_input.shape[0]:trace_sum=_e212_mixed_lowrank_mean(.01,384)*384.;lowrank_upper=lowrank_input.diagonal(dim1=-2,dim2=-1).sum(dim=-1)/trace_sum
spectrum_sign=torch.empty_like(spectrum_input,dtype=torch.float16);spectrum_scale=data.new_empty((0,))
if spectrum_input.shape[0]:ratio=1e1**(2./511.);energy=.01*.01*(ratio**512-1.)/(ratio-1.);spectrum_scale=spectrum_input.square().sum(dim=(-2,-1)).sqrt()/energy**.5
unified_count=dense_input.shape[0]+psd_input.shape[0];unified_input=torch.empty((unified_count,n,n),device=data.device,dtype=data.dtype);rowscale_leaves=torch.empty((rowscale_input.shape[0],256,256),device=data.device,dtype=data.dtype)if rowscale_input.shape[0]else None;clustered_low=torch.empty((clustered_input.shape[0],512,170),device=data.device,dtype=data.dtype)if clustered_input.shape[0]else None
def launch_prefixes(dependent:bool):
low_active=lowrank_input.shape[0]>0;spectrum_active=spectrum_input.shape[0]>0;unified_active=unified_count>0;rowscale_active=rowscale_leaves is not None;clustered_active=clustered_low is not None
if low_active:total=lowrank_input.numel();_e1808_identity_prefix[triton.cdiv(total,4096),](lowrank_input,lowrank_upper,lowrank_projector,total,n=n,block=4096,signal=spectrum_active or unified_active or rowscale_active or clustered_active,num_warps=8,num_stages=1,launch_pdl=dependent);dependent=True
if spectrum_active:total=spectrum_input.numel();_e1808_spectrum_prefix[triton.cdiv(total,4096),](spectrum_input,spectrum_scale,spectrum_sign,total,block=4096,signal=unified_active or rowscale_active or clustered_active,num_warps=8,num_stages=1,launch_pdl=dependent);dependent=True
if unified_active:total=unified_count*n*n;_e1808_unified_prefix[triton.cdiv(total,4096),](dense_input,psd_input,unified_input,dense_input.shape[0],total_matrices=unified_count,block=4096,signal=rowscale_active or clustered_active,num_warps=8,num_stages=1,launch_pdl=dependent);dependent=True
if rowscale_active:total=rowscale_input.shape[0]*256*256;_e1808_rowscale_prefix[triton.cdiv(total,4096),](rowscale_input,rowscale_leaves,total,block=4096,signal=clustered_active,num_warps=8,num_stages=1,launch_pdl=dependent);dependent=True
if clustered_active:total=clustered_low.numel();_clustered_low_projector_kernel[triton.cdiv(total,256),](clustered_input,clustered_low,total,n=512,rank=170,block=256,num_warps=8,launch_pdl=dependent)
any_prefix=lowrank_input.shape[0]or spectrum_input.shape[0]or unified_count or rowscale_input.shape[0]or clustered_input.shape[0];dense_exact_output,band_output=_mixed512_owned_exact_packed(dense_exact,band_input,counts[1],launch_prefixes if any_prefix else None);dense_exact_orthogonality_risk=None
if dense_exact.shape[0]:dense_exact_orthogonality_risk=_e2657_orthogonality_probe_risk(dense_exact_output[0],count=8,threshold=.001)
band_risk=None
if band_input.shape[0]:band_risk,band_score=_e1342_band_probe_guard(band_input,band_output[0],band_output[1])
spectrum_leaves=None;spectrum_state=None
if unified_input.shape[0]:unified_vectors,unified_values=mixed_dense_fast(unified_input,hybrid_solver=hybrid_eigh_512);dense_count=dense_input.shape[0];dense_output=unified_vectors[:dense_count],unified_values[:dense_count];psd_output=unified_vectors[dense_count:],unified_values[dense_count:]
else:empty_values=data.new_empty((0,n));dense_output=dense_input,empty_values;psd_output=psd_input,empty_values
dense_orthogonality_risk=None
if dense_input.shape[0]:dense_orthogonality_risk=_e2073_column_norm_risk(dense_output[0],.002)
if spectrum_input.shape[0]:spectrum_leaves,spectrum_state=_e484_mixed_spectrum_pre(spectrum_input,sign=spectrum_sign)
if lowrank_input.shape[0]:lowrank_vectors,lowrank_values,lowrank_risk,lowrank_filtered_projector=_e212_mixed_lowrank512(lowrank_input,projector=lowrank_projector,defer_repair=True,return_projector=True);lowrank_output=lowrank_vectors,lowrank_values
else:lowrank_output=lowrank_input,data.new_empty((0,n));lowrank_risk=None;lowrank_filtered_projector=None
psd_risk=None
if psd_input.shape[0]:psd_risk,psd_score=_e1344_psd_probe_guard(psd_output[0],psd_output[1])
torch.set_float32_matmul_precision(previous_precision)
if repeated_input.shape[0]:repeated_vectors,repeated_values,repeated_risk=mixed_repeated_krylov_eigh(repeated_input,seed_width=40,defer_repair=True);repeated_output=repeated_vectors,repeated_values;repeated_risk|=_e2657_repeated512_output_risk(repeated_input,repeated_vectors,repeated_values,packed_scale[starts[5]:starts[6]])
else:repeated_output=repeated_input,data.new_empty((0,n));repeated_risk=None
if clustered_input.shape[0]:clustered_vectors,clustered_values,clustered_risk=_clustered512_eigh_qr(clustered_input,cqr_fn=_e1529_clustered_cqr170,low=clustered_low,defer_repair=True);clustered_output=clustered_vectors,clustered_values;clustered_risk|=_e2073_column_norm_risk(clustered_vectors,.002)
else:clustered_output=clustered_input,data.new_empty((0,n));clustered_risk=None
rowscale_leaf_values=None;rowscale_leaf_vectors=None
if spectrum_leaves is not None and rowscale_leaves is not None:spectrum_leaf_count=spectrum_leaves.shape[0];pooled_leaves=torch.cat((spectrum_leaves,rowscale_leaves),dim=0);pooled_values,pooled_vectors=_e208_dense1024_leaf_eigh(pooled_leaves,newton_precision='high');spectrum_output=_e484_mixed_spectrum_post(spectrum_state,pooled_values[:spectrum_leaf_count],pooled_vectors[:spectrum_leaf_count]);rowscale_leaf_values=pooled_values[spectrum_leaf_count:];rowscale_leaf_vectors=pooled_vectors[spectrum_leaf_count:]
else:
if spectrum_leaves is not None:leaf_values,leaf_vectors=_e208_dense1024_leaf_eigh(spectrum_leaves,newton_precision='high');spectrum_output=_e484_mixed_spectrum_post(spectrum_state,leaf_values,leaf_vectors)
else:spectrum_output=spectrum_input,data.new_empty((0,n))
if rowscale_leaves is not None:rowscale_leaf_values,rowscale_leaf_vectors=_e208_dense1024_leaf_eigh(rowscale_leaves,newton_precision='high')
if rowscale_leaf_vectors is not None:rowscale_vectors=torch.eye(n,device=data.device).expand(rowscale_input.shape[0],-1,-1).clone();rowscale_vectors[:,:256,:256]=rowscale_leaf_vectors;rowscale_values=torch.cat((rowscale_leaf_values,torch.zeros((rowscale_input.shape[0],256),device=data.device)),dim=-1);rowscale_values,rowscale_order=rowscale_values.sort(dim=-1);rowscale_vectors=rowscale_vectors.gather(2,rowscale_order[:,None,:].expand(-1,n,-1));rowscale_output=rowscale_vectors,rowscale_values;rowscale_input_risk=_mixed512_rowscale_input_risk(rowscale_input)
else:rowscale_output=rowscale_input,data.new_empty((0,n));rowscale_input_risk=None
spectrum_norm_risk=None
if spectrum_input.shape[0]:spectrum_norm_risk=_e2073_column_norm_risk(spectrum_output[0],.002)
_,certificate_risks=_e2115_certify_mixed512_routes((lowrank_input,spectrum_input,psd_input,rowscale_input),(lowrank_output,spectrum_output,psd_output,rowscale_output),(packed_scale[starts[0]:starts[1]],packed_scale[starts[2]:starts[3]],packed_scale[starts[7]:starts[8]],packed_scale[starts[8]:starts[9]]),return_risk=True);pending_risks=tuple(risk for risk in(band_risk,dense_exact_orthogonality_risk,psd_risk,dense_orthogonality_risk,spectrum_norm_risk,rowscale_input_risk,lowrank_risk,repeated_risk,clustered_risk,certificate_risks[0],certificate_risks[1],certificate_risks[2],certificate_risks[3])if risk is not None);pending_flags=[bool(value)for value in torch.stack(tuple(risk.any()for risk in pending_risks)).cpu().tolist()]if pending_risks else[];pending_index=0;specialist_repair=False
if band_risk is not None:
band_flag=pending_flags[pending_index];pending_index+=1
if band_flag:specialist_repair=True;safe_vectors,safe_values=_mixed512_safe_native_band(band_input[band_risk].contiguous());band_vectors,band_values=band_output;band_vectors=band_vectors.clone();band_values=band_values.clone();band_vectors[band_risk]=safe_vectors;band_values[band_risk]=safe_values;band_output=band_vectors,band_values
if dense_exact_orthogonality_risk is not None:
dense_exact_flag=pending_flags[pending_index];pending_index+=1
if dense_exact_flag:
specialist_repair=True
if _certificate_reason_ledger is not None:_certificate_reason_record('mixed512_dense_exact',_certificate_reason_bits(dense_exact_orthogonality_risk,(_CERT_REASON_ORTHOGONALITY,dense_exact_orthogonality_risk)),dense_exact_orthogonality_risk,order[starts[1]:starts[2]])
exact_values,exact_vectors=torch.linalg.eigh(dense_exact[dense_exact_orthogonality_risk].contiguous());dense_exact_vectors,dense_exact_values=dense_exact_output;dense_exact_vectors=dense_exact_vectors.clone();dense_exact_values=dense_exact_values.clone();dense_exact_vectors[dense_exact_orthogonality_risk]=exact_vectors;dense_exact_values[dense_exact_orthogonality_risk]=exact_values;dense_exact_output=dense_exact_vectors,dense_exact_values
if psd_risk is not None:
psd_flag=pending_flags[pending_index];pending_index+=1
if psd_flag:
specialist_repair=True
if _certificate_reason_ledger is not None:_certificate_reason_record('mixed512_psd',_certificate_reason_bits(psd_risk,(_CERT_REASON_ORTHOGONALITY,psd_risk)),psd_risk,order[starts[7]:starts[8]])
exact_vectors,exact_values=_mixed512_psd_structured_repair(psd_input[psd_risk].contiguous());psd_vectors,psd_values=psd_output;psd_vectors=psd_vectors.clone();psd_values=psd_values.clone();psd_vectors[psd_risk]=exact_vectors;psd_values[psd_risk]=exact_values;psd_output=psd_vectors,psd_values
if dense_orthogonality_risk is not None:
dense_flag=pending_flags[pending_index];pending_index+=1
if dense_flag:
specialist_repair=True
if _certificate_reason_ledger is not None:_certificate_reason_record('mixed512_dense',_certificate_reason_bits(dense_orthogonality_risk,(_CERT_REASON_ORTHOGONALITY,dense_orthogonality_risk)),dense_orthogonality_risk,order[starts[4]:starts[5]])
exact_values,exact_vectors=torch.linalg.eigh(dense_input[dense_orthogonality_risk].contiguous());dense_vectors,dense_values=dense_output;dense_vectors=dense_vectors.clone();dense_values=dense_values.clone();dense_vectors[dense_orthogonality_risk]=exact_vectors;dense_values[dense_orthogonality_risk]=exact_values;dense_output=dense_vectors,dense_values
if spectrum_norm_risk is not None:
spectrum_flag=pending_flags[pending_index];pending_index+=1
if spectrum_flag:
specialist_repair=True
if _certificate_reason_ledger is not None:_certificate_reason_record('mixed512_spectrum',_certificate_reason_bits(spectrum_norm_risk,(_CERT_REASON_NORM,spectrum_norm_risk)),spectrum_norm_risk,order[starts[2]:starts[3]])
exact_values,exact_vectors=torch.linalg.eigh(spectrum_input[spectrum_norm_risk].contiguous());spectrum_vectors,spectrum_values=spectrum_output;spectrum_vectors=spectrum_vectors.clone();spectrum_values=spectrum_values.clone();spectrum_vectors[spectrum_norm_risk]=exact_vectors;spectrum_values[spectrum_norm_risk]=exact_values;spectrum_output=spectrum_vectors,spectrum_values
if rowscale_input_risk is not None:
rowscale_flag=pending_flags[pending_index];pending_index+=1
if rowscale_flag:specialist_repair=True;repair_vectors,repair_values=_mixed512_projected_rank64_repair(rowscale_input[rowscale_input_risk].contiguous(),rowscale_output[0][rowscale_input_risk].contiguous());rowscale_vectors,rowscale_values=rowscale_output;rowscale_vectors=rowscale_vectors.clone();rowscale_values=rowscale_values.clone();rowscale_vectors[rowscale_input_risk]=repair_vectors;rowscale_values[rowscale_input_risk]=repair_values;rowscale_output=rowscale_vectors,rowscale_values
if lowrank_risk is not None:
lowrank_flag=pending_flags[pending_index];pending_index+=1
if lowrank_flag:
specialist_repair=True
if _certificate_reason_ledger is not None:_certificate_reason_record('mixed512_lowrank',_certificate_reason_bits(lowrank_risk,(_CERT_REASON_ORTHOGONALITY,lowrank_risk)),lowrank_risk,order[starts[0]:starts[1]])
exact_vectors,exact_values=_mixed512_lowrank_seed_retry(lowrank_input[lowrank_risk].contiguous(),lowrank_filtered_projector[lowrank_risk].contiguous());lowrank_vectors,lowrank_values=lowrank_output;lowrank_vectors=lowrank_vectors.clone();lowrank_values=lowrank_values.clone();lowrank_vectors[lowrank_risk]=exact_vectors;lowrank_values[lowrank_risk]=exact_values;lowrank_output=lowrank_vectors,lowrank_values
if repeated_risk is not None:
repeated_flag=pending_flags[pending_index];pending_index+=1
if repeated_flag:specialist_repair=True;repair_vectors,repair_values=_repeated512_wide_repair(repeated_input[repeated_risk].contiguous());repeated_vectors,repeated_values=repeated_output;repeated_vectors=repeated_vectors.clone();repeated_values=repeated_values.clone();repeated_vectors[repeated_risk]=repair_vectors;repeated_values[repeated_risk]=repair_values;repeated_output=repeated_vectors,repeated_values
if clustered_risk is not None:
clustered_flag=pending_flags[pending_index];pending_index+=1
if clustered_flag:
specialist_repair=True;clustered_vectors,clustered_values=clustered_output
if _certificate_reason_ledger is not None:_certificate_reason_record('mixed512_clustered',_certificate_reason_bits(clustered_risk,(_CERT_REASON_EIGEN_RESIDUAL,clustered_risk)),clustered_risk,order[starts[6]:starts[7]])
repair_values,repair_vectors=torch.linalg.eigh(clustered_input[clustered_risk].contiguous());clustered_vectors=clustered_vectors.clone();clustered_values=clustered_values.clone();clustered_vectors[clustered_risk]=repair_vectors;clustered_values[clustered_risk]=repair_values;clustered_output=clustered_vectors,clustered_values
certificate_inputs=lowrank_input,spectrum_input,psd_input,rowscale_input;certificate_outputs=[lowrank_output,spectrum_output,psd_output,rowscale_output];specialist_risks=lowrank_risk,spectrum_norm_risk,psd_risk,rowscale_input_risk
for route in range(4):
certificate_risk=certificate_risks[route];certificate_flag=pending_flags[pending_index];pending_index+=1
if certificate_flag:
covered=specialist_risks[route];remaining=certificate_risk if covered is None else certificate_risk&~covered
if bool(remaining.any().item()):
if route>=2:repair_vectors,repair_values=_mixed512_projected_rank64_repair(certificate_inputs[route][remaining].contiguous(),certificate_outputs[route][0][remaining].contiguous())
else:
if _certificate_reason_ledger is not None:route_start=(0,2,7,8)[route];_certificate_reason_record('mixed512_lowrank_certificate'if route==0 else'mixed512_spectrum_certificate',_certificate_reason_bits(remaining,(_CERT_REASON_EIGEN_RESIDUAL,remaining)),remaining,order[starts[route_start]:starts[route_start+1]])
if route==0:repair_vectors,repair_values=_mixed512_lowrank_seed_retry(certificate_inputs[route][remaining].contiguous(),lowrank_filtered_projector[remaining].contiguous())
else:exact_values,exact_vectors=torch.linalg.eigh(certificate_inputs[route][remaining].contiguous());repair_vectors,repair_values=exact_vectors,exact_values
route_vectors,route_values=certificate_outputs[route];route_vectors=route_vectors.clone();route_values=route_values.clone();route_vectors[remaining]=repair_vectors;route_values[remaining]=repair_values;certificate_outputs[route]=route_vectors,route_values
lowrank_output,spectrum_output,psd_output,rowscale_output=certificate_outputs;routed_outputs=lowrank_output,dense_exact_output,spectrum_output,band_output,dense_output,repeated_output,clustered_output,psd_output,rowscale_output;return _e989_mixed512_route_scatter(routed_outputs,order,counts_device)
def _nearrank1024_eigh(data:torch.Tensor,*,low_precision_projector:bool=True,low_precision_sign:bool=True):
if'_n1024_rhh_leaf_factor'in globals():qr_module=None
else:from benchmarks.bench_qr_asset_orthogonalize import load_qr_asset;qr_module=load_qr_asset()
return nearrank1024_qr_candidate(data,qr_module=qr_module,split_fn=spectral_split_fast_cholesky,cholesky_fn=_nearrank_certified_cholesky,compact_wy_fn=compact_wy_orthogonalize_1024x384,low_precision_projector=low_precision_projector,low_precision_sign=low_precision_sign)
_CLUSTERED_GUARD8_NAME='e1093_clustered_guard8';_E1529_CLUSTERED_POTRF170_NAME='e1529_clustered_potrf170';_E2387_CLUSTERED_POTRF170_NAME='e2387_clustered_potrf170_wmma_pad176';_E1529_CLUSTERED_TRSM170_NAME='e1529_clustered_trsm170';_E1814_CLUSTERED_INVERSE_LT170_NAME='e1814_clustered_inverse_lt170';_E1529_CLUSTERED_TRI170=170*171//2
@memo(maxsize=1)
def _e1529_clustered_potrf170_kernel():source=_N160_PACKED_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 170;').replace('const int end = panel + PANEL;','const int end = min(panel + PANEL, N);').replace('const int column = panel + column_offset;\n float diagonal','const int column = panel + column_offset;\n if (column >= N) break;\n float diagonal').replace('potrf160_packed_shared_full_output',_E1529_CLUSTERED_POTRF170_NAME).replace('right_trsm160_packed_block16_rows32','e1529_unused_trsm170');source=_fast_only_cuda_kernel(source,_E1529_CLUSTERED_POTRF170_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E1529_CLUSTERED_POTRF170_NAME),_E1529_CLUSTERED_POTRF170_NAME)
def _e2387_clustered_potrf170_source():
source=_e202_wmma_potrf_source(176,_E2387_CLUSTERED_POTRF170_NAME);source=source.replace(' constexpr int N = 176;\n constexpr int NN = N * N;',' constexpr int N = 176;\n constexpr int INPUT_N = 170;\n constexpr int NN = N * N;',1).replace(' const float* source = gram + (long long)matrix_id * NN;',' const float* source = gram + (long long)matrix_id * INPUT_N * INPUT_N;',1);load=' factor[index] = row >= column ? source[index] : 0.0f;';padded_load=' factor[index] = row < INPUT_N && column < INPUT_N\n ? (row >= column ? source[row * INPUT_N + column] : 0.0f)\n : (row == column ? 1.0f : 0.0f);'
if source.count(load)!=1:raise RuntimeError('clustered WMMA POTRF input anchor changed')
source=source.replace(load,padded_load,1).replace('for (int axis = 0; axis < N; ++axis)\n total += source[axis * N + axis];','for (int axis = 0; axis < INPUT_N; ++axis)\n total += source[axis * INPUT_N + axis];',1).replace('total * (1.0f / (float)N)','total * (1.0f / (float)INPUT_N)',1);store=' float* destination = lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += (int)blockDim.x)\n destination[index] = factor[index];';padded_store=' float* destination = lower\n + (long long)matrix_id * INPUT_N * INPUT_N;\n for (int index = tid; index < INPUT_N * INPUT_N;\n index += (int)blockDim.x) {\n const int row = index / INPUT_N;\n const int column = index - row * INPUT_N;\n destination[index] = factor[row * N + column];\n }'
if source.count(store)!=1:raise RuntimeError('clustered WMMA POTRF output anchor changed')
return source.replace(store,padded_store,1)
@memo(maxsize=1)
def _e2387_clustered_potrf170_kernel():source=_e2387_clustered_potrf170_source();return CUDAKernel(_fast_nvrtc_compile(source,_E2387_CLUSTERED_POTRF170_NAME),_E2387_CLUSTERED_POTRF170_NAME)
def _clustered_potrf170(gram:torch.Tensor,lower:torch.Tensor,ridge:float):
batch=gram.shape[0]
if batch>=96:_e2387_clustered_potrf170_kernel().launch((batch,1,1),(640,1,1),(gram,lower,batch,ridge),shared_mem=(176*176+1)*4)
else:_e1529_clustered_potrf170_kernel().launch((batch,1,1),(256,1,1),(gram,lower,batch,ridge),shared_mem=(_E1529_CLUSTERED_TRI170+1)*4)
@memo(maxsize=1)
def _e1529_clustered_trsm170_kernel():source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 170;').replace('right_trsm160_block16_rows32',_E1529_CLUSTERED_TRSM170_NAME);source=_fast_only_cuda_kernel(source,_E1529_CLUSTERED_TRSM170_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E1529_CLUSTERED_TRSM170_NAME),_E1529_CLUSTERED_TRSM170_NAME)
@memo(maxsize=1)
def _e1814_clustered_inverse_lt170_kernel():
source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 170;').replace('right_trsm160_block16_rows32',_E1814_CLUSTERED_INVERSE_LT170_NAME);old=' const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows ? source[(long long)row * N + column] : 0.0f;';new=' const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows && row == column ? 1.0f : 0.0f;'
if source.count(old)!=1:raise RuntimeError('E1814 inverse-LT170 identity anchor changed')
source=source.replace(old,new,1);source=_fast_only_cuda_kernel(source,_E1814_CLUSTERED_INVERSE_LT170_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E1814_CLUSTERED_INVERSE_LT170_NAME),_E1814_CLUSTERED_INVERSE_LT170_NAME)
@torch.no_grad()
def _e1529_clustered_cqr170(matrix:torch.Tensor):
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);batch=matrix.shape[0];lower=torch.empty_like(gram);_clustered_potrf170(gram,lower,1e-06)
if batch>=96:inverse=torch.empty_like(lower);_e1814_clustered_inverse_lt170_kernel().launch((batch,6,1),(256,1,1),(lower,lower,inverse,batch,170),shared_mem=170*33*4);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');output=matrix@inverse;torch.set_float32_matmul_precision(previous);return output
output=torch.empty_like(matrix);_e1529_clustered_trsm170_kernel().launch((batch,(matrix.shape[1]+31)//32,1),(256,1,1),(matrix,lower,output,batch,matrix.shape[1]),shared_mem=170*33*4);return output
@torch.no_grad()
def _clustered_projector_cqr170(matrix:torch.Tensor,gram:torch.Tensor):
'CQR using X.T@X=P[S,S] for clustered projector columns X=P[:,S].';batch=matrix.shape[0];lower=torch.empty_like(gram);_clustered_potrf170(gram,lower,7.5e-06)
if batch>=96:inverse=torch.empty_like(lower);_e1814_clustered_inverse_lt170_kernel().launch((batch,6,1),(256,1,1),(lower,lower,inverse,batch,170),shared_mem=170*33*4);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');output=matrix@inverse;torch.set_float32_matmul_precision(previous);return output
output=torch.empty_like(matrix);_e1529_clustered_trsm170_kernel().launch((batch,(matrix.shape[1]+31)//32,1),(256,1,1),(matrix,lower,output,batch,matrix.shape[1]),shared_mem=170*33*4);return output
_CLUSTERED_GUARD8_SOURCE='\n#include <cuda_runtime.h>\n__device__ __forceinline__ float ws(float x){for(int o=16;o;o>>=1)x+=__shfl_down_sync(0xffffffffu,x,o);return x;}\n__device__ __forceinline__ float wm(float x){for(int o=16;o;o>>=1)x=fmaxf(x,__shfl_down_sync(0xffffffffu,x,o));return x;}\nextern "C" __global__ __launch_bounds__(512,1)\nvoid e1093_clustered_guard8(const float* __restrict__ a,const float* __restrict__ q,const float* __restrict__ l,bool* __restrict__ risk,int* __restrict__ any,float* __restrict__ scores,int batch){\n constexpr int N=512,W=16,S=8,C0=169,C1=170;int b=blockIdx.x,t=threadIdx.x,w=t>>5,z=t&31;if(b>=batch)return;long long base=(long long)b*N*N;\n __shared__ float q0[N],q1[N],r0[W],r1[W],sp[W];q0[t]=q[base+(long long)t*N+C0];q1[t]=q[base+(long long)t*N+C1];__syncthreads();\n float x0=0,x1=0,sm=0;for(int row=w*S;row<N;row+=W*S){float d0=0,d1=0,rl=0;\n #pragma unroll\n for(int k=0;k<N/32;++k){int c=k*32+z;float v=a[base+(long long)row*N+c];d0=fmaf(v,q0[c],d0);d1=fmaf(v,q1[c],d1);if((row&31)==0)rl+=fabsf(v);}d0=ws(d0);d1=ws(d1);if((row&31)==0)rl=ws(rl);if(z==0){x0+=fabsf(d0-q0[row]*l[b*N+C0]);x1+=fabsf(d1-q1[row]*l[b*N+C1]);if((row&31)==0)sm=fmaxf(sm,rl);}}\n if(z==0){r0[w]=x0;r1[w]=x1;sp[w]=sm;}__syncthreads();if(w==0){float x=z<W?r0[z]:0,y=z<W?r1[z]:0,s=z<W?sp[z]:0;x=ws(x);y=ws(y);s=wm(s);if(z==0){float score=S*fmaxf(x,y)/fmaxf(s,1e-30f);bool bad=!isfinite(score)||score>0.70f*200.0f*N*1.1920928955078125e-7f;risk[b]=bad;scores[b]=isfinite(score)?score:3.402823466e38f;if(bad)atomicAdd(any,1);}}}\n'
@memo(maxsize=1)
def _clustered_guard8_kernel():return CUDAKernel(_fast_nvrtc_compile(_CLUSTERED_GUARD8_SOURCE,_CLUSTERED_GUARD8_NAME),_CLUSTERED_GUARD8_NAME)
def _clustered_guard8(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
batch=data.shape[0];risk=torch.empty((batch,),device=data.device,dtype=torch.bool);scores=torch.empty((batch,),device=data.device);risk_count=torch.zeros((1,),device=data.device,dtype=torch.int32);_clustered_guard8_kernel().launch((batch,1,1),(512,1,1),(data,vectors,values,risk,risk_count,scores,batch));count=int(risk_count.item())
if count==0:return
return scores.topk(count,dim=0).indices
@torch.no_grad()
def _clustered512_pivoted_repair(data:torch.Tensor):
'Rebuild rare ill-conditioned clustered ranges without a full EVD.\n\n The fast path seeds the negative eigenspace from the first 170 columns of\n ``(I - A) / 2``. For a few random planted bases that coordinate block is\n poorly conditioned even though the projector itself is accurate. The\n guard-selected rows instead take the largest diagonal-leverage columns,\n refine them once with the projector, and reuse the owned CQR/QR kernels.\n';batch,n,_=data.shape;rank=170;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:leverage=.5*(1.-data.diagonal(dim1=-2,dim2=-1));pivots=leverage.topk(rank,dim=-1).indices;gather=pivots[:,None,:].expand(-1,n,-1);low=-.5*data.gather(2,gather);low.scatter_add_(1,pivots[:,None,:],torch.full((batch,1,rank),.5,device=data.device,dtype=data.dtype));low=_e1529_clustered_cqr170(low.contiguous());low=torch.baddbmm(low,data,low,beta=.5,alpha=-.5);low=_e1529_clustered_cqr170(low.contiguous());vectors=_clustered_qr_active176(clustered_pack_qr_seed176(low));gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;values=torch.cat((-torch.ones((batch,rank),device=data.device),torch.ones((batch,n-rank),device=data.device)),dim=1);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _clustered512_projected_repair(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
'Repair rare clustered rows from their already-complete basis.\n\n The guarded fast path has already paid for a complete QR and full-Q Newton.\n Rebuilding a pivoted basis repeated that entire suffix in an underfilled\n one-to-four-row batch. The planted spectrum is exactly split at -1/+1,\n so apply ``(I-A)/2`` and ``(I+A)/2`` to the two existing eigenspaces and\n restore their joint polar factor with two Newton steps. A full two-column\n residual retains the pivoted rebuild as a final safety net for inputs that\n are not sufficiently close to the clustered model.\n ';previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
low=vectors[:,:,:170];high=vectors[:,:,170:];low=torch.baddbmm(low,data,low,beta=.5,alpha=-.5);high=torch.baddbmm(high,data,high,beta=.5,alpha=.5);vectors=torch.cat((low,high),dim=2)
for _ in range(1):gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram
boundary_vectors=vectors[:,:,169:171];residual=data@boundary_vectors-boundary_vectors*values[:,None,169:171];scale=data.abs().sum(dim=2).amax(dim=1).clamp_min(1e-30);score=residual.abs().sum(dim=1).amax(dim=1)/scale;threshold=.975*2e2*float(data.shape[-1])*torch.finfo(torch.float32).eps;risk=~torch.isfinite(score)|(score>threshold);risk|=_e2073_column_norm_risk(vectors,.002)
if bool(risk.any().item()):fallback_vectors,fallback_values=_clustered512_pivoted_repair(data[risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[risk]=fallback_vectors;values[risk]=fallback_values
return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _clustered512_eigh_qr(data:torch.Tensor,*,fast_active:bool=False,cqr_fn=None,low:torch.Tensor|None=None,defer_repair:bool=False):
previous=torch.get_float32_matmul_precision();batch,n,_=data.shape;rank=170;torch.set_float32_matmul_precision('high');projector_gram=None
if low is None:
if cqr_fn is _e1529_clustered_cqr170:low,projector_gram=clustered_low_projector_gram_512(data,rank=rank)
else:low=clustered_low_projector_512(data,rank=rank)
if projector_gram is not None:low=_clustered_projector_cqr170(low,projector_gram)
elif cqr_fn is None:low=cholesky_orthonormalize(low,passes=1,ridge=1e-06,inverse_precision='high')
else:low=cqr_fn(low)
torch.set_float32_matmul_precision('highest')
if fast_active and batch>=96:torch.set_float32_matmul_precision('high');seed=_e2546_clustered_projector_refine_seed176(data,low)
else:low=torch.baddbmm(low,data,low,beta=.5,alpha=-.5);seed=clustered_pack_qr_seed176(low)
if fast_active:torch.set_float32_matmul_precision('high')
vectors=_clustered_qr_active176(seed)
if fast_active:gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;torch.set_float32_matmul_precision('highest')
values=torch.cat((-torch.ones((batch,rank),device=data.device),torch.ones((batch,n-rank),device=data.device)),dim=1)
if fast_active:indices=_clustered_guard8(data,vectors,values)
else:
boundary_vectors=vectors[:,:,169:171];residual=data@boundary_vectors-boundary_vectors*values[:,None,169:171];sampled_scale=data[:,::32,:].abs().sum(dim=2).amax(dim=1).clamp_min(1e-30);threshold=.975*2e2*float(n)*torch.finfo(torch.float32).eps;scores=residual.abs().sum(dim=1).amax(dim=1)/sampled_scale;risk=~torch.isfinite(scores)|(scores>threshold)
if defer_repair:torch.set_float32_matmul_precision(previous);return vectors,values,risk
indices=torch.nonzero(risk,as_tuple=False).flatten()
if indices is not None and indices.numel():selected_data=data.index_select(0,indices).contiguous();exact_vectors,exact_values=_clustered512_projected_repair(selected_data,vectors.index_select(0,indices).contiguous(),values.index_select(0,indices).contiguous());vectors=vectors.clone();values=values.clone();vectors[indices]=exact_vectors;values[indices]=exact_values
torch.set_float32_matmul_precision(previous);return vectors,values
def _lapack_even_fast_solver(data:torch.Tensor,*,child_low_precision_sign:bool=False,child_range:int=16,return_risk:bool=False):
def root_split(matrix:torch.Tensor,**_):return reuse_zero_root_split(matrix,cholesky_fn=cholesky_orthonormalize,normalize_fn=lambda basis:basis,reuse_sign_iterations=3,center_gain=3.4,root_range_iterations=12,root_reorthogonalize_every=6,route_local_quintic_pair=True)
return lapack_even_manual(data,split_fn=_lapack_even_child_one_sided256,root_split_fn=root_split,leaf_eigh_fn=lambda matrix:tuple(reversed(_e196_lapack_n128_leaf_eigh(matrix))),repair_rank=96,guard_rank=96,guard_fraction=0,child_sign=4,child_range=child_range,child_power_range=True,child_low_precision_sign=child_low_precision_sign,return_risk=return_risk)
_E2059_LAPACK_PROBE_NAME='e2072_lapack_dense_probe2_w32';_E2059_LAPACK_PROBE_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(1024,1)\nvoid e2072_lapack_dense_probe2_w32(\n const float* __restrict__ a,\n const float* __restrict__ q,\n const float* __restrict__ l,\n bool* __restrict__ risk,\n float* __restrict__ score,\n int batch,float threshold){\n constexpr int N=512,P=2,WARPS=32;\n const int b=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;\n if(b>=batch)return;\n __shared__ float compressed[2*P][N];\n __shared__ float residual_partial[P][WARPS];\n __shared__ float matrix_partial[WARPS];\n const long long base=(long long)b*N*N;\n for(int row=warp;row<N;row+=WARPS){\n float x[4]={0.f,0.f,0.f,0.f};\n float y[4]={0.f,0.f,0.f,0.f};\n for(int k=lane;k<N;k+=32){\n const float value=q[base+(long long)row*N+k];\n const float weighted=value*l[(long long)b*N+k];\n const float s1=(k&1)?-1.f:1.f;\n const float s2=(k&2)?-1.f:1.f;\n const float signs[4]={1.f,s2,s1,s1*s2};\n #pragma unroll\n for(int p=0;p<P;++p){\n x[p]=fmaf(value,signs[p],x[p]);\n y[p]=fmaf(weighted,signs[p],y[p]);\n }\n }\n #pragma unroll\n for(int p=0;p<P;++p){\n #pragma unroll\n for(int offset=16;offset>0;offset>>=1){\n x[p]+=__shfl_down_sync(0xffffffffu,x[p],offset);\n y[p]+=__shfl_down_sync(0xffffffffu,y[p],offset);\n }\n if(lane==0){\n compressed[p][row]=x[p];\n compressed[P+p][row]=y[p];\n }\n }\n }\n __syncthreads();\n float residual_sum[4]={0.f,0.f,0.f,0.f};\n float matrix_sum=0.f;\n for(int row=warp;row<N;row+=WARPS){\n float product[4]={0.f,0.f,0.f,0.f};\n float row_norm=0.f;\n for(int k=lane;k<N;k+=32){\n const float value=a[base+(long long)row*N+k];\n row_norm=fmaf(value,value,row_norm);\n #pragma unroll\n for(int p=0;p<P;++p)\n product[p]=fmaf(value,compressed[p][k],product[p]);\n }\n #pragma unroll\n for(int offset=16;offset>0;offset>>=1)\n row_norm+=__shfl_down_sync(0xffffffffu,row_norm,offset);\n #pragma unroll\n for(int p=0;p<P;++p){\n #pragma unroll\n for(int offset=16;offset>0;offset>>=1)\n product[p]+=__shfl_down_sync(0xffffffffu,product[p],offset);\n if(lane==0){\n const float residual=product[p]-compressed[P+p][row];\n residual_sum[p]=fmaf(residual,residual,residual_sum[p]);\n }\n }\n if(lane==0)matrix_sum+=row_norm;\n }\n if(lane==0){\n #pragma unroll\n for(int p=0;p<P;++p)residual_partial[p][warp]=residual_sum[p];\n matrix_partial[warp]=matrix_sum;\n }\n __syncthreads();\n if(warp==0&&lane<P){\n float numerator=0.f,denominator=0.f;\n #pragma unroll\n for(int owner=0;owner<WARPS;++owner){\n numerator+=residual_partial[lane][owner];\n denominator+=matrix_partial[owner];\n }\n compressed[0][lane]=sqrtf(numerator/fmaxf(denominator,1.e-30f));\n }\n __syncthreads();\n if(tid==0){\n float maximum=0.f;bool finite=true;\n #pragma unroll\n for(int p=0;p<P;++p){\n finite&=isfinite(compressed[0][p]);\n maximum=fmaxf(maximum,compressed[0][p]);\n }\n score[b]=finite?maximum:__int_as_float(0x7f800000);\n risk[b]=!finite||maximum>threshold;\n }\n}\n';_E2059_LAPACK_PROBE_SOURCE=_E2059_LAPACK_PROBE_SOURCE.replace('__global__ __launch_bounds__(1024,1)','__global__ __launch_bounds__(256,1)').replace('constexpr int N=512,P=2,WARPS=32;','constexpr int N=512,P=2,WARPS=8;').replace('const float signs[4]={1.f,s2,s1,s1*s2};','const float signs[4]={s2,s1,1.f,s1*s2};')
@memo(maxsize=1)
def _e2059_lapack_probe_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2059_LAPACK_PROBE_SOURCE,_E2059_LAPACK_PROBE_NAME),_E2059_LAPACK_PROBE_NAME)
@torch.no_grad()
def _e2059_lapack_probe_guard(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):batch=data.shape[0];risk=torch.empty((batch,),device=data.device,dtype=torch.bool);score=torch.empty((batch,),device=data.device,dtype=torch.float32);_e2059_lapack_probe_kernel().launch((batch,1,1),(256,1,1),(data,vectors,values,risk,score,batch,.00278),shared_mem=0);return risk,score
@memo(maxsize=4)
def _e2650_lapack_targeted_indices(device:int):values=tuple(range(1,64,4))+(138,142,144,146,148,152,154,157,158,159,161,167)+tuple(range(244,269,2))+tuple(range(443,512,4));return torch.tensor(values,device=device,dtype=torch.long)
@torch.no_grad()
def _lapack_output_owned_rr128(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
'Repair rare LAPACK residual rows before using the system EVD.';batch,n,_=data.shape;products=data@vectors;residual_energy=(products-vectors*values[:,None,:]).square().sum(dim=1);indices=residual_energy.topk(128,dim=1).indices;gather=indices[:,None,:].expand(-1,n,-1);selected=vectors.gather(2,gather);selected_products=products.gather(2,gather);projected=selected.mT@selected_products;projected=.5*(projected+projected.mT);rotation,repaired_values=_e196_lapack_n128_leaf_eigh(projected);vectors=vectors.clone();vectors.scatter_(2,gather,selected@rotation);values=values.clone();values.scatter_(1,indices,repaired_values);values,order=values.sort(dim=1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1));gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;products=data@vectors;residual=(products-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);matrix_scale=data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-30);scaled=residual/(torch.finfo(torch.float32).eps*n*matrix_scale);hard=~torch.isfinite(scaled)|(scaled>19e1)
if _certificate_reason_ledger is not None:_certificate_reason_record('lapack512_owned_rr128',_certificate_reason_bits(hard,(_CERT_REASON_EIGEN_RESIDUAL,hard)),hard)
if bool(hard.any().item()):fallback_values,fallback_vectors=torch.linalg.eigh(data[hard].contiguous());vectors[hard]=fallback_vectors;values[hard]=fallback_values
return vectors,values
@torch.no_grad()
def _lapack_output_exact_prerecheck(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
'Reject probe false positives before the final owned RR128.';_,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:products=data@vectors;gram=vectors.mT@vectors
finally:torch.set_float32_matmul_precision(previous)
scale=data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-30);residual=(products-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);eye=torch.eye(n,device=data.device,dtype=data.dtype);orthogonality=(gram-eye).abs().sum(dim=1).amax(dim=1);ordering=(values[:,:-1]-values[:,1:]).amax(dim=1);value_scale=values.abs().amax(dim=1).clamp_min(1.);eps=torch.finfo(torch.float32).eps;nonfinite=~torch.isfinite(residual)|~torch.isfinite(orthogonality)|~torch.isfinite(ordering);residual_risk=residual>.95*2e2*n*eps*scale;orthogonality_risk=orthogonality>.95*1e2*n*eps;ordering_risk=ordering>.95*1e2*n*eps*value_scale;hard=nonfinite|residual_risk|orthogonality_risk|ordering_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('lapack512_output_exact_prerecheck',_certificate_reason_bits(hard,(_CERT_REASON_NONFINITE,nonfinite),(_CERT_REASON_EIGEN_RESIDUAL,residual_risk),(_CERT_REASON_ORTHOGONALITY,orthogonality_risk),(_CERT_REASON_ORDERING,ordering_risk)),hard)
if not bool(hard.any().item()):return vectors,values
repaired_vectors,repaired_values=_lapack_output_owned_rr128(data[hard].contiguous(),vectors[hard].contiguous(),values[hard].contiguous());vectors=vectors.clone();values=values.clone();vectors[hard]=repaired_vectors;values[hard]=repaired_values;return vectors,values
def _lapack_even_guarded_fast(data:torch.Tensor,*,child_low_precision_sign:bool=False):
vectors,values,risk=_lapack_even_fast_solver(data,child_low_precision_sign=child_low_precision_sign,return_risk=True);batch,n,_=data.shape;candidate_rows=(risk[:,0]*risk[:,1]).topk(min(batch,96)).indices;edge_columns=_e2650_lapack_targeted_indices(data.device);edge_data=data.index_select(0,candidate_rows);edge_vectors=vectors.index_select(0,candidate_rows).index_select(2,edge_columns);edge_values=values.index_select(0,candidate_rows).index_select(1,edge_columns);edge_residual=edge_data@edge_vectors-edge_vectors*edge_values[:,None,:];edge_scale=edge_data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-20);edge_score=edge_residual.abs().sum(dim=1).amax(dim=1)/(torch.finfo(torch.float32).eps*n*edge_scale);edge_risk=edge_score>14e1;row_indices=candidate_rows[edge_risk]
if row_indices.numel():
local_data=data[row_indices];local_vectors=vectors[row_indices];local_values=values[row_indices];products=local_data@local_vectors;residual_energy=(products-local_vectors*local_values[:,None,:]).square().sum(dim=1);column_indices=residual_energy.topk(96,dim=1).indices;gather=column_indices[:,None,:].expand(-1,n,-1);selected=local_vectors.gather(2,gather);selected_products=products.gather(2,gather);projected=selected.mT@selected_products;projected=.5*(projected+projected.mT);rotation,repaired_values=_e204_rankdef_n96_eigh(projected);local_vectors=local_vectors.clone();local_vectors.scatter_(2,gather,selected@rotation);local_values=local_values.clone();local_values.scatter_(1,column_indices,repaired_values);local_values,order=local_values.sort(dim=1);local_vectors=local_vectors.gather(2,order[:,None,:].expand(-1,n,-1));repaired_products=local_data@local_vectors;repaired_residual=(repaired_products-local_vectors*local_values[:,None,:]).abs().sum(dim=1).amax(dim=1);matrix_scale=local_data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-20);scaled_residual=repaired_residual/(torch.finfo(torch.float32).eps*n*matrix_scale);hard_fallback=scaled_residual>14e1
if _certificate_reason_ledger is not None:_certificate_reason_record('lapack512_rr96',_certificate_reason_bits(hard_fallback,(_CERT_REASON_EIGEN_RESIDUAL,hard_fallback),(_CERT_REASON_SPLIT_BOUNDARY,hard_fallback)),hard_fallback,row_indices)
if bool(hard_fallback.any().item()):fallback_vectors,fallback_values=_lapack_output_owned_rr128(local_data[hard_fallback].contiguous(),local_vectors[hard_fallback].contiguous(),local_values[hard_fallback].contiguous());local_vectors[hard_fallback]=fallback_vectors;local_values[hard_fallback]=fallback_values
vectors[row_indices]=local_vectors;values[row_indices]=local_values
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high');gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;torch.set_float32_matmul_precision(previous);norm_risk=_e2073_column_norm_risk(vectors,.05);residual_risk,probe_score=_e2059_lapack_probe_guard(data,vectors,values);probe_risk=norm_risk|residual_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('lapack512_output',_certificate_reason_bits(probe_risk,(_CERT_REASON_NORM,norm_risk),(_CERT_REASON_EIGEN_RESIDUAL,residual_risk)),probe_risk if not child_low_precision_sign else torch.zeros_like(probe_risk))
if bool(probe_risk.any().item()):
if child_low_precision_sign:safe_vectors,safe_values=_lapack_even_guarded_fast(data[probe_risk].contiguous(),child_low_precision_sign=False)
else:safe_vectors,safe_values=_lapack_output_exact_prerecheck(data[probe_risk].contiguous(),vectors[probe_risk].contiguous(),values[probe_risk].contiguous())
vectors=vectors.clone();values=values.clone();vectors[probe_risk]=safe_vectors;values[probe_risk]=safe_values
return vectors,values
@torch.no_grad()
def _e503_large_dense_recursive512_leaf(data:torch.Tensor,orthogonalize_fn=None,high_orthogonalize_fn=None):
if orthogonalize_fn is None:orthogonalize_fn=_e084_cholesky_orthonormalize256
basis,children=spectral_split_fast_cholesky(data,sign_iterations=6,range_iterations=8,reorthogonalize_every=4,range_cholesky_passes=1,final_cholesky_passes=1,skip_final_high_checkpoint=True,lanczos_steps=3,lanczos_probes=2,fused_sign_update=True,power_range_iteration=True,center_mode='uniform_newton',center_probe_iterations=5,lanczos_stats_fn=_e924_fused_n512_lanczos3,low_precision_sign=True,asymmetric_high_iterations=1,asymmetric_direct_complement=False,asymmetric_normalize_high=False,orthogonalize_fn=orthogonalize_fn,asymmetric_high_orthogonalize_fn=high_orthogonalize_fn);child_values,child_vectors=_e208_dense1024_leaf_eigh(children,final_newton=False);batch=data.shape[0];vectors=torch.empty_like(data);torch.bmm(basis[:,:,:256],child_vectors[:batch],out=vectors[:,:,:256]);torch.bmm(basis[:,:,256:],child_vectors[batch:],out=vectors[:,:,256:]);values=torch.cat((child_values[:batch],child_values[batch:]),dim=-1);values,order=values.sort(dim=-1);vectors=_e160_gather_columns(vectors,order);return vectors,values
@torch.no_grad()
def _dense2048_half_gram(matrix:torch.Tensor):matrix16=matrix.to(torch.float16);return torch.bmm(matrix16.mT,matrix16,out_dtype=torch.float32)
@torch.no_grad()
def _dense2048_cqr1024_half(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,inverse_precision:str='highest',trsm_fn=None,**_):
result=matrix
for pass_index in range(passes):gram=_dense2048_half_gram(result);scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);gram.diagonal(dim1=-2,dim2=-1).add_(pass_ridge*scale[:,None]);lower=torch.linalg.cholesky_ex(gram,check_errors=False)[0];result=block_inverse_triangular_solve(result,lower,precision=inverse_precision,trsm_fn=trsm_fn)
return result
@torch.no_grad()
def _dense2048_cqr512_half(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,potrf_fn=None,trsm_fn=None,**_):
result=matrix
for pass_index in range(passes):gram=_dense2048_half_gram(result);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);result=_e084_solve512(result,_e084_factor512(gram,pass_ridge,potrf_fn=potrf_fn),trsm_fn=trsm_fn)
return result
@torch.no_grad()
def _dense2048_cqr256_half(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,potrf_fn=None,trsm_fn=None,**_):
if trsm_fn is None:trsm_fn=_e084_trsm256_tensor_strided
result=matrix
for pass_index in range(passes):gram=_dense2048_half_gram(result);scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);factor=_e084_potrf256_strided if potrf_fn is None else potrf_fn;lower=factor(gram,scale,leading_dimension=256,row_offset=0,column_offset=0,ridge=pass_ridge);result=trsm_fn(result,lower,source_rows=result.shape[1],source_columns=256,row_offset=0,column_offset=0,rows=result.shape[1])
return result
@torch.no_grad()
def _dense2048_owned_newton_repair(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
'Repair finite output-certificate rows before the system EVD.';batch,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:gram=vectors.mT@vectors;transform=-.5*gram;transform.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@transform;products=data@vectors;gram=vectors.mT@vectors
finally:torch.set_float32_matmul_precision(previous)
matrix_scale=data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-30);residual=(products-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);eye=torch.eye(n,device=data.device,dtype=data.dtype);orthogonality=(gram-eye).abs().sum(dim=1).amax(dim=1);eps=torch.finfo(torch.float32).eps;hard=~torch.isfinite(residual)|~torch.isfinite(orthogonality)|(residual>.95*2e2*n*eps*matrix_scale)|(orthogonality>.95*1e2*n*eps)
if _certificate_reason_ledger is not None:_certificate_reason_record('dense2048_owned_newton',_certificate_reason_bits(hard,(_CERT_REASON_NONFINITE,~torch.isfinite(residual)|~torch.isfinite(orthogonality)),(_CERT_REASON_EIGEN_RESIDUAL,residual>.95*2e2*n*eps*matrix_scale),(_CERT_REASON_ORTHOGONALITY,orthogonality>.95*1e2*n*eps)),hard)
if bool(hard.any().item()):exact_values,exact_vectors=torch.linalg.eigh(data[hard].contiguous());vectors[hard]=exact_vectors;values[hard]=exact_values
return vectors,values
@memo(maxsize=4)
def _dense2048_certificate_indices(device_index:int):device=torch.device('cuda',device_index);orthogonal_centers=torch.tensor((1440,1536,1792,1952),device=device);orthogonal_offsets=torch.arange(-16,17,device=device);orthogonal_indices=(orthogonal_centers[:,None]+orthogonal_offsets[None,:]).reshape(-1);split_centers=torch.tensor((512,1024,1536),device=device);split_offsets=torch.arange(-2,2,device=device);boundary_indices=(split_centers[:,None]+split_offsets[None,:]).reshape(-1);return orthogonal_indices,boundary_indices
def _guarded_large_eigh(data:torch.Tensor,*,low_precision_sign:bool=True):
rank512_calls={'value':0}
def orthogonalize_large(matrix:torch.Tensor,**kwargs):
rank=matrix.shape[-1]
if rank==1024:return _dense2048_cqr1024_half(matrix,trsm_fn=_e2225_tcgen_trsm256,**kwargs)
if rank==512:call=rank512_calls['value'];rank512_calls['value']+=1;return _e3180_panel_cqr512_half(matrix,solve_call_base=2*call,**kwargs)
return cholesky_orthonormalize(matrix,**kwargs)
def leaf_low_orthogonalize(matrix:torch.Tensor,**kwargs):return _e3180_panel_cqr256_half(matrix,x3=False,**kwargs)
def leaf_high_orthogonalize(matrix:torch.Tensor,**kwargs):return _e3180_panel_cqr256_half(matrix,x3=True,**kwargs)
def large_leaf(matrix:torch.Tensor):return _e503_large_dense_recursive512_leaf(matrix,orthogonalize_fn=leaf_low_orthogonalize,high_orthogonalize_fn=leaf_high_orthogonalize)
try:vectors,values=spectral_split_eigh_2048(data,sign_iterations=3,range_iterations=8,reorthogonalize_every=4,lanczos_steps=3,split_levels=2,child_sign_iterations=5,child_range_iterations=8,child_lanczos_steps=3,child_boundary_width=0,boundary_width=0,newton_steps=1,rayleigh_values=False,low_precision_sign=low_precision_sign,asymmetric_high_iterations=1,asymmetric_direct_complement=False,asymmetric_normalize_high=False,orthogonalize_fn=orthogonalize_large,root_lanczos_stats_fn=_e931_fused_n2048_lanczos3,child_lanczos_stats_fn=_e940_fused_n1024_lanczos3,leaf_solver=large_leaf,root_minimax_high_checkpoint=True,diagonal_children=True)
except torch.linalg.LinAlgError:
if _certificate_reason_ledger is not None:factor_risk=torch.ones((data.shape[0],),device=data.device,dtype=torch.bool);_certificate_reason_record('dense2048_factorization',_certificate_reason_bits(factor_risk,(_CERT_REASON_FACTORIZATION,factor_risk)),factor_risk)
exact_values,exact_vectors=torch.linalg.eigh(data);return exact_vectors,exact_values
_,n,_=data.shape;probes=_rademacher_probes(n,16,data.device);probe_rhs=probes[None].expand(data.shape[0],-1,-1);scaled_rhs=values[:,:,None]*probe_rhs;q_products=vectors@torch.cat((probe_rhs,scaled_rhs),dim=2);compressed_vectors=q_products[:,:,:16];compressed_scaled_vectors=q_products[:,:,16:];device_index=data.device.index
if device_index is None:device_index=torch.cuda.current_device()
orthogonal_indices,boundary_indices=_dense2048_certificate_indices(device_index);selected_orthogonal_vectors=vectors.index_select(2,orthogonal_indices);boundary_vectors=vectors.index_select(2,boundary_indices);boundary_values=values.index_select(1,boundary_indices);a_products=data@torch.cat((probe_rhs,compressed_vectors,boundary_vectors),dim=2);source_probe=a_products[:,:,:16];compressed_residual=a_products[:,:,16:32]-compressed_scaled_vectors;boundary_residual=a_products[:,:,32:]-boundary_vectors*boundary_values[:,None,:];q_grams=vectors.mT@torch.cat((compressed_vectors,selected_orthogonal_vectors),dim=2);orthogonality=(q_grams[:,:,:16]-probe_rhs).square().sum(dim=1).sqrt().amax(dim=1);selected_gram=q_grams[:,:,16:];selected_gram[:,orthogonal_indices,torch.arange(orthogonal_indices.numel(),device=data.device)]-=1.;selected_orthogonality=selected_gram.abs().sum(dim=1).amax(dim=1);eigen_score=compressed_residual.square().sum(dim=1).sqrt().amax(dim=1)/source_probe.square().sum(dim=1).sqrt().amax(dim=1).clamp_min(1e-30);matrix_scale=data.abs().sum(dim=1).amax(dim=1).clamp_min(1e-30);boundary_score=boundary_residual.abs().sum(dim=1).amax(dim=1)/matrix_scale;minimum_safe_scale=32.*float(n)*torch.finfo(torch.float32).eps;safe=(eigen_score<=.025)&(orthogonality<=.00066)&(selected_orthogonality<=.95*1e2*float(n)*torch.finfo(torch.float32).eps)&(boundary_score<=.03)&(matrix_scale>=minimum_safe_scale)
if _certificate_reason_ledger is not None:finite_risk=~(torch.isfinite(eigen_score)&torch.isfinite(orthogonality)&torch.isfinite(selected_orthogonality)&torch.isfinite(boundary_score)&torch.isfinite(matrix_scale));_certificate_reason_record('dense2048_output',_certificate_reason_bits(safe,(_CERT_REASON_NONFINITE,finite_risk),(_CERT_REASON_EIGEN_RESIDUAL,eigen_score>.025),(_CERT_REASON_ORTHOGONALITY,(orthogonality>.00066)|(selected_orthogonality>.95*1e2*float(n)*torch.finfo(torch.float32).eps)),(_CERT_REASON_SPLIT_BOUNDARY,boundary_score>.03),(_CERT_REASON_SCALE,matrix_scale<minimum_safe_scale)),~safe)
if bool(safe.all().item()):return vectors,values
fallback=~safe;vectors=vectors.clone();values=values.clone();finite=torch.isfinite(eigen_score)&torch.isfinite(orthogonality)&torch.isfinite(selected_orthogonality)&torch.isfinite(boundary_score)&torch.isfinite(matrix_scale);finite_bad=fallback&finite;nonfinite_bad=fallback&~finite
if bool(finite_bad.any().item()):repair_vectors,repair_values=_dense2048_owned_newton_repair(data[finite_bad].contiguous(),vectors[finite_bad].contiguous(),values[finite_bad].contiguous());vectors[finite_bad]=repair_vectors;values[finite_bad]=repair_values
if bool(nonfinite_bad.any().item()):
if low_precision_sign:repair_vectors,repair_values=_guarded_large_eigh(data[nonfinite_bad].contiguous(),low_precision_sign=False)
else:repair_values,repair_vectors=torch.linalg.eigh(data[nonfinite_bad].contiguous())
vectors[nonfinite_bad]=repair_vectors;values[nonfinite_bad]=repair_values
return vectors,values
@torch.no_grad()
def _rankdef512_owned_newton_repair(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor):
'Repair rare finite rankdef rows before retaining the exact-EVD gate.';_,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:gram=vectors.mT@vectors;transform=-.5*gram;transform.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@transform;products=data@vectors;gram=vectors.mT@vectors
finally:torch.set_float32_matmul_precision(previous)
residual=(products-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);eye=torch.eye(n,device=data.device,dtype=data.dtype);orthogonality=(gram-eye).abs().sum(dim=1).amax(dim=1);eps=torch.finfo(torch.float32).eps;residual_risk=residual>.95*2e2*n*eps*matrix_scale;orthogonality_risk=orthogonality>.95*1e2*n*eps;nonfinite=~torch.isfinite(residual)|~torch.isfinite(orthogonality);hard=nonfinite|residual_risk|orthogonality_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('rankdef512_owned_newton',_certificate_reason_bits(hard,(_CERT_REASON_NONFINITE,nonfinite),(_CERT_REASON_EIGEN_RESIDUAL,residual_risk),(_CERT_REASON_ORTHOGONALITY,orthogonality_risk)),hard)
if bool(hard.any().item()):exact_values,exact_vectors=torch.linalg.eigh(data[hard].contiguous());vectors[hard]=exact_vectors;values[hard]=exact_values
return vectors,values
@torch.no_grad()
def _e204_certify_rankdef512_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor|None=None):
batch,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
windows=(1,4,-5,6),(7,16,-5,6),(5,8,-5,6),(13,16,-5,6);index_chunks=[]
for(numerator,denominator,left,right)in windows:center=numerator*n//denominator;index_chunks.append(torch.arange(max(0,center+left),min(n,center+right),device=data.device))
indices=torch.cat(index_chunks);selected=vectors.index_select(2,indices);selected_values=values.index_select(1,indices);probes=_rademacher_probes(n,4,data.device);probe_rhs=probes[None].expand(batch,-1,-1);scaled_rhs=values[:,:,None]*probe_rhs;q_products=vectors@torch.cat((probe_rhs,scaled_rhs),dim=2);compressed=q_products[:,:,:4];compressed_scaled=q_products[:,:,4:];packed_rhs=torch.cat((probe_rhs,compressed,selected),dim=2);products=data@packed_rhs;source_probe=products[:,:,:4];compressed_residual=products[:,:,4:8]-compressed_scaled;residual=products[:,:,8:]-selected*selected_values[:,None,:]
if matrix_scale is None:matrix_scale=data.abs().sum(dim=1).amax(dim=1)
matrix_scale=matrix_scale.clamp_min(1e-30);eigen_score=residual.abs().sum(dim=1).amax(dim=1)/matrix_scale;value_scale=values.abs().amax(dim=1,keepdim=True).clamp_min(1.);ordering_score=(values[:,:-1]-values[:,1:]).amax(dim=1)/value_scale[:,0];probe_score=compressed_residual.square().sum(dim=1).sqrt().amax(dim=1)/source_probe.square().sum(dim=1).sqrt().amax(dim=1).clamp_min(1e-30);orthogonality_probe_score=(vectors.mT@compressed-probe_rhs).square().sum(dim=1).sqrt().amax(dim=1);eps=torch.finfo(torch.float32).eps;safe=(eigen_score<=.975*2e2*float(n)*eps)&(ordering_score<=.975*1e2*float(n)*eps)&(probe_score<=.005)&(orthogonality_probe_score<=.0003)&torch.isfinite(eigen_score)&torch.isfinite(ordering_score)&torch.isfinite(probe_score)&torch.isfinite(orthogonality_probe_score)
finally:torch.set_float32_matmul_precision(previous)
tail_norm_error=(vectors[:,:,432:512].square().sum(dim=1)-1.).abs().amax(dim=1);safe&=torch.isfinite(tail_norm_error)&(tail_norm_error<=.002)
if bool(safe.all().item()):return vectors,values
repair=~safe;repair_vectors,repair_values=_rankdef512_owned_newton_repair(data[repair].contiguous(),vectors[repair].contiguous(),values[repair].contiguous(),matrix_scale[repair].contiguous());vectors=vectors.clone();values=values.clone();vectors[repair]=repair_vectors;values[repair]=repair_values;return vectors,values
@triton.jit
def _e668_certificate_partial_reduction_kernel(data,selected_vectors,selected_values,products,gram,values,workspace,selected_count:tl.constexpr,n:tl.constexpr,block_cols:tl.constexpr,supplied_scale:tl.constexpr):
batch=tl.program_id(0);column_tile=tl.program_id(1);columns=column_tile*block_cols+tl.arange(0,block_cols);column_mask=columns<n;rows=tl.arange(0,n)
if not supplied_scale:data_tile=tl.load(data+batch*n*n+rows[:,None]*n+columns[None,:],mask=column_mask[None,:],other=.0);matrix_l1=tl.sum(tl.abs(data_tile),axis=0)
selected_mask=columns<selected_count;product_tile=tl.load(products+batch*n*selected_count+rows[:,None]*selected_count+columns[None,:],mask=selected_mask[None,:],other=.0);vector_tile=tl.load(selected_vectors+batch*n*selected_count+rows[:,None]*selected_count+columns[None,:],mask=selected_mask[None,:],other=.0);local_values=tl.load(selected_values+batch*selected_count+columns,mask=selected_mask,other=.0);residual_l1=tl.sum(tl.abs(product_tile-vector_tile*local_values[None,:]),axis=0);gram_rows=tl.arange(0,128);gram_values=tl.load(gram+batch*selected_count*selected_count+gram_rows[:,None]*selected_count+columns[None,:],mask=(gram_rows[:,None]<selected_count)&selected_mask[None,:],other=.0);gram_values-=tl.where((gram_rows[:,None]==columns[None,:])&selected_mask[None,:],1.,.0);orthogonality_l1=tl.sum(tl.abs(gram_values),axis=0);left=tl.load(values+batch*n+columns,mask=column_mask,other=-float('inf'));right=tl.load(values+batch*n+columns+1,mask=columns+1<n,other=float('inf'));ordering_error=left-right;base=batch*4*n+columns
if not supplied_scale:tl.store(workspace+base,matrix_l1,mask=column_mask)
tl.store(workspace+base+n,residual_l1,mask=column_mask);tl.store(workspace+base+2*n,orthogonality_l1,mask=column_mask);tl.store(workspace+base+3*n,ordering_error,mask=column_mask)
@triton.jit
def _e668_certificate_final_reduction_kernel(workspace,values,matrix_scale_input,safe,scores,matrix_scale_stride:tl.constexpr,supplied_scale:tl.constexpr,selected_count:tl.constexpr,eigen_limit:tl.constexpr,ordering_limit:tl.constexpr,n:tl.constexpr):
batch=tl.program_id(0);columns=tl.arange(0,n);base=batch*4*n+columns
if supplied_scale:matrix_scale=tl.load(matrix_scale_input+batch*matrix_scale_stride)
else:matrix_scale=tl.max(tl.load(workspace+base),axis=0)
matrix_scale=tl.maximum(matrix_scale,1e-30)
if supplied_scale:
workspace_base=batch*4*n
if selected_count<=16:selected_columns=tl.arange(0,16);eigen_residual=tl.max(tl.load(workspace+workspace_base+n+selected_columns),axis=0);orthogonality_score=tl.max(tl.load(workspace+workspace_base+2*n+selected_columns),axis=0)
else:selected_columns=tl.arange(0,128);selected_mask=selected_columns<selected_count;eigen_residual=tl.max(tl.load(workspace+workspace_base+n+selected_columns,mask=selected_mask,other=.0),axis=0);orthogonality_score=tl.max(tl.load(workspace+workspace_base+2*n+selected_columns,mask=selected_mask,other=.0),axis=0)
right=tl.load(values+batch*n+columns+1,mask=columns+1<n,other=float('inf'));ordering_error=tl.max(tl.load(values+batch*n+columns)-right,axis=0)
else:eigen_residual=tl.max(tl.load(workspace+base+n),axis=0);orthogonality_score=tl.max(tl.load(workspace+base+2*n),axis=0);ordering_error=tl.max(tl.load(workspace+base+3*n),axis=0)
value_scale=tl.max(tl.abs(tl.load(values+batch*n+columns)),axis=0);value_scale=tl.maximum(value_scale,1.);eigen_score=eigen_residual/matrix_scale;ordering_score=ordering_error/value_scale;decision=(eigen_score<=eigen_limit)&(orthogonality_score<=.002)&(ordering_score<=ordering_limit)&(eigen_score==eigen_score)&(orthogonality_score==orthogonality_score)&(ordering_score==ordering_score)&(tl.abs(eigen_score)<float('inf'))&(tl.abs(orthogonality_score)<float('inf'))&(tl.abs(ordering_score)<float('inf'));tl.store(safe+batch,decision);score_base=batch*4;tl.store(scores+score_base,eigen_score);tl.store(scores+score_base+1,orthogonality_score);tl.store(scores+score_base+2,ordering_score);tl.store(scores+score_base+3,matrix_scale)
_E2115_MIXED_CERTIFICATE_RESIDUAL_NAME='e2115_mixed_certificate_residual';_E2115_MIXED_CERTIFICATE_FINISH_NAME='e2115_mixed_certificate_finish';_E2133_MIXED_CERTIFICATE_B64_NAME='e2133_mixed_certificate_residual_b64';_E2133_MIXED_CERTIFICATE_B96_NAME='e2133_mixed_certificate_residual_b96';_E2115_MIXED_CERTIFICATE_SOURCE='\n#include <mma.h>\n#include <cuda_fp16.h>\nusing namespace nvcuda;\n__device__ __forceinline__ float e2115_tf32(float x){\n unsigned int b=__float_as_uint(x),e=b&0x7f800000u;\n if(e!=0x7f800000u)b=(b+0x00000fffu+((b>>13)&1u))&0xffffe000u;\n return __uint_as_float(b);\n}\n__device__ __forceinline__ float e2115_max(float x){\n #pragma unroll\n for(int o=16;o;o>>=1)x=fmaxf(x,__shfl_down_sync(0xffffffffu,x,o));\n return x;\n}\nextern "C" __global__ __launch_bounds__(128,2)\nvoid e2115_mixed_certificate_residual(\n const float* __restrict__ a,const float* __restrict__ q,\n const float* __restrict__ l,const float* __restrict__ scale,\n const int* __restrict__ indices,\n float* __restrict__ sums,int batch,int count){\n constexpr int N=512,BM=32,BN=32,BK=64;\n const int m=blockIdx.x,row0=blockIdx.y*BM,pos0=blockIdx.z*BN;\n const int t=threadIdx.x,w=t>>5;\n if(m>=batch)return;\n extern __shared__ __align__(1024) unsigned char raw[];\n __half* at=(__half*)raw;__half* bt=at+BM*BK;\n float* out=(float*)(bt+BK*BN);float* residual_tile=out+BM*BN;\n const long long base=(long long)m*N*N;\n const float matrix_scale=fmaxf(scale[m],1.e-30f);\n const float inv_scale=1.f/matrix_scale;\n wmma::fragment<wmma::accumulator,16,16,16,float> acc;\n wmma::fill_fragment(acc,0.f);\n #pragma unroll\n for(int start=0;start<N;start+=BK){\n for(int x=t;x<BM*BK;x+=blockDim.x){\n int r=x/BK,k=x-r*BK;\n at[x]=__float2half_rn(a[base+(long long)(row0+r)*N+start+k]*inv_scale);\n }\n for(int x=t;x<BK*BN;x+=blockDim.x){\n int k=x/BN,j=x-k*BN,p=pos0+j,c=p<count?indices[p]:0;\n bt[x]=p<count?__float2half_rn(q[base+(long long)(start+k)*N+c]):__float2half_rn(0.f);\n }\n __syncthreads();\n const int wm=w>>1,wn=w&1;\n #pragma unroll\n for(int k=0;k<BK;k+=16){\n wmma::fragment<wmma::matrix_a,16,16,16,__half,wmma::row_major> af;\n wmma::fragment<wmma::matrix_b,16,16,16,__half,wmma::row_major> bf;\n wmma::load_matrix_sync(af,at+wm*16*BK+k,BK);\n wmma::load_matrix_sync(bf,bt+k*BN+wn*16,BN);\n wmma::mma_sync(acc,af,bf,acc);\n }\n __syncthreads();\n }\n wmma::store_matrix_sync(out+w*256,acc,16,wmma::mem_row_major);\n __syncthreads();\n for(int x=t;x<BM*BN;x+=blockDim.x){\n int r=x/BN,j=x-r*BN,owner=(r>>4)*2+(j>>4),p=pos0+j;\n float product=out[owner*256+(r&15)*16+(j&15)],residual=0.f;\n if(p<count){\n int c=indices[p];\n residual=fabsf(product-q[base+(long long)(row0+r)*N+c]\n *l[(long long)m*N+c]*inv_scale)*matrix_scale;\n }\n residual_tile[x]=residual;\n }\n __syncthreads();\n if(t<BN&&pos0+t<count){\n float total=0.f;\n #pragma unroll\n for(int r=0;r<BM;++r)total+=residual_tile[r*BN+t];\n atomicAdd(sums+(long long)m*128+pos0+t,total);\n }\n}\nextern "C" __global__ __launch_bounds__(256,2)\nvoid e2115_mixed_certificate_finish(\n const float* __restrict__ sums,const float* __restrict__ values,\n const float* __restrict__ scale,bool* __restrict__ risk,\n int batch,int count,float margin){\n constexpr int N=512;\n int m=blockIdx.x,t=threadIdx.x,w=t>>5,z=t&31;\n if(m>=batch)return;\n float eigen=t<count?sums[(long long)m*128+t]:0.f;\n float ordering=-1.e30f,value_scale=1.f;\n for(int c=t;c<N;c+=256){\n float value=values[(long long)m*N+c];\n value_scale=fmaxf(value_scale,fabsf(value));\n if(c+1<N)ordering=fmaxf(ordering,value-values[(long long)m*N+c+1]);\n }\n eigen=e2115_max(eigen);ordering=e2115_max(ordering);value_scale=e2115_max(value_scale);\n __shared__ float es[8],os[8],vs[8];\n if(z==0){es[w]=eigen;os[w]=ordering;vs[w]=value_scale;}\n __syncthreads();\n if(w==0){\n eigen=e2115_max(z<8?es[z]:0.f);\n ordering=e2115_max(z<8?os[z]:-1.e30f);\n value_scale=e2115_max(z<8?vs[z]:1.f);\n if(z==0){\n float e=eigen/fmaxf(scale[m],1.e-30f),o=ordering/value_scale;\n risk[m]=!isfinite(e)||!isfinite(o)\n ||e>margin*200.f*N*1.1920928955078125e-7f\n ||o>0.975f*100.f*N*1.1920928955078125e-7f;\n }\n }\n}\n'
def _e2133_mixed_certificate_source():
def variant(name:str,bn:int):
source=_E2115_MIXED_CERTIFICATE_SOURCE;replacements=(_E2115_MIXED_CERTIFICATE_RESIDUAL_NAME,name),('__launch_bounds__(128,2)\nvoid '+name+'(',f"__launch_bounds__({bn*4},2)\nvoid {name}("),('constexpr int N=512,BM=32,BN=32,BK=64;',f"constexpr int N=512,BM=32,BN={bn},BK=64;"),('const int wm=w>>1,wn=w&1;',f"const int wm=w/{bn//16},wn=w-wm*{bn//16};"),('owner=(r>>4)*2+(j>>4)',f"owner=(r>>4)*{bn//16}+(j>>4)")
for(old,new)in replacements:
if source.count(old)!=1:raise RuntimeError('E2133 certificate source anchor changed')
source=source.replace(old,new,1)
return source
source64=variant(_E2133_MIXED_CERTIFICATE_B64_NAME,64);source96=_fast_only_cuda_kernel(variant(_E2133_MIXED_CERTIFICATE_B96_NAME,96),_E2133_MIXED_CERTIFICATE_B96_NAME);return source64+source96[source96.find('extern "C"'):]
@memo(maxsize=1)
def _e2115_mixed_certificate_kernels():image=_fast_nvrtc_compile(_e2133_mixed_certificate_source(),_E2133_MIXED_CERTIFICATE_B64_NAME);return CUDAKernel(image,_E2133_MIXED_CERTIFICATE_B64_NAME),CUDAKernel(image,_E2133_MIXED_CERTIFICATE_B96_NAME),CUDAKernel(image,_E2115_MIXED_CERTIFICATE_FINISH_NAME)
@memo(maxsize=8)
def _e2115_mixed_certificate_indices(device:int,route:int):values=(tuple(range(120,136))+tuple(range(220,228))+tuple(range(312,328))+tuple(range(412,420))+tuple(range(472,512)),tuple(range(24))+tuple(range(488,512)),tuple(range(192,224))+tuple(range(264,280)),tuple(range(24))+tuple(range(120,136))+tuple(range(488,512)))[route];return torch.tensor(values,device=device,dtype=torch.int32)
@memo(maxsize=4)
def _e2115_psd_orthogonality_indices(device:int):return torch.arange(192,224,device=device,dtype=torch.int64)
@torch.no_grad()
def _e2115_certify_mixed512_routes(inputs:tuple[torch.Tensor,...],outputs:tuple[output_t,...],scales:tuple[torch.Tensor,...],*,return_risk:bool=False):
counts=tuple(matrix.shape[0]for matrix in inputs);starts=[0]
for count in counts:starts.append(starts[-1]+count)
total=starts[-1]
if total==0:
if return_risk:empty=tuple(torch.empty((0,),device=matrix.device,dtype=torch.bool)for matrix in inputs);return outputs,empty
return outputs
device=inputs[0].device.index
if device is None:device=torch.cuda.current_device()
sums=torch.zeros((total,128),device=inputs[0].device,dtype=torch.float32);risk=torch.empty((total,),device=inputs[0].device,dtype=torch.bool);residual64,residual96,finish=_e2115_mixed_certificate_kernels();shared64=(32*64+64*64)*2+2*32*64*4;shared96=(32*64+64*96)*2+2*32*96*4;margins=.9,.995,.995,.995
for route in range(4):
batch=counts[route]
if batch==0:continue
local=slice(starts[route],starts[route+1]);indices=_e2115_mixed_certificate_indices(device,route);selected_count=indices.numel();vectors,values=outputs[route];residual=residual96 if route==0 else residual64;tile=96 if route==0 else 64;threads=384 if route==0 else 256;shared=shared96 if route==0 else shared64;residual.launch((batch,16,(selected_count+tile-1)//tile),(threads,1,1),(inputs[route],vectors,values,scales[route],indices,sums[local],batch,selected_count),shared_mem=shared);finish.launch((batch,1,1),(256,1,1),(sums[local],values,scales[route],risk[local],batch,selected_count,margins[route]))
if route>=2:risk[local]|=_e2073_column_norm_risk(vectors,.002)
if route==2:orthogonal_vectors=vectors.index_select(2,_e2115_psd_orthogonality_indices(device));gram=orthogonal_vectors.mT@orthogonal_vectors;gram.diagonal(dim1=-2,dim2=-1).sub_(1.);orthogonal_score=gram.abs().sum(dim=1).amax(dim=1);risk[local]|=~torch.isfinite(orthogonal_score)|(orthogonal_score>.002)
if return_risk:return outputs,tuple(risk[starts[route]:starts[route+1]]for route in range(4))
if not bool(risk.any().item()):return outputs
repaired=list(outputs)
for route in range(4):
local_risk=risk[starts[route]:starts[route+1]]
if bool(local_risk.any().item()):exact_values,exact_vectors=torch.linalg.eigh(inputs[route][local_risk].contiguous());vectors,values=outputs[route];vectors=vectors.clone();values=values.clone();vectors[local_risk]=exact_vectors;values[local_risk]=exact_values;repaired[route]=vectors,values
return tuple(repaired)
def _e668_supported_fused_certificate(data:torch.Tensor,windows:tuple[tuple[int,int,int,int],...],adaptive_gap_count:int,eigen_margin:float,residual_probe_count:int,residual_probe_threshold:float,matrix_scale:torch.Tensor|None):return bool(data.shape[-1]==512 and windows==((1,4,-8,8),(7,16,-4,4),(5,8,-8,8),(13,16,-4,4),(1,1,-40,0))and adaptive_gap_count==0 and eigen_margin==.9 and residual_probe_count==0 and residual_probe_threshold==.0 and matrix_scale is None)
@torch.no_grad()
def _e668_certify_mixed512_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,windows:tuple[tuple[int,int,int,int],...],eigen_margin:float):
batch=data.shape[0];previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
index_chunks=[]
for(numerator,denominator,left,right)in windows:center=numerator*512//denominator;index_chunks.append(torch.arange(max(0,center+left),min(512,center+right),device=data.device))
indices=torch.cat(index_chunks);selected_vectors=vectors.index_select(2,indices);selected_values=values.index_select(1,indices);selected_count=selected_vectors.shape[-1];products=data@selected_vectors;gram=selected_vectors.mT@selected_vectors;workspace=torch.empty((batch,4,512),device=data.device,dtype=torch.float32);safe=torch.empty((batch,),device=data.device,dtype=torch.bool);scores=torch.empty((batch,4),device=data.device,dtype=torch.float32);_e668_certificate_partial_reduction_kernel[batch,triton.cdiv(512,8)](data,selected_vectors,selected_values,products,gram,values,workspace,selected_count=selected_count,n=512,block_cols=8,supplied_scale=False,num_warps=8,num_stages=1);eps=torch.finfo(torch.float32).eps;_e668_certificate_final_reduction_kernel[batch,](workspace,values,data,safe,scores,matrix_scale_stride=0,supplied_scale=False,selected_count=selected_count,eigen_limit=eigen_margin*2e2*512.*eps,ordering_limit=.975*1e2*512.*eps,n=512,num_warps=8,num_stages=1)
finally:torch.set_float32_matmul_precision(previous_precision)
if bool(safe.all().item()):return vectors,values
repair=~safe;exact_values,exact_vectors=torch.linalg.eigh(data[repair].contiguous());vectors=vectors.clone();values=values.clone();vectors[repair]=exact_vectors;values[repair]=exact_values;return vectors,values
def _e1163_supported_nearrank_certificate(data:torch.Tensor,windows:tuple[tuple[int,int,int,int],...],adaptive_gap_count:int,eigen_margin:float,residual_probe_count:int,residual_probe_threshold:float,matrix_scale:torch.Tensor|None):return bool(data.shape[-1]==1024 and windows==((1,4,-4,4),)and adaptive_gap_count==8 and eigen_margin==.82 and residual_probe_count==0 and residual_probe_threshold==.0 and matrix_scale is not None)
@torch.no_grad()
def _e1167_certify_selected16_supplied_scale(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,selected_vectors:torch.Tensor,selected_values:torch.Tensor,matrix_scale:torch.Tensor|None,eigen_margin:float):
batch,n,_=data.shape;selected_count=selected_vectors.shape[-1];previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
products=data@selected_vectors;gram=selected_vectors.mT@selected_vectors
if matrix_scale is None:matrix_scale=data.abs().sum(dim=1).amax(dim=1)
workspace=torch.empty((batch,4,n),device=data.device,dtype=torch.float32);safe=torch.empty((batch,),device=data.device,dtype=torch.bool);scores=torch.empty((batch,4),device=data.device,dtype=torch.float32);_e668_certificate_partial_reduction_kernel[batch,triton.cdiv(selected_count,8)](data,selected_vectors,selected_values,products,gram,values,workspace,selected_count=selected_count,n=n,block_cols=8,supplied_scale=True,num_warps=8,num_stages=1);eps=torch.finfo(torch.float32).eps;_e668_certificate_final_reduction_kernel[batch,](workspace,values,matrix_scale,safe,scores,matrix_scale_stride=matrix_scale.stride(0),supplied_scale=True,selected_count=selected_count,eigen_limit=eigen_margin*2e2*float(n)*eps,ordering_limit=.975*1e2*float(n)*eps,n=n,num_warps=8,num_stages=1)
finally:torch.set_float32_matmul_precision(previous_precision)
if bool(safe.all().item()):return vectors,values
repair=~safe;exact_values,exact_vectors=torch.linalg.eigh(data[repair].contiguous());vectors=vectors.clone();values=values.clone();vectors[repair]=exact_vectors;values[repair]=exact_values;return vectors,values
@memo(maxsize=4)
def _e2782_nearrank_certificate_indices(device:int|None):
if device is None:device=torch.cuda.current_device()
return torch.cat((torch.arange(624,632,device=device),torch.arange(816,832,device=device)))
@triton.jit
def _e2782_nearrank_certificate_partial_kernel(inner_vectors,inner_values,outer_values,packed_vectors,packed_products,gram,workspace,packed_count:tl.constexpr,inner_count:tl.constexpr,outer_count:tl.constexpr,n:tl.constexpr,block_cols:tl.constexpr):batch=tl.program_id(0);columns=tl.program_id(1)*block_cols+tl.arange(0,block_cols);rows=tl.arange(0,n);inner_mask=columns<inner_count;inner_products=tl.load(packed_products+batch*n*packed_count+rows[:,None]*packed_count+columns[None,:],mask=inner_mask[None,:],other=.0);inner_q=tl.load(inner_vectors+batch*n*inner_count+rows[:,None]*inner_count+columns[None,:],mask=inner_mask[None,:],other=.0);inner_l=tl.load(inner_values+batch*inner_count+columns,mask=inner_mask,other=.0);inner_residual=tl.sum(tl.abs(inner_products-inner_q*inner_l[None,:]),axis=0);gram_rows=tl.arange(0,16);gram_values=tl.load(gram+batch*inner_count*inner_count+gram_rows[:,None]*inner_count+columns[None,:],mask=(gram_rows[:,None]<inner_count)&inner_mask[None,:],other=.0);gram_values-=tl.where((gram_rows[:,None]==columns[None,:])&inner_mask[None,:],1.,.0);inner_orthogonality=tl.sum(tl.abs(gram_values),axis=0);outer_mask=columns<outer_count;outer_column=inner_count+columns;outer_products=tl.load(packed_products+batch*n*packed_count+rows[:,None]*packed_count+outer_column[None,:],mask=outer_mask[None,:],other=.0);outer_q=tl.load(packed_vectors+batch*n*packed_count+rows[:,None]*packed_count+outer_column[None,:],mask=outer_mask[None,:],other=.0);outer_l=tl.load(outer_values+batch*outer_count+columns,mask=outer_mask,other=.0);outer_residual=tl.sum(tl.abs(outer_products-outer_q*outer_l[None,:]),axis=0);base=batch*4*n+columns;tl.store(workspace+base+n,inner_residual,mask=inner_mask);tl.store(workspace+base+2*n,inner_orthogonality,mask=inner_mask);tl.store(workspace+base+3*n,outer_residual,mask=outer_mask)
@torch.no_grad()
def _n1024_owned_local_rr_repair(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor,start:int,stop:int,stage:str):
'Repair a route-owned invariant window, retaining a full-EVD hard gate.';_,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
local_vectors=vectors[:,:,start:stop];products=data@local_vectors;projected=local_vectors.mT@products;projected=.5*(projected+projected.mT)
if stage=='dense1024_owned_rr256'and stop-start==256:local_values,rotation=_e208_dense1024_leaf_eigh(projected,final_newton=True,newton_precision='highest')
else:local_values,rotation=torch.linalg.eigh(projected)
vectors=vectors.clone();values=values.clone();vectors[:,:,start:stop]=local_vectors@rotation;values[:,start:stop]=local_values;full_products=data@vectors;gram=vectors.mT@vectors
finally:torch.set_float32_matmul_precision(previous)
residual=(full_products-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);eye=torch.eye(n,device=data.device,dtype=data.dtype);orthogonality=(gram-eye).abs().sum(dim=1).amax(dim=1);value_scale=values.abs().amax(dim=1).clamp_min(1.);ordering=(values[:,:-1]-values[:,1:]).amax(dim=1);eps=torch.finfo(torch.float32).eps;nonfinite=~torch.isfinite(residual)|~torch.isfinite(orthogonality)|~torch.isfinite(ordering);residual_risk=residual>.95*2e2*n*eps*matrix_scale;orthogonality_risk=orthogonality>.95*1e2*n*eps;ordering_risk=ordering>.95*1e2*n*eps*value_scale;hard=nonfinite|residual_risk|orthogonality_risk|ordering_risk
if _certificate_reason_ledger is not None:_certificate_reason_record(stage,_certificate_reason_bits(hard,(_CERT_REASON_NONFINITE,nonfinite),(_CERT_REASON_EIGEN_RESIDUAL,residual_risk),(_CERT_REASON_ORTHOGONALITY,orthogonality_risk),(_CERT_REASON_ORDERING,ordering_risk)),hard)
if bool(hard.any().item()):exact_values,exact_vectors=torch.linalg.eigh(data[hard].contiguous());vectors[hard]=exact_vectors;values[hard]=exact_values
return vectors,values
@torch.no_grad()
def _e1163_certify_nearrank1024_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,windows:tuple[tuple[int,int,int,int],...],matrix_scale:torch.Tensor):
batch,n,_=data.shape;gaps=(values[:,1:]-values[:,:-1]).abs();relative_gap=torch.minimum(gaps[:,:-1],gaps[:,1:])/values[:,1:-1].abs().clamp_min(1e-30);adaptive=relative_gap[:,255:].topk(8,dim=1,largest=False).indices+256;fixed=torch.arange(252,260,device=data.device)[None].expand(batch,-1);inner_indices=torch.cat((adaptive,fixed),dim=1);inner_vectors=vectors.gather(2,inner_indices[:,None,:].expand(-1,n,-1));inner_values=values.gather(1,inner_indices);outer_indices=_e2782_nearrank_certificate_indices(data.device.index);outer_vectors=vectors.index_select(2,outer_indices);outer_values=values.index_select(1,outer_indices);previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
packed_vectors=torch.cat((inner_vectors,outer_vectors),dim=2);packed_products=data@packed_vectors;gram=inner_vectors.mT@inner_vectors;workspace=torch.empty((batch,4,n),device=data.device,dtype=torch.float32);safe=torch.empty((batch,),device=data.device,dtype=torch.bool);scores=torch.empty((batch,4),device=data.device,dtype=torch.float32);_e2782_nearrank_certificate_partial_kernel[batch,3](inner_vectors,inner_values,outer_values,packed_vectors,packed_products,gram,workspace,packed_count=40,inner_count=16,outer_count=24,n=1024,block_cols=8,num_warps=8,num_stages=1);eps=torch.finfo(torch.float32).eps;_e668_certificate_final_reduction_kernel[batch,](workspace,values,matrix_scale,safe,scores,matrix_scale_stride=matrix_scale.stride(0),supplied_scale=True,selected_count=16,eigen_limit=.82*2e2*1024.*eps,ordering_limit=.975*1e2*1024.*eps,n=1024,num_warps=8,num_stages=1);outer_score=workspace[:,3,:24].amax(dim=1)/matrix_scale.clamp_min(1e-30);outer_nonfinite_risk=~torch.isfinite(outer_score);outer_residual_risk=outer_score>.95*2e2*1024.*eps;outer_norm_risk=_e2073_column_norm_risk(vectors,.002);effective_nonfinite_risk=safe&outer_nonfinite_risk;effective_residual_risk=safe&outer_residual_risk;effective_norm_risk=safe&outer_norm_risk;effective_outer_risk=effective_nonfinite_risk|effective_residual_risk|effective_norm_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('nearrank1024_output',_certificate_reason_bits(effective_outer_risk,(_CERT_REASON_NONFINITE,effective_nonfinite_risk),(_CERT_REASON_EIGEN_RESIDUAL,effective_residual_risk),(_CERT_REASON_NORM,effective_norm_risk)),effective_outer_risk)
safe&=~(outer_nonfinite_risk|outer_residual_risk|outer_norm_risk)
finally:torch.set_float32_matmul_precision(previous_precision)
if bool(safe.all().item()):return vectors,values
repair=~safe;exact_vectors,exact_values=_n1024_owned_local_rr_repair(data[repair].contiguous(),vectors[repair].contiguous(),values[repair].contiguous(),matrix_scale[repair].contiguous(),240,272,'nearrank1024_owned_rr32');vectors=vectors.clone();values=values.clone();vectors[repair]=exact_vectors;values[repair]=exact_values;return vectors,values
_E2073_COLUMN_NORM_NAME='e2080_column_norm_square';_E2073_COLUMN_NORM_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(512,2)\nvoid e2080_column_norm_square(const float* __restrict__ q,\n int* __restrict__ risk,\n int batch,int n,float threshold){\n const int shards=(n+511)/512;\n const int b=blockIdx.x/shards;\n const int col=(blockIdx.x-b*shards)*512+threadIdx.x;\n if(b>=batch||col>=n)return;\n const long long base=(long long)b*n*n;\n float norm=0.f;\n #pragma unroll 4\n for(int row=0;row<n;++row){\n const float value=q[base+(long long)row*n+col];\n norm=fmaf(value,value,norm);\n }\n if(!isfinite(norm)||fabsf(norm-1.f)>threshold)atomicExch(risk+b,1);\n}\n'
@memo(maxsize=1)
def _e2073_column_norm_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2073_COLUMN_NORM_SOURCE,_E2073_COLUMN_NORM_NAME),_E2073_COLUMN_NORM_NAME)
@torch.no_grad()
def _e2073_column_norm_risk(vectors:torch.Tensor,threshold:float):
batch,_,n=vectors.shape
if n not in(176,352,512,1024):raise ValueError('column norm guard requires n=176, n=352, n=512, or n=1024')
risk=torch.zeros((batch,),device=vectors.device,dtype=torch.int32);shards=(n+511)//512;_e2073_column_norm_kernel().launch((batch*shards,1,1),(512,1,1),(vectors,risk,batch,n,float(threshold)));return risk!=0
def _e1242_supported_dense1024_certificate(data:torch.Tensor,windows:tuple[tuple[int,int,int,int],...],adaptive_gap_count:int,eigen_margin:float,residual_probe_count:int,residual_probe_threshold:float,matrix_scale:torch.Tensor|None):signature=data.shape[-1]==1024 and windows in(((1,4,-16,17),(1,2,-16,17),(3,4,-16,17)),((1,4,-16,17),(1,2,-16,17),(3,4,-16,17),(7,8,-4,5)))or data.shape[-1]==512 and windows==((0,1,24,64),(1,4,-8,9),(3,4,-32,9));return bool(signature and adaptive_gap_count==0 and eigen_margin==.975 and residual_probe_count==0 and residual_probe_threshold==.0 and matrix_scale is not None)
@torch.no_grad()
def _e1242_certify_dense1024_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,windows:tuple[tuple[int,int,int,int],...],matrix_scale:torch.Tensor):
batch,n,_=data.shape;index_chunks=[]
for(numerator,denominator,left,right)in windows:center=numerator*n//denominator;index_chunks.append(torch.arange(max(0,center+left),min(n,center+right),device=data.device))
indices=torch.cat(index_chunks);selected_vectors=vectors.index_select(2,indices);selected_values=values.index_select(1,indices);selected_count=selected_vectors.shape[-1];previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
dense_residual_probe=n==1024
if dense_residual_probe:
probe_count=4 if len(windows)==3 else 8;probe_rhs=_rademacher_probes(n,probe_count,data.device)[None].expand(batch,-1,-1);scaled_rhs=values[:,:,None]*probe_rhs;q_products=vectors@torch.cat((probe_rhs,scaled_rhs),dim=2);compressed=q_products[:,:,:probe_count];compressed_scaled=q_products[:,:,probe_count:]
if len(windows)==3:packed_products=data@torch.cat((probe_rhs,compressed,selected_vectors),dim=2);source_probe=packed_products[:,:,:probe_count];compressed_residual=packed_products[:,:,probe_count:2*probe_count]-compressed_scaled;products=packed_products[:,:,2*probe_count:].contiguous();probe_score=compressed_residual.square().sum(dim=1).sqrt().amax(dim=1)/source_probe.square().sum(dim=1).sqrt().amax(dim=1).clamp_min(1e-30)
else:packed_products=data@torch.cat((compressed,selected_vectors),dim=2);compressed_residual=packed_products[:,:,:probe_count]-compressed_scaled;products=packed_products[:,:,probe_count:].contiguous()
else:products=data@selected_vectors
gram=selected_vectors.mT@selected_vectors;workspace=torch.empty((batch,4,n),device=data.device,dtype=torch.float32);safe=torch.empty((batch,),device=data.device,dtype=torch.bool);scores=torch.empty((batch,4),device=data.device,dtype=torch.float32);_e668_certificate_partial_reduction_kernel[batch,triton.cdiv(selected_count,8)](data,selected_vectors,selected_values,products,gram,values,workspace,selected_count=selected_count,n=n,block_cols=8,supplied_scale=True,num_warps=8,num_stages=1);eps=torch.finfo(torch.float32).eps;_e668_certificate_final_reduction_kernel[batch,](workspace,values,matrix_scale,safe,scores,matrix_scale_stride=matrix_scale.stride(0),supplied_scale=True,selected_count=selected_count,eigen_limit=.975*2e2*float(n)*eps,ordering_limit=.975*1e2*float(n)*eps,n=n,num_warps=8,num_stages=1)
if n==1024 and len(windows)==3:safe&=torch.isfinite(probe_score)&(probe_score<=.03)
if n==1024 and len(windows)==4:
safe&=~_e2073_column_norm_risk(vectors,.05);outer_eigen_score=compressed_residual.abs().sum(dim=1).amax(dim=1)/matrix_scale.clamp_min(1e-30);outer_nonfinite_risk=~torch.isfinite(outer_eigen_score);outer_residual_risk=outer_eigen_score>.018;effective_nonfinite_risk=safe&outer_nonfinite_risk;effective_residual_risk=safe&outer_residual_risk;effective_outer_risk=effective_nonfinite_risk|effective_residual_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('n1024_output',_certificate_reason_bits(effective_outer_risk,(_CERT_REASON_NONFINITE,effective_nonfinite_risk),(_CERT_REASON_EIGEN_RESIDUAL,effective_residual_risk)),effective_outer_risk)
safe&=~(outer_nonfinite_risk|outer_residual_risk)
finally:torch.set_float32_matmul_precision(previous_precision)
if bool(safe.all().item()):return vectors,values
if n==1024 and len(windows)==3:return _dense_1024_specialized(data,low_precision_sign=False,matrix_scale=matrix_scale,certify_output=False)
repair=~safe
if n==512:exact_vectors,exact_values=_lapack_output_owned_rr128(data[repair].contiguous(),vectors[repair].contiguous(),values[repair].contiguous())
else:exact_vectors,exact_values=_n1024_owned_local_rr_repair(data[repair].contiguous(),vectors[repair].contiguous(),values[repair].contiguous(),matrix_scale[repair].contiguous(),768,1024,'dense1024_owned_rr256')
vectors=vectors.clone();values=values.clone();vectors[repair]=exact_vectors;values[repair]=exact_values;return vectors,values
@torch.no_grad()
def _certify_eigh_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,windows:tuple[tuple[int,int,int,int],...]=(),*,adaptive_gap_count:int=0,eigen_margin:float=.975,residual_probe_count:int=0,residual_probe_threshold:float=.0,matrix_scale:torch.Tensor|None=None):
if _e668_supported_fused_certificate(data,windows,adaptive_gap_count,eigen_margin,residual_probe_count,residual_probe_threshold,matrix_scale):return _e668_certify_mixed512_output(data,vectors,values,windows,eigen_margin)
if _e1163_supported_nearrank_certificate(data,windows,adaptive_gap_count,eigen_margin,residual_probe_count,residual_probe_threshold,matrix_scale):return _e1163_certify_nearrank1024_output(data,vectors,values,windows,matrix_scale)
if _e1242_supported_dense1024_certificate(data,windows,adaptive_gap_count,eigen_margin,residual_probe_count,residual_probe_threshold,matrix_scale):return _e1242_certify_dense1024_output(data,vectors,values,windows,matrix_scale)
_,n,_=data.shape;previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
if adaptive_gap_count:
gaps=(values[:,1:]-values[:,:-1]).abs();relative_gap=torch.minimum(gaps[:,:-1],gaps[:,1:])/values[:,1:-1].abs().clamp_min(1e-30);positive_start=n//4;indices=relative_gap[:,positive_start-1:].topk(adaptive_gap_count,dim=1,largest=False).indices+positive_start
if windows:
fixed_chunks=[]
for(numerator,denominator,left,right)in windows:center=numerator*n//denominator;fixed_chunks.append(torch.arange(max(0,center+left),min(n,center+right),device=data.device))
fixed=torch.cat(fixed_chunks)[None].expand(data.shape[0],-1);indices=torch.cat((indices,fixed),dim=1)
selected_vectors=vectors.gather(2,indices[:,None,:].expand(-1,n,-1));selected_values=values.gather(1,indices)
else:
index_chunks=[]
for(numerator,denominator,left,right)in windows:center=numerator*n//denominator;index_chunks.append(torch.arange(max(0,center+left),min(n,center+right),device=data.device))
indices=torch.cat(index_chunks);selected_vectors=vectors.index_select(2,indices);selected_values=values.index_select(1,indices)
residual=data@selected_vectors-selected_vectors*selected_values[:,None,:]
if matrix_scale is None:matrix_scale=data.abs().sum(dim=1).amax(dim=1)
matrix_scale=matrix_scale.clamp_min(1e-30);eigen_score=residual.abs().sum(dim=1).amax(dim=1)/matrix_scale;selected_count=selected_vectors.shape[-1];selected_eye=torch.eye(selected_count,device=data.device,dtype=torch.float32);orthogonality_score=(selected_vectors.mT@selected_vectors-selected_eye).abs().sum(dim=1).amax(dim=1);value_scale=values.abs().amax(dim=1,keepdim=True).clamp_min(1.);ordering_error=values[:,:-1]-values[:,1:];ordering_score=ordering_error.amax(dim=1)/value_scale[:,0];eps=torch.finfo(torch.float32).eps;safe=(eigen_score<=eigen_margin*2e2*float(n)*eps)&(orthogonality_score<=.002)&(ordering_score<=.975*1e2*float(n)*eps)&torch.isfinite(eigen_score)&torch.isfinite(orthogonality_score)&torch.isfinite(ordering_score)
if residual_probe_count:probes=_rademacher_probes(n,residual_probe_count,data.device);source_probe=data@probes;compressed_vectors=vectors@probes;compressed_scaled_vectors=vectors*values[:,None,:]@probes;compressed_residual=data@compressed_vectors-compressed_scaled_vectors;probe_score=compressed_residual.square().sum(dim=1).sqrt().amax(dim=1)/source_probe.square().sum(dim=1).sqrt().amax(dim=1).clamp_min(1e-30);safe=safe&(probe_score<=residual_probe_threshold)&torch.isfinite(probe_score)
finally:torch.set_float32_matmul_precision(previous_precision)
if bool(safe.all().item()):return vectors,values
repair=~safe;exact_values,exact_vectors=torch.linalg.eigh(data[repair].contiguous());vectors=vectors.clone();values=values.clone();vectors[repair]=exact_vectors;values[repair]=exact_values;return vectors,values
_N32_WARP_PAIR_SOURCE='\nextern "C" __global__ __launch_bounds__(512, 1)\nvoid n32_warp_pair_jacobi(\n const float* __restrict__ input,\n float* __restrict__ vectors,\n float* __restrict__ values,\n int batch) {\n constexpr int N = 32;\n constexpr int PITCH = 33;\n constexpr int NN = 1024;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int pair = tid >> 5;\n const int lane = tid & 31;\n if (matrix_id >= batch) return;\n\n __shared__ float a[N * PITCH];\n __shared__ float q[N * PITCH];\n __shared__ int permutation[N];\n __shared__ int converged;\n const float* source = input + (long long)matrix_id * NN;\n #pragma unroll\n for (int linear = tid; linear < NN; linear += 512) {\n const int row = linear >> 5;\n const int col = linear & 31;\n a[row * PITCH + col] = source[linear];\n q[row * PITCH + col] = row == col ? 1.0f : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll\n for (int sweep = 0; sweep < 8; ++sweep) {\n #pragma unroll\n for (int round_id = 0; round_id < 31; ++round_id) {\n const int p = pair == 0 ? round_id : (round_id + pair) % 31;\n const int r = pair == 0 ? 31 : (round_id - pair + 31) % 31;\n float c = 1.0f;\n float s = 0.0f;\n if (lane == 0) {\n const float app = a[p * PITCH + p];\n const float arr = a[r * PITCH + r];\n const float apr = 0.5f * (\n a[p * PITCH + r] + a[r * PITCH + p]);\n float tangent = 0.0f;\n if (fabsf(apr) > 1.0e-30f) {\n const float tau = (arr - app) / (2.0f * apr);\n const float sign = tau >= 0.0f ? 1.0f : -1.0f;\n tangent = sign / (fabsf(tau) + sqrtf(1.0f + tau * tau));\n }\n c = rsqrtf(1.0f + tangent * tangent);\n s = tangent * c;\n }\n c = __shfl_sync(0xffffffff, c, 0);\n s = __shfl_sync(0xffffffff, s, 0);\n\n const float ap = a[lane * PITCH + p];\n const float ar = a[lane * PITCH + r];\n const float qp = q[lane * PITCH + p];\n const float qr = q[lane * PITCH + r];\n a[lane * PITCH + p] = c * ap - s * ar;\n a[lane * PITCH + r] = s * ap + c * ar;\n q[lane * PITCH + p] = c * qp - s * qr;\n q[lane * PITCH + r] = s * qp + c * qr;\n __syncthreads();\n\n const float rp = a[p * PITCH + lane];\n const float rr = a[r * PITCH + lane];\n a[p * PITCH + lane] = c * rp - s * rr;\n a[r * PITCH + lane] = s * rp + c * rr;\n __syncthreads();\n\n }\n\n if (sweep >= 4) {\n if (pair == 0) {\n float block_off = 0.0f;\n float block_scale = 0.0f;\n #pragma unroll\n for (int linear = lane; linear < NN; linear += 32) {\n const int row = linear >> 5;\n const int col = linear & 31;\n const float magnitude = fabsf(a[row * PITCH + col]);\n block_scale = fmaxf(block_scale, magnitude);\n if (row != col) block_off = fmaxf(block_off, magnitude);\n }\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n block_off = fmaxf(block_off,\n __shfl_down_sync(0xffffffff, block_off, delta));\n block_scale = fmaxf(block_scale,\n __shfl_down_sync(0xffffffff, block_scale, delta));\n }\n // The checker is dimension-scaled; 1e-4 remains conservative for\n // n=32 and lets ordinary dense matrices stop after six sweeps.\n if (lane == 0) converged = block_off <= 1.0e-4f * block_scale;\n }\n __syncthreads();\n if (converged) break;\n }\n }\n\n if (tid < N) {\n float key = a[tid * PITCH + tid];\n int key_index = tid;\n #pragma unroll\n for (int width = 2; width <= N; width <<= 1) {\n #pragma unroll\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n const float other = __shfl_xor_sync(\n 0xffffffff, key, stride);\n const int other_index = __shfl_xor_sync(\n 0xffffffff, key_index, stride);\n const bool ascending = (tid & width) == 0;\n const bool lower_lane = (tid & stride) == 0;\n const bool want_minimum = ascending == lower_lane;\n const bool take_other = want_minimum\n ? other < key : other > key;\n if (take_other) {\n key = other;\n key_index = other_index;\n }\n }\n }\n permutation[tid] = key_index;\n values[matrix_id * N + tid] = key;\n }\n __syncthreads();\n float* output = vectors + (long long)matrix_id * NN;\n #pragma unroll\n for (int linear = tid; linear < NN; linear += 512) {\n const int row = linear >> 5;\n const int col = linear & 31;\n output[linear] = q[row * PITCH + permutation[col]];\n }\n}\n'
def _e2474_n32_ranked_tail_source(source:str):
source=source.replace(' __shared__ int permutation[N];\n __shared__ int converged;',' __shared__ int permutation[N];\n __shared__ int converged;\n __shared__ int tail_order[2];\n __shared__ float round_score[32];\n __shared__ float round_partial[16 * 32];',1);old=' #pragma unroll\n for (int sweep = 0; sweep < 8; ++sweep) {\n #pragma unroll\n for (int round_id = 0; round_id < 31; ++round_id) {';new=' #pragma unroll\n for (int sweep = 0; sweep < 8; ++sweep) {\n if (sweep == 5) {\n float score = -1.0f;\n if (lane < 31) {\n const int score_p = pair == 0 ? lane : (lane + pair) % 31;\n const int score_r = pair == 0 ? 31 : (lane - pair + 31) % 31;\n const float value = 0.5f * (\n a[score_p * PITCH + score_r]\n + a[score_r * PITCH + score_p]);\n score = fabsf(value);\n }\n round_partial[pair * 32 + lane] = score;\n __syncthreads();\n if (pair == 0) {\n float maximum = -1.0f;\n #pragma unroll\n for (int local_pair = 0; local_pair < 16; ++local_pair)\n maximum = fmaxf(\n maximum, round_partial[local_pair * 32 + lane]);\n round_score[lane] = maximum;\n __syncwarp();\n #pragma unroll\n for (int selected = 0; selected < 2; ++selected) {\n float best = round_score[lane];\n int best_round = lane;\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n const float other = __shfl_down_sync(0xffffffffu, best, delta);\n const int other_round = __shfl_down_sync(\n 0xffffffffu, best_round, delta);\n if (other > best || (other == best && other_round < best_round)) {\n best = other;\n best_round = other_round;\n }\n }\n if (lane == 0) {\n tail_order[selected] = best_round;\n round_score[best_round] = -1.0f;\n }\n __syncwarp();\n }\n }\n __syncthreads();\n }\n #pragma unroll\n for (int round_id = 0; round_id < (sweep == 5 ? 2 : 31); ++round_id) {\n const int actual_round = sweep == 5 ? tail_order[round_id] : round_id;'
if source.count(old)!=1:raise RuntimeError('n32 ranked-tail loop anchor changed')
source=source.replace(old,new,1);source=source.replace('const int p = pair == 0 ? round_id : (round_id + pair) % 31;\n const int r = pair == 0 ? 31 : (round_id - pair + 31) % 31;','const int p = pair == 0 ? actual_round : (actual_round + pair) % 31;\n const int r = pair == 0 ? 31 : (actual_round - pair + 31) % 31;',1);source=source.replace('block_off <= 1.0e-4f * block_scale','block_off <= 2.5e-4f * block_scale',1);return source
def _e2718_n32_wmma_guard_source(source:str):
source=_e2474_n32_ranked_tail_source(source);source=source.replace('block_off <= 2.5e-4f * block_scale','block_off <= 1.0e-4f * block_scale',1);source='#include <cuda_fp16.h>\n#include <mma.h>\nusing namespace nvcuda;\n'+source;old_shared=' __shared__ int tail_order[2];\n __shared__ float round_score[32];';new_shared=old_shared+'\n __shared__ __align__(32) __half certificate_q[NN];\n __shared__ __align__(32) __half certificate_e[NN];\n __shared__ __align__(32) float certificate_product[4 * 16 * 16];'
if source.count(old_shared)!=1:raise RuntimeError('n32 WMMA certificate shared anchor changed')
source=source.replace(old_shared,new_shared,1);begin=source.index(' if (sweep >= 4) {');end=source.index('\n }\n\n if (tid < N)',begin);old_guard=source[begin:end]
if'block_off <= 1.0e-4f * block_scale'not in old_guard:raise RuntimeError('n32 WMMA convergence anchor changed')
new_guard=' if (sweep >= 4) {\n if (sweep == 5) {\n #pragma unroll\n for (int linear = tid; linear < NN; linear += 512) {\n const int row = linear >> 5;\n const int column = linear & 31;\n certificate_q[linear] = __float2half_rn(q[row * PITCH + column]);\n certificate_e[linear] = __float2half_rn(\n row == column ? 0.0f : a[row * PITCH + column]);\n }\n __syncthreads();\n if (pair < 4) {\n const int tile_row = pair >> 1;\n const int tile_column = pair & 1;\n wmma::fragment<wmma::accumulator, 16, 16, 16, float> accumulator;\n wmma::fill_fragment(accumulator, 0.0f);\n #pragma unroll\n for (int inner = 0; inner < N; inner += 16) {\n wmma::fragment<wmma::matrix_a, 16, 16, 16,\n __half, wmma::row_major> left;\n wmma::fragment<wmma::matrix_b, 16, 16, 16,\n __half, wmma::row_major> right;\n wmma::load_matrix_sync(\n left, certificate_q + tile_row * 16 * N + inner, N);\n wmma::load_matrix_sync(\n right, certificate_e + inner * N + tile_column * 16, N);\n wmma::mma_sync(accumulator, left, right, accumulator);\n }\n wmma::store_matrix_sync(\n certificate_product + pair * 16 * 16,\n accumulator, 16, wmma::mem_row_major);\n }\n __syncthreads();\n if (pair == 0) {\n float residual = 0.0f;\n float scale = 0.0f;\n const int column = lane;\n const int tile_column = column >> 4;\n const int local_column = column & 15;\n #pragma unroll\n for (int row = 0; row < N; ++row) {\n const int tile_row = row >> 4;\n const int local_row = row & 15;\n residual += fabsf(certificate_product[\n (tile_row * 2 + tile_column) * 256\n + local_row * 16 + local_column]);\n scale += fabsf(source[row * N + column]);\n }\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n residual = fmaxf(residual,\n __shfl_down_sync(0xffffffffu, residual, delta));\n scale = fmaxf(scale,\n __shfl_down_sync(0xffffffffu, scale, delta));\n }\n if (lane == 0) {\n // The checker permits 200*N*eps. Use only 70% and disable the\n // FP16 path when tiny residual entries could underflow.\n const float allowed = 0.70f * 200.0f * N\n * 1.1920928955078125e-7f * scale;\n converged = isfinite(residual) && isfinite(scale)\n && scale >= 1.0e-2f && residual <= allowed;\n }\n }\n } else {\n if (pair == 0) {\n float block_off = 0.0f;\n float block_scale = 0.0f;\n #pragma unroll\n for (int linear = lane; linear < NN; linear += 32) {\n const int row = linear >> 5;\n const int col = linear & 31;\n const float magnitude = fabsf(a[row * PITCH + col]);\n block_scale = fmaxf(block_scale, magnitude);\n if (row != col) block_off = fmaxf(block_off, magnitude);\n }\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n block_off = fmaxf(block_off,\n __shfl_down_sync(0xffffffff, block_off, delta));\n block_scale = fmaxf(block_scale,\n __shfl_down_sync(0xffffffff, block_scale, delta));\n }\n if (lane == 0)\n converged = block_off <= 1.0e-4f * block_scale;\n }\n }\n __syncthreads();\n if (converged) break;\n }';return source[:begin]+new_guard+source[end:]
@memo(maxsize=1)
def _small_jacobi_kernel():source=_e2718_n32_wmma_guard_source(_N32_WARP_PAIR_SOURCE);image=_fast_nvrtc_compile(source,'n32_warp_pair_jacobi');return CUDAKernel(image,'n32_warp_pair_jacobi')
@torch.no_grad()
def _small_jacobi_eigh(data:torch.Tensor):vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device,dtype=torch.float32);_small_jacobi_kernel().launch((data.shape[0],1,1),(512,1,1),(data,vectors,values,data.shape[0]));return vectors,values
_N24_BOUNDARY_JACOBI_SOURCE='\nextern "C" __global__ __launch_bounds__(384, 1)\nvoid n24_boundary_jacobi(\n const float* __restrict__ input,\n float* __restrict__ vectors,\n float* __restrict__ values,\n int batch) {\n constexpr int N = 24;\n constexpr int PITCH = 25;\n constexpr int NN = 576;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int pair = tid >> 5;\n const int lane = tid & 31;\n if (matrix_id >= batch) return;\n\n __shared__ float a[N * PITCH];\n __shared__ float q[N * PITCH];\n __shared__ float sorted_values[N];\n __shared__ int permutation[N];\n __shared__ int converged;\n const float* source = input + (long long)matrix_id * NN;\n for (int linear = tid; linear < NN; linear += 384) {\n const int row = linear / N;\n const int column = linear - row * N;\n a[row * PITCH + column] = source[linear];\n q[row * PITCH + column] = row == column ? 1.0f : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll\n for (int sweep = 0; sweep < 8; ++sweep) {\n #pragma unroll\n for (int round_id = 0; round_id < 23; ++round_id) {\n const int p = pair == 0 ? round_id : (round_id + pair) % 23;\n const int r = pair == 0 ? 23 : (round_id - pair + 23) % 23;\n float cosine = 1.0f;\n float sine = 0.0f;\n if (lane == 0) {\n const float app = a[p * PITCH + p];\n const float arr = a[r * PITCH + r];\n const float apr = 0.5f * (\n a[p * PITCH + r] + a[r * PITCH + p]);\n float tangent = 0.0f;\n if (fabsf(apr) > 1.0e-30f) {\n const float tau = (arr - app) / (2.0f * apr);\n const float sign = tau >= 0.0f ? 1.0f : -1.0f;\n tangent = sign / (fabsf(tau) + sqrtf(1.0f + tau * tau));\n }\n cosine = rsqrtf(1.0f + tangent * tangent);\n sine = tangent * cosine;\n }\n cosine = __shfl_sync(0xffffffff, cosine, 0);\n sine = __shfl_sync(0xffffffff, sine, 0);\n\n if (lane < N) {\n const float ap = a[lane * PITCH + p];\n const float ar = a[lane * PITCH + r];\n const float qp = q[lane * PITCH + p];\n const float qr = q[lane * PITCH + r];\n a[lane * PITCH + p] = cosine * ap - sine * ar;\n a[lane * PITCH + r] = sine * ap + cosine * ar;\n q[lane * PITCH + p] = cosine * qp - sine * qr;\n q[lane * PITCH + r] = sine * qp + cosine * qr;\n }\n __syncthreads();\n\n if (lane < N) {\n const float rp = a[p * PITCH + lane];\n const float rr = a[r * PITCH + lane];\n a[p * PITCH + lane] = cosine * rp - sine * rr;\n a[r * PITCH + lane] = sine * rp + cosine * rr;\n }\n __syncthreads();\n }\n\n if (sweep >= 4) {\n if (pair == 0) {\n float block_off = 0.0f;\n float block_scale = 0.0f;\n for (int linear = lane; linear < NN; linear += 32) {\n const int row = linear / N;\n const int column = linear - row * N;\n const float magnitude = fabsf(a[row * PITCH + column]);\n block_scale = fmaxf(block_scale, magnitude);\n if (row != column)\n block_off = fmaxf(block_off, magnitude);\n }\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n block_off = fmaxf(block_off,\n __shfl_down_sync(0xffffffff, block_off, delta));\n block_scale = fmaxf(block_scale,\n __shfl_down_sync(0xffffffff, block_scale, delta));\n }\n if (lane == 0)\n converged = block_off <= 1.0e-5f * block_scale;\n }\n __syncthreads();\n if (converged) break;\n }\n }\n\n if (tid < N) {\n sorted_values[tid] = a[tid * PITCH + tid];\n permutation[tid] = tid;\n }\n __syncthreads();\n if (tid == 0) {\n #pragma unroll\n for (int column = 1; column < N; ++column) {\n const float key = sorted_values[column];\n const int key_index = permutation[column];\n int position = column - 1;\n while (position >= 0 && sorted_values[position] > key) {\n sorted_values[position + 1] = sorted_values[position];\n permutation[position + 1] = permutation[position];\n --position;\n }\n sorted_values[position + 1] = key;\n permutation[position + 1] = key_index;\n }\n }\n __syncthreads();\n if (tid < N)\n values[(long long)matrix_id * N + tid] = sorted_values[tid];\n float* output = vectors + (long long)matrix_id * NN;\n for (int linear = tid; linear < NN; linear += 384) {\n const int row = linear / N;\n const int column = linear - row * N;\n output[linear] = q[row * PITCH + permutation[column]];\n }\n}\n'
@memo(maxsize=1)
def _small_n24_jacobi_kernel():return CUDAKernel(_fast_nvrtc_compile(_N24_BOUNDARY_JACOBI_SOURCE,'n24_boundary_jacobi'),'n24_boundary_jacobi')
@torch.no_grad()
def _small_n24_eigh(data:torch.Tensor):vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device,dtype=torch.float32);_small_n24_jacobi_kernel().launch((data.shape[0],1,1),(384,1,1),(data,vectors,values,data.shape[0]));return vectors,values
_E951_N16_JACOBI_NAME='e951_n16_boundary_jacobi';_E951_N32_JACOBI_NAME='e951_n32_boundary_jacobi_tight'
@memo(maxsize=1)
def _e951_n16_jacobi_kernel():source=_N24_BOUNDARY_JACOBI_SOURCE.replace('__launch_bounds__(384, 1)','__launch_bounds__(256, 1)').replace('n24_boundary_jacobi',_E951_N16_JACOBI_NAME).replace('constexpr int N = 24;','constexpr int N = 16;').replace('constexpr int PITCH = 25;','constexpr int PITCH = 17;').replace('constexpr int NN = 576;','constexpr int NN = 256;').replace('linear += 384','linear += 256').replace('round_id < 23','round_id < 15').replace('? 23 :','? 15 :').replace('% 23','% 15').replace('+ 23','+ 15').replace('1.0e-5f * block_scale','1.0e-6f * block_scale');return CUDAKernel(_fast_nvrtc_compile(source,_E951_N16_JACOBI_NAME),_E951_N16_JACOBI_NAME)
@memo(maxsize=1)
def _e951_n32_jacobi_kernel():source=_N32_WARP_PAIR_SOURCE.replace('n32_warp_pair_jacobi',_E951_N32_JACOBI_NAME).replace('1.0e-4f * block_scale','1.0e-5f * block_scale');return CUDAKernel(_fast_nvrtc_compile(source,_E951_N32_JACOBI_NAME),_E951_N32_JACOBI_NAME)
@torch.no_grad()
def _e951_n16_eigh(data:torch.Tensor):vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device);_e951_n16_jacobi_kernel().launch((data.shape[0],1,1),(256,1,1),(data,vectors,values,data.shape[0]));return vectors,values
@torch.no_grad()
def _e951_n32_eigh(data:torch.Tensor):vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device);_e951_n32_jacobi_kernel().launch((data.shape[0],1,1),(512,1,1),(data,vectors,values,data.shape[0]));return vectors,values
_N160_FULL_CHOLESKY_SOURCE='\n#include <cuda_runtime.h>\n\nextern "C" __global__\nvoid potrf160_ridge(\n const float* __restrict__ gram,\n float* __restrict__ lower,\n int batch,\n float ridge)\n{\n constexpr int N = 160;\n constexpr int NN = N * N;\n const int matrix_id = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n if (matrix_id >= batch) return;\n\n extern __shared__ float storage[];\n float* factor = storage;\n float* diagonal_mean = storage + NN;\n const float* source = gram + (long long)matrix_id * NN;\n\n for (int index = tid; index < NN; index += (int)blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n factor[index] = row >= column ? source[index] : 0.0f;\n }\n if (tid == 0) {\n float total = 0.0f;\n #pragma unroll\n for (int axis = 0; axis < N; ++axis)\n total += source[axis * N + axis];\n *diagonal_mean = total * (1.0f / (float)N);\n }\n __syncthreads();\n\n const float shift = ridge * (*diagonal_mean);\n #pragma unroll 1\n for (int column = 0; column < N; ++column) {\n if (tid == 0) {\n float diagonal = factor[column * N + column] + shift;\n #pragma unroll 1\n for (int k = 0; k < column; ++k) {\n const float value = factor[column * N + k];\n diagonal = fmaf(-value, value, diagonal);\n }\n factor[column * N + column] = sqrtf(fmaxf(diagonal, 1.0e-30f));\n }\n __syncthreads();\n const float inverse_diagonal = 1.0f / factor[column * N + column];\n for (int row = column + 1 + tid; row < N; row += (int)blockDim.x) {\n float value = factor[row * N + column];\n #pragma unroll 1\n for (int k = 0; k < column; ++k)\n value = fmaf(\n -factor[row * N + k], factor[column * N + k], value);\n factor[row * N + column] = value * inverse_diagonal;\n }\n __syncthreads();\n }\n\n float* destination = lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += (int)blockDim.x)\n destination[index] = factor[index];\n}\n\n// Right-looking blocked Cholesky. A 16-column panel is factored by lane 0,\n// all rows below it solve independently, then the CTA applies one rank-16\n// update to the remaining lower triangle. This reduces the synchronization\n// count from two per scalar column to three per panel.\nextern "C" __global__\nvoid potrf160_block16_ridge(\n const float* __restrict__ gram,\n float* __restrict__ lower,\n int batch,\n float ridge)\n{\n constexpr int N = 160;\n constexpr int NN = N * N;\n constexpr int PANEL = 16;\n const int matrix_id = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n if (matrix_id >= batch) return;\n\n extern __shared__ float storage[];\n float* factor = storage;\n float* diagonal_mean = storage + NN;\n const float* source = gram + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += (int)blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n factor[index] = row >= column ? source[index] : 0.0f;\n }\n if (tid == 0) {\n float total = 0.0f;\n #pragma unroll\n for (int axis = 0; axis < N; ++axis)\n total += source[axis * N + axis];\n *diagonal_mean = total * (1.0f / (float)N);\n }\n __syncthreads();\n const float shift = ridge * (*diagonal_mean);\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = min(panel + PANEL, N);\n if (tid == 0) {\n #pragma unroll\n for (int column_offset = 0; column_offset < PANEL; ++column_offset) {\n const int column = panel + column_offset;\n if (column >= N) break;\n float diagonal = factor[column * N + column] + shift;\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n const float value = factor[column * N + k];\n diagonal = fmaf(-value, value, diagonal);\n }\n factor[column * N + column] =\n sqrtf(fmaxf(diagonal, 1.0e-30f));\n for (int row = column + 1; row < end; ++row) {\n float value = factor[row * N + column];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[row * N + column] =\n value / factor[column * N + column];\n }\n }\n }\n __syncthreads();\n\n // L21 <- A21 inv(L11^T); one thread owns one complete row.\n for (int row = end + tid; row < N; row += (int)blockDim.x) {\n #pragma unroll\n for (int column_offset = 0; column_offset < PANEL; ++column_offset) {\n const int column = panel + column_offset;\n if (column >= N) break;\n float value = factor[row * N + column];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[row * N + column] =\n value / factor[column * N + column];\n }\n }\n __syncthreads();\n\n // A22 <- A22 - L21 L21^T.\n for (int index = tid; index < NN; index += (int)blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n if (row >= end && column >= end && row >= column) {\n float value = factor[index];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= N) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[index] = value;\n }\n }\n __syncthreads();\n }\n\n float* destination = lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += (int)blockDim.x)\n destination[index] = factor[index];\n}\n\nextern "C" __global__\nvoid right_trsm160_rows64(\n const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch,\n int rows)\n{\n constexpr int N = 160;\n constexpr int NN = N * N;\n constexpr int ROWS = 64;\n const int matrix_id = (int)blockIdx.x;\n const int local_row = (int)threadIdx.x;\n const int row = (int)blockIdx.y * ROWS + local_row;\n if (matrix_id >= batch || local_row >= ROWS || row >= rows) return;\n\n extern __shared__ float solved[]; // transposed [N][ROWS]\n const float* source = matrix +\n ((long long)matrix_id * rows + row) * N;\n const float* factor = lower + (long long)matrix_id * NN;\n float* destination = output +\n ((long long)matrix_id * rows + row) * N;\n\n #pragma unroll 1\n for (int column = 0; column < N; ++column) {\n float value = source[column];\n #pragma unroll 1\n for (int k = 0; k < column; ++k)\n value = fmaf(-factor[column * N + k], solved[k * ROWS + local_row], value);\n value /= factor[column * N + column];\n solved[column * ROWS + local_row] = value;\n destination[column] = value;\n }\n}\n\n// Block forward solve for 32 rows. The RHS tile is transposed in shared\n// memory so a warp owns one future column: L loads broadcast while the 32 RHS\n// row values are bank-conflict-free. Rank-16 trailing updates use the whole\n// CTA instead of leaving each row thread with a long serial dot product.\nextern "C" __global__\nvoid right_trsm160_block16_rows32(\n const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch,\n int rows)\n{\n constexpr int N = 160;\n constexpr int NN = N * N;\n constexpr int ROWS = 32;\n constexpr int PITCH = 33;\n constexpr int PANEL = 16;\n const int matrix_id = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n const int row_base = (int)blockIdx.y * ROWS;\n if (matrix_id >= batch) return;\n\n extern __shared__ float rhs[]; // transposed [N][ROWS]\n const float* source = matrix + (long long)matrix_id * rows * N;\n const float* factor = lower + (long long)matrix_id * NN;\n float* destination = output + (long long)matrix_id * rows * N;\n for (int index = tid; index < N * ROWS; index += (int)blockDim.x) {\n const int local_row = index / N;\n const int column = index - local_row * N;\n const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows ? source[(long long)row * N + column] : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = min(panel + PANEL, N);\n if (tid < ROWS) {\n const int local_row = tid;\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n if (column >= N) break;\n float value = rhs[column * PITCH + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[column * N + k],\n rhs[k * PITCH + local_row], value);\n }\n rhs[column * PITCH + local_row] =\n value / factor[column * N + column];\n }\n }\n __syncthreads();\n const int remaining = (N - end) * ROWS;\n for (int index = tid; index < remaining; index += (int)blockDim.x) {\n const int column = end + index / ROWS;\n const int local_row = index % ROWS;\n float value = rhs[column * PITCH + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= N) break;\n value = fmaf(-factor[column * N + k],\n rhs[k * PITCH + local_row], value);\n }\n rhs[column * PITCH + local_row] = value;\n }\n __syncthreads();\n }\n\n for (int index = tid; index < N * ROWS; index += (int)blockDim.x) {\n const int local_row = index / N;\n const int column = index - local_row * N;\n const int row = row_base + local_row;\n if (row < rows)\n destination[(long long)row * N + column] =\n rhs[column * PITCH + local_row];\n }\n}\n';_N160_PACKED_FULL_CHOLESKY_SOURCE='\n#include <cuda_runtime.h>\n\nconstexpr int N = 160;\nconstexpr int NN = N * N;\nconstexpr int TRI = N * (N + 1) / 2;\nconstexpr int PANEL = 16;\n\n__device__ __forceinline__ int pidx(int row, int column) {\n return row * (row + 1) / 2 + column;\n}\n\nextern "C" __global__ __launch_bounds__(256, 1)\nvoid potrf160_packed_shared_full_output(\n const float* __restrict__ gram,\n float* __restrict__ packed_lower,\n int batch,\n float ridge) {\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n if (matrix_id >= batch) return;\n extern __shared__ float storage[];\n float* factor = storage;\n float* diagonal_mean = factor + TRI;\n const float* source = gram + (long long)matrix_id * NN;\n\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n if (row >= column) factor[pidx(row, column)] = source[index];\n }\n if (tid == 0) {\n float total = 0.0f;\n #pragma unroll\n for (int axis = 0; axis < N; ++axis)\n total += source[axis * N + axis];\n *diagonal_mean = total * (1.0f / N);\n }\n __syncthreads();\n const float shift = ridge * (*diagonal_mean);\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = panel + PANEL;\n if (tid == 0) {\n #pragma unroll\n for (int column_offset = 0; column_offset < PANEL; ++column_offset) {\n const int column = panel + column_offset;\n float diagonal = factor[pidx(column, column)] + shift;\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n const float value = factor[pidx(column, k)];\n diagonal = fmaf(-value, value, diagonal);\n }\n factor[pidx(column, column)] =\n sqrtf(fmaxf(diagonal, 1.0e-30f));\n for (int row = column + 1; row < end; ++row) {\n float value = factor[pidx(row, column)];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[pidx(row, k)],\n factor[pidx(column, k)], value);\n }\n factor[pidx(row, column)] =\n value / factor[pidx(column, column)];\n }\n }\n }\n __syncthreads();\n\n for (int row = end + tid; row < N; row += blockDim.x) {\n #pragma unroll\n for (int column_offset = 0; column_offset < PANEL; ++column_offset) {\n const int column = panel + column_offset;\n float value = factor[pidx(row, column)];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[pidx(row, k)],\n factor[pidx(column, k)], value);\n }\n factor[pidx(row, column)] =\n value / factor[pidx(column, column)];\n }\n }\n __syncthreads();\n\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n if (row >= end && column >= end && row >= column) {\n float value = factor[pidx(row, column)];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n value = fmaf(-factor[pidx(row, k)],\n factor[pidx(column, k)], value);\n }\n factor[pidx(row, column)] = value;\n }\n }\n __syncthreads();\n }\n\n float* destination = packed_lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n destination[index] = row >= column ? factor[pidx(row, column)] : 0.0f;\n }\n}\n\nextern "C" __global__ __launch_bounds__(256, 1)\nvoid right_trsm160_packed_block16_rows32(\n const float* __restrict__ matrix,\n const float* __restrict__ packed_lower,\n float* __restrict__ output,\n int batch,\n int rows) {\n constexpr int ROWS = 32;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int row_base = blockIdx.y * ROWS;\n if (matrix_id >= batch) return;\n extern __shared__ float rhs[];\n const float* source = matrix + (long long)matrix_id * rows * N;\n const float* factor = packed_lower + (long long)matrix_id * TRI;\n float* destination = output + (long long)matrix_id * rows * N;\n\n for (int index = tid; index < N * ROWS; index += blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n rhs[index] = row < rows ? source[(long long)row * N + column] : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = panel + PANEL;\n if (tid < ROWS) {\n const int local_row = tid;\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[pidx(column, k)],\n rhs[k * ROWS + local_row], value);\n }\n rhs[column * ROWS + local_row] =\n value / factor[pidx(column, column)];\n }\n }\n __syncthreads();\n const int remaining = (N - end) * ROWS;\n for (int index = tid; index < remaining; index += blockDim.x) {\n const int column = end + index / ROWS;\n const int local_row = index - (column - end) * ROWS;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n value = fmaf(-factor[pidx(column, k)],\n rhs[k * ROWS + local_row], value);\n }\n rhs[column * ROWS + local_row] = value;\n }\n __syncthreads();\n }\n\n for (int index = tid; index < N * ROWS; index += blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n if (row < rows)\n destination[(long long)row * N + column] = rhs[index];\n }\n}\n';_N160_EIGH_SOURCE='\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;\n\nconstexpr int N = 160;\nconstexpr int HALF = 80;\nconstexpr int SORT_N = 256;\nconstexpr int TRI = N * (N + 1) / 2;\n\nextern "C" __global__\n__cluster_dims__(2, 1, 1)\n__launch_bounds__(512, 4)\nvoid tridiagonal160_column_t512_g4(\n const float* __restrict__ input,\n float* __restrict__ output_q,\n float* __restrict__ output_l,\n float* __restrict__ saved_reflectors,\n unsigned long long* __restrict__ timers,\n float ql_epsilon) {\n cg::cluster_group cluster = cg::this_cluster();\n const int rank = cluster.block_rank();\n const int matrix_id = blockIdx.x >> 1;\n const int tid = threadIdx.x;\n const int first_row = rank * HALF;\n float* saved_matrix = saved_reflectors + (long long)matrix_id * N * N;\n\n extern __shared__ __align__(16) unsigned char storage[];\n float* a = reinterpret_cast<float*>(storage);\n // Packed A and the row half of Q have disjoint lifetimes.\n float* q = a;\n float* v = a + TRI;\n float* w = v + N;\n float* partial = w + N;\n float* scratch = partial + N;\n float* diagonal = scratch + 16;\n float* off_diagonal = diagonal + N;\n\n float* a0 = cluster.map_shared_rank(a, 0);\n float* v0 = cluster.map_shared_rank(v, 0);\n float* w0 = cluster.map_shared_rank(w, 0);\n float* partial0 = cluster.map_shared_rank(partial, 0);\n float* diagonal0 = cluster.map_shared_rank(diagonal, 0);\n float* off_diagonal0 = cluster.map_shared_rank(off_diagonal, 0);\n\n if (rank == 0 && tid == 0) timers[matrix_id * 5] = clock64();\n\n if (rank == 0) {\n const float* source = input + (long long)matrix_id * N * N;\n for (int row = 0; row < N; ++row)\n if (tid <= row)\n a[row * (row + 1) / 2 + tid] = source[row * N + tid];\n }\n cluster.sync();\n\n // CTA 0 performs the right-looking symmetric Householder reduction.\n // Warp-shuffle block reductions use two CTA barriers instead of the nine\n // barriers in a shared-memory binary tree. The normalized reflector norm\n // is known analytically from ||x|| and x_0, avoiding a second reduction.\n if (rank == 0) {\n for (int column = 0; column < N - 2; ++column) {\n const int begin = column + 1;\n const int length = N - begin;\n float local = 0.0f;\n for (int i = tid; i < length; i += blockDim.x) {\n const float x = a[(begin + i) * (begin + i + 1) / 2 + column];\n local += x * x;\n }\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n local += __shfl_down_sync(0xffffffff, local, offset);\n if ((tid & 31) == 0) scratch[tid >> 5] = local;\n __syncthreads();\n float total = tid < 8 ? scratch[tid] : 0.0f;\n if (tid < 32) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n total += __shfl_down_sync(0xffffffff, total, offset);\n if (tid == 0) scratch[0] = total;\n }\n __syncthreads();\n\n const float norm = sqrtf(scratch[0]);\n const float first = a[begin * (begin + 1) / 2 + column];\n const float tridiagonal_value = -copysignf(norm, first);\n if (tid == 0) off_diagonal[column] = tridiagonal_value;\n const float reflector_norm = sqrtf(\n fmaxf(2.0f * norm * (norm + fabsf(first)), 0.0f));\n const float inverse = reflector_norm > 0.0f\n ? 1.0f / reflector_norm : 0.0f;\n for (int i = tid; i < length; i += blockDim.x) {\n float value = a[(begin + i) * (begin + i + 1) / 2 + column];\n if (i == 0) value -= tridiagonal_value;\n v[i] = value * inverse;\n }\n __syncthreads();\n\n constexpr int MATVEC_GROUP = 4;\n constexpr int MATVEC_ROWS = 128;\n const int matvec_lane = tid & (MATVEC_GROUP - 1);\n const int matvec_local_row = tid / MATVEC_GROUP;\n for (int row_base = 0; row_base < length; row_base += MATVEC_ROWS) {\n const int local_row = row_base + matvec_local_row;\n float product = 0.0f;\n if (local_row < length) {\n const int row = begin + local_row;\n for (int j = matvec_lane; j < length; j += MATVEC_GROUP) {\n const float value = j <= local_row\n ? a[row * (row + 1) / 2 + begin + j]\n : a[(begin + j) * (begin + j + 1) / 2 + row];\n product += value * v[j];\n }\n }\n #pragma unroll\n for (int offset = MATVEC_GROUP / 2; offset > 0; offset >>= 1)\n product += __shfl_down_sync(0xffffffff, product, offset);\n if (matvec_lane == 0 && local_row < length)\n w[local_row] = product;\n }\n __syncthreads();\n\n local = 0.0f;\n for (int i = tid; i < length; i += blockDim.x)\n local += v[i] * w[i];\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n local += __shfl_down_sync(0xffffffff, local, offset);\n if ((tid & 31) == 0) scratch[tid >> 5] = local;\n __syncthreads();\n total = tid < 8 ? scratch[tid] : 0.0f;\n if (tid < 32) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n total += __shfl_down_sync(0xffffffff, total, offset);\n if (tid == 0) scratch[0] = total;\n }\n __syncthreads();\n const float projection = scratch[0];\n for (int i = tid; i < length; i += blockDim.x)\n w[i] = 2.0f * (w[i] - projection * v[i]);\n __syncthreads();\n\n // Update only the lower triangle. At each row, participating threads\n // write a contiguous segment and avoid dynamic integer division.\n constexpr int UPDATE_ROWS = 32;\n constexpr int UPDATE_COLUMNS = 16;\n const int tile_row = tid / UPDATE_COLUMNS;\n const int tile_column = tid - tile_row * UPDATE_COLUMNS;\n for (int row_base = 0; row_base < length; row_base += UPDATE_ROWS) {\n for (int column_base = 0; column_base < row_base + UPDATE_ROWS;\n column_base += UPDATE_COLUMNS) {\n const int i = row_base + tile_row;\n const int j = column_base + tile_column;\n if (i < length && j < length && j <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n v[i] * w[j] + w[i] * v[j];\n }\n }\n __syncthreads();\n for (int i = tid; i < length; i += blockDim.x)\n a[(begin + i) * (begin + i + 1) / 2 + column] = v[i];\n __syncthreads();\n }\n if (tid == 0) off_diagonal[N - 2] = a[(N - 1) * N / 2 + N - 2];\n if (tid < N) {\n diagonal[tid] = a[tid * (tid + 1) / 2 + tid];\n if (tid == N - 1) off_diagonal[tid] = 0.0;\n }\n }\n // Spill saved unit reflectors before q aliases packed A.\n if (rank == 0) {\n for (int column = 0; column < N - 2; ++column)\n for (int row = column + 1 + tid; row < N; row += blockDim.x)\n saved_matrix[column * N + row] =\n a[row * (row + 1) / 2 + column];\n // cluster.sync() orders DSM, but the reflector spill is device memory.\n // Every producing thread must publish its own global stores before rank 1\n // later reloads the compact-WY panels.\n __threadfence();\n }\n cluster.sync();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 1] = clock64();\n // Replicate the tiny tridiagonal representation. Without this copy CTA 1\n // performs every Sturm/LDL recurrence through remote DSM loads.\n if (rank == 1 && tid < N) {\n diagonal[tid] = diagonal0[tid];\n off_diagonal[tid] = off_diagonal0[tid];\n }\n cluster.sync();\n\n // Parallel Sturm bisection followed by one shifted inverse iteration per\n // eigenpair. CTA ranks split the 176 eigenpairs 88/88. ``q`` is reused as\n // per-eigenpair LDL diagonal workspace, while output_q temporarily stores\n // the modified right-hand sides and resulting tridiagonal eigenvectors.\n if (rank == 0 && tid == 0) {\n float lower = diagonal[0] - fabsf(off_diagonal[0]);\n float upper = diagonal[0] + fabsf(off_diagonal[0]);\n for (int i = 1; i < N; ++i) {\n const float radius = fabsf(off_diagonal[i - 1]) +\n (i + 1 < N ? fabsf(off_diagonal[i]) : 0.0f);\n lower = fminf(lower, diagonal[i] - radius);\n upper = fmaxf(upper, diagonal[i] + radius);\n }\n partial[0] = lower;\n partial[1] = upper;\n }\n cluster.sync();\n if (tid < HALF) {\n const int eigen_index = first_row + tid;\n float lower = partial0[0];\n float upper = partial0[1];\n #pragma unroll\n for (int step = 0; step < 24; ++step) {\n const float midpoint = 0.5f * (lower + upper);\n float pivot = diagonal[0] - midpoint;\n int count = pivot < 0.0f;\n for (int i = 1; i < N; ++i) {\n if (fabsf(pivot) < 1.0e-12f)\n pivot = copysignf(1.0e-12f, pivot == 0.0f ? -1.0f : pivot);\n pivot = diagonal[i] - midpoint -\n off_diagonal[i - 1] * off_diagonal[i - 1] / pivot;\n count += pivot < 0.0f;\n }\n if (count <= eigen_index) lower = midpoint;\n else upper = midpoint;\n }\n const float eigenvalue = 0.5f * (lower + upper);\n float* left_diagonal = q + tid * N;\n float* right_diagonal = output_q + matrix_id * N * N + eigen_index;\n float pivot = diagonal[0] - eigenvalue;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_diagonal[0] = pivot;\n for (int i = 1; i < N; ++i) {\n pivot = diagonal[i] - eigenvalue -\n off_diagonal[i - 1] * off_diagonal[i - 1] / left_diagonal[i - 1];\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_diagonal[i] = pivot;\n }\n pivot = diagonal[N - 1] - eigenvalue;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n right_diagonal[(N - 1) * N] = pivot;\n for (int i = N - 2; i >= 0; --i) {\n pivot = diagonal[i] - eigenvalue -\n off_diagonal[i] * off_diagonal[i] / right_diagonal[(i + 1) * N];\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n right_diagonal[i * N] = pivot;\n }\n int twist = 0;\n float minimum_gamma = fabsf(\n diagonal[0] - eigenvalue -\n off_diagonal[0] * off_diagonal[0] / right_diagonal[N]);\n for (int i = 1; i < N; ++i) {\n float gamma = diagonal[i] - eigenvalue -\n off_diagonal[i - 1] * off_diagonal[i - 1] / left_diagonal[i - 1];\n if (i + 1 < N)\n gamma -= off_diagonal[i] * off_diagonal[i] /\n right_diagonal[(i + 1) * N];\n if (fabsf(gamma) < minimum_gamma) {\n minimum_gamma = fabsf(gamma);\n twist = i;\n }\n }\n right_diagonal[twist * N] = 1.0f;\n for (int i = twist - 1; i >= 0; --i)\n right_diagonal[i * N] = -off_diagonal[i] / left_diagonal[i] *\n right_diagonal[(i + 1) * N];\n for (int i = twist + 1; i < N; ++i)\n right_diagonal[i * N] = -off_diagonal[i - 1] /\n right_diagonal[i * N] * right_diagonal[(i - 1) * N];\n float norm2 = 0.0f;\n for (int i = 0; i < N; ++i)\n norm2 += right_diagonal[i * N] * right_diagonal[i * N];\n const float inverse_norm = rsqrtf(norm2);\n for (int i = 0; i < N; ++i) {\n right_diagonal[i * N] *= inverse_norm;\n }\n output_l[matrix_id * N + eigen_index] = eigenvalue;\n }\n cluster.sync();\n // Preserve the tridiagonal solver\'s natural ownership: each CTA keeps\n // 88 complete eigenvector columns. Within a CTA Q is row-major so a warp\n // reading one row across local eigenvectors sees conflict-free banks.\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int row = index / HALF;\n const int local_eigen = index - row * HALF;\n const int eigen_column = first_row + local_eigen;\n q[row * HALF + local_eigen] =\n output_q[(matrix_id * N + row) * N + eigen_column];\n }\n __syncthreads();\n\n // Ordered MGS is local for all but the single rank boundary. The two CTAs\n // process their independent contiguous ranges concurrently.\n for (int local_eigen = 1; local_eigen < HALF; ++local_eigen) {\n const int eigen_column = first_row + local_eigen;\n const float gap = output_l[matrix_id * N + eigen_column] -\n output_l[matrix_id * N + eigen_column - 1];\n if (gap < 5.0e-3f) {\n if (tid == 0) {\n float dot = 0.0f;\n for (int row = 0; row < N; ++row)\n dot += q[row * HALF + local_eigen - 1] *\n q[row * HALF + local_eigen];\n partial[0] = dot;\n }\n __syncthreads();\n const float dot = partial[0];\n const float inverse_norm = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));\n for (int row = tid; row < N; row += blockDim.x)\n q[row * HALF + local_eigen] =\n (q[row * HALF + local_eigen] -\n dot * q[row * HALF + local_eigen - 1]) * inverse_norm;\n __syncthreads();\n }\n }\n\n // Repair the only adjacent pair split across ranks when it is tight.\n float* q0 = cluster.map_shared_rank(q, 0);\n float* q1 = cluster.map_shared_rank(q, 1);\n const float boundary_gap = output_l[matrix_id * N + HALF] -\n output_l[matrix_id * N + HALF - 1];\n if (boundary_gap < 5.0e-3f) {\n if (rank == 0 && tid == 0) {\n float dot = 0.0f;\n for (int row = 0; row < N; ++row)\n dot += q0[row * HALF + HALF - 1] * q1[row * HALF];\n partial[0] = dot;\n }\n cluster.sync();\n const float dot = partial0[0];\n const float inverse_norm = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));\n if (rank == 1)\n for (int row = tid; row < N; row += blockDim.x)\n q[row * HALF] = (q[row * HALF] -\n dot * q0[row * HALF + HALF - 1]) * inverse_norm;\n cluster.sync();\n }\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 2] = clock64();\n\n // Apply normalized reflectors in reverse blocks of WY_BLOCK. For a\n // block V=[v_0,...,v_b-1], first form S=V^T Q. The coefficients of\n // H_{b-1}...H_0 Q are then obtained by the descending recurrence\n // c_i = 2 * (s_i - sum_{j>i} (v_i^T v_j) c_j).\n // This is the compact-WY triangular solve without explicitly materializing\n // T. It preserves the exact reflector order but amortizes cluster barriers.\n constexpr int WY_BLOCK = 8;\n float* wy_s = v; // WY_BLOCK x HALF, CTA-local\n float* wy_c = wy_s + WY_BLOCK * HALF; // WY_BLOCK x HALF, CTA-local\n float* wy_g = wy_c + WY_BLOCK * HALF; // WY_BLOCK x WY_BLOCK, CTA-local\n\n for (int block_end = N - 3; block_end >= 0; block_end -= WY_BLOCK) {\n const int block_start = max(0, block_end - WY_BLOCK + 1);\n const int block_count = block_end - block_start + 1;\n\n // Replicate the small reflector Gram matrix. This removes the DSM\n // barrier from the critical per-panel loop.\n if (tid < WY_BLOCK * WY_BLOCK) {\n const int i = tid / WY_BLOCK;\n const int j = tid - i * WY_BLOCK;\n float dot = 0.0f;\n if (i < block_count && j < block_count) {\n const int column_i = block_start + i;\n const int column_j = block_start + j;\n const int begin = max(column_i, column_j) + 1;\n for (int row = begin; row < N; ++row)\n dot += saved_matrix[column_i * N + row] *\n saved_matrix[column_j * N + row];\n }\n wy_g[tid] = dot;\n }\n __syncthreads();\n\n // Every CTA has complete vectors, so V.T @ Q is purely local.\n for (int index = tid; index < block_count * HALF; index += blockDim.x) {\n const int i = index / HALF;\n const int local_eigen = index - i * HALF;\n const int reflector_column = block_start + i;\n float dot = 0.0f;\n for (int row = reflector_column + 1; row < N; ++row)\n dot += saved_matrix[reflector_column * N + row] *\n q[row * HALF + local_eigen];\n wy_s[index] = dot;\n }\n __syncthreads();\n\n if (tid < HALF) {\n for (int i = block_count - 1; i >= 0; --i) {\n float coefficient = wy_s[i * HALF + tid];\n for (int j = i + 1; j < block_count; ++j)\n coefficient -= wy_g[i * WY_BLOCK + j] * wy_c[j * HALF + tid];\n wy_c[i * HALF + tid] = 2.0f * coefficient;\n }\n }\n __syncthreads();\n\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int row = index / HALF;\n const int local_eigen = index - row * HALF;\n float update = 0.0f;\n #pragma unroll\n for (int i = 0; i < WY_BLOCK; ++i) {\n if (i < block_count) {\n const int reflector_column = block_start + i;\n if (row > reflector_column)\n update += saved_matrix[reflector_column * N + row] *\n wy_c[i * HALF + local_eigen];\n }\n }\n q[row * HALF + local_eigen] -= update;\n }\n __syncthreads();\n }\n cluster.sync();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 3] = clock64();\n\n // Sturm bisection emits ascending eigenvalues. Store each CTA\'s complete\n // column tile directly; no cross-CTA permutation is required.\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int row = index / HALF;\n const int local_eigen = index - row * HALF;\n const int eigen_column = first_row + local_eigen;\n output_q[(matrix_id * N + row) * N + eigen_column] =\n q[row * HALF + local_eigen];\n }\n cluster.sync();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 4] = clock64();\n}\n';_E084_POTRF256_SOURCE='\n#include <cuda_runtime.h>\n#include <mma.h>\nusing namespace nvcuda;\n\nconstexpr int N = 256;\nconstexpr int NN = N * N;\nconstexpr int TRI = N * (N + 1) / 2;\nconstexpr int PANEL = 16;\nconstexpr int WARPS = 19;\nconstexpr int TILE_FLOATS = 256;\n\n__device__ __forceinline__ int pidx(int row, int column) {\n return row * (row + 1) / 2 + column;\n}\n\nextern "C" __global__ __launch_bounds__(608, 1)\nvoid potrf256_strided_tf32x3_w19(\n const float* __restrict__ gram,\n const float* __restrict__ diagonal_scale,\n float* __restrict__ lower,\n int batch, int leading_dimension, int row_offset, int column_offset,\n float ridge) {\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int warp = tid >> 5;\n const int lane = tid & 31;\n if (matrix_id >= batch) return;\n\n extern __shared__ __align__(32) float storage[];\n float* factor = storage;\n float* diagonal_mean = factor + TRI;\n // Round the packed-factor tail up to a 32-byte boundary for WMMA loads.\n float* stage = factor + TRI + 8;\n float* a_hi = stage + (warp * 5 + 0) * TILE_FLOATS;\n float* a_lo = stage + (warp * 5 + 1) * TILE_FLOATS;\n float* b_hi = stage + (warp * 5 + 2) * TILE_FLOATS;\n float* b_lo = stage + (warp * 5 + 3) * TILE_FLOATS;\n float* c_tile = stage + (warp * 5 + 4) * TILE_FLOATS;\n const float* source = gram + (long long)matrix_id * leading_dimension * leading_dimension;\n\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n if (row >= column) factor[pidx(row, column)] = source[(long long)(row_offset + row) * leading_dimension + column_offset + column];\n }\n if (tid == 0) *diagonal_mean = diagonal_scale[matrix_id];\n __syncthreads();\n const float shift = ridge * (*diagonal_mean);\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = panel + PANEL;\n if (tid == 0) {\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n float diagonal = factor[pidx(column, column)] + shift;\n #pragma unroll\n for (int k = panel; k < column; ++k) {\n const float value = factor[pidx(column, k)];\n diagonal = fmaf(-value, value, diagonal);\n }\n factor[pidx(column, column)] = sqrtf(fmaxf(diagonal, 1.0e-30f));\n #pragma unroll\n for (int row = column + 1; row < end; ++row) {\n float value = factor[pidx(row, column)];\n #pragma unroll\n for (int k = panel; k < column; ++k)\n value = fmaf(-factor[pidx(row, k)], factor[pidx(column, k)], value);\n factor[pidx(row, column)] = value / factor[pidx(column, column)];\n }\n }\n }\n __syncthreads();\n\n // Panel TRSM: a thread owns one complete row below the panel.\n for (int row = end + tid; row < N; row += blockDim.x) {\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n float value = factor[pidx(row, column)];\n #pragma unroll\n for (int k = panel; k < column; ++k)\n value = fmaf(-factor[pidx(row, k)], factor[pidx(column, k)], value);\n factor[pidx(row, column)] = value / factor[pidx(column, column)];\n }\n }\n __syncthreads();\n\n const int trailing_tiles = (N - end) / PANEL;\n const int jobs = trailing_tiles * (trailing_tiles + 1) / 2;\n for (int job = warp; job < jobs; job += WARPS) {\n int tile_row = 0;\n int tile_column = job;\n while (tile_column > tile_row) {\n tile_column -= tile_row + 1;\n ++tile_row;\n }\n const int row_base = end + tile_row * PANEL;\n const int column_base = end + tile_column * PANEL;\n for (int index = lane; index < TILE_FLOATS; index += 32) {\n const int i = index >> 4;\n const int j = index & 15;\n const float av = factor[pidx(row_base + i, panel + j)];\n const float bv = factor[pidx(column_base + j, panel + i)];\n const float ah = wmma::__float_to_tf32(av);\n const float bh = wmma::__float_to_tf32(bv);\n a_hi[index] = ah;\n a_lo[index] = wmma::__float_to_tf32(av - ah);\n b_hi[index] = -bh;\n b_lo[index] = -wmma::__float_to_tf32(bv - bh);\n const int global_row = row_base + i;\n const int global_column = column_base + j;\n c_tile[index] = global_row >= global_column\n ? factor[pidx(global_row, global_column)]\n : factor[pidx(global_column, global_row)];\n }\n __syncwarp();\n\n wmma::fragment<wmma::accumulator, 16, 16, 8, float> c;\n wmma::load_matrix_sync(c, c_tile, 16, wmma::mem_row_major);\n #pragma unroll\n for (int k0 = 0; k0 < PANEL; k0 += 8) {\n wmma::fragment<wmma::matrix_a, 16, 16, 8,\n wmma::precision::tf32, wmma::row_major> ah, al;\n wmma::fragment<wmma::matrix_b, 16, 16, 8,\n wmma::precision::tf32, wmma::row_major> bh, bl;\n wmma::load_matrix_sync(ah, a_hi + k0, 16);\n wmma::load_matrix_sync(al, a_lo + k0, 16);\n wmma::load_matrix_sync(bh, b_hi + k0 * 16, 16);\n wmma::load_matrix_sync(bl, b_lo + k0 * 16, 16);\n wmma::mma_sync(c, ah, bh, c);\n wmma::mma_sync(c, al, bh, c);\n wmma::mma_sync(c, ah, bl, c);\n }\n wmma::store_matrix_sync(c_tile, c, 16, wmma::mem_row_major);\n __syncwarp();\n for (int index = lane; index < TILE_FLOATS; index += 32) {\n const int row = row_base + (index >> 4);\n const int column = column_base + (index & 15);\n if (row >= column) factor[pidx(row, column)] = c_tile[index];\n }\n }\n __syncthreads();\n }\n\n float* destination = lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n destination[index] = row >= column ? factor[pidx(row, column)] : 0.0f;\n }\n}\n';_E084_SIMT_TRSM256_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__\nvoid right_trsm256_strided_panel32_rows32(\n const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch, int source_rows, int source_columns, int row_offset, int column_offset, int rows)\n{\n constexpr int N = 256;\n constexpr int NN = N * N;\n constexpr int ROWS = 32;\n constexpr int PANEL = 32;\n const int matrix_id = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n const int row_base = (int)blockIdx.y * ROWS;\n if (matrix_id >= batch) return;\n\n extern __shared__ float rhs[]; // transposed [N][ROWS]\n const float* source = matrix + (long long)matrix_id * source_rows * source_columns;\n const float* factor = lower + (long long)matrix_id * NN;\n float* destination = output + (long long)matrix_id * rows * N;\n for (int index = tid; index < N * ROWS; index += (int)blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n rhs[index] = row < rows ? source[(long long)(row_offset + row) * source_columns + column_offset + column] : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = min(panel + PANEL, N);\n if (tid < ROWS) {\n const int local_row = tid;\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n if (column >= N) break;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[column * N + k],\n rhs[k * ROWS + local_row], value);\n }\n rhs[column * ROWS + local_row] =\n value / factor[column * N + column];\n }\n }\n __syncthreads();\n const int remaining = (N - end) * ROWS;\n for (int index = tid; index < remaining; index += (int)blockDim.x) {\n const int column = end + index / ROWS;\n const int local_row = index % ROWS;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= N) break;\n value = fmaf(-factor[column * N + k],\n rhs[k * ROWS + local_row], value);\n }\n rhs[column * ROWS + local_row] = value;\n }\n __syncthreads();\n }\n\n for (int index = tid; index < N * ROWS; index += (int)blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n if (row < rows)\n destination[(long long)row * N + column] = rhs[index];\n }\n}\n';_E084_TENSOR_TRSM256_SOURCE='\n#include <cuda_runtime.h>\n#include <mma.h>\nusing namespace nvcuda;\n\nconstexpr int N = 256;\nconstexpr int NN = N * N;\nconstexpr int ROWS = 32;\nconstexpr int PANEL = 16;\nconstexpr int WARPS = 4;\nconstexpr int TILE = 256;\n\nextern "C" __global__ __launch_bounds__(512, 1)\nvoid right_trsm256_wmma_tf32x1_w4(\n const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch, int source_rows, int source_columns,\n int row_offset, int column_offset, int rows,\n int destination_columns, int destination_column_offset) {\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int warp = tid >> 5;\n const int lane = tid & 31;\n const int row_base = blockIdx.y * ROWS;\n if (matrix_id >= batch) return;\n\n extern __shared__ __align__(32) float storage[];\n float* rhs = storage; // transposed [N][ROWS]\n float* stage = rhs + N * ROWS;\n float* a_hi = stage + (warp * 3 + 0) * TILE;\n float* b0_hi = stage + (warp * 3 + 1) * TILE;\n float* b1_hi = stage + (warp * 3 + 2) * TILE;\n const float* source = matrix +\n (long long)matrix_id * source_rows * source_columns;\n const float* factor = lower + (long long)matrix_id * NN;\n float* destination = output\n + (long long)matrix_id * rows * destination_columns;\n\n for (int index = tid; index < N * ROWS; index += blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n rhs[index] = row < rows\n ? source[(long long)(row_offset + row) * source_columns +\n column_offset + column]\n : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = panel + PANEL;\n if (tid < ROWS) {\n const int local_row = tid;\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k = panel; k < column; ++k)\n value = fmaf(-factor[column * N + k],\n rhs[k * ROWS + local_row], value);\n rhs[column * ROWS + local_row] =\n value / factor[column * N + column];\n }\n }\n __syncthreads();\n\n const int column_tiles = (N - end) / 16;\n for (int column_tile = warp; column_tile < column_tiles;\n column_tile += WARPS) {\n const int output_column = end + column_tile * 16;\n for (int index = lane; index < TILE; index += 32) {\n const int i = index >> 4;\n const int j = index & 15;\n const float av = factor[(output_column + i) * N + panel + j];\n const float b0 = rhs[(panel + i) * ROWS + j];\n const float b1 = rhs[(panel + i) * ROWS + 16 + j];\n a_hi[index] = wmma::__float_to_tf32(av);\n b0_hi[index] = -wmma::__float_to_tf32(b0);\n b1_hi[index] = -wmma::__float_to_tf32(b1);\n }\n __syncwarp();\n wmma::fragment<wmma::accumulator, 16, 16, 8, float> c0, c1;\n wmma::load_matrix_sync(c0, rhs + output_column * ROWS,\n ROWS, wmma::mem_row_major);\n wmma::load_matrix_sync(c1, rhs + output_column * ROWS + 16,\n ROWS, wmma::mem_row_major);\n #pragma unroll\n for (int k0 = 0; k0 < PANEL; k0 += 8) {\n wmma::fragment<wmma::matrix_a, 16, 16, 8,\n wmma::precision::tf32, wmma::row_major> ah;\n wmma::fragment<wmma::matrix_b, 16, 16, 8,\n wmma::precision::tf32, wmma::row_major> b0h, b1h;\n wmma::load_matrix_sync(ah, a_hi + k0, 16);\n wmma::load_matrix_sync(b0h, b0_hi + k0 * 16, 16);\n wmma::load_matrix_sync(b1h, b1_hi + k0 * 16, 16);\n wmma::mma_sync(c0, ah, b0h, c0);\n wmma::mma_sync(c1, ah, b1h, c1);\n }\n wmma::store_matrix_sync(rhs + output_column * ROWS,\n c0, ROWS, wmma::mem_row_major);\n wmma::store_matrix_sync(rhs + output_column * ROWS + 16,\n c1, ROWS, wmma::mem_row_major);\n }\n __syncthreads();\n }\n\n for (int index = tid; index < N * ROWS; index += blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n if (row < rows)\n destination[(long long)row * destination_columns\n + destination_column_offset + column] = rhs[index];\n }\n}\n';_N96_PACKED_TILED_SOURCE='\n\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;\n\nconstexpr int N = 96;\nconstexpr int HALF = 96;\nconstexpr int SORT_N = 256;\nconstexpr int TRI = N * (N + 1) / 2;\n\nextern "C" __global__\n__launch_bounds__(256, 5)\nvoid tridiagonal96_packed_tiled_singlecta(\n const float* __restrict__ input,\n float* __restrict__ output_q,\n float* __restrict__ output_l,\n float* __restrict__ saved_reflectors,\n unsigned long long* __restrict__ timers,\n float ql_epsilon) {\n constexpr int rank = 0;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n constexpr int first_row = 0;\n float* saved_matrix = saved_reflectors + (long long)matrix_id * N * N;\n\n extern __shared__ __align__(16) unsigned char storage[];\n float* a = reinterpret_cast<float*>(storage);\n // Packed A and the row half of Q have disjoint lifetimes.\n float* q = a;\n float* v = a + N * N;\n float* w = v + N;\n float* partial = w + N;\n float* rotation_c = partial + N;\n float* rotation_s = rotation_c + N;\n float* sort_values = rotation_s + N;\n int* sort_indices = reinterpret_cast<int*>(sort_values + SORT_N);\n float* scratch = reinterpret_cast<float*>(sort_indices + SORT_N);\n float* diagonal = scratch + 256;\n float* off_diagonal = diagonal + N;\n int* control = reinterpret_cast<int*>(off_diagonal + N);\n\n float* a0 = a;\n float* v0 = v;\n float* w0 = w;\n float* partial0 = partial;\n float* rotation_c0 = rotation_c;\n float* rotation_s0 = rotation_s;\n float* diagonal0 = diagonal;\n float* off_diagonal0 = off_diagonal;\n int* control0 = control;\n\n if (rank == 0 && tid == 0) timers[matrix_id * 5] = clock64();\n\n if (rank == 0) {\n const float* source = input + (long long)matrix_id * N * N;\n for (int row = 0; row < N; ++row)\n if (tid <= row)\n a[row * (row + 1) / 2 + tid] = source[row * N + tid];\n }\n __syncthreads();\n\n // CTA 0 performs the right-looking symmetric Householder reduction.\n // Warp-shuffle block reductions use two CTA barriers instead of the nine\n // barriers in a shared-memory binary tree. The normalized reflector norm\n // is known analytically from ||x|| and x_0, avoiding a second reduction.\n if (rank == 0) {\n for (int column = 0; column < N - 2; ++column) {\n const int begin = column + 1;\n const int length = N - begin;\n float local = 0.0f;\n for (int i = tid; i < length; i += blockDim.x) {\n const float x = a[(begin + i) * (begin + i + 1) / 2 + column];\n local += x * x;\n }\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n local += __shfl_down_sync(0xffffffff, local, offset);\n if ((tid & 31) == 0) scratch[tid >> 5] = local;\n __syncthreads();\n float total = tid < 8 ? scratch[tid] : 0.0f;\n if (tid < 32) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n total += __shfl_down_sync(0xffffffff, total, offset);\n if (tid == 0) scratch[0] = total;\n }\n __syncthreads();\n\n const float norm = sqrtf(scratch[0]);\n const float first = a[begin * (begin + 1) / 2 + column];\n const float tridiagonal_value = -copysignf(norm, first);\n if (tid == 0) off_diagonal[column] = tridiagonal_value;\n const float reflector_norm = sqrtf(\n fmaxf(2.0f * norm * (norm + fabsf(first)), 0.0f));\n const float inverse = reflector_norm > 0.0f\n ? 1.0f / reflector_norm : 0.0f;\n for (int i = tid; i < length; i += blockDim.x) {\n float value = a[(begin + i) * (begin + i + 1) / 2 + column];\n if (i == 0) value -= tridiagonal_value;\n v[i] = value * inverse;\n }\n __syncthreads();\n\n if (tid < length) {\n const int row = begin + tid;\n float product = 0.0f;\n // Only the lower active triangle is current. Reflect it on load.\n for (int j = 0; j <= tid; ++j)\n product += a[row * (row + 1) / 2 + begin + j] * v[j];\n for (int j = tid + 1; j < length; ++j)\n product += a[(begin + j) * (begin + j + 1) / 2 + row] * v[j];\n w[tid] = product;\n }\n __syncthreads();\n\n local = 0.0f;\n for (int i = tid; i < length; i += blockDim.x)\n local += v[i] * w[i];\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n local += __shfl_down_sync(0xffffffff, local, offset);\n if ((tid & 31) == 0) scratch[tid >> 5] = local;\n __syncthreads();\n total = tid < 8 ? scratch[tid] : 0.0f;\n if (tid < 32) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n total += __shfl_down_sync(0xffffffff, total, offset);\n if (tid == 0) scratch[0] = total;\n }\n __syncthreads();\n const float projection = scratch[0];\n for (int i = tid; i < length; i += blockDim.x)\n w[i] = 2.0f * (w[i] - projection * v[i]);\n __syncthreads();\n\n // Update only the lower triangle. At each row, participating threads\n // write a contiguous segment and avoid dynamic integer division.\n constexpr int UPDATE_TILE = 16;\n const int tile_row = tid >> 4;\n const int tile_column = tid & 15;\n for (int row_base = 0; row_base < length; row_base += UPDATE_TILE) {\n for (int column_base = 0; column_base <= row_base;\n column_base += UPDATE_TILE) {\n const int i = row_base + tile_row;\n const int j = column_base + tile_column;\n if (i < length && j < length && j <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n v[i] * w[j] + w[i] * v[j];\n }\n }\n __syncthreads();\n for (int i = tid; i < length; i += blockDim.x)\n a[(begin + i) * (begin + i + 1) / 2 + column] = v[i];\n __syncthreads();\n }\n if (tid == 0) off_diagonal[N - 2] = a[(N - 1) * N / 2 + N - 2];\n if (tid < N) {\n diagonal[tid] = a[tid * (tid + 1) / 2 + tid];\n if (tid == N - 1) off_diagonal[tid] = 0.0;\n }\n }\n // Spill saved unit reflectors before q aliases packed A.\n if (rank == 0) {\n for (int column = 0; column < N - 2; ++column)\n for (int row = column + 1 + tid; row < N; row += blockDim.x)\n saved_matrix[column * N + row] =\n a[row * (row + 1) / 2 + column];\n // cluster.sync() orders DSM, but the reflector spill is device memory.\n // Every producing thread must publish its own global stores before rank 1\n // later reloads the compact-WY panels.\n __threadfence();\n }\n __syncthreads();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 1] = clock64();\n // Replicate the tiny tridiagonal representation. Without this copy CTA 1\n // performs every Sturm/LDL recurrence through remote DSM loads.\n if (rank == 1 && tid < N) {\n diagonal[tid] = diagonal0[tid];\n off_diagonal[tid] = off_diagonal0[tid];\n }\n __syncthreads();\n\n // Parallel Sturm bisection followed by one shifted inverse iteration per\n // eigenpair. CTA ranks split the 176 eigenpairs 88/88. ``q`` is reused as\n // per-eigenpair LDL diagonal workspace, while output_q temporarily stores\n // the modified right-hand sides and resulting tridiagonal eigenvectors.\n if (rank == 0 && tid == 0) {\n float lower = diagonal[0] - fabsf(off_diagonal[0]);\n float upper = diagonal[0] + fabsf(off_diagonal[0]);\n for (int i = 1; i < N; ++i) {\n const float radius = fabsf(off_diagonal[i - 1]) +\n (i + 1 < N ? fabsf(off_diagonal[i]) : 0.0f);\n lower = fminf(lower, diagonal[i] - radius);\n upper = fmaxf(upper, diagonal[i] + radius);\n }\n partial[0] = lower;\n partial[1] = upper;\n }\n __syncthreads();\n if (tid < HALF) {\n const int eigen_index = first_row + tid;\n float lower = partial0[0];\n float upper = partial0[1];\n #pragma unroll\n for (int step = 0; step < 27; ++step) {\n const float midpoint = 0.5f * (lower + upper);\n float pivot = diagonal[0] - midpoint;\n int count = pivot < 0.0f;\n for (int i = 1; i < N; ++i) {\n if (fabsf(pivot) < 1.0e-12f)\n pivot = copysignf(1.0e-12f, pivot == 0.0f ? -1.0f : pivot);\n pivot = diagonal[i] - midpoint -\n off_diagonal[i - 1] * off_diagonal[i - 1] / pivot;\n count += pivot < 0.0f;\n }\n if (count <= eigen_index) lower = midpoint;\n else upper = midpoint;\n }\n const float eigenvalue = 0.5f * (lower + upper);\n float* left_diagonal = q + tid * N;\n float* right_diagonal = output_q + matrix_id * N * N + eigen_index;\n float pivot = diagonal[0] - eigenvalue;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_diagonal[0] = pivot;\n for (int i = 1; i < N; ++i) {\n pivot = diagonal[i] - eigenvalue -\n off_diagonal[i - 1] * off_diagonal[i - 1] / left_diagonal[i - 1];\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_diagonal[i] = pivot;\n }\n pivot = diagonal[N - 1] - eigenvalue;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n right_diagonal[(N - 1) * N] = pivot;\n for (int i = N - 2; i >= 0; --i) {\n pivot = diagonal[i] - eigenvalue -\n off_diagonal[i] * off_diagonal[i] / right_diagonal[(i + 1) * N];\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n right_diagonal[i * N] = pivot;\n }\n int twist = 0;\n float minimum_gamma = fabsf(\n diagonal[0] - eigenvalue -\n off_diagonal[0] * off_diagonal[0] / right_diagonal[N]);\n for (int i = 1; i < N; ++i) {\n float gamma = diagonal[i] - eigenvalue -\n off_diagonal[i - 1] * off_diagonal[i - 1] / left_diagonal[i - 1];\n if (i + 1 < N)\n gamma -= off_diagonal[i] * off_diagonal[i] /\n right_diagonal[(i + 1) * N];\n if (fabsf(gamma) < minimum_gamma) {\n minimum_gamma = fabsf(gamma);\n twist = i;\n }\n }\n right_diagonal[twist * N] = 1.0f;\n for (int i = twist - 1; i >= 0; --i)\n right_diagonal[i * N] = -off_diagonal[i] / left_diagonal[i] *\n right_diagonal[(i + 1) * N];\n for (int i = twist + 1; i < N; ++i)\n right_diagonal[i * N] = -off_diagonal[i - 1] /\n right_diagonal[i * N] * right_diagonal[(i - 1) * N];\n float norm2 = 0.0f;\n for (int i = 0; i < N; ++i)\n norm2 += right_diagonal[i * N] * right_diagonal[i * N];\n const float inverse_norm = rsqrtf(norm2);\n for (int i = 0; i < N; ++i) {\n right_diagonal[i * N] *= inverse_norm;\n }\n output_l[matrix_id * N + eigen_index] = eigenvalue;\n }\n __syncthreads();\n // Convert the global column workspace to the row-distributed shared Q.\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int local_row = index / N;\n const int eigen_column = index - local_row * N;\n q[index] = output_q[(matrix_id * N + first_row + local_row) * N + eigen_column];\n }\n __syncthreads();\n // Independent inverse iterations only lose visible orthogonality for very\n // tight adjacent pairs. A local ordered MGS pass repairs those pairs while\n // rotating by an amount far below the checker tolerance for their gap.\n float* partial1 = partial;\n for (int eigen_column = 1; eigen_column < N; ++eigen_column) {\n const float gap = output_l[matrix_id * N + eigen_column] -\n output_l[matrix_id * N + eigen_column - 1];\n if (gap < 5.0e-3f) {\n if (tid == 0) {\n float dot = 0.0f;\n for (int local_row = 0; local_row < HALF; ++local_row)\n dot += q[local_row * N + eigen_column - 1] *\n q[local_row * N + eigen_column];\n partial[0] = dot;\n }\n __syncthreads();\n const float dot = partial[0];\n const float inverse_norm = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));\n if (tid < HALF)\n q[tid * N + eigen_column] =\n (q[tid * N + eigen_column] -\n dot * q[tid * N + eigen_column - 1]) * inverse_norm;\n __syncthreads();\n }\n }\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 2] = clock64();\n\n // Apply normalized reflectors in reverse blocks of WY_BLOCK. For a\n // block V=[v_0,...,v_b-1], first form S=V^T Q. The coefficients of\n // H_{b-1}...H_0 Q are then obtained by the descending recurrence\n // c_i = 2 * (s_i - sum_{j>i} (v_i^T v_j) c_j).\n // This is the compact-WY triangular solve without explicitly materializing\n // T. It preserves the exact reflector order but amortizes cluster barriers.\n constexpr int WY_BLOCK = 10;\n float* wy_s = v; // WY_BLOCK x N, CTA-local\n float* wy_c = wy_s + WY_BLOCK * N; // WY_BLOCK x N, CTA-local\n float* wy_g = wy_c + WY_BLOCK * N; // WY_BLOCK x WY_BLOCK, rank 0\n float* wy_s0 = wy_s;\n float* wy_s1 = wy_s;\n float* wy_g0 = wy_g;\n\n for (int block_end = N - 3; block_end >= 0; block_end -= WY_BLOCK) {\n const int block_start = max(0, block_end - WY_BLOCK + 1);\n const int block_count = block_end - block_start + 1;\n\n // Gram matrix of the stored normalized reflectors. Only rank 0 writes it;\n // both ranks consume it after the cluster barrier.\n if (rank == 0 && tid < WY_BLOCK * WY_BLOCK) {\n const int i = tid / WY_BLOCK;\n const int j = tid - i * WY_BLOCK;\n float dot = 0.0f;\n if (i < block_count && j < block_count) {\n const int column_i = block_start + i;\n const int column_j = block_start + j;\n const int begin = max(column_i, column_j) + 1;\n for (int row = begin; row < N; ++row)\n dot += saved_matrix[column_i * N + row] * saved_matrix[column_j * N + row];\n }\n wy_g[tid] = dot;\n }\n __syncthreads();\n\n // Each CTA forms its half-row contribution to V^T Q.\n for (int index = tid; index < block_count * N; index += blockDim.x) {\n const int i = index / N;\n const int eigen_column = index - i * N;\n const int reflector_column = block_start + i;\n const int begin = reflector_column + 1;\n float dot = 0.0f;\n for (int local_row = 0; local_row < HALF; ++local_row) {\n const int global_row = first_row + local_row;\n if (global_row >= begin)\n dot += saved_matrix[reflector_column * N + global_row] *\n q[local_row * N + eigen_column];\n }\n wy_s[index] = dot;\n }\n __syncthreads();\n\n // Both CTAs redundantly solve the tiny upper-triangular recurrence. The\n // duplicate work is tiny and keeps the subsequent Q update CTA-local.\n if (tid < N) {\n for (int i = block_count - 1; i >= 0; --i) {\n float coefficient = wy_s[i * N + tid];\n for (int j = i + 1; j < block_count; ++j)\n coefficient -= wy_g0[i * WY_BLOCK + j] * wy_c[j * N + tid];\n wy_c[i * N + tid] = 2.0f * coefficient;\n }\n }\n __syncthreads();\n\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int local_row = index / N;\n const int eigen_column = index - local_row * N;\n const int global_row = first_row + local_row;\n float update = 0.0f;\n #pragma unroll\n for (int i = 0; i < WY_BLOCK; ++i) {\n if (i < block_count) {\n const int reflector_column = block_start + i;\n if (global_row > reflector_column)\n update += saved_matrix[reflector_column * N + global_row] *\n wy_c[i * N + eigen_column];\n }\n }\n q[index] -= update;\n }\n // The next iteration\'s Gram barrier also guarantees both CTAs completed\n // their local Q update before either computes the next V^T Q.\n }\n __syncthreads();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 3] = clock64();\n\n // Independent bitonic sorts avoid another permutation exchange. Bisection\n // already orders the columns, but Rayleigh correction can swap a close pair.\n if (tid < N) {\n sort_values[tid] = output_l[matrix_id * N + tid];\n sort_indices[tid] = tid;\n } else {\n sort_values[tid] = __int_as_float(0x7f800000);\n sort_indices[tid] = tid;\n }\n __syncthreads();\n for (int width = 2; width <= SORT_N; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n const int other = tid ^ stride;\n const float self_value = sort_values[tid];\n const float other_value = sort_values[other];\n const int self_index = sort_indices[tid];\n const int other_index = sort_indices[other];\n const bool ascending = (tid & width) == 0;\n const bool take_minimum = ascending == (tid < other);\n const bool self_is_minimum = self_value < other_value ||\n (self_value == other_value && self_index < other_index);\n if (take_minimum != self_is_minimum) {\n sort_values[tid] = other_value;\n sort_indices[tid] = other_index;\n }\n __syncthreads();\n }\n }\n if (rank == 0 && tid < N)\n output_l[matrix_id * N + tid] = sort_values[tid];\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int local_row = index / N;\n const int sorted_column = index - local_row * N;\n output_q[(matrix_id * N + first_row + local_row) * N + sorted_column] =\n q[local_row * N + sort_indices[sorted_column]];\n }\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 4] = clock64();\n}\n';_N128_PACKED_TILED_SOURCE='\n\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;\n\nconstexpr int N = 128;\nconstexpr int HALF = 128;\nconstexpr int SORT_N = 256;\nconstexpr int TRI = N * (N + 1) / 2;\n\nextern "C" __global__\n__launch_bounds__(256, 2)\nvoid tridiagonal128_packed_tiled_singlecta(\n const float* __restrict__ input,\n float* __restrict__ output_q,\n float* __restrict__ output_l,\n float* __restrict__ saved_reflectors,\n unsigned long long* __restrict__ timers,\n float ql_epsilon) {\n constexpr int rank = 0;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n constexpr int first_row = 0;\n float* saved_matrix = saved_reflectors + (long long)matrix_id * N * N;\n\n extern __shared__ __align__(16) unsigned char storage[];\n float* a = reinterpret_cast<float*>(storage);\n // Packed A and the row half of Q have disjoint lifetimes.\n float* q = a;\n float* v = a + N * N;\n float* w = v + N;\n float* partial = w + N;\n float* rotation_c = partial + N;\n float* rotation_s = rotation_c + N;\n float* sort_values = rotation_s + N;\n int* sort_indices = reinterpret_cast<int*>(sort_values + SORT_N);\n float* scratch = reinterpret_cast<float*>(sort_indices + SORT_N);\n float* diagonal = scratch + 256;\n float* off_diagonal = diagonal + N;\n int* control = reinterpret_cast<int*>(off_diagonal + N);\n\n float* a0 = a;\n float* v0 = v;\n float* w0 = w;\n float* partial0 = partial;\n float* rotation_c0 = rotation_c;\n float* rotation_s0 = rotation_s;\n float* diagonal0 = diagonal;\n float* off_diagonal0 = off_diagonal;\n int* control0 = control;\n\n if (rank == 0 && tid == 0) timers[matrix_id * 5] = clock64();\n\n if (rank == 0) {\n const float* source = input + (long long)matrix_id * N * N;\n for (int row = 0; row < N; ++row)\n if (tid <= row)\n a[row * (row + 1) / 2 + tid] = source[row * N + tid];\n }\n __syncthreads();\n\n // CTA 0 performs the right-looking symmetric Householder reduction.\n // Warp-shuffle block reductions use two CTA barriers instead of the nine\n // barriers in a shared-memory binary tree. The normalized reflector norm\n // is known analytically from ||x|| and x_0, avoiding a second reduction.\n if (rank == 0) {\n for (int column = 0; column < N - 2; ++column) {\n const int begin = column + 1;\n const int length = N - begin;\n float local = 0.0f;\n for (int i = tid; i < length; i += blockDim.x) {\n const float x = a[(begin + i) * (begin + i + 1) / 2 + column];\n local += x * x;\n }\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n local += __shfl_down_sync(0xffffffff, local, offset);\n if ((tid & 31) == 0) scratch[tid >> 5] = local;\n __syncthreads();\n float total = tid < 8 ? scratch[tid] : 0.0f;\n if (tid < 32) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n total += __shfl_down_sync(0xffffffff, total, offset);\n if (tid == 0) scratch[0] = total;\n }\n __syncthreads();\n\n const float norm = sqrtf(scratch[0]);\n const float first = a[begin * (begin + 1) / 2 + column];\n const float tridiagonal_value = -copysignf(norm, first);\n if (tid == 0) off_diagonal[column] = tridiagonal_value;\n const float reflector_norm = sqrtf(\n fmaxf(2.0f * norm * (norm + fabsf(first)), 0.0f));\n const float inverse = reflector_norm > 0.0f\n ? 1.0f / reflector_norm : 0.0f;\n for (int i = tid; i < length; i += blockDim.x) {\n float value = a[(begin + i) * (begin + i + 1) / 2 + column];\n if (i == 0) value -= tridiagonal_value;\n v[i] = value * inverse;\n }\n __syncthreads();\n\n if (tid < length) {\n const int row = begin + tid;\n float product = 0.0f;\n // Only the lower active triangle is current. Reflect it on load.\n for (int j = 0; j <= tid; ++j)\n product += a[row * (row + 1) / 2 + begin + j] * v[j];\n for (int j = tid + 1; j < length; ++j)\n product += a[(begin + j) * (begin + j + 1) / 2 + row] * v[j];\n w[tid] = product;\n }\n __syncthreads();\n\n local = 0.0f;\n for (int i = tid; i < length; i += blockDim.x)\n local += v[i] * w[i];\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n local += __shfl_down_sync(0xffffffff, local, offset);\n if ((tid & 31) == 0) scratch[tid >> 5] = local;\n __syncthreads();\n total = tid < 8 ? scratch[tid] : 0.0f;\n if (tid < 32) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n total += __shfl_down_sync(0xffffffff, total, offset);\n if (tid == 0) scratch[0] = total;\n }\n __syncthreads();\n const float projection = scratch[0];\n for (int i = tid; i < length; i += blockDim.x)\n w[i] = 2.0f * (w[i] - projection * v[i]);\n __syncthreads();\n\n // Update only the lower triangle. At each row, participating threads\n // write a contiguous segment and avoid dynamic integer division.\n constexpr int UPDATE_TILE = 16;\n const int tile_row = tid >> 4;\n const int tile_column = tid & 15;\n for (int row_base = 0; row_base < length; row_base += UPDATE_TILE) {\n for (int column_base = 0; column_base <= row_base;\n column_base += UPDATE_TILE) {\n const int i = row_base + tile_row;\n const int j = column_base + tile_column;\n if (i < length && j < length && j <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n v[i] * w[j] + w[i] * v[j];\n }\n }\n __syncthreads();\n for (int i = tid; i < length; i += blockDim.x)\n a[(begin + i) * (begin + i + 1) / 2 + column] = v[i];\n __syncthreads();\n }\n if (tid == 0) off_diagonal[N - 2] = a[(N - 1) * N / 2 + N - 2];\n if (tid < N) {\n diagonal[tid] = a[tid * (tid + 1) / 2 + tid];\n if (tid == N - 1) off_diagonal[tid] = 0.0;\n }\n }\n // Spill saved unit reflectors before q aliases packed A.\n if (rank == 0) {\n for (int column = 0; column < N - 2; ++column)\n for (int row = column + 1 + tid; row < N; row += blockDim.x)\n saved_matrix[column * N + row] =\n a[row * (row + 1) / 2 + column];\n // cluster.sync() orders DSM, but the reflector spill is device memory.\n // Every producing thread must publish its own global stores before rank 1\n // later reloads the compact-WY panels.\n __threadfence();\n }\n __syncthreads();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 1] = clock64();\n // Replicate the tiny tridiagonal representation. Without this copy CTA 1\n // performs every Sturm/LDL recurrence through remote DSM loads.\n if (rank == 1 && tid < N) {\n diagonal[tid] = diagonal0[tid];\n off_diagonal[tid] = off_diagonal0[tid];\n }\n __syncthreads();\n\n // Parallel Sturm bisection followed by one shifted inverse iteration per\n // eigenpair. CTA ranks split the 176 eigenpairs 88/88. ``q`` is reused as\n // per-eigenpair LDL diagonal workspace, while output_q temporarily stores\n // the modified right-hand sides and resulting tridiagonal eigenvectors.\n if (rank == 0 && tid == 0) {\n float lower = diagonal[0] - fabsf(off_diagonal[0]);\n float upper = diagonal[0] + fabsf(off_diagonal[0]);\n for (int i = 1; i < N; ++i) {\n const float radius = fabsf(off_diagonal[i - 1]) +\n (i + 1 < N ? fabsf(off_diagonal[i]) : 0.0f);\n lower = fminf(lower, diagonal[i] - radius);\n upper = fmaxf(upper, diagonal[i] + radius);\n }\n partial[0] = lower;\n partial[1] = upper;\n }\n __syncthreads();\n if (tid < HALF) {\n const int eigen_index = first_row + tid;\n float lower = partial0[0];\n float upper = partial0[1];\n #pragma unroll\n for (int step = 0; step < 27; ++step) {\n const float midpoint = 0.5f * (lower + upper);\n float pivot = diagonal[0] - midpoint;\n int count = pivot < 0.0f;\n for (int i = 1; i < N; ++i) {\n if (fabsf(pivot) < 1.0e-12f)\n pivot = copysignf(1.0e-12f, pivot == 0.0f ? -1.0f : pivot);\n pivot = diagonal[i] - midpoint -\n off_diagonal[i - 1] * off_diagonal[i - 1] / pivot;\n count += pivot < 0.0f;\n }\n if (count <= eigen_index) lower = midpoint;\n else upper = midpoint;\n }\n const float eigenvalue = 0.5f * (lower + upper);\n float* left_diagonal = q + tid * N;\n float* right_diagonal = output_q + matrix_id * N * N + eigen_index;\n float pivot = diagonal[0] - eigenvalue;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_diagonal[0] = pivot;\n for (int i = 1; i < N; ++i) {\n pivot = diagonal[i] - eigenvalue -\n off_diagonal[i - 1] * off_diagonal[i - 1] / left_diagonal[i - 1];\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_diagonal[i] = pivot;\n }\n pivot = diagonal[N - 1] - eigenvalue;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n right_diagonal[(N - 1) * N] = pivot;\n for (int i = N - 2; i >= 0; --i) {\n pivot = diagonal[i] - eigenvalue -\n off_diagonal[i] * off_diagonal[i] / right_diagonal[(i + 1) * N];\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n right_diagonal[i * N] = pivot;\n }\n int twist = 0;\n float minimum_gamma = fabsf(\n diagonal[0] - eigenvalue -\n off_diagonal[0] * off_diagonal[0] / right_diagonal[N]);\n for (int i = 1; i < N; ++i) {\n float gamma = diagonal[i] - eigenvalue -\n off_diagonal[i - 1] * off_diagonal[i - 1] / left_diagonal[i - 1];\n if (i + 1 < N)\n gamma -= off_diagonal[i] * off_diagonal[i] /\n right_diagonal[(i + 1) * N];\n if (fabsf(gamma) < minimum_gamma) {\n minimum_gamma = fabsf(gamma);\n twist = i;\n }\n }\n right_diagonal[twist * N] = 1.0f;\n for (int i = twist - 1; i >= 0; --i)\n right_diagonal[i * N] = -off_diagonal[i] / left_diagonal[i] *\n right_diagonal[(i + 1) * N];\n for (int i = twist + 1; i < N; ++i)\n right_diagonal[i * N] = -off_diagonal[i - 1] /\n right_diagonal[i * N] * right_diagonal[(i - 1) * N];\n float norm2 = 0.0f;\n for (int i = 0; i < N; ++i)\n norm2 += right_diagonal[i * N] * right_diagonal[i * N];\n const float inverse_norm = rsqrtf(norm2);\n for (int i = 0; i < N; ++i) {\n right_diagonal[i * N] *= inverse_norm;\n }\n output_l[matrix_id * N + eigen_index] = eigenvalue;\n }\n __syncthreads();\n // Convert the global column workspace to the row-distributed shared Q.\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int local_row = index / N;\n const int eigen_column = index - local_row * N;\n q[index] = output_q[(matrix_id * N + first_row + local_row) * N + eigen_column];\n }\n __syncthreads();\n // Independent inverse iterations only lose visible orthogonality for very\n // tight adjacent pairs. A local ordered MGS pass repairs those pairs while\n // rotating by an amount far below the checker tolerance for their gap.\n float* partial1 = partial;\n for (int eigen_column = 1; eigen_column < N; ++eigen_column) {\n const float gap = output_l[matrix_id * N + eigen_column] -\n output_l[matrix_id * N + eigen_column - 1];\n if (gap < 5.0e-3f) {\n if (tid == 0) {\n float dot = 0.0f;\n for (int local_row = 0; local_row < HALF; ++local_row)\n dot += q[local_row * N + eigen_column - 1] *\n q[local_row * N + eigen_column];\n partial[0] = dot;\n }\n __syncthreads();\n const float dot = partial[0];\n const float inverse_norm = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));\n if (tid < HALF)\n q[tid * N + eigen_column] =\n (q[tid * N + eigen_column] -\n dot * q[tid * N + eigen_column - 1]) * inverse_norm;\n __syncthreads();\n }\n }\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 2] = clock64();\n\n // Apply normalized reflectors in reverse blocks of WY_BLOCK. For a\n // block V=[v_0,...,v_b-1], first form S=V^T Q. The coefficients of\n // H_{b-1}...H_0 Q are then obtained by the descending recurrence\n // c_i = 2 * (s_i - sum_{j>i} (v_i^T v_j) c_j).\n // This is the compact-WY triangular solve without explicitly materializing\n // T. It preserves the exact reflector order but amortizes cluster barriers.\n constexpr int WY_BLOCK = 10;\n float* wy_s = v; // WY_BLOCK x N, CTA-local\n float* wy_c = wy_s + WY_BLOCK * N; // WY_BLOCK x N, CTA-local\n float* wy_g = wy_c + WY_BLOCK * N; // WY_BLOCK x WY_BLOCK, rank 0\n float* wy_s0 = wy_s;\n float* wy_s1 = wy_s;\n float* wy_g0 = wy_g;\n\n for (int block_end = N - 3; block_end >= 0; block_end -= WY_BLOCK) {\n const int block_start = max(0, block_end - WY_BLOCK + 1);\n const int block_count = block_end - block_start + 1;\n\n // Gram matrix of the stored normalized reflectors. Only rank 0 writes it;\n // both ranks consume it after the cluster barrier.\n if (rank == 0 && tid < WY_BLOCK * WY_BLOCK) {\n const int i = tid / WY_BLOCK;\n const int j = tid - i * WY_BLOCK;\n float dot = 0.0f;\n if (i < block_count && j < block_count) {\n const int column_i = block_start + i;\n const int column_j = block_start + j;\n const int begin = max(column_i, column_j) + 1;\n for (int row = begin; row < N; ++row)\n dot += saved_matrix[column_i * N + row] * saved_matrix[column_j * N + row];\n }\n wy_g[tid] = dot;\n }\n __syncthreads();\n\n // Each CTA forms its half-row contribution to V^T Q.\n for (int index = tid; index < block_count * N; index += blockDim.x) {\n const int i = index / N;\n const int eigen_column = index - i * N;\n const int reflector_column = block_start + i;\n const int begin = reflector_column + 1;\n float dot = 0.0f;\n for (int local_row = 0; local_row < HALF; ++local_row) {\n const int global_row = first_row + local_row;\n if (global_row >= begin)\n dot += saved_matrix[reflector_column * N + global_row] *\n q[local_row * N + eigen_column];\n }\n wy_s[index] = dot;\n }\n __syncthreads();\n\n // Both CTAs redundantly solve the tiny upper-triangular recurrence. The\n // duplicate work is tiny and keeps the subsequent Q update CTA-local.\n if (tid < N) {\n for (int i = block_count - 1; i >= 0; --i) {\n float coefficient = wy_s[i * N + tid];\n for (int j = i + 1; j < block_count; ++j)\n coefficient -= wy_g0[i * WY_BLOCK + j] * wy_c[j * N + tid];\n wy_c[i * N + tid] = 2.0f * coefficient;\n }\n }\n __syncthreads();\n\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int local_row = index / N;\n const int eigen_column = index - local_row * N;\n const int global_row = first_row + local_row;\n float update = 0.0f;\n #pragma unroll\n for (int i = 0; i < WY_BLOCK; ++i) {\n if (i < block_count) {\n const int reflector_column = block_start + i;\n if (global_row > reflector_column)\n update += saved_matrix[reflector_column * N + global_row] *\n wy_c[i * N + eigen_column];\n }\n }\n q[index] -= update;\n }\n // The next iteration\'s Gram barrier also guarantees both CTAs completed\n // their local Q update before either computes the next V^T Q.\n }\n __syncthreads();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 3] = clock64();\n\n // Independent bitonic sorts avoid another permutation exchange. Bisection\n // already orders the columns, but Rayleigh correction can swap a close pair.\n if (tid < N) {\n sort_values[tid] = output_l[matrix_id * N + tid];\n sort_indices[tid] = tid;\n } else {\n sort_values[tid] = __int_as_float(0x7f800000);\n sort_indices[tid] = tid;\n }\n __syncthreads();\n for (int width = 2; width <= SORT_N; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n const int other = tid ^ stride;\n const float self_value = sort_values[tid];\n const float other_value = sort_values[other];\n const int self_index = sort_indices[tid];\n const int other_index = sort_indices[other];\n const bool ascending = (tid & width) == 0;\n const bool take_minimum = ascending == (tid < other);\n const bool self_is_minimum = self_value < other_value ||\n (self_value == other_value && self_index < other_index);\n if (take_minimum != self_is_minimum) {\n sort_values[tid] = other_value;\n sort_indices[tid] = other_index;\n }\n __syncthreads();\n }\n }\n if (rank == 0 && tid < N)\n output_l[matrix_id * N + tid] = sort_values[tid];\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int local_row = index / N;\n const int sorted_column = index - local_row * N;\n output_q[(matrix_id * N + first_row + local_row) * N + sorted_column] =\n q[local_row * N + sort_indices[sorted_column]];\n }\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 4] = clock64();\n}\n';_TRIDIAG_TWISTED512_SOURCE='\n\n#include <cuda_runtime.h>\n\nconstexpr int N = 512;\n\nextern "C" __global__\n__launch_bounds__(512, 1)\nvoid tridiag_twisted512(\n const float* __restrict__ matrix,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace) {\n const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;\n __syncthreads();\n\n if (eigen == 0) {\n float lower = diagonal[0] - fabsf(off[0]);\n float upper = diagonal[0] + fabsf(off[0]);\n for (int i = 1; i < N; ++i) {\n const float radius = fabsf(off[i - 1])\n + (i + 1 < N ? fabsf(off[i]) : 0.0f);\n lower = fminf(lower, diagonal[i] - radius);\n upper = fmaxf(upper, diagonal[i] + radius);\n }\n bounds[0] = lower;\n bounds[1] = upper;\n }\n __syncthreads();\n\n float lower = bounds[0];\n float upper = bounds[1];\n #pragma unroll 1\n for (int step = 0; step < 30; ++step) {\n const float shift = 0.5f * (lower + upper);\n float pivot = diagonal[0] - shift;\n int count = pivot < 0.0f;\n #pragma unroll 1\n for (int row = 1; row < N; ++row) {\n if (fabsf(pivot) < 1.0e-12f)\n pivot = copysignf(1.0e-12f, pivot == 0.0f ? -1.0f : pivot);\n pivot = diagonal[row] - shift\n - off[row - 1] * off[row - 1] / pivot;\n count += pivot < 0.0f;\n }\n if (count <= eigen) lower = shift;\n else upper = shift;\n }\n const float lambda = 0.5f * (lower + upper);\n values[(long long)batch * N + eigen] = lambda;\n\n // Forward and backward LDL diagonals, laid out [row,eigen] so a warp\'s\n // stores are coalesced despite every thread owning one recurrence.\n const long long wb = (long long)batch * N * N;\n float pivot = diagonal[0] - lambda;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_workspace[wb + eigen] = pivot;\n for (int row = 1; row < N; ++row) {\n const float previous = left_workspace[wb + (long long)(row - 1) * N + eigen];\n pivot = diagonal[row] - lambda\n - off[row - 1] * off[row - 1] / previous;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_workspace[wb + (long long)row * N + eigen] = pivot;\n }\n pivot = diagonal[N - 1] - lambda;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n vectors[mb + (long long)(N - 1) * N + eigen] = pivot;\n for (int row = N - 2; row >= 0; --row) {\n const float next = vectors[mb + (long long)(row + 1) * N + eigen];\n pivot = diagonal[row] - lambda - off[row] * off[row] / next;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n vectors[mb + (long long)row * N + eigen] = pivot;\n }\n\n int twist = 0;\n float gamma0 = diagonal[0] - lambda;\n if (N > 1)\n gamma0 -= off[0] * off[0] / vectors[mb + N + eigen];\n float minimum = fabsf(gamma0);\n for (int row = 1; row < N; ++row) {\n float gamma = diagonal[row] - lambda\n - off[row - 1] * off[row - 1]\n / left_workspace[wb + (long long)(row - 1) * N + eigen];\n if (row + 1 < N)\n gamma -= off[row] * off[row]\n / vectors[mb + (long long)(row + 1) * N + eigen];\n if (fabsf(gamma) < minimum) {\n minimum = fabsf(gamma);\n twist = row;\n }\n }\n vectors[mb + (long long)twist * N + eigen] = 1.0f;\n for (int row = twist - 1; row >= 0; --row) {\n vectors[mb + (long long)row * N + eigen] =\n -off[row]\n / left_workspace[wb + (long long)row * N + eigen]\n * vectors[mb + (long long)(row + 1) * N + eigen];\n }\n for (int row = twist + 1; row < N; ++row) {\n vectors[mb + (long long)row * N + eigen] =\n -off[row - 1]\n / vectors[mb + (long long)row * N + eigen]\n * vectors[mb + (long long)(row - 1) * N + eigen];\n }\n float norm2 = 0.0f;\n for (int row = 0; row < N; ++row) {\n const float x = vectors[mb + (long long)row * N + eigen];\n norm2 += x * x;\n }\n const float inverse = rsqrtf(norm2);\n for (int row = 0; row < N; ++row)\n vectors[mb + (long long)row * N + eigen] *= inverse;\n}\n';_N176_CHOLESKYQR_SOURCE='\n#include <cuda_runtime.h>\n\nextern "C" __global__\nvoid potrf176_block16_ridge(\n const float* __restrict__ gram,\n float* __restrict__ lower,\n int batch,\n float ridge)\n{\n constexpr int N = 176;\n constexpr int NN = N * N;\n constexpr int PANEL = 16;\n const int matrix_id = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n if (matrix_id >= batch) return;\n\n extern __shared__ float storage[];\n float* factor = storage;\n float* diagonal_mean = storage + NN;\n const float* source = gram + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += (int)blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n factor[index] = row >= column ? source[index] : 0.0f;\n }\n if (tid == 0) {\n float total = 0.0f;\n #pragma unroll\n for (int axis = 0; axis < N; ++axis)\n total += source[axis * N + axis];\n *diagonal_mean = total * (1.0f / (float)N);\n }\n __syncthreads();\n const float shift = ridge * (*diagonal_mean);\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = min(panel + PANEL, N);\n if (tid == 0) {\n #pragma unroll\n for (int column_offset = 0; column_offset < PANEL; ++column_offset) {\n const int column = panel + column_offset;\n if (column >= N) break;\n float diagonal = factor[column * N + column] + shift;\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n const float value = factor[column * N + k];\n diagonal = fmaf(-value, value, diagonal);\n }\n factor[column * N + column] =\n sqrtf(fmaxf(diagonal, 1.0e-30f));\n for (int row = column + 1; row < end; ++row) {\n float value = factor[row * N + column];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[row * N + column] =\n value / factor[column * N + column];\n }\n }\n }\n __syncthreads();\n\n for (int row = end + tid; row < N; row += (int)blockDim.x) {\n #pragma unroll\n for (int column_offset = 0; column_offset < PANEL; ++column_offset) {\n const int column = panel + column_offset;\n if (column >= N) break;\n float value = factor[row * N + column];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[row * N + column] =\n value / factor[column * N + column];\n }\n }\n __syncthreads();\n\n for (int index = tid; index < NN; index += (int)blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n if (row >= end && column >= end && row >= column) {\n float value = factor[index];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= N) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[index] = value;\n }\n }\n __syncthreads();\n }\n\n float* destination = lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += (int)blockDim.x)\n destination[index] = factor[index];\n}\n\nextern "C" __global__\nvoid right_trsm176_block16_rows32(\n const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch,\n int rows)\n{\n constexpr int N = 176;\n constexpr int NN = N * N;\n constexpr int ROWS = 32;\n constexpr int PANEL = 16;\n const int matrix_id = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n const int row_base = (int)blockIdx.y * ROWS;\n if (matrix_id >= batch) return;\n\n extern __shared__ float rhs[];\n const float* source = matrix + (long long)matrix_id * rows * N;\n const float* factor = lower + (long long)matrix_id * NN;\n float* destination = output + (long long)matrix_id * rows * N;\n for (int index = tid; index < N * ROWS; index += (int)blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n rhs[index] = row < rows\n ? source[(long long)row * N + column] : 0.0f;\n }\n __syncthreads();\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = min(panel + PANEL, N);\n if (tid < ROWS) {\n const int local_row = tid;\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n if (column >= N) break;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[column * N + k],\n rhs[k * ROWS + local_row], value);\n }\n rhs[column * ROWS + local_row] =\n value / factor[column * N + column];\n }\n }\n __syncthreads();\n const int remaining = (N - end) * ROWS;\n for (int index = tid; index < remaining; index += (int)blockDim.x) {\n const int column = end + index / ROWS;\n const int local_row = index % ROWS;\n float value = rhs[column * ROWS + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= N) break;\n value = fmaf(-factor[column * N + k],\n rhs[k * ROWS + local_row], value);\n }\n rhs[column * ROWS + local_row] = value;\n }\n __syncthreads();\n }\n\n for (int index = tid; index < N * ROWS; index += (int)blockDim.x) {\n const int column = index / ROWS;\n const int local_row = index - column * ROWS;\n const int row = row_base + local_row;\n if (row < rows)\n destination[(long long)row * N + column] = rhs[index];\n }\n}\n';_N176_SHARED_BYTES=80*1024;_N96_SHARED_BYTES=44*1024;_N128_SHARED_BYTES=80*1024
@memo(maxsize=None)
def _n176_choleskyqr_kernel(name:str):
source=_N176_CHOLESKYQR_SOURCE
if name=='right_trsm176_block16_rows16':source=source.replace('void right_trsm176_block16_rows32(','void right_trsm176_block16_rows16(').replace('constexpr int ROWS = 32;','constexpr int ROWS = 16;')
source=_fast_only_cuda_kernel(source,name);image=_fast_nvrtc_compile(source,name);return CUDAKernel(image,name)
def _paired_packed_update_source(source:str,*,wide:bool):
if wide:scalar=' constexpr int UPDATE_ROWS = 32;\n constexpr int UPDATE_COLUMNS = 16;\n const int tile_row = tid / UPDATE_COLUMNS;\n const int tile_column = tid - tile_row * UPDATE_COLUMNS;\n for (int row_base = 0; row_base < length; row_base += UPDATE_ROWS) {\n for (int column_base = 0; column_base < row_base + UPDATE_ROWS;\n column_base += UPDATE_COLUMNS) {\n const int i = row_base + tile_row;\n const int j = column_base + tile_column;\n if (i < length && j < length && j <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n v[i] * w[j] + w[i] * v[j];\n }\n }';paired=' constexpr int UPDATE_ROWS = 32;\n constexpr int UPDATE_PAIRS = 16;\n const int tile_row = tid / UPDATE_PAIRS;\n const int tile_pair = tid - tile_row * UPDATE_PAIRS;\n for (int row_base = 0; row_base < length; row_base += UPDATE_ROWS) {\n for (int column_base = 0; column_base < row_base + UPDATE_ROWS;\n column_base += 2 * UPDATE_PAIRS) {\n const int i = row_base + tile_row;\n const int j = column_base + 2 * tile_pair;\n if (i < length && j < length && j <= i) {\n const float vi = v[i];\n const float wi = w[i];\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n vi * w[j] + wi * v[j];\n if (j + 1 < length && j + 1 <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j + 1] -=\n vi * w[j + 1] + wi * v[j + 1];\n }\n }\n }'
else:scalar=' constexpr int UPDATE_TILE = 16;\n const int tile_row = tid >> 4;\n const int tile_column = tid & 15;\n for (int row_base = 0; row_base < length; row_base += UPDATE_TILE) {\n for (int column_base = 0; column_base <= row_base;\n column_base += UPDATE_TILE) {\n const int i = row_base + tile_row;\n const int j = column_base + tile_column;\n if (i < length && j < length && j <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n v[i] * w[j] + w[i] * v[j];\n }\n }';paired=' constexpr int UPDATE_ROWS = 16;\n constexpr int UPDATE_PAIRS = 16;\n const int tile_row = tid >> 4;\n const int tile_pair = tid & 15;\n for (int row_base = 0; row_base < length; row_base += UPDATE_ROWS) {\n for (int column_base = 0; column_base <= row_base;\n column_base += 2 * UPDATE_PAIRS) {\n const int i = row_base + tile_row;\n const int j = column_base + 2 * tile_pair;\n if (i < length && j < length && j <= i) {\n const float vi = v[i];\n const float wi = w[i];\n a[(begin + i) * (begin + i + 1) / 2 + begin + j] -=\n vi * w[j] + wi * v[j];\n if (j + 1 < length && j + 1 <= i)\n a[(begin + i) * (begin + i + 1) / 2 + begin + j + 1] -=\n vi * w[j + 1] + wi * v[j + 1];\n }\n }\n }'
if scalar not in source:raise RuntimeError('packed update template changed')
return source.replace(scalar,paired,1)
_E148_SOLVE176_NAME='e148_solve176_cluster2';_E148_N176_TRI=176*177//2;_E148_SOLVE176_SHARED_BYTES=(_E148_N176_TRI+5*176+16)*4;_E185_N176_BLOCK4_REDUCE_NAME='e185_reduce176_block4';_E185_N176_BLOCK4_SHARED_BYTES=(176*180+2*4*176+176+32+2*4)*4;_E185_N176_BLOCK4_REDUCE_SOURCE='\n#include <cuda_runtime.h>\nconstexpr int N=176,TRI=N*(N+1)/2,B=4;\nextern "C" __global__ __launch_bounds__(768,1)\nvoid e185_reduce176_block4(const float*__restrict__ input,\n float*__restrict__ saved,float*__restrict__ dout,float*__restrict__ eout,\n unsigned long long*__restrict__ timers){\n int mid=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;\n extern __shared__ float sm[];\n float*a=sm,*V=a+TRI,*W=V+B*N,*xvec=W+B*N,*scratch=xvec+N;\n float*coeff_a=scratch+32,*coeff_b=coeff_a+B;\n const float*src=input+(long long)mid*N*N;\n float*sv=saved+(long long)mid*N*N;\n for(int row=0;row<N;++row)if(tid<=row)\n a[row*(row+1)/2+tid]=src[row*N+tid];\n __syncthreads();\n if(tid==0)timers[mid*5]=clock64();\n for(int panel=0;panel<N-2;panel+=B){\n int count=min(B,N-2-panel);\n #pragma unroll\n for(int p=0;p<B;++p){\n if(p>=count)break;\n int col=panel+p,begin=col+1,length=N-begin;\n float local=0.f;\n for(int i=tid;i<length;i+=blockDim.x){\n int row=begin+i;\n float value=a[row*(row+1)/2+col];\n #pragma unroll\n for(int h=0;h<B;++h)if(h<p)\n value-=V[h*N+row]*W[h*N+col]+W[h*N+row]*V[h*N+col];\n xvec[i]=value;\n local=fmaf(value,value,local);\n }\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n local+=__shfl_down_sync(0xffffffff,local,sh);\n if(lane==0)scratch[warp]=local;\n __syncthreads();\n float total=tid<24?scratch[tid]:0.f;\n if(tid<32){\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n total+=__shfl_down_sync(0xffffffff,total,sh);\n if(tid==0)scratch[0]=total;\n }\n __syncthreads();\n float norm=sqrtf(scratch[0]),first=xvec[0];\n float beta=-copysignf(norm,first);\n float rn=sqrtf(fmaxf(2.f*norm*(norm+fabsf(first)),0.f));\n float inv=rn>0.f?1.f/rn:0.f;\n if(tid==0){\n eout[mid*N+col]=beta;\n float dv=a[col*(col+1)/2+col];\n #pragma unroll\n for(int h=0;h<B;++h)if(h<p)\n dv-=2.f*V[h*N+col]*W[h*N+col];\n dout[mid*N+col]=dv;\n }\n for(int i=tid;i<length;i+=blockDim.x){\n float value=xvec[i];\n if(i==0)value-=beta;\n value*=inv;\n V[p*N+begin+i]=value;\n sv[col*N+begin+i]=value;\n }\n __syncthreads();\n if(warp<p){\n float da=0.f,db=0.f;\n for(int i=lane;i<length;i+=32){\n int row=begin+i;\n da=fmaf(W[warp*N+row],V[p*N+row],da);\n db=fmaf(V[warp*N+row],V[p*N+row],db);\n }\n #pragma unroll\n for(int sh=16;sh;sh>>=1){\n da+=__shfl_down_sync(0xffffffff,da,sh);\n db+=__shfl_down_sync(0xffffffff,db,sh);\n }\n if(lane==0){coeff_a[warp]=da;coeff_b[warp]=db;}\n }\n __syncthreads();\n constexpr int G=4;\n int mlane=tid&(G-1),row=tid/G;\n float product=0.f;\n if(row<length){\n int gr=begin+row;\n for(int j=mlane;j<length;j+=G){\n int gc=begin+j,rr=gr>=gc?gr:gc,cc=gr>=gc?gc:gr;\n product=fmaf(a[rr*(rr+1)/2+cc],V[p*N+gc],product);\n }\n }\n #pragma unroll\n for(int sh=G/2;sh;sh>>=1)\n product+=__shfl_down_sync(0xffffffff,product,sh);\n float lp=0.f;\n if(mlane==0&&row<length){\n int gr=begin+row;\n #pragma unroll\n for(int h=0;h<B;++h)if(h<p)\n product-=V[h*N+gr]*coeff_a[h]+W[h*N+gr]*coeff_b[h];\n W[p*N+gr]=product;\n lp=V[p*N+gr]*product;\n }\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n lp+=__shfl_down_sync(0xffffffff,lp,sh);\n if(lane==0)scratch[warp]=lp;\n __syncthreads();\n total=tid<24?scratch[tid]:0.f;\n if(tid<32){\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n total+=__shfl_down_sync(0xffffffff,total,sh);\n if(tid==0)scratch[0]=total;\n }\n __syncthreads();\n float projection=scratch[0];\n for(int i=tid;i<length;i+=blockDim.x){\n int gr=begin+i;\n W[p*N+gr]=2.f*(W[p*N+gr]-projection*V[p*N+gr]);\n }\n __syncthreads();\n }\n int begin=panel+count,length=N-begin;\n constexpr int PAIRS=16,ROWS=48;\n int tr=tid/PAIRS,tp=tid-tr*PAIRS;\n for(int rb=0;rb<length;rb+=ROWS)\n for(int cb=0;cb<rb+ROWS;cb+=2*PAIRS){\n int i=rb+tr,j=cb+2*tp;\n if(i<length&&j<length&&j<=i){\n int gr=begin+i,gc=begin+j;\n int packed=gr*(gr+1)/2+gc;\n float value0=a[packed];\n bool paired=j+1<length&&j+1<=i;\n float value1=paired?a[packed+1]:0.f;\n #pragma unroll\n for(int h=0;h<B;++h)if(h<count){\n float vi=V[h*N+gr],wi=W[h*N+gr];\n value0-=vi*W[h*N+gc]+wi*V[h*N+gc];\n if(paired)value1-=vi*W[h*N+gc+1]+wi*V[h*N+gc+1];\n }\n a[packed]=value0;\n if(paired)a[packed+1]=value1;\n }\n }\n __syncthreads();\n }\n if(tid==0){\n eout[mid*N+N-2]=a[(N-1)*N/2+N-2];\n eout[mid*N+N-1]=0.f;\n dout[mid*N+N-2]=a[(N-2)*(N-1)/2+N-2];\n dout[mid*N+N-1]=a[(N-1)*N/2+N-1];\n timers[mid*5+1]=clock64();\n }\n}\n'
def _e185_reduce176_block4_source():
source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=176,STRIDE=180,TRI=N*STRIDE,B=4;',1).replace(' for(int row=0;row<N;++row)if(tid<=row)\n a[row*(row+1)/2+tid]=src[row*N+tid];',' for(int index=tid;index<N*N;index+=blockDim.x){\n int row=index/N,col=index-row*N;\n if(col<=row)a[row*STRIDE+col]=src[index];\n }',1)
for(old,new)in(('a[row*(row+1)/2+col]','a[row*STRIDE+col]'),('a[col*(col+1)/2+col]','a[col*STRIDE+col]'),('a[rr*(rr+1)/2+cc]','a[rr*STRIDE+cc]'),('a[packed]','a[gr*STRIDE+gc]'),('a[packed+1]','a[gr*STRIDE+gc+1]'),('a[(N-1)*N/2+N-2]','a[(N-1)*STRIDE+N-2]'),('a[(N-2)*(N-1)/2+N-2]','a[(N-2)*STRIDE+N-2]'),('a[(N-1)*N/2+N-1]','a[(N-1)*STRIDE+N-1]')):source=source.replace(old,new)
source=source.replace(' int packed=gr*(gr+1)/2+gc;\n','');return source
@memo(maxsize=1)
def _e185_reduce176_block4_kernel():return CUDAKernel(_fast_nvrtc_compile(_e185_reduce176_block4_source(),_E185_N176_BLOCK4_REDUCE_NAME),_E185_N176_BLOCK4_REDUCE_NAME)
_E1630_REDUCE176_NAME='e1630_reduce176_fused_conditional_scale'
def _e2775_chain_reduce176_source(source:str):"Keep the current panel reflector's norm/projection on warp zero.";norm_begin=source.index(' float local=0.f;\n for(int i=tid;i<length;i+=blockDim.x){');norm_end_marker=' __syncthreads();\n if(warp<p){';norm_end=source.index(norm_end_marker,norm_begin);warp_norm=' if(warp==0){\n float local=0.f;\n for(int i=lane;i<length;i+=32){\n int row=begin+i;\n float value=a[row*STRIDE+col];\n #pragma unroll\n for(int h=0;h<B;++h)if(h<p)\n value-=V[h*N+row]*W[h*N+col]+W[h*N+row]*V[h*N+col];\n xvec[i]=value;\n local=fmaf(value,value,local);\n }\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n local+=__shfl_down_sync(0xffffffff,local,sh);\n float total=__shfl_sync(0xffffffff,local,0);\n float norm=sqrtf(total),first=xvec[0];\n float beta=-copysignf(norm,first);\n float rn=sqrtf(fmaxf(2.f*norm*(norm+fabsf(first)),0.f));\n float inv=rn>0.f?1.f/rn:0.f;\n if(lane==0){\n eout[mid*N+col]=beta;\n float dv=a[col*STRIDE+col];\n #pragma unroll\n for(int h=0;h<B;++h)if(h<p)\n dv-=2.f*V[h*N+col]*W[h*N+col];\n dout[mid*N+col]=dv;\n }\n for(int i=lane;i<length;i+=32){\n float value=xvec[i];\n if(i==0)value-=beta;\n value*=inv;\n V[p*N+begin+i]=value;\n sv[col*N+begin+i]=value;\n }\n }\n __syncthreads();';source=source[:norm_begin]+warp_norm+source[norm_end+len(' __syncthreads();'):];projection_begin=source.index(' #pragma unroll\n for(int sh=16;sh;sh>>=1)\n lp+=__shfl_down_sync(0xffffffff,lp,sh);');projection_end=source.index(' }\n int begin=panel',projection_begin);chain_projection=' #pragma unroll\n for(int sh=16;sh;sh>>=1)\n lp+=__shfl_down_sync(0xffffffff,lp,sh);\n if(lane==0)scratch[warp]=lp;\n __syncthreads();\n if(warp==0){\n float projection=0.f;\n #pragma unroll\n for(int item=0;item<24;++item)projection+=scratch[item];\n for(int i=lane;i<length;i+=32){\n int gr=begin+i;\n W[p*N+gr]=2.f*(W[p*N+gr]-projection*V[p*N+gr]);\n }\n }\n // Warp zero immediately consumes W for the next reflector. The full CTA\n // only rendezvous at the panel boundary before the trailing update.\n if(p==count-1)__syncthreads();\n';source=source[:projection_begin]+chain_projection+source[projection_end:];source=source.replace(' __syncthreads();\n constexpr int G=4;',' constexpr int G=4;',1);source=source.replace(' float lp=0.f;\n if(mlane==0&&row<length){',' __syncthreads();\n float lp=0.f;\n if(mlane==0&&row<length){',1);return source
def _e1630_reduce176_source():
source=_e185_reduce176_block4_source().replace(_E185_N176_BLOCK4_REDUCE_NAME,_E1630_REDUCE176_NAME,1);old_signature=' float*__restrict__ saved,float*__restrict__ dout,float*__restrict__ eout,\n unsigned long long*__restrict__ timers){';new_signature=' float*__restrict__ saved,float*__restrict__ dout,float*__restrict__ eout,\n float*__restrict__ scales,unsigned long long*__restrict__ timers){'
if source.count(old_signature)!=1:raise RuntimeError('E1630 reducer signature anchor changed')
source=source.replace(old_signature,new_signature,1);old_load=' for(int index=tid;index<N*N;index+=blockDim.x){\n int row=index/N,col=index-row*N;\n if(col<=row)a[row*STRIDE+col]=src[index];\n }\n __syncthreads();';new_load=' float local_maximum=0.f;\n for(int index=tid;index<N*N;index+=blockDim.x){\n int row=index/N,col=index-row*N;\n float value=src[index];\n local_maximum=fmaxf(local_maximum,fabsf(value));\n if(col<=row)a[row*STRIDE+col]=value;\n }\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n local_maximum=fmaxf(local_maximum,\n __shfl_down_sync(0xffffffff,local_maximum,sh));\n if(lane==0)scratch[warp]=local_maximum;\n __syncthreads();\n float maximum=tid<24?scratch[tid]:0.f;\n if(tid<32){\n #pragma unroll\n for(int sh=16;sh;sh>>=1)\n maximum=fmaxf(maximum,__shfl_down_sync(0xffffffff,maximum,sh));\n if(tid==0){\n float scale=maximum>0.f&&(maximum<1.0e-4f||maximum>1.0e4f)\n ?maximum:1.f;\n scratch[0]=1.f/scale;\n scales[mid]=scale;\n }\n }\n __syncthreads();\n float inverse_scale=scratch[0];\n for(int index=tid;index<N*N;index+=blockDim.x){\n int row=index/N,col=index-row*N;\n if(col<=row)a[row*STRIDE+col]*=inverse_scale;\n }\n __syncthreads();'
if source.count(old_load)!=1:raise RuntimeError('E1630 reducer input anchor changed')
source=source.replace(old_load,new_load,1);source=_e2775_chain_reduce176_source(source);coefficient=' constexpr int G=4;'
if source.count(coefficient)!=1:raise RuntimeError('n176 safe-chain coefficient barrier changed')
source=source.replace(coefficient,' __syncthreads();\n constexpr int G=4;',1);late=' __syncthreads();\n float lp=0.f;'
if source.count(late)!=1:raise RuntimeError('n176 safe-chain late barrier changed')
return source.replace(late,' float lp=0.f;',1)
@memo(maxsize=1)
def _e1630_reduce176_kernel():return CUDAKernel(_fast_nvrtc_compile(_e1630_reduce176_source(),_E1630_REDUCE176_NAME),_E1630_REDUCE176_NAME)
def _parallelize_n176_adjacent_mgs(source:str):
old=' // Ordered MGS is local for all but the single rank boundary. The two CTAs\n // process their independent contiguous ranges concurrently.\n for (int local_eigen = 1; local_eigen < HALF; ++local_eigen) {\n const int eigen_column = first_row + local_eigen;\n const float gap = output_l[matrix_id * N + eigen_column] -\n output_l[matrix_id * N + eigen_column - 1];\n if (gap < 5.0e-3f) {\n if (tid == 0) {\n float dot = 0.0f;\n for (int row = 0; row < N; ++row)\n dot += q[row * HALF + local_eigen - 1] *\n q[row * HALF + local_eigen];\n partial[0] = dot;\n }\n __syncthreads();\n const float dot = partial[0];\n const float inverse_norm = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));\n for (int row = tid; row < N; row += blockDim.x)\n q[row * HALF + local_eigen] =\n (q[row * HALF + local_eigen] -\n dot * q[row * HALF + local_eigen - 1]) * inverse_norm;\n __syncthreads();\n }\n }\n';new=' // Adjacent close pairs form a path. Its two parity colors are\n // independent, so the active warps repair disjoint pairs concurrently\n // instead of serializing one scalar dot and two CTA barriers per accepted\n // gap.\n const int local_warp = tid >> 5;\n const int local_lane = tid & 31;\n #pragma unroll\n for (int parity = 0; parity < 2; ++parity) {\n for (int local_eigen = 1 + parity + 2 * local_warp;\n local_eigen < HALF; local_eigen += 16) {\n const int eigen_column = first_row + local_eigen;\n const float gap = output_l[matrix_id * N + eigen_column] -\n output_l[matrix_id * N + eigen_column - 1];\n if (gap < 5.0e-3f) {\n float dot = 0.0f;\n for (int row = local_lane; row < N; row += 32)\n dot += q[row * HALF + local_eigen - 1] *\n q[row * HALF + local_eigen];\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n dot += __shfl_down_sync(0xffffffff, dot, offset);\n dot = __shfl_sync(0xffffffff, dot, 0);\n const float inverse_norm = rsqrtf(\n fmaxf(1.0f - dot * dot, 1.0e-12f));\n for (int row = local_lane; row < N; row += 32)\n q[row * HALF + local_eigen] =\n (q[row * HALF + local_eigen] -\n dot * q[row * HALF + local_eigen - 1]) * inverse_norm;\n }\n }\n __syncthreads();\n }\n'
if source.count(old)!=1:raise RuntimeError('n176 adjacent MGS template changed')
return source.replace(old,new,1)
def _e2337_cluster_coarse_sturm(source:str,*,steps:int,probes:int,fine_steps:int,levels:int):
'Build one monotone Sturm table per two-CTA solve cluster.';cluster_barrier='cluster.sync();'if'cg::cluster_group cluster'in source else'asm volatile("barrier.cluster.arrive.aligned;" ::: "memory");\n asm volatile("barrier.cluster.wait.aligned;" ::: "memory");';old=f""" if (tid < HALF) {{
const int eigen_index = first_row + tid;
float lower = partial0[0];
float upper = partial0[1];
#pragma unroll
for (int step = 0; step < {steps}; ++step) {{""";coarse=f""" const float global_lower = partial0[0];
const float global_upper = partial0[1];
const float coarse_step =
(global_upper - global_lower) * (1.0f / {probes+1}.0f);
if (rank == 0 && tid < {probes}) {{
const float coarse_shift = global_lower + (tid + 1) * coarse_step;
float coarse_pivot = diagonal[0] - coarse_shift;
int coarse_count = coarse_pivot < 0.0f;
#pragma unroll 1
for (int row = 1; row < N; ++row) {{
if (fabsf(coarse_pivot) < 1.0e-12f)
coarse_pivot = copysignf(
1.0e-12f, coarse_pivot == 0.0f ? -1.0f : coarse_pivot);
coarse_pivot = diagonal[row] - coarse_shift
- off_diagonal[row - 1] * off_diagonal[row - 1] / coarse_pivot;
coarse_count += coarse_pivot < 0.0f;
}}
partial[2 + tid] = (float)coarse_count;
}}
{cluster_barrier}
if (tid < HALF) {{
const int eigen_index = first_row + tid;
int low_probe = -1;
int high_probe = {probes};
#pragma unroll
for (int level = 0; level < {levels}; ++level) {{
if (high_probe - low_probe > 1) {{
const int middle = (low_probe + high_probe) >> 1;
if ((int)partial0[2 + middle] <= eigen_index) low_probe = middle;
else high_probe = middle;
}}
}}
float lower = low_probe < 0 ? global_lower
: global_lower + (low_probe + 1) * coarse_step;
float upper = high_probe >= {probes} ? global_upper
: global_lower + (high_probe + 1) * coarse_step;
#pragma unroll
for (int step = 0; step < {fine_steps}; ++step) {{"""
if source.count(old)!=1:raise RuntimeError('E2337 cluster coarse Sturm anchor changed')
return source.replace(old,coarse,1)
_E2775_N176_DIRECT_STORE=" // Sturm bisection emits ascending eigenvalues. Store each CTA's complete\n // column tile directly; no cross-CTA permutation is required.\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int row = index / HALF;\n const int local_eigen = index - row * HALF;\n const int eigen_column = first_row + local_eigen;\n output_q[(matrix_id * N + row) * N + eigen_column] =\n q[row * HALF + local_eigen];\n }\n cluster.sync();\n if (rank == 0 && tid == 0) timers[matrix_id * 5 + 4] = clock64();\n}";_E2775_N176_CERTIFIED_STORE=" // Sturm bisection emits ascending eigenvalues. Store each CTA's complete\n // column tile directly; no cross-CTA permutation is required.\n for (int index = tid; index < HALF * N; index += blockDim.x) {\n const int row = index / HALF;\n const int local_eigen = index - row * HALF;\n const int eigen_column = first_row + local_eigen;\n output_q[(matrix_id * N + row) * N + eigen_column] =\n q[row * HALF + local_eigen];\n }\n\n if (tid == 0) *reinterpret_cast<int*>(partial) = 0;\n __syncthreads();\n if (tid < HALF) {\n float norm = 0.0f;\n float adjacent = 0.0f;\n #pragma unroll 1\n for (int row = 0; row < N; ++row) {\n const float x = q[row * HALF + tid];\n norm = fmaf(x, x, norm);\n if (tid + 1 < HALF)\n adjacent = fmaf(x, q[row * HALF + tid + 1], adjacent);\n }\n if (!isfinite(norm) || fabsf(norm - 1.0f) > 2.0e-3f\n || !isfinite(adjacent) || fabsf(adjacent) > 1.2e-3f)\n atomicExch(reinterpret_cast<int*>(partial), 1);\n }\n __syncthreads();\n __threadfence();\n cluster.sync();\n\n if (rank == 0 && tid == 0) {\n bool safe = *reinterpret_cast<int*>(partial) == 0;\n const int remote_unsafe = *reinterpret_cast<int*>(\n cluster.map_shared_rank(partial, 1));\n safe = safe && remote_unsafe == 0;\n int run = 0, maximum_run = 0;\n float previous = output_l[matrix_id * N];\n float spectrum_scale = fabsf(previous);\n safe = safe && isfinite(previous);\n #pragma unroll\n for (int index = 1; index < N; ++index) {\n const float current = output_l[matrix_id * N + index];\n const float gap = current - previous;\n safe = safe && isfinite(current) && gap >= 0.0f;\n spectrum_scale = fmaxf(spectrum_scale, fabsf(current));\n run = gap < 5.0e-3f ? run + 1 : 0;\n maximum_run = max(maximum_run, run);\n previous = current;\n }\n float boundary = 0.0f;\n #pragma unroll 1\n for (int row = 0; row < N; ++row)\n boundary = fmaf(output_q[(matrix_id * N + row) * N + HALF - 1],\n output_q[(matrix_id * N + row) * N + HALF], boundary);\n safe = safe && maximum_run <= 6\n && spectrum_scale >= 1.0e-8f && spectrum_scale <= 1.0e8f\n && isfinite(boundary) && fabsf(boundary) <= 1.2e-3f;\n flags[matrix_id] = safe ? 0 : 1;\n timers[matrix_id * 5 + 4] = clock64();\n }\n // Rank 1 must retain its DSM allocation until rank 0 consumed the bit.\n cluster.sync();\n}"
def _e2783_n176_central_refinement(source:str):
"Spend the 23rd Sturm bit only on the spectrum's failure-bearing middle.";old=' #pragma unroll\n for (int step = 0; step < 16; ++step) {\n const float midpoint = 0.5f * (lower + upper);\n float pivot = diagonal[0] - midpoint;\n int count = pivot < 0.0f;\n for (int i = 1; i < N; ++i) {\n if (fabsf(pivot) < 1.0e-12f)\n pivot = copysignf(1.0e-12f, pivot == 0.0f ? -1.0f : pivot);\n pivot = diagonal[i] - midpoint -\n off_diagonal[i - 1] * off_diagonal[i - 1] / pivot;\n count += pivot < 0.0f;\n }\n if (count <= eigen_index) lower = midpoint;\n else upper = midpoint;\n }';new=' #pragma unroll\n for (int step = 0; step < 17; ++step) {\n if (step < 16 || (eigen_index >= 48 && eigen_index < 128)) {\n const float midpoint = 0.5f * (lower + upper);\n float pivot = diagonal[0] - midpoint;\n int count = pivot < 0.0f;\n for (int i = 1; i < N; ++i) {\n if (fabsf(pivot) < 1.0e-12f)\n pivot = copysignf(1.0e-12f, pivot == 0.0f ? -1.0f : pivot);\n pivot = diagonal[i] - midpoint -\n off_diagonal[i - 1] * off_diagonal[i - 1] / pivot;\n count += pivot < 0.0f;\n }\n if (count <= eigen_index) lower = midpoint;\n else upper = midpoint;\n }\n }'
if source.count(old)!=1:raise RuntimeError('E2783 central Sturm refinement changed')
return source.replace(old,new,1)
@memo(maxsize=None)
def _e1762_solve176_pdl_kernel(sturm_steps:int=23,fused_certificate:bool=False):
suffix='_fused_certificate'if fused_certificate else'';name=f"e1762_solve176_step{sturm_steps}_pdl_wy_builder{suffix}";source=_e107_solve160_source().replace('constexpr int N = 160;','constexpr int N = 176;').replace('constexpr int HALF = 80;','constexpr int HALF = 88;').replace(_E107_SOLVE160_NAME,name).replace('__launch_bounds__(384, 4)','__launch_bounds__(256, 1)')
if source.count('step < 24')!=1:raise RuntimeError('E1762 Sturm schedule changed')
source=source.replace('step < 24',f"step < {sturm_steps}",1);source=_parallelize_n176_adjacent_mgs(source);anchor=' const int matrix_id = blockIdx.x >> 1;\n const int tid = threadIdx.x;';replacement=' const int matrix_id = blockIdx.x >> 1;\n const int tid = threadIdx.x;\n if (tid == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);'
if source.count(anchor)!=1:raise RuntimeError('E1762 solve entry changed')
source=source.replace(anchor,replacement,1);source=_e2337_cluster_coarse_sturm(source,steps=sturm_steps,probes=64,fine_steps=sturm_steps-6,levels=7)
if fused_certificate:
if sturm_steps!=22:raise ValueError('fused n176 certificate expects the 22+central schedule')
source=_e2783_n176_central_refinement(source);signature=' unsigned long long* __restrict__ timers) {';replacement=' unsigned long long* __restrict__ timers,\n int* __restrict__ flags) {'
if source.count(signature)!=1:raise RuntimeError('E2775 n176 solve signature changed')
source=source.replace(signature,replacement,1)
if source.count(_E2775_N176_DIRECT_STORE)!=1:raise RuntimeError('E2775 n176 solve store changed')
source=source.replace(_E2775_N176_DIRECT_STORE,_E2775_N176_CERTIFIED_STORE,1)
return CUDAKernel(_fast_nvrtc_compile(source,name),name)
_E1630_N176_JACOBI_NAME='e1630_n176_device_jacobi_s8_scaled';_E1630_N176_JACOBI_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(256, 1)\nvoid e1630_n176_device_jacobi_s8_scaled(\n const float* __restrict__ diagonal,\n const float* __restrict__ off,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ q_workspace,\n const float* __restrict__ scales,\n int batch) {\n constexpr int N = 176;\n constexpr int PAIRS = N / 2;\n constexpr int SWEEPS = 8;\n const int matrix = blockIdx.x;\n const int tid = threadIdx.x;\n if (matrix >= batch) return;\n __shared__ int unsafe;\n __shared__ int order[N];\n __shared__ int next_order[N];\n __shared__ int sorted[N];\n __shared__ float cosine[PAIRS];\n __shared__ float sine[PAIRS];\n extern __shared__ float a[];\n const long long vector_base = (long long)matrix * N * N;\n const long long value_base = (long long)matrix * N;\n const float matrix_scale = scales[matrix];\n\n if (tid == 0) {\n bool safe = true;\n int run = 0, maximum_run = 0;\n float previous = values[value_base];\n float spectrum_scale = fabsf(previous);\n safe = safe && isfinite(previous);\n #pragma unroll\n for (int index = 1; index < N; ++index) {\n const float current = values[value_base + index];\n const float gap = current - previous;\n safe = safe && isfinite(current) && gap >= 0.0f;\n spectrum_scale = fmaxf(spectrum_scale, fabsf(current));\n run = gap < 5.0e-3f ? run + 1 : 0;\n maximum_run = max(maximum_run, run);\n previous = current;\n }\n safe = safe && maximum_run <= 6\n && spectrum_scale >= 1.0e-8f && spectrum_scale <= 1.0e8f;\n unsafe = !safe;\n }\n __syncthreads();\n if (!unsafe) {\n if (tid < N) values[value_base + tid] *= matrix_scale;\n return;\n }\n\n for (int index = tid; index < N * N; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n float value = 0.0f;\n if (row == column) value = diagonal[value_base + row];\n else if (row + 1 == column) value = off[value_base + row];\n else if (column + 1 == row) value = off[value_base + column];\n a[index] = value;\n q_workspace[vector_base + index] = row == column ? 1.0f : 0.0f;\n }\n if (tid < N) order[tid] = tid;\n __syncthreads();\n\n #pragma unroll 1\n for (int sweep = 0; sweep < SWEEPS; ++sweep) {\n #pragma unroll 1\n for (int round = 0; round < N - 1; ++round) {\n if (tid < PAIRS) {\n const int p = order[tid];\n const int q = order[N - 1 - tid];\n const float app = a[p * N + p];\n const float aqq = a[q * N + q];\n const float apq = 0.5f * (a[p * N + q] + a[q * N + p]);\n float c = 1.0f, s = 0.0f;\n if (apq != 0.0f && isfinite(apq)) {\n const float tau = (aqq - app) / (2.0f * apq);\n const float t = copysignf(1.0f, tau)\n / (fabsf(tau) + sqrtf(1.0f + tau * tau));\n c = rsqrtf(1.0f + t * t);\n s = t * c;\n }\n cosine[tid] = c;\n sine[tid] = s;\n }\n __syncthreads();\n for (int index = tid; index < N * PAIRS; index += blockDim.x) {\n const int row = index / PAIRS;\n const int pair = index - row * PAIRS;\n const int p = order[pair];\n const int q = order[N - 1 - pair];\n const float c = cosine[pair], s = sine[pair];\n const float x = a[row * N + p], y = a[row * N + q];\n a[row * N + p] = fmaf(-s, y, c * x);\n a[row * N + q] = fmaf( s, x, c * y);\n }\n __syncthreads();\n for (int index = tid; index < N * PAIRS; index += blockDim.x) {\n const int column = index / PAIRS;\n const int pair = index - column * PAIRS;\n const int p = order[pair];\n const int q = order[N - 1 - pair];\n const float c = cosine[pair], s = sine[pair];\n float x = a[p * N + column], y = a[q * N + column];\n a[p * N + column] = fmaf(-s, y, c * x);\n a[q * N + column] = fmaf( s, x, c * y);\n x = q_workspace[vector_base + (long long)column * N + p];\n y = q_workspace[vector_base + (long long)column * N + q];\n q_workspace[vector_base + (long long)column * N + p]\n = fmaf(-s, y, c * x);\n q_workspace[vector_base + (long long)column * N + q]\n = fmaf( s, x, c * y);\n }\n __syncthreads();\n if (tid < N) {\n if (tid == 0) next_order[tid] = order[tid];\n else if (tid == 1) next_order[tid] = order[N - 1];\n else next_order[tid] = order[tid - 1];\n }\n __syncthreads();\n if (tid < N) order[tid] = next_order[tid];\n __syncthreads();\n }\n }\n\n if (tid == 0) {\n for (int index = 0; index < N; ++index) sorted[index] = index;\n for (int index = 1; index < N; ++index) {\n const int item = sorted[index];\n const float value = a[item * N + item];\n int position = index;\n while (position > 0\n && a[sorted[position - 1] * N + sorted[position - 1]] > value) {\n sorted[position] = sorted[position - 1];\n --position;\n }\n sorted[position] = item;\n }\n }\n __syncthreads();\n if (tid < N)\n values[value_base + tid]\n = a[sorted[tid] * N + sorted[tid]] * matrix_scale;\n for (int index = tid; index < N * N; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n vectors[vector_base + index]\n = q_workspace[vector_base + (long long)row * N + sorted[column]];\n }\n}\n'
@memo(maxsize=1)
def _e1630_n176_jacobi_kernel():
source=_E1630_N176_JACOBI_SOURCE;signature=' int batch) {';replacement=' int batch,\n const int* __restrict__ flags) {'
if source.count(signature)!=1:raise RuntimeError('E2775 Jacobi signature changed')
source=source.replace(signature,replacement,1);decision_begin=source.index(' if (tid == 0) {\n bool safe = true;');decision=' __syncthreads();\n if (!unsafe) {';decision_end=source.index(decision,decision_begin);repair_decision=' if (tid == 0) unsafe = flags[matrix];\n __syncthreads();\n if (!unsafe) {';source=source[:decision_begin]+repair_decision+source[decision_end+len(decision):];return CUDAKernel(_fast_nvrtc_compile(source,_E1630_N176_JACOBI_NAME),_E1630_N176_JACOBI_NAME)
@torch.no_grad()
def _e185_guarded_n176_eigh(data:torch.Tensor):batch=data.shape[0];vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device,dtype=torch.float32);saved=torch.empty_like(data);diagonal=torch.empty_like(values);off_diagonal=torch.empty_like(values);scales=torch.empty((batch,),device=data.device);timers=torch.empty((batch,5),device=data.device,dtype=torch.int64);flags=torch.empty((batch,),device=data.device,dtype=torch.int32);panel_total=(176-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=data.device,dtype=torch.float32);_e1630_reduce176_kernel().launch((batch,1,1),(768,1,1),(data,saved,diagonal,off_diagonal,scales,timers),shared_mem=_E185_N176_BLOCK4_SHARED_BYTES);_e1762_solve176_pdl_kernel(22,True).launch((batch*2,1,1),(256,1,1),(saved,diagonal,off_diagonal,vectors,values,timers,flags),shared_mem=_E148_SOLVE176_SHARED_BYTES);_householder_compact_t16_kernel[batch,panel_total](saved,triangular,176,panel_total,panel_width=32,block_rows=32,num_warps=4,num_stages=2,launch_pdl=True);workspace=torch.empty_like(vectors);_e1630_n176_jacobi_kernel().launch((batch,1,1),(256,1,1),(diagonal,off_diagonal,vectors,values,workspace,scales,batch,flags),shared_mem=176*176*4);vectors=_e1762_tensor_wy_apply_saved_(vectors,saved,triangular,block_cols=32);return vectors,values
@memo(maxsize=1)
def _small_n96_kernel():source=_paired_packed_update_source(_N96_PACKED_TILED_SOURCE,wide=False).replace('step < 27','step < 18');wy_begin=source.index(' // Apply normalized reflectors in reverse blocks of WY_BLOCK.');direct_store=' // Sturm bisection already emits ascending columns.\n for (int index = tid; index < N * N; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n output_q[(matrix_id * N + row) * N + column] = q[row * N + column];\n }\n if (tid == 0) timers[matrix_id * 5 + 4] = clock64();\n}\n';source=source[:wy_begin]+direct_store;name='e114_tridiagonal96_tensor_wy';source=source.replace('tridiagonal96_packed_tiled_singlecta',name);image=_fast_nvrtc_compile(source,name);return CUDAKernel(image,name)
@torch.no_grad()
def _small_n96_eigh(data:torch.Tensor):vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device,dtype=torch.float32);saved_reflectors=torch.empty_like(data);timers=torch.empty((data.shape[0],5),device=data.device,dtype=torch.int64);_small_n96_kernel().launch((data.shape[0],1,1),(256,1,1),(data,vectors,values,saved_reflectors,timers,2e-07),shared_mem=_N96_SHARED_BYTES);return _tensor_wy_backtransform_saved_(vectors,saved_reflectors),values
_E204_RANKDEF_N96_REDUCE_NAME='e3038_rankdef_n96_reduce_b4_t128_g1';_E204_RANKDEF_N96_REDUCE_SHARED=(96*97//2+2*4*96+96+32+2*4)*4
@memo(maxsize=1)
def _e204_rankdef_n96_reduce_kernel():source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=96,TRI=N*(N+1)/2,B=4;').replace(_E185_N176_BLOCK4_REDUCE_NAME,_E204_RANKDEF_N96_REDUCE_NAME).replace('__launch_bounds__(768,1)','__launch_bounds__(128,1)').replace('total=tid<24?scratch[tid]:0.f','total=tid<4?scratch[tid]:0.f').replace('constexpr int G=4;','constexpr int G=1;').replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=4,ROWS=32;');return CUDAKernel(_fast_nvrtc_compile(source,_E204_RANKDEF_N96_REDUCE_NAME),_E204_RANKDEF_N96_REDUCE_NAME)
@torch.no_grad()
def _e204_rankdef_n96_eigh(data:torch.Tensor):batch=data.shape[0];saved=torch.empty_like(data);diagonal=torch.empty(data.shape[:-1],device=data.device);off_diagonal=torch.empty_like(diagonal);vectors=torch.empty_like(data);values=torch.empty_like(diagonal);workspace=torch.empty_like(data);timers=torch.empty((batch,5),device=data.device,dtype=torch.int64);_e204_rankdef_n96_reduce_kernel().launch((batch,1,1),(128,1,1),(data,saved,diagonal,off_diagonal,timers),shared_mem=_E204_RANKDEF_N96_REDUCE_SHARED);_e1778_n96_solve_pdl_kernel().launch((batch,1,1),(96,1,1),(diagonal,off_diagonal,vectors,values,workspace));panel_total=(96-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=data.device,dtype=data.dtype);_householder_compact_t16_kernel[batch,panel_total](saved,triangular,96,panel_total,panel_width=32,block_rows=32,num_warps=1,num_stages=2,launch_pdl=True);_e1284_n96_mgs_kernel().launch((batch,1,1),(384,1,1),(vectors,values,batch),shared_mem=_E1284_N96_MGS_SHARED_BYTES);return _e1762_tensor_wy_apply_saved_(vectors,saved,triangular,block_cols=128),values
@memo(maxsize=1)
def _small_n128_kernel():
source=_N128_PACKED_TILED_SOURCE.replace('step < 27','step < 18');scalar_matvec=' if (tid < length) {\n const int row = begin + tid;\n float product = 0.0f;\n // Only the lower active triangle is current. Reflect it on load.\n for (int j = 0; j <= tid; ++j)\n product += a[row * (row + 1) / 2 + begin + j] * v[j];\n for (int j = tid + 1; j < length; ++j)\n product += a[(begin + j) * (begin + j + 1) / 2 + row] * v[j];\n w[tid] = product;\n }\n __syncthreads();';pair_matvec=' // Adjacent lanes share one row of A*v. This halves\n // the longest FMA dependency chain while retaining all 256 threads.\n const int matvec_row = tid >> 1;\n const int matvec_lane = tid & 1;\n float product = 0.0f;\n if (matvec_row < length) {\n const int row = begin + matvec_row;\n for (int j = matvec_lane; j <= matvec_row; j += 2)\n product += a[row * (row + 1) / 2 + begin + j] * v[j];\n for (int j = matvec_row + 1 + matvec_lane; j < length; j += 2)\n product += a[(begin + j) * (begin + j + 1) / 2 + row] * v[j];\n }\n product += __shfl_xor_sync(0xffffffff, product, 1);\n if (matvec_row < length && matvec_lane == 0)\n w[matvec_row] = product;\n __syncthreads();'
if source.count(scalar_matvec)!=1:raise RuntimeError('n128 packed matvec template changed')
source=source.replace(scalar_matvec,pair_matvec);wy_begin=source.index(' // Apply normalized reflectors in reverse blocks of WY_BLOCK.');direct_store=' // Sturm bisection already emits ascending columns.\n for (int index = tid; index < N * N; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n output_q[(matrix_id * N + row) * N + column] = q[row * N + column];\n }\n if (tid == 0) timers[matrix_id * 5 + 4] = clock64();\n}\n';source=source[:wy_begin]+direct_store;name='e113_tridiagonal128_tensor_wy';source=source.replace('tridiagonal128_packed_tiled_singlecta',name);image=_fast_nvrtc_compile(source,name);return CUDAKernel(image,name)
@torch.no_grad()
def _small_n128_eigh(data:torch.Tensor):vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device,dtype=torch.float32);saved_reflectors=torch.empty_like(data);timers=torch.empty((data.shape[0],5),device=data.device,dtype=torch.int64);_small_n128_kernel().launch((data.shape[0],1,1),(256,1,1),(data,vectors,values,saved_reflectors,timers,2e-07),shared_mem=_N128_SHARED_BYTES);return _tensor_wy_backtransform_saved_(vectors,saved_reflectors),values
_E196_N128_REDUCE_NAME='e3036_lapack_n128_reduce_b4_t128';_E196_N128_REDUCE_SHARED_BYTES=(128*129//2+2*4*128+128+32+2*4)*4
@memo(maxsize=1)
def _e196_lapack_n128_reduce_kernel():
source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=128,TRI=N*(N+1)/2,B=4;').replace(_E185_N176_BLOCK4_REDUCE_NAME,_E196_N128_REDUCE_NAME).replace('__launch_bounds__(768,1)','__launch_bounds__(128,1)').replace('total=tid<24?scratch[tid]:0.f','total=tid<4?scratch[tid]:0.f').replace('constexpr int G=4;','constexpr int G=1;').replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=4,ROWS=32;');entry=' int mid=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;'
if source.count(entry)!=1:raise RuntimeError('n128 reducer entry template changed')
clock_start=' if(tid==0)timers[mid*5]=clock64();'
if source.count(clock_start)!=1:raise RuntimeError('n128 reducer start timer changed')
source=source.replace(clock_start,'',1);source=source.replace(entry,entry+'\n if(tid==0){\n'+' timers[mid]=0ULL;\n'+' __threadfence();\n'+' asm volatile("griddepcontrol.launch_dependents;":::);\n'+' }',1);tail=' timers[mid*5+1]=clock64();\n }\n}';tail_ready=' }\n __threadfence();\n __syncthreads();\n if(tid==0) timers[mid]=1ULL;\n}'
if source.count(tail)!=1:raise RuntimeError('n128 reducer tail template changed')
source=source.replace(tail,tail_ready,1);return CUDAKernel(_fast_nvrtc_compile(source,_E196_N128_REDUCE_NAME),_E196_N128_REDUCE_NAME)
_E1282_N128_SOLVE_NAME='e1282_n128_twisted18_mb12';_E1776_N128_SOLVE_PDL_NAME='e1776_n128_twisted18_pdl_wy';_E1282_N128_MGS_NAME='e1282_n128_adjacent_two_color_mgs'
def _e2333_small_leaf_coarse_sturm(source:str,n:int,levels:int):
source=source.replace(' __shared__ float bounds[2];\n',f" __shared__ float bounds[2];\n __shared__ int coarse_counts[{n}];\n",1);old_bisection=' float lower = bounds[0];\n float upper = bounds[1];\n #pragma unroll 1\n for (int step = 0; step < 18; ++step) {';coarse_bisection=f""" const float global_lower = bounds[0];
const float global_upper = bounds[1];
const float coarse_step =
(global_upper - global_lower) * (1.0f / {n+1}.0f);
const float coarse_shift =
global_lower + (eigen + 1) * coarse_step;
float coarse_pivot = diagonal[0] - coarse_shift;
int coarse_count = coarse_pivot < 0.0f;
#pragma unroll 1
for (int row = 1; row < N; ++row) {{
if (fabsf(coarse_pivot) < 1.0e-12f)
coarse_pivot = copysignf(
1.0e-12f, coarse_pivot == 0.0f ? -1.0f : coarse_pivot);
coarse_pivot = diagonal[row] - coarse_shift
- off[row - 1] * off[row - 1] / coarse_pivot;
coarse_count += coarse_pivot < 0.0f;
}}
coarse_counts[eigen] = coarse_count;
__syncthreads();
int low_probe = -1;
int high_probe = {n};
#pragma unroll
for (int level = 0; level < {levels}; ++level) {{
if (high_probe - low_probe > 1) {{
const int middle = (low_probe + high_probe) >> 1;
if (coarse_counts[middle] <= eigen) low_probe = middle;
else high_probe = middle;
}}
}}
float lower = low_probe < 0 ? global_lower
: global_lower + (low_probe + 1) * coarse_step;
float upper = high_probe >= {n} ? global_upper
: global_lower + (high_probe + 1) * coarse_step;
#pragma unroll 1
for (int step = 0; step < 11; ++step) {{"""
if source.count(old_bisection)!=1:raise RuntimeError('E2333 small-leaf coarse Sturm anchor changed')
return source.replace(old_bisection,coarse_bisection,1)
def _e1282_n128_solve_source(min_blocks:int=16):
source=_TRIDIAG_TWISTED512_SOURCE.replace('constexpr int N = 512;','constexpr int N = 128;',1).replace('__launch_bounds__(512, 1)',f"__launch_bounds__(128, {min_blocks})",1).replace('tridiag_twisted512',_E1282_N128_SOLVE_NAME,1).replace('for (int step = 0; step < 30; ++step)','for (int step = 0; step < 18; ++step)',1);signature=' const float* __restrict__ matrix,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace) {';replacement=' const float* __restrict__ diagonal_input,\n const float* __restrict__ off_input,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace,\n const unsigned long long* __restrict__ timers) {'
if source.count(signature)!=1:raise RuntimeError('n128 split solve signature changed')
source=source.replace(signature,replacement,1);header=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;\n __syncthreads();';new_header=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n if (eigen == 0) {\n volatile const unsigned long long* ready = timers + batch;\n while (*ready == 0ULL) __nanosleep(64);\n }\n __syncthreads();\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float off_squared[N];\n __shared__ float bounds[2];\n diagonal[eigen] = diagonal_input[(long long)batch * N + eigen];\n off[eigen] = off_input[(long long)batch * N + eigen];\n off_squared[eigen] = off[eigen] * off[eigen];\n __syncthreads();'
if source.count(header)!=1:raise RuntimeError('n128 split solve input changed')
source=source.replace(header,new_header,1)
for(old,new)in(('off[row - 1] * off[row - 1]','off_squared[row - 1]'),('off[0] * off[0]','off_squared[0]'),('off[row] * off[row]','off_squared[row]')):
if source.count(old)==0:raise RuntimeError(f"n128 off-square expression changed: {old}")
source=source.replace(old,new)
old_ldl=" // Forward and backward LDL diagonals, laid out [row,eigen] so a warp's\n // stores are coalesced despite every thread owning one recurrence.\n const long long wb = (long long)batch * N * N;\n float pivot = diagonal[0] - lambda;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_workspace[wb + eigen] = pivot;\n for (int row = 1; row < N; ++row) {\n const float previous = left_workspace[wb + (long long)(row - 1) * N + eigen];\n pivot = diagonal[row] - lambda\n - off_squared[row - 1] / previous;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n left_workspace[wb + (long long)row * N + eigen] = pivot;\n }\n pivot = diagonal[N - 1] - lambda;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n vectors[mb + (long long)(N - 1) * N + eigen] = pivot;\n for (int row = N - 2; row >= 0; --row) {\n const float next = vectors[mb + (long long)(row + 1) * N + eigen];\n pivot = diagonal[row] - lambda - off_squared[row] / next;\n if (fabsf(pivot) < 1.0e-10f) pivot = copysignf(1.0e-10f, pivot);\n vectors[mb + (long long)row * N + eigen] = pivot;\n }\n";new_ldl=' // Advance the independent forward/backward LDL chains together so\n // division latency is hidden and each preceding pivot stays in a register.\n const long long wb = (long long)batch * N * N;\n float left_pivot = diagonal[0] - lambda;\n float right_pivot = diagonal[N - 1] - lambda;\n if (fabsf(left_pivot) < 1.0e-10f)\n left_pivot = copysignf(1.0e-10f, left_pivot);\n if (fabsf(right_pivot) < 1.0e-10f)\n right_pivot = copysignf(1.0e-10f, right_pivot);\n left_workspace[wb + eigen] = left_pivot;\n vectors[mb + (long long)(N - 1) * N + eigen] = right_pivot;\n #pragma unroll 1\n for (int step = 1; step < N; ++step) {\n const int left_row = step;\n const int right_row = N - 1 - step;\n left_pivot = diagonal[left_row] - lambda\n - off_squared[left_row - 1] / left_pivot;\n right_pivot = diagonal[right_row] - lambda\n - off_squared[right_row] / right_pivot;\n if (fabsf(left_pivot) < 1.0e-10f)\n left_pivot = copysignf(1.0e-10f, left_pivot);\n if (fabsf(right_pivot) < 1.0e-10f)\n right_pivot = copysignf(1.0e-10f, right_pivot);\n left_workspace[wb + (long long)left_row * N + eigen] = left_pivot;\n vectors[mb + (long long)right_row * N + eigen] = right_pivot;\n }\n'
if source.count(old_ldl)!=1:raise RuntimeError('n128 forward/backward LDL block changed')
source=source.replace(old_ldl,new_ldl,1);gamma_scan=' for (int row = 1; row < N; ++row) {\n float gamma =';sparse_gamma_scan=' // Either parity is a valid twisted-factorization origin in exact\n // arithmetic. Sampling one parity halves the division-heavy stability\n // scan; the outer residual/orthogonality certificates retain ownership of\n // unusual recurrences and their existing repair paths.\n for (int row = 1; row < N; row += 2) {\n float gamma ='
if source.count(gamma_scan)!=1:raise RuntimeError('n128 gamma scan anchor changed')
return source.replace(gamma_scan,sparse_gamma_scan,1)
@memo(maxsize=2)
def _e1776_n128_solve_pdl_kernel(min_blocks:int=16):
name=f"{_E1776_N128_SOLVE_PDL_NAME}_mb{min_blocks}";source=_e2333_small_leaf_coarse_sturm(_e1282_n128_solve_source(min_blocks),128,8).replace(_E1282_N128_SOLVE_NAME,name,1);ready=' diagonal[eigen] = diagonal_input[(long long)batch * N + eigen];\n off[eigen] = off_input[(long long)batch * N + eigen];\n off_squared[eigen] = off[eigen] * off[eigen];\n __syncthreads();';signal=ready+'\n if (eigen == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);'
if source.count(ready)!=1:raise RuntimeError('n128 PDL post-ready template changed')
source=source.replace(ready,signal,1);return CUDAKernel(_fast_nvrtc_compile(source,name),name)
_E2984_N128_SOLVE_MGS_NAME='e2984_n128_twisted18_mb8_pdl_fused_mgs'
@memo(maxsize=1)
def _e2984_n128_solve_mgs_kernel():
source=_e2333_small_leaf_coarse_sturm(_e1282_n128_solve_source(8),128,8).replace(_E1282_N128_SOLVE_NAME,_E2984_N128_SOLVE_MGS_NAME,1);ready=' diagonal[eigen] = diagonal_input[(long long)batch * N + eigen];\n off[eigen] = off_input[(long long)batch * N + eigen];\n off_squared[eigen] = off[eigen] * off[eigen];\n __syncthreads();';signal=ready+'\n if (eigen == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);'
if source.count(ready)!=1:raise RuntimeError('n128 fused solve/MGS PDL anchor changed')
source=source.replace(ready,signal,1);tail=' for (int row = 0; row < N; ++row)\n vectors[mb + (long long)row * N + eigen] *= inverse;\n}';fused_tail=' for (int row = 0; row < N; ++row)\n vectors[mb + (long long)row * N + eigen] *= inverse;\n __syncthreads();\n const int warp = eigen >> 5;\n const int lane = eigen & 31;\n #pragma unroll\n for (int parity = 0; parity < 2; ++parity) {\n for (int column = 1 + parity + 2 * warp;\n column < N; column += 8) {\n const float gap = values[(long long)batch * N + column]\n - values[(long long)batch * N + column - 1];\n if (gap < 5.0e-3f) {\n float dot = 0.0f;\n #pragma unroll\n for (int row = lane; row < N; row += 32)\n dot = fmaf(\n vectors[mb + (long long)row * N + column - 1],\n vectors[mb + (long long)row * N + column], dot);\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n dot += __shfl_down_sync(0xffffffffu, dot, offset);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n const float repair_inverse = rsqrtf(\n fmaxf(1.0f - dot * dot, 1.0e-12f));\n #pragma unroll\n for (int row = lane; row < N; row += 32)\n vectors[mb + (long long)row * N + column] =\n (vectors[mb + (long long)row * N + column]\n - dot * vectors[mb + (long long)row * N + column - 1])\n * repair_inverse;\n }\n }\n __syncthreads();\n }\n}'
if source.count(tail)!=1:raise RuntimeError('n128 fused solve/MGS tail changed')
source=source.replace(tail,fused_tail,1);return CUDAKernel(_fast_nvrtc_compile(source,_E2984_N128_SOLVE_MGS_NAME),_E2984_N128_SOLVE_MGS_NAME)
_E1282_N128_MGS_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(256,3)\nvoid e1282_n128_adjacent_two_color_mgs(\n float*__restrict__ q,const float*__restrict__ values,int batch){\n constexpr int N=128,PITCH=129;\n int matrix=blockIdx.x;if(matrix>=batch)return;\n int tid=threadIdx.x,warp=tid>>5,lane=tid&31;\n float*global_q=q+(long long)matrix*N*N;\n const float*local_l=values+(long long)matrix*N;\n extern __shared__ float local_q[];\n for(int index=tid;index<N*N;index+=256){\n int row=index/N,column=index-row*N;\n local_q[row*PITCH+column]=global_q[index];\n }\n __syncthreads();\n #pragma unroll\n for(int parity=0;parity<2;++parity){\n for(int column=1+parity+2*warp;column<N;column+=16){\n float gap=local_l[column]-local_l[column-1];\n if(gap<5.0e-3f){\n float dot=0.f;\n #pragma unroll\n for(int row=lane;row<N;row+=32)\n dot=fmaf(local_q[row*PITCH+column-1],local_q[row*PITCH+column],dot);\n #pragma unroll\n for(int offset=16;offset>0;offset>>=1)\n dot+=__shfl_down_sync(0xffffffffu,dot,offset);\n dot=__shfl_sync(0xffffffffu,dot,0);\n float inverse=rsqrtf(fmaxf(1.f-dot*dot,1.0e-12f));\n #pragma unroll\n for(int row=lane;row<N;row+=32)\n local_q[row*PITCH+column]=(local_q[row*PITCH+column]\n -dot*local_q[row*PITCH+column-1])*inverse;\n }\n }\n __syncthreads();\n }\n for(int index=tid;index<N*N;index+=256){\n int row=index/N,column=index-row*N;\n global_q[index]=local_q[row*PITCH+column];\n }\n}\n';_E1282_N128_MGS_SHARED_BYTES=128*129*4;_E3039_N128_MGS_NAME='e3039_n128_mgs_t512'
@memo(maxsize=1)
def _e1282_n128_mgs_kernel():source=_E1282_N128_MGS_SOURCE.replace(_E1282_N128_MGS_NAME,_E3039_N128_MGS_NAME,1).replace('__launch_bounds__(256,3)','__launch_bounds__(512,3)',1).replace('index+=256','index+=512').replace('column+=16','column+=32',1);return CUDAKernel(_fast_nvrtc_compile(source,_E3039_N128_MGS_NAME),_E3039_N128_MGS_NAME)
_E1284_N96_SOLVE_NAME='e1284_n96_twisted18_mb12';_E1778_N96_SOLVE_PDL_NAME='e1778_n96_twisted18_pdl_wy';_E1284_N96_MGS_NAME='e1284_n96_adjacent_two_color_mgs';_E3043_N96_MGS_NAME='e3043_n96_mgs_t384'
def _e1284_n96_solve_source():
source=_TRIDIAG_TWISTED512_SOURCE.replace('constexpr int N = 512;','constexpr int N = 96;',1).replace('__launch_bounds__(512, 1)','__launch_bounds__(96, 12)',1).replace('tridiag_twisted512',_E1284_N96_SOLVE_NAME,1).replace('for (int step = 0; step < 30; ++step)','for (int step = 0; step < 18; ++step)',1);signature=' const float* __restrict__ matrix,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace) {';replacement=' const float* __restrict__ diagonal_input,\n const float* __restrict__ off_input,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace) {'
if source.count(signature)!=1:raise RuntimeError('n96 split solve signature changed')
source=source.replace(signature,replacement,1);header=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;\n __syncthreads();';new_header=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = diagonal_input[(long long)batch * N + eigen];\n off[eigen] = off_input[(long long)batch * N + eigen];\n __syncthreads();'
if source.count(header)!=1:raise RuntimeError('n96 split solve input changed')
return source.replace(header,new_header,1)
@memo(maxsize=1)
def _e1778_n96_solve_pdl_kernel():
source=_e2333_small_leaf_coarse_sturm(_e1284_n96_solve_source(),96,7).replace(_E1284_N96_SOLVE_NAME,_E1778_N96_SOLVE_PDL_NAME,1);anchor=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;';replacement=' const int batch = blockIdx.x;\n if (threadIdx.x == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);\n const int eigen = threadIdx.x;'
if source.count(anchor)!=1:raise RuntimeError('n96 PDL solve entry changed')
source=source.replace(anchor,replacement,1);return CUDAKernel(_fast_nvrtc_compile(source,_E1778_N96_SOLVE_PDL_NAME),_E1778_N96_SOLVE_PDL_NAME)
@memo(maxsize=1)
def _e1284_n96_mgs_kernel():source=_E1282_N128_MGS_SOURCE.replace(_E1282_N128_MGS_NAME,_E1284_N96_MGS_NAME,1).replace('__launch_bounds__(256,3)','__launch_bounds__(256,5)',1).replace('constexpr int N=128,PITCH=129;','constexpr int N=96,PITCH=97;',1).replace(_E1284_N96_MGS_NAME,_E3043_N96_MGS_NAME,1).replace('__launch_bounds__(256,5)','__launch_bounds__(384,5)',1).replace('index+=256','index+=384').replace('column+=16','column+=24',1);return CUDAKernel(_fast_nvrtc_compile(source,_E3043_N96_MGS_NAME),_E3043_N96_MGS_NAME)
_E1284_N96_MGS_SHARED_BYTES=96*97*4
@torch.no_grad()
def _e196_lapack_n128_leaf_eigh(data:torch.Tensor,*,fuse_mgs:bool=False,compact_t_tf32:bool=False):
batch=data.shape[0];vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device,dtype=torch.float32);saved_reflectors=torch.empty_like(data);diagonal=torch.empty_like(values);off_diagonal=torch.empty_like(values);workspace=torch.empty_like(data);timers=torch.empty((batch,),device=data.device,dtype=torch.int64);_e196_lapack_n128_reduce_kernel().launch((batch,1,1),(128,1,1),(data,saved_reflectors,diagonal,off_diagonal,timers),shared_mem=_E196_N128_REDUCE_SHARED_BYTES);solve_min_blocks=8 if batch<=640 else 16;solve_kernel=_e2984_n128_solve_mgs_kernel()if fuse_mgs and batch<=640 else _e1776_n128_solve_pdl_kernel(solve_min_blocks);solve_kernel.launch_pdl((batch,1,1),(128,1,1),(diagonal,off_diagonal,vectors,values,workspace,timers));panel_total=(128-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=data.device,dtype=data.dtype);_householder_compact_t16_kernel[batch,panel_total](saved_reflectors,triangular,128,panel_total,panel_width=32,block_rows=32,use_tf32=compact_t_tf32,num_warps=1,num_stages=2,launch_pdl=True)
if not(fuse_mgs and batch<=640):_e1282_n128_mgs_kernel().launch((batch,1,1),(512,1,1),(vectors,values,batch),shared_mem=_E1282_N128_MGS_SHARED_BYTES)
return _e1762_tensor_wy_apply_saved_(vectors,saved_reflectors,triangular,block_cols=128),values
_E187_MIXED_TRI24_SHARDS=16;_E187_MIXED_TRI24_EIGEN_PER_SHARD=512//_E187_MIXED_TRI24_SHARDS;_E187_MIXED_TRI24_NAME='e677_mixed_tridiag_twisted512_s24_split16';_E1782_MIXED_TRI24_SHARDS=8;_E1782_MIXED_TRI24_EIGEN_PER_SHARD=512//_E1782_MIXED_TRI24_SHARDS;_E1782_MIXED_TRI24_PDL_NAME='e1782_mixed_tridiag_twisted512_s24_split8_pdl_wy'
@memo(maxsize=1)
def _e187_mixed_tridiag24_kernel():
source=_TRIDIAG_TWISTED512_SOURCE.replace('for (int step = 0; step < 30; ++step)','for (int step = 0; step < 24; ++step)').replace('__launch_bounds__(512, 1)',f"__launch_bounds__({_E187_MIXED_TRI24_EIGEN_PER_SHARD}, 1)").replace('tridiag_twisted512',_E187_MIXED_TRI24_NAME);original=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;\n __syncthreads();\n\n if (eigen == 0) {';split=f""" constexpr int SHARDS = {_E187_MIXED_TRI24_SHARDS};
constexpr int EIGEN_PER_SHARD = {_E187_MIXED_TRI24_EIGEN_PER_SHARD};
const int batch = blockIdx.x / SHARDS;
const int shard = blockIdx.x - batch * SHARDS;
const int eigen = shard * EIGEN_PER_SHARD + threadIdx.x;
const long long mb = (long long)batch * N * N;
__shared__ float diagonal[N];
__shared__ float off[N];
__shared__ float bounds[2];
for (int item = threadIdx.x; item < N; item += blockDim.x) {{
diagonal[item] = matrix[mb + (long long)item * N + item];
off[item] = item + 1 < N
? matrix[mb + (long long)(item + 1) * N + item] : 0.0f;
}}
__syncthreads();
if (threadIdx.x == 0) {{"""
if source.count(original)!=1:raise RuntimeError('E677 n512 tridiagonal split source anchor changed')
source=source.replace(original,split,1);return CUDAKernel(_fast_nvrtc_compile(source,_E187_MIXED_TRI24_NAME),_E187_MIXED_TRI24_NAME)
@memo(maxsize=1)
def _e1782_mixed_tridiag24_pdl_kernel():
source=_TRIDIAG_TWISTED512_SOURCE.replace('for (int step = 0; step < 30; ++step)','for (int step = 0; step < 24; ++step)').replace('__launch_bounds__(512, 1)',f"__launch_bounds__({_E1782_MIXED_TRI24_EIGEN_PER_SHARD}, 1)").replace('tridiag_twisted512',_E1782_MIXED_TRI24_PDL_NAME);original=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;\n __syncthreads();\n\n if (eigen == 0) {';split=f''' constexpr int SHARDS = {_E1782_MIXED_TRI24_SHARDS};
constexpr int EIGEN_PER_SHARD = {_E1782_MIXED_TRI24_EIGEN_PER_SHARD};
const int batch = blockIdx.x / SHARDS;
const int shard = blockIdx.x - batch * SHARDS;
if (threadIdx.x == 0)
asm volatile("griddepcontrol.launch_dependents;" :::);
const int eigen = shard * EIGEN_PER_SHARD + threadIdx.x;
const long long mb = (long long)batch * N * N;
__shared__ float diagonal[N];
__shared__ float off[N];
__shared__ float bounds[2];
for (int item = threadIdx.x; item < N; item += blockDim.x) {{
diagonal[item] = matrix[mb + (long long)item * N + item];
off[item] = item + 1 < N
? matrix[mb + (long long)(item + 1) * N + item] : 0.0f;
}}
__syncthreads();
if (threadIdx.x == 0) {{'''
if source.count(original)!=1:raise RuntimeError('E1782 n512 tridiagonal split source changed')
source=source.replace(original,split,1);source=source.replace(' __shared__ float bounds[2];\n',' __shared__ float bounds[2];\n __shared__ int coarse_counts[EIGEN_PER_SHARD];\n',1);old_bisection=' float lower = bounds[0];\n float upper = bounds[1];\n #pragma unroll 1\n for (int step = 0; step < 24; ++step) {';coarse_bisection=' const float global_lower = bounds[0];\n const float global_upper = bounds[1];\n const float coarse_step =\n (global_upper - global_lower) * (1.0f / 65.0f);\n const float coarse_shift =\n global_lower + (threadIdx.x + 1) * coarse_step;\n float coarse_pivot = diagonal[0] - coarse_shift;\n int coarse_count = coarse_pivot < 0.0f;\n #pragma unroll 1\n for (int row = 1; row < N; ++row) {\n if (fabsf(coarse_pivot) < 1.0e-12f)\n coarse_pivot = copysignf(\n 1.0e-12f, coarse_pivot == 0.0f ? -1.0f : coarse_pivot);\n coarse_pivot = diagonal[row] - coarse_shift\n - off[row - 1] * off[row - 1] / coarse_pivot;\n coarse_count += coarse_pivot < 0.0f;\n }\n coarse_counts[threadIdx.x] = coarse_count;\n __syncthreads();\n\n int low_probe = -1;\n int high_probe = EIGEN_PER_SHARD;\n #pragma unroll\n for (int level = 0; level < 7; ++level) {\n if (high_probe - low_probe > 1) {\n const int middle = (low_probe + high_probe) >> 1;\n if (coarse_counts[middle] <= eigen) low_probe = middle;\n else high_probe = middle;\n }\n }\n float lower = low_probe < 0 ? global_lower\n : global_lower + (low_probe + 1) * coarse_step;\n float upper = high_probe >= EIGEN_PER_SHARD ? global_upper\n : global_lower + (high_probe + 1) * coarse_step;\n #pragma unroll 1\n for (int step = 0; step < 18; ++step) {'
if source.count(old_bisection)!=1:raise RuntimeError('E2323 n512 coarse Sturm anchor changed')
source=source.replace(old_bisection,coarse_bisection,1);return CUDAKernel(_fast_nvrtc_compile(source,_E1782_MIXED_TRI24_PDL_NAME),_E1782_MIXED_TRI24_PDL_NAME)
@triton.jit
def _dense_guard_column_stats_kernel(matrix,vectors,aq,values,residual_energy,residual_l1,matrix_l1,n:tl.constexpr,block_cols:tl.constexpr):batch=tl.program_id(0);group=tl.program_id(1);rows=tl.arange(0,n)[:,None];column_ids=group*block_cols+tl.arange(0,block_cols);columns=column_ids[None,:];column_mask=column_ids<n;mask=column_mask[None,:];offsets=batch*n*n+rows*n+columns;q=tl.load(vectors+offsets,mask=mask,other=.0).to(tl.float32);product=tl.load(aq+offsets,mask=mask,other=.0).to(tl.float32);rayleigh=tl.sum(q*product,axis=0);residual=product-q*rayleigh[None,:];column_offsets=batch*n+column_ids;tl.store(values+column_offsets,rayleigh,mask=column_mask);tl.store(residual_energy+column_offsets,tl.sum(residual*residual,axis=0),mask=column_mask);tl.store(residual_l1+column_offsets,tl.sum(tl.abs(residual),axis=0),mask=column_mask);original=tl.load(matrix+offsets,mask=mask,other=.0).to(tl.float32);tl.store(matrix_l1+column_offsets,tl.sum(tl.abs(original),axis=0),mask=column_mask)
@triton.jit
def _dense_guard_reduce_kernel(residual_l1,matrix_l1,risk,threshold:tl.constexpr,n:tl.constexpr):batch=tl.program_id(0);columns=tl.arange(0,n);residual=tl.load(residual_l1+batch*n+columns);scale=tl.sum(tl.load(matrix_l1+batch*n+columns),axis=0);score=tl.max(residual,axis=0)/tl.maximum(scale,1e-30);tl.store(risk+batch,score>threshold)
@torch.no_grad()
def _dense_guard_statistics(data:torch.Tensor,vectors:torch.Tensor,aq:torch.Tensor,threshold:float):batch,n,_=data.shape;values=torch.empty((batch,n),device=data.device,dtype=torch.float32);residual_energy=torch.empty_like(values);residual_l1=torch.empty_like(values);matrix_l1=torch.empty_like(values);block_cols=8;_dense_guard_column_stats_kernel[batch,triton.cdiv(n,block_cols)](data,vectors,aq,values,residual_energy,residual_l1,matrix_l1,n=n,block_cols=block_cols,num_warps=8,num_stages=1);risk=torch.empty((batch,),device=data.device,dtype=torch.bool);_dense_guard_reduce_kernel[batch,](residual_l1,matrix_l1,risk,threshold=threshold,n=n,num_warps=8,num_stages=1);return values,residual_energy,torch.nonzero(risk,as_tuple=False).flatten()
@triton.jit
def _small_dense_classifier_kernel(matrix,flags,n:tl.constexpr,block_n:tl.constexpr):
batch=tl.program_id(0);column=tl.arange(0,block_n);maximum=.0;minimum=float('inf');maximum_off_diagonal=.0
for sample in range(64):row=sample*(n-1)//63;value=tl.load(matrix+batch*n*n+row*n+column,mask=column<n,other=.0).to(tl.float32);energy=tl.sum(value*value,axis=0);diagonal_energy=tl.sum(tl.where(column==row,value*value,.0),axis=0);maximum=tl.maximum(maximum,energy);minimum=tl.minimum(minimum,energy);maximum_off_diagonal=tl.maximum(maximum_off_diagonal,energy-diagonal_energy)
dynamic=maximum/tl.maximum(minimum,1e-30);tl.store(flags+batch,(dynamic>32.)&(dynamic<256.)&(maximum_off_diagonal>1e-20*maximum))
@torch.no_grad()
def _is_small_dense_cond1(data:torch.Tensor):flags=torch.empty((data.shape[0],),device=data.device,dtype=torch.bool);_small_dense_classifier_kernel[data.shape[0],](data,flags,n=data.shape[-1],block_n=triton.next_power_of_2(data.shape[-1]),num_warps=8,num_stages=1);return bool(flags.all().item())
@triton.jit
def _large_dense_trace_frob_partials_kernel(matrix,partials,n:tl.constexpr,tiles:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);tile=tl.program_id(1);local=tl.arange(0,BLOCK);offsets=tile*BLOCK+local;total=n*n;mask=offsets<total;values=tl.load(matrix+batch*total+offsets,mask=mask,other=.0).to(tl.float32);local_rows=local//n;columns=local-local_rows*n;row0=tile*(BLOCK//n);row1=row0+1;energy0=tl.sum(tl.where(local_rows==0,values*values,.0),axis=0);energy1=tl.sum(tl.where(local_rows==1,values*values,.0),axis=0);diagonal_value0=tl.sum(tl.where((local_rows==0)&(columns==row0),values,.0),axis=0);diagonal_value1=tl.sum(tl.where((local_rows==1)&(columns==row1),values,.0),axis=0);sample0=(row0*63+n-2)//(n-1);sample1=(row1*63+n-2)//(n-1);sampled0=row0==sample0*(n-1)//63;sampled1=row1==sample1*(n-1)//63;sampled_maximum=tl.maximum(tl.where(sampled0,energy0,.0),tl.where(sampled1,energy1,.0));sampled_minimum=tl.minimum(tl.where(sampled0,energy0,float('inf')),tl.where(sampled1,energy1,float('inf')));sampled_off_diagonal=tl.maximum(tl.where(sampled0,energy0-diagonal_value0*diagonal_value0,.0),tl.where(sampled1,energy1-diagonal_value1*diagonal_value1,.0));base=(batch*tiles+tile)*5;tl.store(partials+base,energy0+energy1);tl.store(partials+base+1,diagonal_value0+diagonal_value1);tl.store(partials+base+2,sampled_maximum);tl.store(partials+base+3,sampled_minimum);tl.store(partials+base+4,sampled_off_diagonal)
@triton.jit
def _large_dense_cond1_finish_kernel(partials,output,batch:tl.constexpr,tiles:tl.constexpr,BLOCK:tl.constexpr):
tile=tl.arange(0,BLOCK);mask=tile<tiles;safe=1
for matrix in range(batch):base=(matrix*tiles+tile)*5;frobenius_squared=tl.sum(tl.load(partials+base,mask=mask,other=.0),axis=0);trace=tl.sum(tl.load(partials+base+1,mask=mask,other=.0),axis=0);maximum=tl.max(tl.load(partials+base+2,mask=mask,other=.0),axis=0);minimum=tl.min(tl.load(partials+base+3,mask=mask,other=float('inf')),axis=0);maximum_off_diagonal=tl.max(tl.load(partials+base+4,mask=mask,other=.0),axis=0);dynamic=maximum/tl.maximum(minimum,1e-30);sampled=(dynamic>32.)&(dynamic<256.)&(maximum_off_diagonal>1e-20*maximum);safe&=sampled&(tl.abs(trace)<.5*tl.sqrt(tl.maximum(frobenius_squared,1e-30)))
tl.store(output,safe)
@torch.no_grad()
def _is_large_dense_cond1(data:torch.Tensor):batch,n,_=data.shape;block=4096;tiles=triton.cdiv(n*n,block);partials=torch.empty((batch,tiles,5),device=data.device);_large_dense_trace_frob_partials_kernel[batch,tiles](data,partials,n=n,tiles=tiles,BLOCK=block,num_warps=8,num_stages=1);output=torch.empty((),device=data.device,dtype=torch.bool);_large_dense_cond1_finish_kernel[1,](partials,output,batch=batch,tiles=tiles,BLOCK=triton.next_power_of_2(tiles),num_warps=8,num_stages=1);return bool(output.item())
_E096_LANCZOS3_NAME='e096_n352_fused_lanczos3';_E096_LANCZOS3_SOURCE='\n#include <cuda_runtime.h>\n\nstatic __device__ __forceinline__ float e096_warp_sum(float value) {\n value += __shfl_down_sync(0xffffffffu, value, 16);\n value += __shfl_down_sync(0xffffffffu, value, 8);\n value += __shfl_down_sync(0xffffffffu, value, 4);\n value += __shfl_down_sync(0xffffffffu, value, 2);\n value += __shfl_down_sync(0xffffffffu, value, 1);\n return value;\n}\n\nextern "C" __global__ __launch_bounds__(512, 1)\nvoid e096_n352_fused_lanczos3(\n const float* __restrict__ matrix,\n float* __restrict__ median,\n float* __restrict__ lower,\n float* __restrict__ upper,\n bool* __restrict__ dense_flags) {\n constexpr int N = 352;\n constexpr int WARPS = 16;\n const int batch = blockIdx.x;\n const int tid = threadIdx.x;\n const int lane = tid & 31;\n const int warp = tid >> 5;\n const long long base = (long long)batch * N * N;\n __shared__ float previous[N];\n __shared__ float current[N];\n __shared__ float product[N];\n __shared__ float warp_partial[WARPS];\n __shared__ float scalar;\n __shared__ float diagonal[3];\n __shared__ float off[2];\n __shared__ float row_energy[N];\n __shared__ float row_off[N];\n __shared__ float warp_maximum[WARPS];\n __shared__ float warp_minimum[WARPS];\n __shared__ float warp_off[WARPS];\n\n for (int row = tid; row < N; row += blockDim.x) {\n previous[row] = 0.0f;\n current[row] = 0.05330017908890261f;\n }\n __syncthreads();\n float beta = 0.0f;\n\n #pragma unroll\n for (int step = 0; step < 3; ++step) {\n for (int row = warp; row < N; row += WARPS) {\n float value = 0.0f;\n float energy = 0.0f;\n float diagonal_energy = 0.0f;\n for (int column = lane; column < N; column += 32) {\n const float a = matrix[base + (long long)row * N + column];\n value = fmaf(a, current[column], value);\n if (step == 0) {\n energy = fmaf(a, a, energy);\n if (column == row) diagonal_energy = a * a;\n }\n }\n value = e096_warp_sum(value);\n if (step == 0) {\n energy = e096_warp_sum(energy);\n diagonal_energy = e096_warp_sum(diagonal_energy);\n }\n if (lane == 0) {\n product[row] = value - beta * previous[row];\n if (step == 0) {\n row_energy[row] = energy;\n row_off[row] = energy - diagonal_energy;\n }\n }\n }\n __syncthreads();\n\n // The old route launched a second kernel that sampled 64 rows for this\n // gate. The first Lanczos matvec already reads every matrix element, so\n // accumulate the same row-energy criterion from those resident values.\n if (step == 0) {\n float local_maximum = 0.0f;\n float local_minimum = 3.4e38f;\n float local_off = 0.0f;\n for (int row = tid; row < N; row += blockDim.x) {\n local_maximum = fmaxf(local_maximum, row_energy[row]);\n local_minimum = fminf(local_minimum, row_energy[row]);\n local_off = fmaxf(local_off, row_off[row]);\n }\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n local_maximum = fmaxf(local_maximum,\n __shfl_down_sync(0xffffffffu, local_maximum, delta));\n local_minimum = fminf(local_minimum,\n __shfl_down_sync(0xffffffffu, local_minimum, delta));\n local_off = fmaxf(local_off,\n __shfl_down_sync(0xffffffffu, local_off, delta));\n }\n if (lane == 0) {\n warp_maximum[warp] = local_maximum;\n warp_minimum[warp] = local_minimum;\n warp_off[warp] = local_off;\n }\n __syncthreads();\n if (warp == 0) {\n local_maximum = lane < WARPS ? warp_maximum[lane] : 0.0f;\n local_minimum = lane < WARPS ? warp_minimum[lane] : 3.4e38f;\n local_off = lane < WARPS ? warp_off[lane] : 0.0f;\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1) {\n local_maximum = fmaxf(local_maximum,\n __shfl_down_sync(0xffffffffu, local_maximum, delta));\n local_minimum = fminf(local_minimum,\n __shfl_down_sync(0xffffffffu, local_minimum, delta));\n local_off = fmaxf(local_off,\n __shfl_down_sync(0xffffffffu, local_off, delta));\n }\n if (lane == 0) {\n const float dynamic = local_maximum\n / fmaxf(local_minimum, 1.0e-30f);\n dense_flags[batch] = dynamic > 32.0f && dynamic < 256.0f\n && local_off > 1.0e-20f * local_maximum;\n }\n }\n __syncthreads();\n }\n\n float partial = 0.0f;\n for (int row = tid; row < N; row += blockDim.x)\n partial = fmaf(current[row], product[row], partial);\n partial = e096_warp_sum(partial);\n if (lane == 0) warp_partial[warp] = partial;\n __syncthreads();\n if (warp == 0) {\n float value = lane < WARPS ? warp_partial[lane] : 0.0f;\n value = e096_warp_sum(value);\n if (lane == 0) scalar = value;\n }\n __syncthreads();\n const float alpha = scalar;\n if (tid == 0) diagonal[step] = alpha;\n\n partial = 0.0f;\n for (int row = tid; row < N; row += blockDim.x) {\n const float value = product[row] - alpha * current[row];\n product[row] = value;\n partial = fmaf(value, value, partial);\n }\n partial = e096_warp_sum(partial);\n if (lane == 0) warp_partial[warp] = partial;\n __syncthreads();\n if (warp == 0) {\n float value = lane < WARPS ? warp_partial[lane] : 0.0f;\n value = e096_warp_sum(value);\n if (lane == 0) scalar = sqrtf(fmaxf(value, 1.0e-40f));\n }\n __syncthreads();\n const float next_beta = scalar;\n if (step < 2 && tid == 0) off[step] = next_beta;\n for (int row = tid; row < N; row += blockDim.x) {\n previous[row] = current[row];\n current[row] = product[row] / next_beta;\n }\n beta = next_beta;\n __syncthreads();\n }\n\n if (tid == 0) {\n const float a = diagonal[0];\n const float d = diagonal[1];\n const float f = diagonal[2];\n const float b = off[0];\n const float e = off[1];\n const float q = (a + d + f) * (1.0f / 3.0f);\n const float aa = a - q;\n const float dd = d - q;\n const float ff = f - q;\n const float p = sqrtf(fmaxf(\n (aa * aa + dd * dd + ff * ff + 2.0f * (b * b + e * e))\n * (1.0f / 6.0f), 1.0e-30f));\n const float ia = aa / p;\n const float id = dd / p;\n const float iff = ff / p;\n const float ib = b / p;\n const float ie = e / p;\n const float determinant = ia * (id * iff - ie * ie) - ib * ib * iff;\n const float r = fminf(1.0f, fmaxf(-1.0f, 0.5f * determinant));\n const float phi = acosf(r) * (1.0f / 3.0f);\n const float maximum = q + 2.0f * p * cosf(phi);\n const float minimum = q + 2.0f * p * cosf(phi + 2.0943951023931953f);\n lower[batch] = minimum;\n upper[batch] = maximum;\n median[batch] = 3.0f * q - minimum - maximum;\n }\n}\n'
def _e2993_n352_probe_scale_source():
source=_E096_LANCZOS3_SOURCE;replacements=(' bool* __restrict__ dense_flags) {',' bool* __restrict__ dense_flags,\n float* __restrict__ matrix_scale) {'),(' __shared__ float row_off[N];',' __shared__ float row_off[N];\n __shared__ float row_l1[N];'),(' __shared__ float warp_off[WARPS];',' __shared__ float warp_off[WARPS];\n __shared__ float warp_l1[WARPS];'),(' float diagonal_energy = 0.0f;',' float diagonal_energy = 0.0f;\n float l1 = 0.0f;'),(' energy = fmaf(a, a, energy);',' energy = fmaf(a, a, energy);\n l1 += fabsf(a);'),(' diagonal_energy = e096_warp_sum(diagonal_energy);',' diagonal_energy = e096_warp_sum(diagonal_energy);\n l1 = e096_warp_sum(l1);'),(' row_off[row] = energy - diagonal_energy;',' row_off[row] = energy - diagonal_energy;\n row_l1[row] = l1;'),(' float local_off = 0.0f;\n for (int row = tid; row < N; row += blockDim.x) {',' float local_off = 0.0f;\n float local_l1 = 0.0f;\n for (int row = tid; row < N; row += blockDim.x) {'),(' local_off = fmaxf(local_off, row_off[row]);',' local_off = fmaxf(local_off, row_off[row]);\n local_l1 = fmaxf(local_l1, row_l1[row]);'),(' local_off = fmaxf(local_off,\n __shfl_down_sync(0xffffffffu, local_off, delta));',' local_off = fmaxf(local_off,\n __shfl_down_sync(0xffffffffu, local_off, delta));\n local_l1 = fmaxf(local_l1,\n __shfl_down_sync(0xffffffffu, local_l1, delta));'),(' warp_off[warp] = local_off;',' warp_off[warp] = local_off;\n warp_l1[warp] = local_l1;'),(' local_off = lane < WARPS ? warp_off[lane] : 0.0f;',' local_off = lane < WARPS ? warp_off[lane] : 0.0f;\n local_l1 = lane < WARPS ? warp_l1[lane] : 0.0f;'),(' local_off = fmaxf(local_off,\n __shfl_down_sync(0xffffffffu, local_off, delta));',' local_off = fmaxf(local_off,\n __shfl_down_sync(0xffffffffu, local_off, delta));\n local_l1 = fmaxf(local_l1,\n __shfl_down_sync(0xffffffffu, local_l1, delta));'),(' && local_off > 1.0e-20f * local_maximum;',' && local_off > 1.0e-20f * local_maximum;\n matrix_scale[batch] = local_l1;')
for(old,new)in replacements:
if source.count(old)!=1:raise RuntimeError('n352 probe-scale template changed')
source=source.replace(old,new,1)
return source
@memo(maxsize=1)
def _e096_lanczos3_kernel():return CUDAKernel(_fast_nvrtc_compile(_e2993_n352_probe_scale_source(),_E096_LANCZOS3_NAME),_E096_LANCZOS3_NAME)
@torch.no_grad()
def _e096_fused_lanczos3(matrix:torch.Tensor,*,steps:int=3,probes:int=1,dense_flags:torch.Tensor|None=None,return_scale:bool=False):
if matrix.shape[-1]!=352 or steps!=3 or probes!=1:return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
median=torch.empty((matrix.shape[0],),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);matrix_scale=torch.empty_like(median)
if dense_flags is None:dense_flags=torch.empty((matrix.shape[0],),device=matrix.device,dtype=torch.bool)
_e096_lanczos3_kernel().launch((matrix.shape[0],1,1),(512,1,1),(matrix,median,lower,upper,dense_flags,matrix_scale));stats=median,lower,upper;return(stats,matrix_scale)if return_scale else stats
@torch.no_grad()
def _e133_n352_probe(matrix:torch.Tensor):dense_flags=torch.empty((matrix.shape[0],),device=matrix.device,dtype=torch.bool);stats,matrix_scale=_e096_fused_lanczos3(matrix,dense_flags=dense_flags,return_scale=True);return stats,dense_flags,matrix_scale
@triton.jit
def _e202_n352_symmetric_square_kernel(matrix,output,n:tl.constexpr,block:tl.constexpr):
batch=tl.program_id(0);tile=tl.program_id(1);tile_m=tl.where(tile<6,0,tl.where(tile<11,1,tl.where(tile<15,2,tl.where(tile<18,3,tl.where(tile<20,4,5)))));base=tl.where(tile_m==0,0,tl.where(tile_m==1,6,tl.where(tile_m==2,11,tl.where(tile_m==3,15,tl.where(tile_m==4,18,20)))));tile_n=tile_m+tile-base;rows=tile_m*block+tl.arange(0,block);cols=tile_n*block+tl.arange(0,block);accumulator=tl.zeros((block,block),tl.float32)
for start in tl.range(0,n,block,num_stages=3):inner=start+tl.arange(0,block);left=tl.load(matrix+batch*n*n+rows[:,None]*n+inner[None,:],mask=(rows[:,None]<n)&(inner[None,:]<n),other=.0);right=tl.load(matrix+batch*n*n+inner[:,None]*n+cols[None,:],mask=(inner[:,None]<n)&(cols[None,:]<n),other=.0);accumulator+=tl.dot(left,right)
if tile_n==tile_m:accumulator=.5*(accumulator+tl.trans(accumulator))
mask=(rows[:,None]<n)&(cols[None,:]<n);tl.store(output+batch*n*n+rows[:,None]*n+cols[None,:],accumulator,mask=mask);tl.store(output+batch*n*n+cols[:,None]*n+rows[None,:],tl.trans(accumulator),mask=tl.trans(mask))
@triton.jit
def _e202_n352_symmetric_sign_update_kernel(matrix,square,output,n:tl.constexpr,block:tl.constexpr,alpha:tl.constexpr):
batch=tl.program_id(0);tile=tl.program_id(1);tile_m=tl.where(tile<6,0,tl.where(tile<11,1,tl.where(tile<15,2,tl.where(tile<18,3,tl.where(tile<20,4,5)))));base=tl.where(tile_m==0,0,tl.where(tile_m==1,6,tl.where(tile_m==2,11,tl.where(tile_m==3,15,tl.where(tile_m==4,18,20)))));tile_n=tile_m+tile-base;rows=tile_m*block+tl.arange(0,block);cols=tile_n*block+tl.arange(0,block);accumulator=tl.zeros((block,block),tl.float32)
for start in tl.range(0,n,block,num_stages=3):inner=start+tl.arange(0,block);left=tl.load(matrix+batch*n*n+rows[:,None]*n+inner[None,:],mask=(rows[:,None]<n)&(inner[None,:]<n),other=.0);right=tl.load(square+batch*n*n+inner[:,None]*n+cols[None,:],mask=(inner[:,None]<n)&(cols[None,:]<n),other=.0);accumulator+=tl.dot(left,right)
old=tl.load(matrix+batch*n*n+rows[:,None]*n+cols[None,:],mask=(rows[:,None]<n)&(cols[None,:]<n),other=.0).to(tl.float32);result=alpha*old+(1.-alpha)*accumulator
if tile_n==tile_m:result=.5*(result+tl.trans(result))
mask=(rows[:,None]<n)&(cols[None,:]<n);tl.store(output+batch*n*n+rows[:,None]*n+cols[None,:],result,mask=mask);tl.store(output+batch*n*n+cols[:,None]*n+rows[None,:],tl.trans(result),mask=tl.trans(mask))
@torch.no_grad()
def _e202_n352_symmetric_sign(sign:torch.Tensor,steps:int,*,growth_alpha:float=1.5,growth_steps:int=0):
grid=sign.shape[0],21,1;square=torch.empty_like(sign);buffer=torch.empty_like(sign)
if steps==7 and growth_alpha==1.5 and growth_steps==0:schedule=1.25,2.1,2.1,2.1,1.5,1.5
else:schedule=tuple(growth_alpha if iteration<growth_steps else 1.5 for iteration in range(steps))
for alpha in schedule:_e202_n352_symmetric_square_kernel[grid](sign,square,n=352,block=64,num_warps=4,num_stages=1);_e202_n352_symmetric_sign_update_kernel[grid](sign,square,buffer,n=352,block=64,alpha=alpha,num_warps=4,num_stages=1);sign,buffer=buffer,sign
return sign
@torch.no_grad()
def _e202_n352_symmetric_square(matrix:torch.Tensor):output=torch.empty_like(matrix);_e202_n352_symmetric_square_kernel[matrix.shape[0],21,1](matrix,output,n=352,block=64,num_warps=4,num_stages=1);return output
_E202_N352_POTRF176_NAME='e202_n352_potrf176_wmma_tf32x3'
def _e202_wmma_potrf_source(n:int,name:str):
source=_N176_CHOLESKYQR_SOURCE.replace('#include <cuda_runtime.h>','#include <cuda_runtime.h>\n#include <mma.h>\nnamespace wmma=nvcuda::wmma;',1).replace('potrf176_block16_ridge',name,1)
if n!=176:source=source.replace('constexpr int N = 176;',f"constexpr int N = {n};",1)
old=' for (int index = tid; index < NN; index += (int)blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n if (row >= end && column >= end && row >= column) {\n float value = factor[index];\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= N) break;\n value = fmaf(-factor[row * N + k],\n factor[column * N + k], value);\n }\n factor[index] = value;\n }\n }\n __syncthreads();';new=' const int warp = tid >> 5;\n const int trailing_tiles = (N - end) / 16;\n for (int tile = warp; tile < trailing_tiles * trailing_tiles;\n tile += blockDim.x / 32) {\n const int tile_row = tile / trailing_tiles;\n const int tile_col = tile - tile_row * trailing_tiles;\n if (tile_row >= tile_col) {\n const int row = end + 16 * tile_row;\n const int col = end + 16 * tile_col;\n wmma::fragment<wmma::accumulator,16,16,8,float> acc;\n wmma::load_matrix_sync(acc, factor + row * N + col,\n N, wmma::mem_row_major);\n #pragma unroll\n for (int k0 = 0; k0 < 16; k0 += 8) {\n wmma::fragment<wmma::matrix_a,16,16,8,\n wmma::precision::tf32,wmma::row_major> a;\n wmma::fragment<wmma::matrix_a,16,16,8,\n wmma::precision::tf32,wmma::row_major> al;\n wmma::fragment<wmma::matrix_b,16,16,8,\n wmma::precision::tf32,wmma::col_major> b;\n wmma::fragment<wmma::matrix_b,16,16,8,\n wmma::precision::tf32,wmma::col_major> bl;\n wmma::load_matrix_sync(a, factor + row * N + panel + k0, N);\n wmma::load_matrix_sync(b, factor + col * N + panel + k0, N);\n #pragma unroll\n for (int item = 0; item < a.num_elements; ++item) {\n const float value = a.x[item];\n const float high = wmma::__float_to_tf32(value);\n a.x[item] = -high;\n al.x[item] = -wmma::__float_to_tf32(value - high);\n }\n #pragma unroll\n for (int item = 0; item < b.num_elements; ++item) {\n const float value = b.x[item];\n const float high = wmma::__float_to_tf32(value);\n b.x[item] = high;\n bl.x[item] = wmma::__float_to_tf32(value - high);\n }\n wmma::mma_sync(acc, a, b, acc);\n wmma::mma_sync(acc, al, b, acc);\n wmma::mma_sync(acc, a, bl, acc);\n }\n wmma::store_matrix_sync(\n factor + row * N + col, acc, N, wmma::mem_row_major);\n }\n }\n __syncthreads();'
if source.count(old)!=1:raise RuntimeError('WMMA POTRF trailing template changed')
return source.replace(old,new,1)
@memo(maxsize=1)
def _e202_n352_potrf176_kernel():source=_e202_wmma_potrf_source(176,_E202_N352_POTRF176_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E202_N352_POTRF176_NAME),_E202_N352_POTRF176_NAME)
@torch.no_grad()
def _e202_n352_cqr176(matrix:torch.Tensor,precision:str='highest',trsm_fn=None):
batch,rows,_=matrix.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(precision);gram=matrix.mT@matrix;torch.set_float32_matmul_precision(previous);lower=gram;_e202_n352_potrf176_kernel().launch((batch,1,1),(640,1,1),(gram,lower,batch,1e-06),shared_mem=(176*176+1)*4)
if trsm_fn is not None:return trsm_fn(matrix,lower)
_n176_choleskyqr_kernel('right_trsm176_block16_rows16').launch((batch,(rows+15)//16,1),(256,1,1),(matrix,lower,matrix,batch,rows),shared_mem=176*16*4);return matrix
@torch.no_grad()
def _e202_n352_child176_eigh(data:torch.Tensor):batch=data.shape[0];vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device);saved=torch.empty_like(data);diagonal=torch.empty_like(values);off_diagonal=torch.empty_like(values);timers=torch.empty((batch,5),device=data.device,dtype=torch.int64);panel_total=(176-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=data.device,dtype=torch.float32);_e185_reduce176_block4_kernel().launch((batch,1,1),(768,1,1),(data,saved,diagonal,off_diagonal,timers),shared_mem=_E185_N176_BLOCK4_SHARED_BYTES);_e1762_solve176_pdl_kernel(20).launch((batch*2,1,1),(256,1,1),(saved,diagonal,off_diagonal,vectors,values,timers),shared_mem=_E148_SOLVE176_SHARED_BYTES);_householder_compact_t16_kernel[batch,panel_total](saved,triangular,176,panel_total,panel_width=32,block_rows=16,num_warps=2,num_stages=2,launch_pdl=True);return _e1762_tensor_wy_apply_saved_(vectors,saved,triangular),values
@torch.no_grad()
def _e202_n352_low_only_split(data:torch.Tensor,lanczos_stats:tuple[torch.Tensor,torch.Tensor,torch.Tensor]):batch,n,_=data.shape;rank=n//2;eye=torch.eye(n,device=data.device).expand(batch,-1,-1);center=data.diagonal(dim1=-2,dim2=-1).mean(dim=-1);_,lower,upper=lanczos_stats;radius=torch.maximum((lower-center).abs(),(upper-center).abs()).clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(data,center,radius);sign=_e202_n352_symmetric_sign(sign,7);low_projector=_e483_half_sign_projector(sign);operator=_e202_n352_symmetric_square(low_projector);operator=_e202_n352_symmetric_square(operator);low=operator[:,:,:rank].contiguous();low=_e202_n352_cqr176(low);low=operator@low;low=_e202_n352_cqr176(low);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');high=eye[:,:,rank:]-low@low[:,rank:,:].mT;high=torch.baddbmm(high,low,low.mT@high,beta=1.,alpha=-1.);torch.set_float32_matmul_precision('high');high=_e202_n352_cqr176(high);basis=torch.cat((low,high),dim=-1);transformed=basis.mT@data@basis;torch.set_float32_matmul_precision(previous);children=torch.cat((transformed[:,:rank,:rank],transformed[:,rank:,rank:]),dim=0).contiguous();return basis,children
@torch.no_grad()
def _e202_small_n352_eigh(data:torch.Tensor,*,lanczos_stats:tuple[torch.Tensor,torch.Tensor,torch.Tensor]|None=None):
batch,n,_=data.shape;rank=n//2;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
if lanczos_stats is None:lanczos_stats=_e096_fused_lanczos3(data,steps=3,probes=1)
basis,children=_e202_n352_low_only_split(data,lanczos_stats);child_vectors,child_values=_e202_n352_child176_eigh(children);vectors=torch.empty_like(data);torch.bmm(basis[:,:,:rank],child_vectors[:batch],out=vectors[:,:,:rank]);torch.bmm(basis[:,:,rank:],child_vectors[batch:],out=vectors[:,:,rank:]);values=torch.cat((child_values[:batch],child_values[batch:]),dim=-1);begin,end=rank-12,rank+12;window=vectors[:,:,begin:end];rotation,local_values=_small_n24_eigh(window.mT@data@window);vectors[:,:,begin:end]=window@rotation;values[:,begin:end]=local_values;values,order=values.sort(dim=-1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1));torch.set_float32_matmul_precision('highest');gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);return vectors,values
finally:torch.set_float32_matmul_precision(previous)
_E2676_N352_CERTIFICATE_RESIDUAL_NAME='e2676_n352_residual_m176n16';_E2676_N352_CERTIFICATE_FINISH_NAME='e2676_n352_finish16'
@memo(maxsize=1)
def _e2676_n352_certificate_kernels():residual_source=_e2167_dense_certificate_residual_source().replace(_E2167_DENSE_CERTIFICATE_RESIDUAL_NAME,_E2676_N352_CERTIFICATE_RESIDUAL_NAME).replace('constexpr int N=512,BM=256,BN=16,BK=64;','constexpr int N=352,BM=176,BN=16,BK=64;').replace('at[x]=__float2half_rn(a[base+(long long)(row0+r)*N+start+k]*inv_scale);','at[x]=start+k<N?__float2half_rn(a[base+(long long)(row0+r)*N+start+k]*inv_scale):__float2half_rn(0.f);').replace('bt[x]=p<count?__float2half_rn(q[base+(long long)(start+k)*N+c]):__float2half_rn(0.f);','bt[x]=p<count&&start+k<N?__float2half_rn(q[base+(long long)(start+k)*N+c]):__float2half_rn(0.f);');finish_source=_E2167_DENSE_CERTIFICATE_FINISH_SOURCE.replace(_E2167_DENSE_CERTIFICATE_FINISH_NAME,_E2676_N352_CERTIFICATE_FINISH_NAME).replace('constexpr int N=512;','constexpr int N=352;');return CUDAKernel(_fast_nvrtc_compile(residual_source,_E2676_N352_CERTIFICATE_RESIDUAL_NAME),_E2676_N352_CERTIFICATE_RESIDUAL_NAME),CUDAKernel(_fast_nvrtc_compile(finish_source,_E2676_N352_CERTIFICATE_FINISH_NAME),_E2676_N352_CERTIFICATE_FINISH_NAME)
@memo(maxsize=4)
def _e2676_n352_certificate_indices(device:int):return torch.cat((torch.arange(160,168,device=device),torch.arange(344,352,device=device))).to(torch.int32)
_E2992_N352_NORM_GUARD_NAME='e2992_n352_inplace_norm_guard';_E2992_N352_NORM_GUARD_SOURCE='\n#include <cuda_runtime.h>\n__device__ __forceinline__ float e2992_max(float value) {\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n value = fmaxf(\n value, __shfl_down_sync(0xffffffffu, value, offset));\n return value;\n}\nextern "C" __global__ __launch_bounds__(352, 1)\nvoid e2992_n352_inplace_norm_guard(\n const float* __restrict__ vectors,\n const float* __restrict__ values,\n const float* __restrict__ scale,\n const float* __restrict__ sums,\n const bool* __restrict__ deferred,\n bool* __restrict__ risk,\n const bool* __restrict__ route_flags,\n int batch,\n int enforce_route,\n int fuse_finish,\n float threshold) {\n constexpr int N = 352;\n constexpr int WARPS = 11;\n const int matrix = blockIdx.x;\n const int column = threadIdx.x;\n const int lane = column & 31;\n const int warp = column >> 5;\n if (matrix >= batch) return;\n const long long base = (long long)matrix * N * N;\n __shared__ float eigen_warp[4];\n __shared__ float norm_warp[4];\n __shared__ float adjacent_warp[4];\n __shared__ float ordering_warp[4];\n __shared__ float scale_warp[4];\n if (fuse_finish) {\n constexpr int COUNT = 16;\n const float inf = __int_as_float(0x7f800000);\n float eigen = column < COUNT\n ? sums[(long long)matrix * 256 + column] : 0.0f;\n float selected_norm = column < COUNT\n ? fabsf(sums[(long long)matrix * 256 + COUNT + column] - 1.0f)\n : 0.0f;\n float adjacent = column < COUNT\n ? fabsf(sums[(long long)matrix * 256 + 2 * COUNT + column])\n : 0.0f;\n eigen = isfinite(eigen) ? eigen : inf;\n selected_norm = isfinite(selected_norm) ? selected_norm : inf;\n adjacent = isfinite(adjacent) ? adjacent : inf;\n eigen = e2992_max(eigen);\n selected_norm = e2992_max(selected_norm);\n adjacent = e2992_max(adjacent);\n float ordering = -1.0e30f;\n float value_scale = 1.0f;\n if (column < 128) {\n for (int item = column; item < N; item += 128) {\n const float value = values[(long long)matrix * N + item];\n if (!isfinite(value)) {\n ordering = inf;\n value_scale = inf;\n } else {\n value_scale = fmaxf(value_scale, fabsf(value));\n if (item + 1 < N) {\n const float next = values[(long long)matrix * N + item + 1];\n ordering = !isfinite(next)\n ? inf : fmaxf(ordering, value - next);\n }\n }\n }\n }\n ordering = e2992_max(ordering);\n value_scale = e2992_max(value_scale);\n if (lane == 0 && warp < 4) {\n eigen_warp[warp] = eigen;\n norm_warp[warp] = selected_norm;\n adjacent_warp[warp] = adjacent;\n ordering_warp[warp] = ordering;\n scale_warp[warp] = value_scale;\n }\n __syncthreads();\n if (warp == 0) {\n eigen = e2992_max(lane < 4 ? eigen_warp[lane] : 0.0f);\n selected_norm = e2992_max(lane < 4 ? norm_warp[lane] : 0.0f);\n adjacent = e2992_max(lane < 4 ? adjacent_warp[lane] : 0.0f);\n ordering = e2992_max(lane < 4 ? ordering_warp[lane] : -1.0e30f);\n value_scale = e2992_max(lane < 4 ? scale_warp[lane] : 1.0f);\n if (lane == 0) {\n const float eigen_score = eigen / fmaxf(scale[matrix], 1.0e-30f);\n const float ordering_score = ordering / value_scale;\n risk[matrix] = !isfinite(eigen_score) || !isfinite(ordering_score)\n || !isfinite(selected_norm) || !isfinite(adjacent)\n || eigen_score > 0.95f * 200.0f * N * 1.1920928955078125e-7f\n || selected_norm > 0.002f || adjacent > 0.002f\n || ordering_score > 0.975f * 100.0f * N\n * 1.1920928955078125e-7f;\n }\n }\n __syncthreads();\n }\n float norm = 0.0f;\n #pragma unroll 4\n for (int row = 0; row < N; ++row) {\n const float value = vectors[base + (long long)row * N + column];\n norm = fmaf(value, value, norm);\n }\n int bad = !isfinite(norm) || fabsf(norm - 1.0f) > threshold;\n #pragma unroll\n for (int offset = 16; offset > 0; offset >>= 1)\n bad |= __shfl_down_sync(0xffffffffu, bad, offset);\n __shared__ int warp_bad[WARPS];\n if (lane == 0) warp_bad[warp] = bad;\n __syncthreads();\n if (column == 0) {\n int any = 0;\n #pragma unroll\n for (int item = 0; item < WARPS; ++item) any |= warp_bad[item];\n int route_bad = 0;\n if (enforce_route) {\n for (int item = 0; item < batch; ++item)\n route_bad |= !route_flags[item];\n }\n if (any || route_bad) risk[matrix] = true;\n }\n}\n'
@memo(maxsize=1)
def _e2992_n352_norm_guard_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2992_N352_NORM_GUARD_SOURCE,_E2992_N352_NORM_GUARD_NAME),_E2992_N352_NORM_GUARD_NAME)
@torch.no_grad()
def _e202_guarded_small_n352_eigh(data:torch.Tensor,*,lanczos_stats:tuple[torch.Tensor,torch.Tensor,torch.Tensor]|None=None,matrix_scale:torch.Tensor|None=None,route_flags:torch.Tensor|None=None):
vectors,values=_e202_small_n352_eigh(data,lanczos_stats=lanczos_stats);batch=data.shape[0]
if matrix_scale is None:matrix_scale=data.abs().sum(dim=1).amax(dim=1)
sums=torch.zeros((batch,256),device=data.device);risk=torch.empty((batch,),device=data.device,dtype=torch.bool);deferred=torch.zeros_like(risk)if _certificate_reason_ledger is not None else risk;device=data.device.index
if device is None:device=torch.cuda.current_device()
indices=_e2676_n352_certificate_indices(device);residual,finish=_e2676_n352_certificate_kernels();residual.launch((batch,2,1),(352,1,1),(data,vectors,values,matrix_scale,indices,sums,batch,16),shared_mem=47104)
if _certificate_reason_ledger is not None:finish.launch((batch,1,1),(128,1,1),(values,matrix_scale,sums,deferred,risk,batch,16,.95))
if _certificate_reason_ledger is None:_e2992_n352_norm_guard_kernel().launch((batch,1,1),(352,1,1),(vectors,values,matrix_scale,sums,deferred,risk,risk if route_flags is None else route_flags,batch,int(route_flags is not None),1,.002));norm_risk=None
else:norm_risk=_e2073_column_norm_risk(vectors,.002);risk|=norm_risk;route_risk=torch.zeros_like(risk)if route_flags is None else~route_flags.all().expand_as(risk);risk|=route_risk
if _certificate_reason_ledger is not None:eps=torch.finfo(torch.float32).eps;eigen_score=sums[:,:16].amax(dim=1)/matrix_scale.clamp_min(1e-30);selected_norm=(sums[:,16:32]-1.).abs().amax(dim=1);selected_adjacent=sums[:,32:48].abs().amax(dim=1);value_scale=values.abs().amax(dim=1).clamp_min(1.);ordering_score=(values[:,:-1]-values[:,1:]).amax(dim=1)/value_scale;finite_risk=~(torch.isfinite(eigen_score)&torch.isfinite(selected_norm)&torch.isfinite(selected_adjacent)&torch.isfinite(ordering_score));eigen_risk=eigen_score>.95*2e2*352.*eps;selected_orthogonal_risk=(selected_norm>.002)|(selected_adjacent>.002);ordering_risk=ordering_score>.975*1e2*352.*eps;classified=finite_risk|eigen_risk|selected_orthogonal_risk|ordering_risk|norm_risk;unknown_risk=risk&~(classified|route_risk);_certificate_reason_record('n352_output',_certificate_reason_bits(risk,(_CERT_REASON_NONFINITE,finite_risk),(_CERT_REASON_EIGEN_RESIDUAL,eigen_risk),(_CERT_REASON_ORTHOGONALITY,selected_orthogonal_risk),(_CERT_REASON_ORDERING,ordering_risk),(_CERT_REASON_NORM,norm_risk),(_CERT_REASON_ROUTE,route_risk),(_CERT_REASON_ROUTE,unknown_risk)),risk)
if bool(risk.any().item()):exact_values,exact_vectors=torch.linalg.eigh(data[risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[risk]=exact_vectors;values[risk]=exact_values
return vectors,values
N=256;BLOCK=32
@memo(maxsize=64)
def _e549_upper_tile_pairs(device_index:int,tiles:int,first_column_tile:int=0):
rows=[];columns=[]
for row in range(tiles):
for column in range(max(row,first_column_tile),tiles):rows.append(row);columns.append(column)
device=torch.device('cuda',device_index);return torch.tensor(rows,device=device,dtype=torch.int32),torch.tensor(columns,device=device,dtype=torch.int32)
@triton.jit
def _e549_symmetric_update_upper_kernel(matrix,panel,companion,output,pair_rows,pair_columns,n:tl.constexpr,active_start,panel_rows,block:tl.constexpr,use_tf32x3:tl.constexpr):
batch=tl.program_id(0);pair=tl.program_id(1);tile_m=tl.load(pair_rows+pair);tile_n=tl.load(pair_columns+pair);rows=tile_m*block+tl.arange(0,block);cols=tile_n*block+tl.arange(0,block);rank=tl.arange(0,32);row_mask=rows<n;col_mask=cols<n;old=tl.load(matrix+batch*n*n+rows[:,None]*n+cols[None,:],mask=row_mask[:,None]&col_mask[None,:],other=.0);local_rows=rows-active_start;local_cols=cols-active_start;vr=tl.load(panel+batch*panel_rows*32+local_rows[:,None]+rank[None,:]*panel_rows,mask=(local_rows[:,None]>=0)&(local_rows[:,None]<panel_rows),other=.0);vc=tl.load(panel+batch*panel_rows*32+local_cols[:,None]+rank[None,:]*panel_rows,mask=(local_cols[:,None]>=0)&(local_cols[:,None]<panel_rows),other=.0);wr=tl.load(companion+batch*n*32+rows[:,None]*32+rank[None,:],mask=row_mask[:,None],other=.0);wc=tl.load(companion+batch*n*32+cols[:,None]*32+rank[None,:],mask=col_mask[:,None],other=.0)
if use_tf32x3:update=tl.dot(vr,tl.trans(wc),input_precision='tf32x3');update+=tl.dot(wr,tl.trans(vc),input_precision='tf32x3')
else:update=tl.dot(vr,tl.trans(wc),input_precision='tf32');update+=tl.dot(wr,tl.trans(vc),input_precision='tf32')
result=old-update;tl.store(output+batch*n*n+rows[:,None]*n+cols[None,:],result,mask=row_mask[:,None]&col_mask[None,:]);tl.store(output+batch*n*n+cols[:,None]*n+rows[None,:],tl.trans(result),mask=(tile_m!=tile_n)&col_mask[:,None]&row_mask[None,:])
def _e549_symmetric_compact_block_update(matrix:torch.Tensor,panel:torch.Tensor,triangular:torch.Tensor,*,active_start:int,precision:str):
batch=matrix.shape[0];n=matrix.shape[-1];panel_rows=n-active_start;tiles=triton.cdiv(n,64);product=torch.empty((batch,n,BLOCK),device=matrix.device);_compact_matrix_panel_kernel[batch,tiles](matrix,panel,product,n,active_start,panel_rows,block_m=64,block_k=64,use_tf32x3=precision=='tf32x3',num_warps=8,num_stages=3);companion=torch.empty_like(product);_compact_companion_kernel[batch,tiles](product,panel,triangular,companion,n,active_start,panel_rows,block_m=64,block_k=64,num_warps=8,num_stages=2);device_index=matrix.device.index
if device_index is None:device_index=torch.cuda.current_device()
pair_rows,pair_columns=_e549_upper_tile_pairs(device_index,tiles,active_start//64);_e549_symmetric_update_upper_kernel[batch,pair_rows.numel()](matrix,panel,companion,matrix,pair_rows,pair_columns,n,active_start,panel_rows,block=64,use_tf32x3=precision=='tf32x3',num_warps=8,num_stages=2);return matrix
T32_NAME='blocked_build_t32';T32_SOURCE=_N352_GAU_T_SOURCE+'\nextern "C" __global__ __launch_bounds__(512, 1)\nvoid blocked_build_t32(const float* gram, const float* tau,\n float* output, long long tau_stride) {\n constexpr int K = 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const float* gram_b = gram + batch * K * K;\n const float* tau_b = tau + batch * tau_stride;\n float* out_b = output + batch * K * K;\n extern __shared__ float storage[];\n float* lower = storage;\n float* inverse = lower + K * K;\n float* mid = inverse + K * K;\n\n if (tid < K * K / 8) {\n const int row = tid / (K / 8);\n const int col = (tid - row * (K / 8)) * 8;\n float values[8];\n ldg_f32<8>(values, gram_b + row * K + col);\n #pragma unroll\n for (int item = 0; item < 8; ++item)\n lower[row * K + col + item] = col + item < row ? values[item] : 0.0f;\n }\n __syncthreads();\n build_t32_inverse_block<K>(lower, tau_b, inverse, mid);\n __syncthreads();\n\n if (tid < K * K / 8) {\n const int row = tid / (K / 8);\n const int col = (tid - row * (K / 8)) * 8;\n float values[8];\n #pragma unroll\n for (int item = 0; item < 8; ++item)\n values[item] = col + item <= row ? inverse[row * K + col + item] : 0.0f;\n stg_f32<8>(out_b + row * K + col, values);\n }\n}\n'
@memo(maxsize=1)
def _t32_kernel():return CUDAKernel(_fast_nvrtc_compile(T32_SOURCE,T32_NAME),T32_NAME)
_E197_FUSED_GRAM_T32_NAME='e197_fused_gram_t32';_E197_FUSED_GRAM_T32_SOURCE=_N352_GAU_T_SOURCE+'\nextern "C" __global__ __launch_bounds__(512, 1)\nvoid e197_fused_gram_t32(\n const float* __restrict__ panel,\n const float* __restrict__ tau,\n float* __restrict__ output,\n int panel_rows,\n long long tau_stride) {\n constexpr int K = 32;\n const int tid = threadIdx.x;\n const int batch = blockIdx.x;\n const long long panel_base = (long long)batch * panel_rows * K;\n const float* tau_b = tau + (long long)batch * tau_stride;\n float* output_b = output + (long long)batch * K * K;\n extern __shared__ float storage[];\n float* lower = storage;\n float* inverse = lower + K * K;\n float* mid = inverse + K * K;\n\n for (int index = tid; index < K * K; index += blockDim.x) {\n const int row = index / K;\n const int col = index - row * K;\n float value = 0.0f;\n if (col < row) {\n for (int k = 0; k < panel_rows; ++k)\n value = fmaf(\n panel[panel_base + k + (long long)row * panel_rows],\n panel[panel_base + k + (long long)col * panel_rows],\n value);\n }\n lower[index] = value;\n }\n __syncthreads();\n build_t32_inverse_block<K>(lower, tau_b, inverse, mid);\n __syncthreads();\n if (tid < K * K / 8) {\n const int row = tid / (K / 8);\n const int col = (tid - row * (K / 8)) * 8;\n float values[8];\n #pragma unroll\n for (int item = 0; item < 8; ++item)\n values[item] = col + item <= row\n ? inverse[row * K + col + item] : 0.0f;\n stg_f32<8>(output_b + row * K + col, values);\n }\n}\n'
@memo(maxsize=1)
def _e197_fused_gram_t32_kernel():return CUDAKernel(_fast_nvrtc_compile(_E197_FUSED_GRAM_T32_SOURCE,_E197_FUSED_GRAM_T32_NAME),_E197_FUSED_GRAM_T32_NAME)
_E834_ENTRY_SHARDED_GRAM_T32_NAME='e834_entry_sharded_gram_t32';_E834_ENTRY_SHARDED_GRAM_T32_SOURCE=_N352_GAU_T_SOURCE.replace('#include <cuda_fp16.h>','#include <cuda_fp16.h>\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;',1)+'\nextern "C" __global__ __cluster_dims__(4, 1, 1) __launch_bounds__(512, 1)\nvoid e834_entry_sharded_gram_t32(\n const float* __restrict__ panel, const float* __restrict__ tau,\n float* __restrict__ output, int panel_rows, long long tau_stride) {\n constexpr int K = 32;\n constexpr int CLUSTER = 4;\n constexpr int ENTRIES = K * (K - 1) / 2;\n constexpr int ENTRIES_PER_CTA = 128;\n const int tid = threadIdx.x;\n cg::cluster_group cluster = cg::this_cluster();\n const int rank = cluster.block_rank();\n const int batch = blockIdx.x / CLUSTER;\n const long long panel_base = (long long)batch * panel_rows * K;\n const float* tau_b = tau + (long long)batch * tau_stride;\n float* output_b = output + (long long)batch * K * K;\n\n extern __shared__ float storage[];\n float* entries = storage;\n float* lower = entries + ENTRIES;\n float* inverse = lower + K * K;\n float* mid = inverse + K * K;\n\n const int ordinal = rank * ENTRIES_PER_CTA + tid;\n if (tid < ENTRIES_PER_CTA && ordinal < ENTRIES) {\n const int row = (1 + (int)sqrtf(1.0f + 8.0f * ordinal)) / 2;\n const int col = ordinal - row * (row - 1) / 2;\n float value = 0.0f;\n for (int k = 0; k < panel_rows; ++k) {\n value = fmaf(\n panel[panel_base + k + (long long)row * panel_rows],\n panel[panel_base + k + (long long)col * panel_rows], value);\n }\n entries[ordinal] = value;\n }\n cluster.sync();\n\n if (rank == 0) {\n for (int index = tid; index < K * K; index += blockDim.x)\n lower[index] = 0.0f;\n __syncthreads();\n for (int item = tid; item < ENTRIES; item += blockDim.x) {\n const int owner = item / ENTRIES_PER_CTA;\n const float* remote = cluster.map_shared_rank(entries, owner);\n const int row = (1 + (int)sqrtf(1.0f + 8.0f * item)) / 2;\n const int col = item - row * (row - 1) / 2;\n lower[row * K + col] = remote[item];\n }\n __syncthreads();\n build_t32_inverse_block<K>(lower, tau_b, inverse, mid);\n __syncthreads();\n if (tid < K * K / 8) {\n const int pack = tid;\n const int row = pack / (K / 8);\n const int col = (pack - row * (K / 8)) * 8;\n float values[8];\n #pragma unroll\n for (int item = 0; item < 8; ++item) {\n const int item_col = col + item;\n values[item] = item_col <= row\n ? inverse[row * K + item_col] : 0.0f;\n }\n stg_f32<8>(output_b + row * K + col, values);\n }\n }\n cluster.sync();\n}\n'
def _e2020_direct_entry_gram_cluster_source():
helper='\ntemplate <typename T>\n__device__ __forceinline__ T* e2020_map_shared_rank(T* pointer, int rank) {\n unsigned long long remote;\n asm volatile("mapa.u64 %0, %1, %2;"\n : "=l"(remote)\n : "l"((unsigned long long)pointer), "r"(rank));\n return reinterpret_cast<T*>(remote);\n}\n';source=_E834_ENTRY_SHARDED_GRAM_T32_SOURCE;rewrites=('#include <cooperative_groups.h>\n',helper),('namespace cg = cooperative_groups;\n',''),(' cg::cluster_group cluster = cg::this_cluster();\n',''),(' const int rank = cluster.block_rank();',' const int rank = blockIdx.x & 3;'),('cluster.map_shared_rank(entries, owner)','e2020_map_shared_rank(entries, owner)'),(' cluster.sync();',' asm volatile("barrier.cluster.arrive.aligned;" ::: "memory");\n asm volatile("barrier.cluster.wait.aligned;" ::: "memory");');expected=1,1,1,1,1,2
for((old,new),count)in zip(rewrites,expected):
if source.count(old)!=count:raise RuntimeError('E2020 direct Gram cluster PTX anchor changed')
source=source.replace(old,new)
return source
@memo(maxsize=1)
def _e834_entry_sharded_gram_t32_kernel():source=_e2020_direct_entry_gram_cluster_source();return CUDAKernel(_fast_nvrtc_compile(source,_E834_ENTRY_SHARDED_GRAM_T32_NAME),_E834_ENTRY_SHARDED_GRAM_T32_NAME)
_OWNED_N256_COLUMN_BASE_SOURCE='\n#include <cuda_runtime.h>\n\nstatic __device__ __forceinline__ float warp_sum(float value) {\n value += __shfl_down_sync(0xffffffffu, value, 16);\n value += __shfl_down_sync(0xffffffffu, value, 8);\n value += __shfl_down_sync(0xffffffffu, value, 4);\n value += __shfl_down_sync(0xffffffffu, value, 2);\n value += __shfl_down_sync(0xffffffffu, value, 1);\n return value;\n}\n\nextern "C" __global__ void column_tiled_band_to_tridiagonal_512(\n float* __restrict__ matrices,\n float* __restrict__ reflectors\n) {\n constexpr int n = 512;\n constexpr int bandwidth = 63;\n constexpr int storage_band = 64;\n constexpr int max_blocks = 9;\n const int tid = threadIdx.x;\n const int lane = tid & 31;\n const int warp = tid >> 5;\n const int warps = blockDim.x / 32;\n const int batch = blockIdx.x;\n const long long matrix_base = (long long)batch * n * n;\n const long long reflector_base =\n (long long)batch * n * max_blocks * storage_band;\n\n __shared__ float vectors[max_blocks][storage_band];\n __shared__ float left_projection[storage_band];\n __shared__ float right_projection[storage_band];\n __shared__ float scalar;\n\n for (int output_column = 0; output_column < n - 2; ++output_column) {\n const int remaining = n - output_column - 2;\n const int block_count = (remaining + bandwidth - 1) / bandwidth;\n const int first_support = output_column + 1;\n\n // Generate v_0 directly from the target band column.\n if (warp == 0) {\n const int length = min(bandwidth, n - first_support);\n const int i0 = lane;\n const int i1 = lane + 32;\n float x0 = i0 < length ? matrices[\n matrix_base + (long long)(first_support + i0) * n\n + output_column] : 0.0f;\n float x1 = i1 < length ? matrices[\n matrix_base + (long long)(first_support + i1) * n\n + output_column] : 0.0f;\n float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);\n x0 *= inverse_norm;\n x1 *= inverse_norm;\n vectors[0][i0] = x0;\n vectors[0][i1] = x1;\n const long long destination = reflector_base\n + (long long)output_column * max_blocks * storage_band;\n reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;\n }\n __syncthreads();\n\n // The first column of H_k is all that is needed to create the next\n // bulge: x_{k+1} = A[S_{k+1},S_k] H_k e_0.\n for (int block = 1; block < block_count; ++block) {\n const int previous_start = first_support + (block - 1) * bandwidth;\n const int support_start = previous_start + bandwidth;\n const int length = min(bandwidth, n - support_start);\n for (int row = warp; row < length; row += warps) {\n float dot = 0.0f;\n for (int j = lane; j < bandwidth; j += 32) {\n const float h_column = (j == 0 ? 1.0f : 0.0f)\n - 2.0f * vectors[block - 1][j]\n * vectors[block - 1][0];\n dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + previous_start + j] * h_column;\n }\n dot = warp_sum(dot);\n if (lane == 0) vectors[block][row] = dot;\n }\n __syncthreads();\n if (warp == 0) {\n const int i0 = lane;\n const int i1 = lane + 32;\n float x0 = i0 < length ? vectors[block][i0] : 0.0f;\n float x1 = i1 < length ? vectors[block][i1] : 0.0f;\n float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);\n x0 *= inverse_norm;\n x1 *= inverse_norm;\n vectors[block][i0] = x0;\n vectors[block][i1] = x1;\n const long long destination = reflector_base\n + ((long long)output_column * max_blocks + block)\n * storage_band;\n reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;\n }\n __syncthreads();\n }\n\n // Transform the one structurally nonzero prefix row against S_0.\n if (warp == 0) {\n const int prefix_row = output_column;\n const int first_length = min(bandwidth, n - first_support);\n float dot = 0.0f;\n for (int j = lane; j < first_length; j += 32) {\n dot += matrices[matrix_base + (long long)prefix_row * n\n + first_support + j] * vectors[0][j];\n }\n dot = warp_sum(dot);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n for (int j = lane; j < first_length; j += 32) {\n const float updated = matrices[\n matrix_base + (long long)prefix_row * n\n + first_support + j] - 2.0f * dot * vectors[0][j];\n matrices[matrix_base + (long long)prefix_row * n\n + first_support + j] = updated;\n matrices[matrix_base + (long long)(first_support + j) * n\n + prefix_row] = updated;\n }\n }\n __syncthreads();\n\n for (int block = 0; block < block_count; ++block) {\n const int support_start = first_support + block * bandwidth;\n const int length = min(bandwidth, n - support_start);\n\n // Diagonal tile H_k A_kk H_k.\n for (int row = warp; row < length; row += warps) {\n float dot = 0.0f;\n for (int j = lane; j < length; j += 32) {\n dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j] * vectors[block][j];\n }\n dot = warp_sum(dot);\n if (lane == 0) right_projection[row] = dot;\n }\n __syncthreads();\n if (warp == 0) {\n float projection = 0.0f;\n if (lane < length) {\n projection += vectors[block][lane]\n * right_projection[lane];\n }\n if (lane + 32 < length) {\n projection += vectors[block][lane + 32]\n * right_projection[lane + 32];\n }\n projection = warp_sum(projection);\n if (lane == 0) scalar = projection;\n }\n __syncthreads();\n const float diagonal_scalar = scalar;\n for (int row = warp; row < length; row += warps) {\n const float u_row = vectors[block][row];\n const float w_row = 2.0f\n * (right_projection[row] - u_row * diagonal_scalar);\n for (int j = lane; j < length; j += 32) {\n const float w_j = 2.0f\n * (right_projection[j]\n - vectors[block][j] * diagonal_scalar);\n matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j]\n -= w_row * vectors[block][j] + u_row * w_j;\n }\n }\n __syncthreads();\n\n if (block + 1 >= block_count) {\n // When the active range is an exact multiple of 63, one\n // trailing scalar remains outside the final reflector. It\n // is an identity block, but the cross tile still needs the\n // left/right action of H_k. Omitting this case first shows\n // up at output column 6 and then rapidly pollutes the chain.\n const int tail_start = support_start + length;\n if (tail_start < n) {\n if (warp == 0) {\n float dot = 0.0f;\n for (int j = lane; j < length; j += 32) {\n dot += matrices[matrix_base\n + (long long)tail_start * n\n + support_start + j] * vectors[block][j];\n }\n dot = warp_sum(dot);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n for (int j = lane; j < length; j += 32) {\n const float updated = matrices[matrix_base\n + (long long)tail_start * n\n + support_start + j]\n - 2.0f * dot * vectors[block][j];\n matrices[matrix_base + (long long)tail_start * n\n + support_start + j] = updated;\n matrices[matrix_base\n + (long long)(support_start + j) * n\n + tail_start] = updated;\n }\n }\n __syncthreads();\n }\n continue;\n }\n const int next_start = support_start + bandwidth;\n const int next_length = min(bandwidth, n - next_start);\n\n // Both projections of the neighboring off-diagonal tile. Read\n // upper and lower copies row-wise so every transaction coalesces.\n for (int row = warp; row < length; row += warps) {\n float dot = 0.0f;\n for (int j = lane; j < next_length; j += 32) {\n dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j] * vectors[block + 1][j];\n }\n dot = warp_sum(dot);\n if (lane == 0) right_projection[row] = dot;\n }\n for (int row = warp; row < next_length; row += warps) {\n float dot = 0.0f;\n for (int j = lane; j < length; j += 32) {\n dot += matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j] * vectors[block][j];\n }\n dot = warp_sum(dot);\n if (lane == 0) left_projection[row] = dot;\n }\n __syncthreads();\n if (warp == 0) {\n float cross = 0.0f;\n if (lane < length) {\n cross += vectors[block][lane] * right_projection[lane];\n }\n if (lane + 32 < length) {\n cross += vectors[block][lane + 32]\n * right_projection[lane + 32];\n }\n cross = warp_sum(cross);\n if (lane == 0) scalar = cross;\n }\n __syncthreads();\n const float cross_scalar = scalar;\n\n for (int row = warp; row < length; row += warps) {\n for (int j = lane; j < next_length; j += 32) {\n const float updated = matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j]\n - 2.0f * vectors[block][row] * left_projection[j]\n - 2.0f * right_projection[row] * vectors[block + 1][j]\n + 4.0f * vectors[block][row] * cross_scalar\n * vectors[block + 1][j];\n matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j] = updated;\n }\n }\n for (int row = warp; row < next_length; row += warps) {\n for (int j = lane; j < length; j += 32) {\n const float updated = matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j]\n - 2.0f * left_projection[row] * vectors[block][j]\n - 2.0f * vectors[block + 1][row] * right_projection[j]\n + 4.0f * vectors[block + 1][row] * cross_scalar\n * vectors[block][j];\n matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j] = updated;\n }\n }\n __syncthreads();\n }\n }\n}\n';TRIDIAG_NAME='tridiag_twisted256';TRIDIAG_SOURCE=_TRIDIAG_TWISTED512_SOURCE.replace('constexpr int N = 512;','constexpr int N = 256;').replace('__launch_bounds__(512, 1)','__launch_bounds__(256, 1)').replace('tridiag_twisted512',TRIDIAG_NAME);REPAIR_NAME='repair_twisted256_adjacent';REPAIR_SOURCE='\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;\n\nextern "C" __global__ __cluster_dims__(4, 1, 1)\n__launch_bounds__(128, 1)\nvoid repair_twisted256_adjacent(\n const float* __restrict__ input,\n const float* __restrict__ values,\n float* __restrict__ output,\n int batch,\n float gap_threshold) {\n constexpr int N = 256;\n constexpr int ROWS = 64;\n constexpr int PITCH = 257;\n cg::cluster_group cluster = cg::this_cluster();\n const int rank = cluster.block_rank();\n const int matrix = blockIdx.x >> 2;\n const int tid = threadIdx.x;\n if (matrix >= batch) return;\n const int first_row = rank * ROWS;\n const long long base = (long long)matrix * N * N;\n extern __shared__ float storage[];\n float* q = storage;\n float* partial = q + ROWS * PITCH;\n float* partial0 = cluster.map_shared_rank(partial, 0);\n\n for (int index = tid; index < ROWS * N; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n q[row * PITCH + column] =\n input[base + (long long)(first_row + row) * N + column];\n }\n cluster.sync();\n\n // Ordered adjacent MGS repairs the isolated inverse-iteration loss. The\n // gap gate keeps every rotation inside a numerically tight eigenspace.\n for (int column = 1; column < N; ++column) {\n const float gap = values[(long long)matrix * N + column]\n - values[(long long)matrix * N + column - 1];\n const float span = values[(long long)matrix * N + N - 1]\n - values[(long long)matrix * N];\n if (gap < gap_threshold * fmaxf(span, 1.0e-30f)) {\n if (tid == 0) {\n float dot = 0.0f;\n for (int row = 0; row < ROWS; ++row)\n dot += q[row * PITCH + column - 1]\n * q[row * PITCH + column];\n partial[0] = dot;\n }\n cluster.sync();\n if (rank == 0 && tid == 0) {\n float dot = 0.0f;\n #pragma unroll\n for (int owner = 0; owner < 4; ++owner) {\n float* remote = cluster.map_shared_rank(partial, owner);\n dot += remote[0];\n }\n partial[0] = dot;\n }\n cluster.sync();\n const float dot = partial0[0];\n const float inverse = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));\n if (tid < ROWS)\n q[tid * PITCH + column] =\n (q[tid * PITCH + column]\n - dot * q[tid * PITCH + column - 1]) * inverse;\n cluster.sync();\n }\n }\n\n for (int index = tid; index < ROWS * N; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n output[base + (long long)(first_row + row) * N + column] =\n q[row * PITCH + column];\n }\n}\n';REPAIR_SHARED_BYTES=(64*257+1)*4
def _e2021_direct_repair_cluster_source(source:str,rank_mask:int,sync_count:int):
helper='\ntemplate <typename T>\n__device__ __forceinline__ T* e2021_map_shared_rank(T* pointer, int rank) {\n unsigned long long remote;\n asm volatile("mapa.u64 %0, %1, %2;"\n : "=l"(remote)\n : "l"((unsigned long long)pointer), "r"(rank));\n return reinterpret_cast<T*>(remote);\n}\n';rewrites=('#include <cooperative_groups.h>\n',helper),('namespace cg = cooperative_groups;\n',''),(' cg::cluster_group cluster = cg::this_cluster();\n',''),(' const int rank = cluster.block_rank();',f" const int rank = blockIdx.x & {rank_mask};"),('cluster.map_shared_rank(','e2021_map_shared_rank('),('cluster.sync();','asm volatile("barrier.cluster.arrive.aligned;" ::: "memory");\n asm volatile("barrier.cluster.wait.aligned;" ::: "memory");');expected=1,1,1,1,2,sync_count
for((old,new),count)in zip(rewrites,expected):
if source.count(old)!=count:raise RuntimeError('E2021 direct repair cluster PTX anchor changed')
source=source.replace(old,new)
return source
@memo(maxsize=1)
def _repair_kernel():source=_e2021_direct_repair_cluster_source(REPAIR_SOURCE,3,4);return CUDAKernel(_fast_nvrtc_compile(source,REPAIR_NAME),REPAIR_NAME)
@torch.no_grad()
def repair_twisted256_adjacent(q:torch.Tensor,values:torch.Tensor,*,gap_threshold:float=.001):output=torch.empty_like(q);_repair_kernel().launch(grid=(q.shape[0]*4,1,1),block=(128,1,1),shared_mem=REPAIR_SHARED_BYTES,args=(q,values,output,q.shape[0],gap_threshold));return output
@torch.no_grad()
def dense_panel_backtransform(q:torch.Tensor,panels:list[tuple[int,torch.Tensor,torch.Tensor]],*,precision:str='high'):
if precision in('highest','tf32')and q.is_contiguous():return dense_panel_backtransform_fused(q,panels,use_tf32=precision=='tf32')
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(precision)
try:
for(active_start,v,triangular)in reversed(panels):active=q[:,active_start:,:];updated=active-v@(triangular.mT@(v.mT@active));q=q.clone();q[:,active_start:,:]=updated
finally:torch.set_float32_matmul_precision(previous)
return q
@triton.jit
def _compact_wy_backtransform_kernel(q,panel,triangular,n:tl.constexpr,active_start,panel_rows,panel_batch_stride,panel_offset,panel_column_stride,panel_count,support_start,triangular_batch_stride,triangular_offset,panel_width:tl.constexpr,masked_support:tl.constexpr,projection_stages:tl.constexpr,block_rows:tl.constexpr,block_cols:tl.constexpr,use_tf32:tl.constexpr):
batch=tl.program_id(0);tile_col=tl.program_id(1);rank=tl.arange(0,panel_width);cols=tile_col*block_cols+tl.arange(0,block_cols);column_mask=cols<n;projection=tl.zeros((panel_width,block_cols),tl.float32)
for current in tl.range(0,panel_rows,block_rows,num_stages=projection_stages):
local_rows=current+tl.arange(0,block_rows);rows=active_start+local_rows;panel_mask=(local_rows[None,:]<panel_rows)&(rank[:,None]<panel_count)
if masked_support:panel_mask&=local_rows[None,:]>=support_start+rank[:,None]
v_transpose=tl.load(panel+batch*panel_batch_stride+panel_offset+local_rows[None,:]+rank[:,None]*panel_column_stride,mask=panel_mask,other=.0);q_tile=tl.load(q+batch*n*n+rows[:,None]*n+cols[None,:],mask=(local_rows[:,None]<panel_rows)&column_mask[None,:],other=.0)
if use_tf32:projection+=tl.dot(v_transpose,q_tile,input_precision='tf32')
else:projection+=tl.dot(v_transpose,q_tile,input_precision='ieee')
triangular_transpose=tl.load(triangular+batch*triangular_batch_stride+triangular_offset+rank[None,:]*panel_width+rank[:,None])
if use_tf32:transformed=tl.dot(triangular_transpose,projection,input_precision='tf32')
else:transformed=tl.dot(triangular_transpose,projection,input_precision='ieee')
for current in tl.range(0,panel_rows,block_rows):
local_rows=current+tl.arange(0,block_rows);rows=active_start+local_rows;panel_mask=(local_rows[:,None]<panel_rows)&(rank[None,:]<panel_count)
if masked_support:panel_mask&=local_rows[:,None]>=support_start+rank[None,:]
v=tl.load(panel+batch*panel_batch_stride+panel_offset+local_rows[:,None]+rank[None,:]*panel_column_stride,mask=panel_mask,other=.0);q_tile=tl.load(q+batch*n*n+rows[:,None]*n+cols[None,:],mask=(local_rows[:,None]<panel_rows)&column_mask[None,:],other=.0)
if use_tf32:q_tile-=tl.dot(v,transformed,input_precision='tf32')
else:q_tile-=tl.dot(v,transformed,input_precision='ieee')
tl.store(q+batch*n*n+rows[:,None]*n+cols[None,:],q_tile,mask=(local_rows[:,None]<panel_rows)&column_mask[None,:])
@triton.jit
def _householder_compact_t16_kernel(saved,triangular,n:tl.constexpr,panel_count_total:tl.constexpr,panel_width:tl.constexpr,block_rows:tl.constexpr,use_tf32:tl.constexpr=False):
batch=tl.program_id(0);panel_id=tl.program_id(1);rank=tl.arange(0,panel_width);panel_start=panel_id*panel_width;active=rank<tl.minimum(panel_width,n-2-panel_start);gram=tl.zeros((panel_width,panel_width),tl.float32)
for current in tl.range(panel_start+1,n,block_rows,num_stages=2):
rows=current+tl.arange(0,block_rows);support=rows[None,:]>panel_start+rank[:,None];vectors=tl.load(saved+batch*n*n+(panel_start+rank[:,None])*n+rows[None,:],mask=(rows[None,:]<n)&support&active[:,None],other=.0)
if use_tf32:gram+=tl.dot(vectors,tl.trans(vectors),input_precision='tf32')
else:gram+=tl.dot(vectors,tl.trans(vectors),input_precision='ieee')
row=rank[:,None];column=rank[None,:];control=tl.where(row==column,2.,.0).to(tl.float32)
for current_row in tl.static_range(panel_width-1,-1,-1):gram_row=tl.sum(tl.where(row==current_row,gram,.0),axis=0);correction=tl.sum(tl.where(rank[:,None]>current_row,gram_row[:,None]*control,.0),axis=0);replacement=2.*(rank==current_row).to(tl.float32)-2.*correction;control=tl.where(row==current_row,replacement[None,:],control)
tl.store(triangular+batch*panel_count_total*panel_width*panel_width+panel_id*panel_width*panel_width+column*panel_width+row,control)
@torch.no_grad()
def _tensor_wy_backtransform_saved_(vectors:torch.Tensor,saved:torch.Tensor,*,factor_warps:int=4,block_cols:int=64,use_tf32:bool|None=None):
batch,n,_=vectors.shape;panel_width=32;panel_total=(n-2+panel_width-1)//panel_width;triangular=torch.empty((batch,panel_total,panel_width,panel_width),device=vectors.device,dtype=vectors.dtype);_householder_compact_t16_kernel[batch,panel_total](saved,triangular,n,panel_total,panel_width=panel_width,block_rows=32,num_warps=factor_warps,num_stages=2)
for panel_id in range(panel_total-1,-1,-1):panel_start=panel_id*panel_width;panel_count=min(panel_width,n-2-panel_start);active_start=panel_start+1;panel_rows=n-active_start;_compact_wy_backtransform_kernel[batch,triton.cdiv(n,block_cols)](vectors,saved,triangular,n,active_start,panel_rows,n*n,panel_start*n+active_start,n,panel_count,0,panel_total*panel_width*panel_width,panel_id*panel_width*panel_width,panel_width=panel_width,masked_support=True,projection_stages=1,block_rows=32,block_cols=block_cols,use_tf32=n in(256,320)if use_tf32 is None else use_tf32,num_warps=4,num_stages=1)
return vectors
@torch.no_grad()
def _e1762_tensor_wy_apply_saved_(vectors:torch.Tensor,saved:torch.Tensor,triangular:torch.Tensor,*,block_cols:int=64,use_tf32:bool|None=None):
batch,n,_=vectors.shape;panel_width=32;panel_total=triangular.shape[1]
for panel_id in range(panel_total-1,-1,-1):panel_start=panel_id*panel_width;panel_count=min(panel_width,n-2-panel_start);active_start=panel_start+1;panel_rows=n-active_start;_compact_wy_backtransform_kernel[batch,triton.cdiv(n,block_cols)](vectors,saved,triangular,n,active_start,panel_rows,n*n,panel_start*n+active_start,n,panel_count,0,panel_total*panel_width*panel_width,panel_id*panel_width*panel_width,panel_width=panel_width,masked_support=True,projection_stages=1,block_rows=32,block_cols=block_cols,use_tf32=n in(256,320)if use_tf32 is None else use_tf32,num_warps=4,num_stages=1)
return vectors
@torch.no_grad()
def dense_panel_backtransform_fused(q:torch.Tensor,panels:list[tuple[int,torch.Tensor,torch.Tensor]],*,block_cols:int=32,use_tf32:bool=False):
n=q.shape[-1];block_rows=64 if use_tf32 else 32;replay_warps=2 if use_tf32 else 4
for(active_start,panel,triangular)in reversed(panels):panel_rows=n-active_start;_compact_wy_backtransform_kernel[q.shape[0],triton.cdiv(n,block_cols)](q,panel,triangular,n,active_start,panel_rows,panel_rows*32,0,panel_rows,32,0,32*32,0,panel_width=32,masked_support=False,projection_stages=2,block_rows=block_rows,block_cols=block_cols,use_tf32=use_tf32,num_warps=replay_warps,num_stages=2)
return q
@triton.jit
def _compact_matrix_panel_kernel(matrix,panel,output,n:tl.constexpr,active_start,panel_rows,block_m:tl.constexpr,block_k:tl.constexpr,use_tf32x3:tl.constexpr):
batch=tl.program_id(0);tile_m=tl.program_id(1);rows=tile_m*block_m+tl.arange(0,block_m);cols=tl.arange(0,32);accumulator=tl.zeros((block_m,32),tl.float32)
for current_k in tl.range(active_start,n,block_k):
inner=current_k+tl.arange(0,block_k);a=tl.load(matrix+batch*n*n+rows[:,None]*n+inner[None,:],mask=(rows[:,None]<n)&(inner[None,:]<n),other=.0);local=inner-active_start;v=tl.load(panel+batch*panel_rows*32+local[:,None]+cols[None,:]*panel_rows,mask=(local[:,None]>=0)&(local[:,None]<panel_rows),other=.0)
if use_tf32x3:accumulator+=tl.dot(a,v,input_precision='tf32x3')
else:accumulator+=tl.dot(a,v,input_precision='tf32')
tl.store(output+batch*n*32+rows[:,None]*32+cols[None,:],accumulator,mask=rows[:,None]<n)
@triton.jit
def _compact_companion_kernel(product,panel,triangular,output,n:tl.constexpr,active_start,panel_rows,block_m:tl.constexpr,block_k:tl.constexpr):
batch=tl.program_id(0);tile_m=tl.program_id(1);rows=tile_m*block_m+tl.arange(0,block_m);rank=tl.arange(0,32);h=tl.zeros((32,32),tl.float32)
for current_k in tl.range(0,panel_rows,block_k):inner=current_k+tl.arange(0,block_k);v=tl.load(panel+batch*panel_rows*32+inner[None,:]+rank[:,None]*panel_rows,mask=inner[None,:]<panel_rows,other=.0);p=tl.load(product+batch*n*32+(active_start+inner[:,None])*32+rank[None,:],mask=inner[:,None]<panel_rows,other=.0);h+=tl.dot(v,p,input_precision='ieee')
t=tl.load(triangular+batch*32*32+rank[:,None]*32+rank[None,:]);correction=tl.dot(t,h,input_precision='ieee');correction=tl.dot(correction,tl.trans(t),input_precision='ieee');p_rows=tl.load(product+batch*n*32+rows[:,None]*32+rank[None,:],mask=rows[:,None]<n,other=.0);y=tl.dot(p_rows,tl.trans(t),input_precision='ieee');local_rows=rows-active_start;v_rows=tl.load(panel+batch*panel_rows*32+local_rows[:,None]+rank[None,:]*panel_rows,mask=(local_rows[:,None]>=0)&(local_rows[:,None]<panel_rows),other=.0);companion=y-.5*tl.dot(v_rows,correction,input_precision='ieee');tl.store(output+batch*n*32+rows[:,None]*32+rank[None,:],companion,mask=rows[:,None]<n)
@triton.jit
def _compact_control_once_kernel(product,panel,triangular,control,n:tl.constexpr,active_start,panel_rows,block_k:tl.constexpr):
batch=tl.program_id(0);rank=tl.arange(0,32);h=tl.zeros((32,32),tl.float32)
for current_k in tl.range(0,panel_rows,block_k):inner=current_k+tl.arange(0,block_k);v=tl.load(panel+batch*panel_rows*32+inner[None,:]+rank[:,None]*panel_rows,mask=inner[None,:]<panel_rows,other=.0);p=tl.load(product+batch*n*32+(active_start+inner[:,None])*32+rank[None,:],mask=inner[:,None]<panel_rows,other=.0);h+=tl.dot(v,p,input_precision='ieee')
t=tl.load(triangular+batch*32*32+rank[:,None]*32+rank[None,:]);correction=tl.dot(t,h,input_precision='ieee');correction=tl.dot(correction,tl.trans(t),input_precision='ieee');tl.store(control+batch*32*32+rank[:,None]*32+rank[None,:],correction)
@triton.jit
def _compact_companion_apply_kernel(product,panel,triangular,control,output,n:tl.constexpr,active_start,panel_rows,block_m:tl.constexpr):batch=tl.program_id(0);tile_m=tl.program_id(1);rows=tile_m*block_m+tl.arange(0,block_m);rank=tl.arange(0,32);correction=tl.load(control+batch*32*32+rank[:,None]*32+rank[None,:]);t=tl.load(triangular+batch*32*32+rank[:,None]*32+rank[None,:]);p_rows=tl.load(product+batch*n*32+rows[:,None]*32+rank[None,:],mask=rows[:,None]<n,other=.0);y=tl.dot(p_rows,tl.trans(t),input_precision='ieee');local_rows=rows-active_start;v_rows=tl.load(panel+batch*panel_rows*32+local_rows[:,None]+rank[None,:]*panel_rows,mask=(local_rows[:,None]>=0)&(local_rows[:,None]<panel_rows),other=.0);companion=y-.5*tl.dot(v_rows,correction,input_precision='ieee');tl.store(output+batch*n*32+rows[:,None]*32+rank[None,:],companion,mask=rows[:,None]<n)
@triton.jit
def _compact_symmetric_update_kernel(matrix,panel,companion,output,n:tl.constexpr,active_start,panel_rows,block:tl.constexpr,use_tf32x3:tl.constexpr):
batch=tl.program_id(0);tile_m=tl.program_id(1);tile_n=tl.program_id(2);rows=tile_m*block+tl.arange(0,block);cols=tile_n*block+tl.arange(0,block);rank=tl.arange(0,32);row_mask=rows<n;col_mask=cols<n;a=tl.load(matrix+batch*n*n+rows[:,None]*n+cols[None,:],mask=row_mask[:,None]&col_mask[None,:],other=.0);local_rows=rows-active_start;local_cols=cols-active_start;vr=tl.load(panel+batch*panel_rows*32+local_rows[:,None]+rank[None,:]*panel_rows,mask=(local_rows[:,None]>=0)&(local_rows[:,None]<panel_rows),other=.0);vc=tl.load(panel+batch*panel_rows*32+local_cols[:,None]+rank[None,:]*panel_rows,mask=(local_cols[:,None]>=0)&(local_cols[:,None]<panel_rows),other=.0);wr=tl.load(companion+batch*n*32+rows[:,None]*32+rank[None,:],mask=row_mask[:,None],other=.0);wc=tl.load(companion+batch*n*32+cols[:,None]*32+rank[None,:],mask=col_mask[:,None],other=.0)
if use_tf32x3:update=tl.dot(vr,tl.trans(wc),input_precision='tf32x3');update+=tl.dot(wr,tl.trans(vc),input_precision='tf32x3')
else:update=tl.dot(vr,tl.trans(wc),input_precision='tf32');update+=tl.dot(wr,tl.trans(vc),input_precision='tf32')
tl.store(output+batch*n*n+rows[:,None]*n+cols[None,:],a-update,mask=row_mask[:,None]&col_mask[None,:])
def _compact_block_update(matrix:torch.Tensor,panel:torch.Tensor,triangular:torch.Tensor,*,active_start:int,precision:str):batch=matrix.shape[0];n=matrix.shape[-1];panel_rows=n-active_start;tiles=triton.cdiv(n,64);product=torch.empty((batch,n,BLOCK),device=matrix.device);_compact_matrix_panel_kernel[batch,tiles](matrix,panel,product,n,active_start,panel_rows,block_m=64,block_k=64,use_tf32x3=precision=='tf32x3',num_warps=8,num_stages=3);control=torch.empty((batch,32,32),device=matrix.device,dtype=torch.float32);_compact_control_once_kernel[batch,](product,panel,triangular,control,n,active_start,panel_rows,block_k=64,num_warps=8,num_stages=2);companion=torch.empty_like(product);_compact_companion_apply_kernel[batch,tiles](product,panel,triangular,control,companion,n,active_start,panel_rows,block_m=64,num_warps=8,num_stages=2);output=torch.empty_like(matrix);_compact_symmetric_update_kernel[batch,tiles,tiles](matrix,panel,companion,output,n,active_start,panel_rows,block=64,use_tf32x3=precision=='tf32x3',num_warps=8,num_stages=2);return output
def _panel_source(n:int=N):
names=[];wrappers=[]
for active_start in range(BLOCK,n,BLOCK):rows=n-active_start;name=f"blocked_qr_n{n}_r{rows}_b32";names.append(name);wrappers.append(f'''
extern "C" __global__ __cluster_dims__(2, 1, 1)
__launch_bounds__(128, 1)
void {name}(const float* input, float* output, float* tau,
float* v_fp32, __half* v_fp16) {{
qr2_gau_panel_body<{rows}, 32, {n}, 4>(
input, output, tau, v_fp32, v_fp16);
}}
''')
names_tuple=tuple(names);source=_N1024_GAU_PANEL_SOURCE+''.join(wrappers);return _fast_only_cuda_kernels(source,names_tuple),names_tuple
@memo(maxsize=3)
def _panel_kernels(n:int=N):source,names=_panel_source(n);image=_fast_nvrtc_compile(source,names[0]);return tuple(CUDAKernel(image,name)for name in names)
@triton.jit
def _e1332_reconstruct_final_panel(factored,panel,n:tl.constexpr):program=tl.program_id(0);batch=program//4;item=(program-batch*4)*256+tl.arange(0,256);row=item//32;column=item-row*32;value=tl.load(factored+batch*n*n+row*n+column,mask=item<1024,other=.0);value=tl.where(row>column,value,tl.where(row==column,1.,.0));tl.store(panel+batch*1024+row+column*32,value,mask=item<1024)
@torch.no_grad()
def _e1332_dense_to_band32(matrix:torch.Tensor,*,precision:str,save_panels:bool):
batch,n,_=matrix.shape;row_counts=tuple(range(n-BLOCK,0,-BLOCK));offsets=[];total=0
for rows in row_counts:offsets.append(total);total+=batch*rows*BLOCK
vectors=torch.empty(total,device=matrix.device,dtype=torch.float32);half_scratch=torch.empty(batch*row_counts[0]*BLOCK,device=matrix.device,dtype=torch.float16);triangulars=torch.empty((len(row_counts),batch,BLOCK,BLOCK),device=matrix.device,dtype=torch.float32);product=torch.empty((batch,n,BLOCK),device=matrix.device,dtype=torch.float32);companion=torch.empty_like(product);reduced=matrix;qr_output=torch.empty_like(matrix);tau=torch.empty((batch,n),device=matrix.device,dtype=torch.float32);panels=[];tiles=triton.cdiv(n,64);device_index=matrix.device.index
if device_index is None:device_index=torch.cuda.current_device()
for(panel_id,panel_start)in enumerate(range(0,n-BLOCK,BLOCK)):
active_start=panel_start+BLOCK;rows=n-active_start;v32=torch.as_strided(vectors,(batch,rows,BLOCK),(rows*BLOCK,1,rows),offsets[panel_id]);v16=torch.as_strided(half_scratch,(batch,rows,BLOCK),(rows*BLOCK,1,rows),0);_panel_kernels(n)[panel_id].launch(grid=(batch*2,1,1),block=(128,1,1),shared_mem=(rows*(BLOCK//2)+BLOCK)*4+BLOCK*8,args=(reduced[:,active_start:,panel_start:panel_start+BLOCK],qr_output[:,active_start:,panel_start:panel_start+BLOCK],tau[:,panel_start:panel_start+BLOCK],v32,v16))
if rows==BLOCK:_e1332_reconstruct_final_panel[batch*4,](qr_output[:,active_start:,panel_start:panel_start+BLOCK],v32,n,num_warps=4,num_stages=1)
triangular=triangulars[panel_id];_e834_entry_sharded_gram_t32_kernel().launch(grid=(batch*4,1,1),block=(512,1,1),shared_mem=(BLOCK*(BLOCK-1)//2+2*BLOCK*BLOCK+16*16)*4,args=(v32,tau[:,panel_start:panel_start+BLOCK],triangular,rows,n));_compact_matrix_panel_kernel[batch,tiles](reduced,v32,product,n,active_start,rows,block_m=64,block_k=64,use_tf32x3=precision=='tf32x3',num_warps=8,num_stages=3);_compact_companion_kernel[batch,tiles](product,v32,triangular,companion,n,active_start,rows,block_m=64,block_k=64,num_warps=8,num_stages=2);pair_rows,pair_columns=_e549_upper_tile_pairs(device_index,tiles,active_start//64);_e549_symmetric_update_upper_kernel[batch,pair_rows.numel()](reduced,v32,companion,reduced,pair_rows,pair_columns,n,active_start,rows,block=64,use_tf32x3=precision=='tf32x3',num_warps=8,num_stages=2)
if save_panels:panels.append((active_start,v32,triangular))
return reduced,panels
@torch.no_grad()
def dense_to_band32(matrix:torch.Tensor,*,precision:str='tf32x3',save_panels:bool=True,symmetric_update:bool=False,cluster_gram:bool=False):
if matrix.ndim!=3 or matrix.shape[-1]!=matrix.shape[-2]:raise ValueError('expected batched square matrices')
if matrix.dtype!=torch.float32 or not matrix.is_cuda:raise ValueError('expected CUDA FP32 input')
if not matrix.is_contiguous():matrix=matrix.contiguous()
batch,n,_=matrix.shape
if n%BLOCK:raise ValueError('band32 reduction requires a multiple-of-32 dimension')
if symmetric_update and cluster_gram and n==1024:return _e1332_dense_to_band32(matrix,precision=precision,save_panels=save_panels)
reduced=matrix.clone();qr_output=torch.empty_like(matrix);tau=torch.empty((batch,n),device=matrix.device,dtype=torch.float32);panels=[]
for(panel_id,panel_start)in enumerate(range(0,n-BLOCK,BLOCK)):
active_start=panel_start+BLOCK;rows=n-active_start;v32=torch.empty((batch,BLOCK,rows),device=matrix.device,dtype=torch.float32).transpose(1,2);v16=torch.empty((batch,BLOCK,rows),device=matrix.device,dtype=torch.float16).transpose(1,2);_panel_kernels(n)[panel_id].launch(grid=(batch*2,1,1),block=(128,1,1),shared_mem=(rows*(BLOCK//2)+BLOCK)*4+BLOCK*8,args=(reduced[:,active_start:,panel_start:panel_start+BLOCK],qr_output[:,active_start:,panel_start:panel_start+BLOCK],tau[:,panel_start:panel_start+BLOCK],v32,v16))
if rows==BLOCK:factored=qr_output[:,active_start:,panel_start:panel_start+BLOCK];identity=torch.eye(BLOCK,device=matrix.device).expand(batch,-1,-1);reconstructed=torch.tril(factored,diagonal=-1)+identity;v32=reconstructed.mT.contiguous().mT
triangular=torch.empty((batch,BLOCK,BLOCK),device=matrix.device,dtype=torch.float32);gram_kernel=_e834_entry_sharded_gram_t32_kernel()if cluster_gram else _e197_fused_gram_t32_kernel();gram_kernel.launch(grid=(batch*4 if cluster_gram else batch,1,1),block=(512,1,1),shared_mem=(BLOCK*(BLOCK-1)//2+2*BLOCK*BLOCK+16*16)*4 if cluster_gram else(2*BLOCK*BLOCK+16*16)*4,args=(v32,tau[:,panel_start:panel_start+BLOCK],triangular,rows,n));update_fn=_e549_symmetric_compact_block_update if symmetric_update else _compact_block_update;reduced=update_fn(reduced,v32,triangular,active_start=active_start,precision=precision)
if save_panels:panels.append((active_start,v32,triangular))
return reduced,panels
REPAIR_NAME_320='repair_twisted320_adjacent';REPAIR_SOURCE_320=REPAIR_SOURCE.replace('repair_twisted256_adjacent',REPAIR_NAME_320).replace('constexpr int N = 256;','constexpr int N = 320;').replace('constexpr int PITCH = 257;','constexpr int PITCH = 321;').replace('__cluster_dims__(4, 1, 1)','__cluster_dims__(5, 1, 1)').replace('const int matrix = blockIdx.x >> 2;','const int matrix = blockIdx.x / 5;').replace('owner < 4','owner < 5');REPAIR_SOURCE_320=REPAIR_SOURCE_320.replace(' if (tid == 0) {\n float dot = 0.0f;\n for (int row = 0; row < ROWS; ++row)\n dot += q[row * PITCH + column - 1]\n * q[row * PITCH + column];\n partial[0] = dot;\n }',' if (tid < 32) {\n float dot = q[tid * PITCH + column - 1]\n * q[tid * PITCH + column]\n + q[(tid + 32) * PITCH + column - 1]\n * q[(tid + 32) * PITCH + column];\n dot += __shfl_down_sync(0xffffffffu, dot, 16);\n dot += __shfl_down_sync(0xffffffffu, dot, 8);\n dot += __shfl_down_sync(0xffffffffu, dot, 4);\n dot += __shfl_down_sync(0xffffffffu, dot, 2);\n dot += __shfl_down_sync(0xffffffffu, dot, 1);\n if (tid == 0) partial[0] = dot;\n }').replace(' if (tid < ROWS)\n q[tid * PITCH + column] =\n (q[tid * PITCH + column]\n - dot * q[tid * PITCH + column - 1]) * inverse;\n cluster.sync();',' if (tid < ROWS)\n q[tid * PITCH + column] =\n (q[tid * PITCH + column]\n - dot * q[tid * PITCH + column - 1]) * inverse;\n __syncthreads();');REPAIR_SOURCE_320=REPAIR_SOURCE_320.replace(' float* partial = q + ROWS * PITCH;\n float* partial0 = cluster.map_shared_rank(partial, 0);',' float* partial = q + ROWS * PITCH;\n float* repair_flags = partial + 1;\n float* partial0 = cluster.map_shared_rank(partial, 0);').replace(' }\n cluster.sync();\n\n // Ordered adjacent MGS repairs',' }\n const float spectrum_span = values[(long long)matrix * N + N - 1]\n - values[(long long)matrix * N];\n for (int column = tid + 1; column < N; column += blockDim.x) {\n const float gap = values[(long long)matrix * N + column]\n - values[(long long)matrix * N + column - 1];\n repair_flags[column] =\n gap < gap_threshold * fmaxf(spectrum_span, 1.0e-30f);\n }\n cluster.sync();\n\n // Ordered adjacent MGS repairs').replace(' const float gap = values[(long long)matrix * N + column]\n - values[(long long)matrix * N + column - 1];\n const float span = values[(long long)matrix * N + N - 1]\n - values[(long long)matrix * N];\n if (gap < gap_threshold * fmaxf(span, 1.0e-30f)) {',' if (repair_flags[column] != 0.0f) {');_BAND16_REPAIR512_NAME='band16_repair_twisted512_cluster8';_BAND16_REPAIR512_SOURCE=REPAIR_SOURCE_320.replace(REPAIR_NAME_320,_BAND16_REPAIR512_NAME).replace('constexpr int N = 320;','constexpr int N = 512;').replace('constexpr int PITCH = 321;','constexpr int PITCH = 513;').replace('__cluster_dims__(5, 1, 1)','__cluster_dims__(8, 1, 1)').replace('const int matrix = blockIdx.x / 5;','const int matrix = blockIdx.x >> 3;').replace('owner < 5','owner < 8').replace('float* repair_flags = partial + 1;','float* repair_flags = partial + 2;').replace(' partial[0] = dot;',' partial[1] = dot;').replace(' const float dot = partial0[0];',' const float dot = partial0[1];');_BAND16_REPAIR512_SHARED_BYTES=(64*513+2+512)*4
@memo(maxsize=1)
def _band16_repair512_kernel():source=_e2021_direct_repair_cluster_source(_BAND16_REPAIR512_SOURCE,7,3);return CUDAKernel(_fast_nvrtc_compile(source,_BAND16_REPAIR512_NAME),_BAND16_REPAIR512_NAME)
@torch.no_grad()
def _band16_repair_twisted512(vectors:torch.Tensor,values:torch.Tensor,*,gap:float):output=torch.empty_like(vectors);_band16_repair512_kernel().launch((vectors.shape[0]*8,1,1),(128,1,1),(vectors,values,output,vectors.shape[0],float(gap)),shared_mem=_BAND16_REPAIR512_SHARED_BYTES);return output
_E146_DIRECT320_N=320;_E146_DIRECT320_TRI=_E146_DIRECT320_N*(_E146_DIRECT320_N+1)//2;_E147_REDUCE320_NAME='e147_reduce320_singlecta';_E147_SOLVE320_NAME='e147_solve320_cluster2';_E1769_SOLVE320_PDL_NAME='e1769_solve320_cluster2_pdl_wy';_E147_REDUCE320_SHARED_BYTES=(_E146_DIRECT320_TRI+2*8*_E146_DIRECT320_N+_E146_DIRECT320_N+32)*4;_E147_SOLVE320_SHARED_BYTES=(_E146_DIRECT320_TRI+5*_E146_DIRECT320_N+16)*4
def _e147_reduce320_source():
source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=320,TRI=N*(N+1)/2,B=8;').replace('__launch_bounds__(768,1)','__launch_bounds__(640,1)').replace(_E102_REDUCE160_NAME,_E147_REDUCE320_NAME).replace(_E185_N176_BLOCK4_REDUCE_NAME,_E147_REDUCE320_NAME).replace('float total=tid<24?scratch[tid]:0.f;','float total=tid<20?scratch[tid]:0.f;').replace('total=tid<24?scratch[tid]:0.f;','total=tid<20?scratch[tid]:0.f;').replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=16,ROWS=40;').replace('constexpr int G=4;','constexpr int G=2;').replace('float*coeff_a=scratch+32,*coeff_b=coeff_a+B;','float*coeff_a=xvec,*coeff_b=coeff_a+B;');entry=' int mid=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;'
if source.count(entry)!=1:raise RuntimeError('n320 reducer entry template changed')
source=source.replace(entry,entry+'\n if(tid==0) asm volatile("griddepcontrol.launch_dependents;":::);',1);tail=' timers[mid*5+1]=clock64();\n }\n}';tail_ready=' timers[mid*5+1]=clock64();\n }\n __threadfence();\n __syncthreads();\n if(tid==0) timers[mid*5+4]=1ULL;\n}'
if source.count(tail)!=1:raise RuntimeError('n320 reducer tail template changed')
return source.replace(tail,tail_ready,1)
def _e147_solve320_source():
source=_e107_solve160_source().replace('constexpr int N = 160;','constexpr int N = 320;').replace('constexpr int HALF = 80;','constexpr int HALF = 160;').replace('constexpr int SORT_N = 256;','constexpr int SORT_N = 512;').replace('__launch_bounds__(384, 4)','__launch_bounds__(512, 1)').replace(_E107_SOLVE160_NAME,_E147_SOLVE320_NAME).replace('step < 24','step < 23');anchor=' const int tid = threadIdx.x;\n const int first_row = rank * HALF;';wait=' const int tid = threadIdx.x;\n if (rank == 0 && tid == 0) {\n volatile unsigned long long* ready = timers + matrix_id * 5 + 4;\n while (*ready == 0ULL) __nanosleep(64);\n }\n cluster.sync();\n const int first_row = rank * HALF;'
if source.count(anchor)!=1:raise RuntimeError('n320 solve ready template changed')
return source.replace(anchor,wait,1)
@memo(maxsize=1)
def _e147_reduce320_kernel():return CUDAKernel(_fast_nvrtc_compile(_e147_reduce320_source(),_E147_REDUCE320_NAME),_E147_REDUCE320_NAME)
@memo(maxsize=1)
def _e1769_solve320_pdl_kernel():
source=_parallelize_n176_adjacent_mgs(_e147_solve320_source())
if source.count('local_eigen += 16')!=1:raise RuntimeError('n320 PDL parallel MGS stride template changed')
source=source.replace('local_eigen += 16','local_eigen += 32',1);anchor=' cluster.sync();\n const int first_row = rank * HALF;';replacement=' cluster.sync();\n if (tid == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);\n const int first_row = rank * HALF;'
if source.count(anchor)!=1:raise RuntimeError('n320 PDL post-ready template changed')
source=source.replace(anchor,replacement,1).replace(_E147_SOLVE320_NAME,_E1769_SOLVE320_PDL_NAME,1);source=_e2023_direct_solve_cluster_source(source);source=_e2337_cluster_coarse_sturm(source,steps=23,probes=128,fine_steps=16,levels=8);return CUDAKernel(_fast_nvrtc_compile(source,_E1769_SOLVE320_PDL_NAME),_E1769_SOLVE320_PDL_NAME)
@torch.no_grad()
def _e146_direct_eigh320(matrix:torch.Tensor):batch=matrix.shape[0];vectors=torch.empty_like(matrix);values=torch.empty((batch,_E146_DIRECT320_N),device=matrix.device,dtype=torch.float32);saved_reflectors=torch.empty_like(matrix);diagonal=torch.empty_like(values);off_diagonal=torch.empty_like(values);timers=torch.zeros((batch,5),device=matrix.device,dtype=torch.int64);_e147_reduce320_kernel().launch((batch,1,1),(640,1,1),(matrix,saved_reflectors,diagonal,off_diagonal,timers),shared_mem=_E147_REDUCE320_SHARED_BYTES);_e1769_solve320_pdl_kernel().launch_pdl((batch*2,1,1),(512,1,1),(saved_reflectors,diagonal,off_diagonal,vectors,values,timers),shared_mem=_E147_SOLVE320_SHARED_BYTES);panel_total=(_E146_DIRECT320_N-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=matrix.device,dtype=matrix.dtype);_householder_compact_t16_kernel[batch,panel_total](saved_reflectors,triangular,_E146_DIRECT320_N,panel_total,panel_width=32,block_rows=32,num_warps=4,num_stages=2,launch_pdl=True);return _e1762_tensor_wy_apply_saved_(vectors,saved_reflectors,triangular),values
@torch.no_grad()
def owned_eigh_n320(matrix:torch.Tensor):return _e146_direct_eigh320(matrix)
@torch.no_grad()
def truncated_range_eigh(matrix:torch.Tensor,*,rank:int=320,power_steps:int=2,range_passes:int=1,tail_values:str='rayleigh',projected_solver:str='exact',split_sign:int=5,split_range:int=8,final_newton:bool=True,newton_precision:str='highest',range_orthogonalization:str='cqr'):
batch,n,_=matrix.shape
if n!=512 or rank!=320:raise ValueError('prototype currently supports n=512, rank=320')
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
eye=torch.eye(n,device=matrix.device).expand(batch,-1,-1)
if power_steps:
basis=matrix@matrix[:,:,:rank]
if range_orthogonalization=='normalize':basis=normalize_columns_(basis)
elif range_orthogonalization=='cqr':basis=cholesky_orthonormalize(basis,passes=range_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision='high')
else:raise ValueError("range_orthogonalization must be 'cqr' or 'normalize'")
for _ in range(1,power_steps):
basis=matrix@(matrix@basis)
if range_orthogonalization=='normalize':basis=normalize_columns_(basis)
else:basis=cholesky_orthonormalize(basis,passes=range_passes,ridge=1e-06,final_ridge=1e-08,inverse_precision='high')
else:basis=eye[:,:,:rank].clone()
complete=mixed_qr_active(basis.contiguous(),complete=True,module=None);top=complete[:,:,:rank];tail=complete[:,:,rank:];projected=top.mT@matrix@top
if projected_solver=='exact':top_values,rotation=torch.linalg.eigh(projected)
elif projected_solver=='owned320':rotation,top_values=owned_eigh_n320(projected)
elif projected_solver=='split160':rotation,top_values=spectral_split_eigh_2048(projected,sign_iterations=split_sign,range_iterations=split_range,reorthogonalize_every=4,lanczos_steps=3,boundary_width=16,newton_steps=1,newton_precision='highest',rayleigh_values=True,split_levels=1)
elif projected_solver=='manual80':rotation,top_values=dense512_manual_eigh(projected,root_sign=split_sign,root_range=split_range,root_reorth=3,child_sign=split_sign+2,child_range=split_range+2,child_reorth=4,boundary=20,root_lanczos=3,child_lanczos=5,power_range=False,child_power_range=False,final_newton_steps=1,leaf_solver='eigh')
else:raise ValueError(projected_solver)
top=top@rotation
if tail_values=='rayleigh':rest_values=(tail*(matrix@tail)).sum(dim=1);vectors=torch.cat((top,tail),dim=-1);values=torch.cat((top_values,rest_values),dim=-1)
elif tail_values=='zero':vectors,values=_e159_dense512_merge(top,tail,top_values)
else:raise ValueError(tail_values)
if final_newton:torch.set_float32_matmul_precision(newton_precision);gram=vectors.mT@vectors;gram.mul_(-.505);gram.diagonal(dim1=-2,dim2=-1).add_(1.505);vectors=vectors@gram;torch.set_float32_matmul_precision('high')
if tail_values!='zero':values,order=values.sort(dim=-1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1))
return vectors,values
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def mixed_dense_fast(data:torch.Tensor,*,hybrid_solver=None):return truncated_range_eigh(data,power_steps=1,projected_solver='owned320',tail_values='zero',final_newton=True,newton_precision='high',range_orthogonalization='normalize')
_E208_REDUCE256_NAME='e208_dense1024_reduce256_block4';_E208_SOLVE256_NAME='e208_dense1024_twisted256_diagonal_s26';_E208_REDUCE256_SHARED_BYTES=(256*257//2+2*4*256+256+32+2*4)*4
@memo(maxsize=1)
def _e208_reduce256_kernel():source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=256,TRI=N*(N+1)/2,B=4;').replace('__launch_bounds__(768,1)','__launch_bounds__(1024,1)').replace(_E185_N176_BLOCK4_REDUCE_NAME,_E208_REDUCE256_NAME).replace('float total=tid<24?scratch[tid]:0.f;','float total=tid<32?scratch[tid]:0.f;').replace('total=tid<24?scratch[tid]:0.f;','total=tid<32?scratch[tid]:0.f;').replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=16,ROWS=64;');return CUDAKernel(_fast_nvrtc_compile(source,_E208_REDUCE256_NAME),_E208_REDUCE256_NAME)
_E1521_REDUCE256_NAME='e1521_dense1024_reduce256_halfstore_b4';_E1521_REDUCE256_SHARED_BYTES=256*257//2*2+(2*8*256+256+32+2*8)*4
@memo(maxsize=1)
def _e1521_reduce256_kernel():
source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('#include <cuda_runtime.h>','#include <cuda_runtime.h>\n#include <cuda_fp16.h>',1).replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=256,TRI=N*(N+1)/2,B=8;',1).replace('__launch_bounds__(768,1)','__launch_bounds__(512,2)',1).replace(_E185_N176_BLOCK4_REDUCE_NAME,_E1521_REDUCE256_NAME,1).replace('float total=tid<24?scratch[tid]:0.f;','float total=tid<16?scratch[tid]:0.f;',1).replace('total=tid<24?scratch[tid]:0.f;','total=tid<16?scratch[tid]:0.f;',1).replace('constexpr int G=4;','constexpr int G=2;',1).replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=16,ROWS=32;',1).replace(' extern __shared__ float sm[];\n float*a=sm,*V=a+TRI,*W=V+B*N,*xvec=W+B*N,*scratch=xvec+N;',' extern __shared__ __align__(16) unsigned char sm[];\n __half*a=(__half*)sm;\n float*V=(float*)(a+TRI),*W=V+B*N,*xvec=W+B*N,*scratch=xvec+N;',1)
for(old,new)in(('a[row*(row+1)/2+tid]=src[row*N+tid];','a[row*(row+1)/2+tid]=__float2half_rn(src[row*N+tid]);'),('float value=a[row*(row+1)/2+col];','float value=__half2float(a[row*(row+1)/2+col]);'),('float dv=a[col*(col+1)/2+col];','float dv=__half2float(a[col*(col+1)/2+col]);'),('product=fmaf(a[rr*(rr+1)/2+cc],V[p*N+gc],product);','product=fmaf(__half2float(a[rr*(rr+1)/2+cc]),V[p*N+gc],product);'),('float value0=a[packed];','float value0=__half2float(a[packed]);'),('float value1=paired?a[packed+1]:0.f;','float value1=paired?__half2float(a[packed+1]):0.f;'),('a[packed]=value0;','a[packed]=__float2half_rn(value0);'),('if(paired)a[packed+1]=value1;','if(paired)a[packed+1]=__float2half_rn(value1);'),('eout[mid*N+N-2]=a[(N-1)*N/2+N-2];','eout[mid*N+N-2]=__half2float(a[(N-1)*N/2+N-2]);'),('dout[mid*N+N-2]=a[(N-2)*(N-1)/2+N-2];','dout[mid*N+N-2]=__half2float(a[(N-2)*(N-1)/2+N-2]);'),('dout[mid*N+N-1]=a[(N-1)*N/2+N-1];','dout[mid*N+N-1]=__half2float(a[(N-1)*N/2+N-1]);')):source=source.replace(old,new,1)
return CUDAKernel(_fast_nvrtc_compile(source,_E1521_REDUCE256_NAME),_E1521_REDUCE256_NAME)
def _e208_solve256_source():source=TRIDIAG_SOURCE.replace(TRIDIAG_NAME,_E208_SOLVE256_NAME).replace('step < 30','step < 26');source=source.replace(' const float* __restrict__ matrix,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace) {',' const float* __restrict__ diagonal_input,\n const float* __restrict__ off_input,\n float* __restrict__ vectors,\n float* __restrict__ values,\n float* __restrict__ left_workspace) {',1);return source.replace(' diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;',' diagonal[eigen] = diagonal_input[(long long)batch * N + eigen];\n off[eigen] = off_input[(long long)batch * N + eigen];',1)
@memo(maxsize=1)
def _e208_solve256_kernel():return CUDAKernel(_fast_nvrtc_compile(_e208_solve256_source(),_E208_SOLVE256_NAME),_E208_SOLVE256_NAME)
_E1773_DENSE1024_SOLVE256_PDL_NAME='e1773_dense1024_twisted256_s20_fused_repair_pdl_wy';_E1789_DENSE2048_SOLVE256_PDL_NAME='e1789_twisted256_s26_fused_repair_pdl_wy'
def _e1773_fused_repair256_source(name:str,steps:int):
source=_e208_solve256_source().replace(_E208_SOLVE256_NAME,name).replace('step < 26',f"step < {steps}");anchor=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;';replacement=' const int batch = blockIdx.x;\n if (threadIdx.x == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);\n const int eigen = threadIdx.x;'
if source.count(anchor)!=1:raise RuntimeError('E1773 n256 solve entry changed')
source=source.replace(anchor,replacement,1);tail=' for (int row = 0; row < N; ++row)\n vectors[mb + (long long)row * N + eigen] *= inverse;\n}';fused=' for (int row = 0; row < N; ++row)\n vectors[mb + (long long)row * N + eigen] *= inverse;\n __syncthreads();\n int* repair_columns = reinterpret_cast<int*>(diagonal);\n if (eigen == 0) {\n const float span = values[(long long)batch * N + N - 1]\n - values[(long long)batch * N];\n int count = 0;\n for (int column = 1; column < N; ++column) {\n const float gap = values[(long long)batch * N + column]\n - values[(long long)batch * N + column - 1];\n if (gap < 1.0e-4f * fmaxf(span, 1.0e-30f))\n repair_columns[count++] = column;\n }\n bounds[0] = __int_as_float(count);\n }\n __syncthreads();\n const int repair_count = __float_as_int(bounds[0]);\n for (int item = 0; item < repair_count; ++item) {\n const int column = repair_columns[item];\n if (eigen < 4) {\n float dot = 0.0f;\n const int begin = eigen * 64;\n #pragma unroll 1\n for (int row = begin; row < begin + 64; ++row)\n dot += vectors[mb + (long long)row * N + column - 1]\n * vectors[mb + (long long)row * N + column];\n off[eigen] = dot;\n }\n __syncthreads();\n if (eigen == 0) {\n float dot = 0.0f;\n #pragma unroll\n for (int owner = 0; owner < 4; ++owner) dot += off[owner];\n bounds[1] = dot;\n }\n __syncthreads();\n const float dot = bounds[1];\n const float inverse_norm = rsqrtf(\n fmaxf(1.0f - dot * dot, 1.0e-12f));\n vectors[mb + (long long)eigen * N + column] =\n (vectors[mb + (long long)eigen * N + column]\n - dot * vectors[mb + (long long)eigen * N + column - 1])\n * inverse_norm;\n __syncthreads();\n }\n}'
if source.count(tail)!=1:raise RuntimeError('E1773 n256 solve tail changed')
source=source.replace(tail,fused,1);return source
@memo(maxsize=1)
def _e1773_dense1024_solve256_pdl_kernel():
source=_e1773_fused_repair256_source(_E1773_DENSE1024_SOLVE256_PDL_NAME,20);source=source.replace(' __shared__ float bounds[2];\n',' __shared__ float bounds[2];\n __shared__ int coarse_counts[64];\n',1);old_bisection=' float lower = bounds[0];\n float upper = bounds[1];\n #pragma unroll 1\n for (int step = 0; step < 20; ++step) {';coarse_bisection=' const float global_lower = bounds[0];\n const float global_upper = bounds[1];\n const float coarse_step =\n (global_upper - global_lower) * (1.0f / 65.0f);\n if (eigen < 64) {\n const float coarse_shift =\n global_lower + (eigen + 1) * coarse_step;\n float coarse_pivot = diagonal[0] - coarse_shift;\n int coarse_count = coarse_pivot < 0.0f;\n #pragma unroll 1\n for (int row = 1; row < N; ++row) {\n if (fabsf(coarse_pivot) < 1.0e-12f)\n coarse_pivot = copysignf(\n 1.0e-12f,\n coarse_pivot == 0.0f ? -1.0f : coarse_pivot);\n coarse_pivot = diagonal[row] - coarse_shift\n - off[row - 1] * off[row - 1] / coarse_pivot;\n coarse_count += coarse_pivot < 0.0f;\n }\n coarse_counts[eigen] = coarse_count;\n }\n __syncthreads();\n int low_probe = -1;\n int high_probe = 64;\n #pragma unroll\n for (int level = 0; level < 7; ++level) {\n if (high_probe - low_probe > 1) {\n const int middle = (low_probe + high_probe) >> 1;\n if (coarse_counts[middle] <= eigen) low_probe = middle;\n else high_probe = middle;\n }\n }\n float lower = low_probe < 0 ? global_lower\n : global_lower + (low_probe + 1) * coarse_step;\n float upper = high_probe >= 64 ? global_upper\n : global_lower + (high_probe + 1) * coarse_step;\n #pragma unroll 1\n for (int step = 0; step < 14; ++step) {'
if source.count(old_bisection)!=1:raise RuntimeError('E2327 dense1024 coarse Sturm anchor changed')
source=source.replace(old_bisection,coarse_bisection,1);return CUDAKernel(_fast_nvrtc_compile(source,_E1773_DENSE1024_SOLVE256_PDL_NAME),_E1773_DENSE1024_SOLVE256_PDL_NAME)
@memo(maxsize=1)
def _e1789_dense2048_solve256_pdl_kernel():source=_e1773_fused_repair256_source(_E1789_DENSE2048_SOLVE256_PDL_NAME,26);return CUDAKernel(_fast_nvrtc_compile(source,_E1789_DENSE2048_SOLVE256_PDL_NAME),_E1789_DENSE2048_SOLVE256_PDL_NAME)
@torch.no_grad()
def _e208_dense1024_leaf_eigh(leaves:torch.Tensor,*,final_newton:bool=True,newton_precision:str='highest'):
batch=leaves.shape[0];saved=torch.empty_like(leaves);diagonal=torch.empty((batch,256),device=leaves.device);off_diagonal=torch.empty_like(diagonal);timers=torch.empty((batch,5),device=leaves.device,dtype=torch.int64);_e208_reduce256_kernel().launch((batch,1,1),(1024,1,1),(leaves,saved,diagonal,off_diagonal,timers),shared_mem=_E208_REDUCE256_SHARED_BYTES);vectors=torch.empty_like(leaves);values=torch.empty_like(diagonal);workspace=torch.empty_like(leaves)
if final_newton:_e208_solve256_kernel().launch((batch,1,1),(256,1,1),(diagonal,off_diagonal,vectors,values,workspace));vectors=repair_twisted256_adjacent(vectors,values,gap_threshold=.0001);vectors=_tensor_wy_backtransform_saved_(vectors,saved)
else:panel_total=(256-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=leaves.device,dtype=torch.float32);_e1789_dense2048_solve256_pdl_kernel().launch((batch,1,1),(256,1,1),(diagonal,off_diagonal,vectors,values,workspace));_householder_compact_t16_kernel[batch,panel_total](saved,triangular,256,panel_total,panel_width=32,block_rows=32,num_warps=4,num_stages=2,launch_pdl=True);vectors=_e1762_tensor_wy_apply_saved_(vectors,saved,triangular)
if final_newton:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(newton_precision);gram=vectors.mT@vectors;gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;torch.set_float32_matmul_precision(previous)
return values,vectors
@torch.no_grad()
def _dense1024_leaf_eigh(leaves:torch.Tensor):return _e208_dense1024_leaf_eigh(leaves)
@torch.no_grad()
def _e507_dense1024_leaf_eigh_skip_newton(leaves:torch.Tensor):batch=leaves.shape[0];saved=torch.empty_like(leaves);diagonal=torch.empty((batch,256),device=leaves.device);off_diagonal=torch.empty_like(diagonal);timers=torch.empty((batch,5),device=leaves.device,dtype=torch.int64);_e1521_reduce256_kernel().launch((batch,1,1),(512,1,1),(leaves,saved,diagonal,off_diagonal,timers),shared_mem=_E1521_REDUCE256_SHARED_BYTES);vectors=torch.empty_like(leaves);values=torch.empty_like(diagonal);workspace=torch.empty_like(leaves);panel_total=(256-2+31)//32;triangular=torch.empty((batch,panel_total,32,32),device=leaves.device,dtype=torch.float32);_e1773_dense1024_solve256_pdl_kernel().launch((batch,1,1),(256,1,1),(diagonal,off_diagonal,vectors,values,workspace));_householder_compact_t16_kernel[batch,panel_total](saved,triangular,256,panel_total,panel_width=32,block_rows=32,num_warps=4,num_stages=2,launch_pdl=True);vectors=_e1762_tensor_wy_apply_saved_(vectors,saved,triangular,use_tf32=True);return values,vectors
@torch.no_grad()
def _dense_1024_specialized(data:torch.Tensor,*,low_precision_sign:bool=False,matrix_scale:torch.Tensor|None=None,leaf_eigh_fn=None,child_complete_qr:bool=False,root_lanczos_stats=None,certify_output:bool=True,nonic_sign:bool=False):
previous_precision=torch.get_float32_matmul_precision()
try:
if leaf_eigh_fn is None:leaf_eigh_fn=_dense1024_leaf_eigh
pure_dense_route=leaf_eigh_fn is _e507_dense1024_leaf_eigh_skip_newton;root_cqr_call=0;child_cqr_call=0
def root_orthogonalize(matrix:torch.Tensor,**kwargs):
nonlocal root_cqr_call;kwargs['gram_precision']='highest'if root_cqr_call==0 else'high';root_cqr_call+=1
if pure_dense_route:return _e1074_cholesky_orthonormalize512_pdl(matrix,trsm_fn=_e2225_tcgen_trsm256,**kwargs)
return _e084_cholesky_orthonormalize512(matrix,trsm_fn=_e2225_tcgen_trsm256,**kwargs)
def child_orthogonalize(matrix:torch.Tensor,**kwargs):nonlocal child_cqr_call;kwargs['gram_precision']='highest'if child_cqr_call==0 else'high';child_cqr_call+=1;return _e084_cholesky_orthonormalize256(matrix,trsm_fn=_e2225_tcgen_trsm256,**kwargs)
root_stats_fn=_e189_fused_n1024_lanczos3
if root_lanczos_stats is not None:
def root_stats_fn(matrix:torch.Tensor,*,steps:int=3,probes:int=2):return root_lanczos_stats
vectors,values=dense512_manual_eigh(data,root_sign=6,root_range=6,root_reorth=3,child_sign=8,child_range=8,child_reorth=8,boundary=48,root_lanczos=3,child_lanczos=5,power_range=False,child_power_range=False,final_newton_steps=1,leaf_eigh_fn=leaf_eigh_fn,boundary_eigh_fn=_e204_rankdef_n96_eigh,root_orthogonalize_fn=root_orthogonalize,child_orthogonalize_fn=child_orthogonalize,root_lanczos_stats_fn=root_stats_fn,child_lanczos_stats_fn=_e192_fused_n512_lanczos5,root_low_precision_sign=low_precision_sign,child_low_precision_sign=low_precision_sign,root_asymmetric_high_iterations=1,root_asymmetric_direct_complement=False,root_asymmetric_normalize_high=True,root_asymmetric_complete_qr=True,root_asymmetric_skip_final_low_cqr=pure_dense_route,root_asymmetric_direct_qr_workspace=pure_dense_route,child_asymmetric_high_iterations=0,child_asymmetric_direct_complement=True,child_asymmetric_normalize_high=False,child_asymmetric_complete_qr=child_complete_qr,child_asymmetric_skip_final_low_cqr=pure_dense_route,root_nonic_sign=nonic_sign,child_nonic_sign=nonic_sign);certificate_windows=(1,4,-16,17),(1,2,-16,17),(3,4,-16,17)
if pure_dense_route:certificate_windows+=(7,8,-4,5),
if certify_output:return _certify_eigh_output(data,vectors,values,certificate_windows,matrix_scale=matrix_scale)
return vectors,values
finally:torch.set_float32_matmul_precision(previous_precision)
def _e1055_apply_b32_diamond(source:str):
shared=' __shared__ float chain_tile[2][32][32];\n';shared_new=shared+' __shared__ volatile int ready[32];\n volatile int* ready0 = (volatile int*)cluster.map_shared_rank(\n (int*)ready, 0);\n if (tid < 32) ready[tid] = 0;\n cluster.sync();\n'
if source.count(shared)!=1:raise RuntimeError('B32 diamond shared anchor changed')
source=source.replace(shared,shared_new,1);column_header=' const int remaining = n - output_column - 2;\n const int block_count = (remaining + bandwidth - 1) / bandwidth;\n const int first_support = output_column + 1;\n';column_header_new=column_header+' const bool diamond_first =\n output_column < 1021\n && output_column + 1 < n - 2\n && block_count >= 3;\n const int previous_blocks =\n (n - (output_column - 1) - 2 + bandwidth - 1) / bandwidth;\n const bool diamond_second =\n output_column > 0\n && output_column < 1022\n && previous_blocks >= 3;\n const int epoch = output_column;\n'
if source.count(column_header)!=1:raise RuntimeError('B32 diamond column anchor changed')
source=source.replace(column_header,column_header_new,1);chain_begin=source.index(' // Generate v_0 directly from the target band column.\n');chain_edge=' cluster.sync();\n\n // Only rank 0 owns the prefix row/column.\n';chain_end=source.index(chain_edge,chain_begin);chain=source[chain_begin:chain_end];block_loop=' for (int block = 1; block < block_count; ++block) {';waited_loop=block_loop+'\n const int dependency = min(block + 1, previous_blocks - 1);\n if (tid == 0) {\n while (atomicAdd((int*)&ready0[dependency], 0) < epoch)\n __nanosleep(64);\n }\n __syncthreads();'
if chain.count(block_loop)!=1:raise RuntimeError('B32 diamond chain anchor changed')
waited_chain=chain.replace(block_loop,waited_loop,1);chain_replacement=' if (!diamond_second) {\n'+chain+' cluster.sync();\n } else {\n if (rank == 0) {\n if (tid == 0) {\n while (atomicAdd((int*)&ready0[1], 0) < epoch)\n __nanosleep(64);\n }\n __syncthreads();\n'+waited_chain+' }\n cluster.sync();\n for (int index = tid; index < block_count * storage_band;\n index += blockDim.x) {\n const int block = index / storage_band;\n const int item = index - block * storage_band;\n vectors[block][item] = reflectors[\n reflector_base\n + ((long long)output_column * max_blocks + block)\n * storage_band + item];\n }\n __syncthreads();\n }\n\n // Only rank 0 owns the prefix row/column.\n';source=source[:chain_begin]+chain_replacement+source[chain_end+len(chain_edge):];old_header=' for (int block_wave = rank; block_wave < block_count; block_wave += (block_wave + 13 < block_count ? 16 : 4)) {\n const bool quad = block_wave + 13 < block_count;\n const int subgroup = quad ? (warp >> 3) : 0;\n const int local_warp = quad ? (warp & 7) : warp;\n const int local_warps = quad ? 8 : warps;\n const int block = block_wave + subgroup * 4;';new_header=' const int full_groups = (block_count - 1) >> 2;\n const int tail_count = block_count - (full_groups << 2);\n const int consumer_groups = full_groups - 2;\n const bool tail_diamond_first =\n diamond_first && block_count < 17;\n for (int phase = 0; phase < 6; ++phase) {\n bool active = false;\n bool quad = false;\n bool pair = false;\n bool strided = false;\n int block_base = 0;\n if (tail_diamond_first) {\n const int interior_quads = (block_count - 1) >> 2;\n const int scalar_tail =\n block_count - (interior_quads << 2);\n if (rank > 0 && phase == 0) {\n const int group = rank - 1;\n active = group < interior_quads;\n quad = active;\n block_base = group << 2;\n } else if (rank > 0 && phase == 1) {\n const int item = rank - 1;\n active = item < scalar_tail;\n block_base = (interior_quads << 2) + item;\n } else if (rank == 1 && phase == 2) {\n active = scalar_tail == 4;\n block_base = (interior_quads << 2) + 3;\n }\n } else if (diamond_first) {\n if (phase == 0) {\n active = true;\n pair = true;\n strided = true;\n block_base = rank;\n } else if (rank > 0) {\n if (phase <= 2) {\n const int group =\n 2 + (phase - 1) * 3 + (rank - 1);\n active = group < full_groups;\n quad = active;\n block_base = group << 2;\n }\n if (!active) {\n int item = -1;\n if (consumer_groups == 5 && rank == 3)\n item = phase - 2;\n else if (consumer_groups == 4 && rank >= 2)\n item = 2 * (phase - 2) + (rank - 2);\n else if (consumer_groups == 3)\n item = 3 * (phase - 2) + (rank - 1);\n else if (consumer_groups == 2 && rank == 3)\n item = phase - 1;\n active = item >= 0 && item < tail_count;\n block_base = (full_groups << 2) + item;\n }\n }\n } else {\n int production_wave = rank;\n #pragma unroll\n for (int prior = 0; prior < phase; ++prior) {\n if (production_wave < block_count)\n production_wave +=\n production_wave + 13 < block_count ? 16 : 4;\n }\n active = production_wave < block_count;\n quad = active && production_wave + 13 < block_count;\n strided = quad;\n block_base = production_wave;\n }\n if (!active) continue;\n const int subgroup =\n quad ? (warp >> 3) : (pair ? (warp >> 4) : 0);\n const int local_warp =\n quad ? (warp & 7) : (pair ? (warp & 15) : warp);\n const int local_warps = quad ? 8 : (pair ? 16 : warps);\n const int tile_copies = quad ? 4 : (pair ? 2 : 1);\n const int block = strided\n ? block_base + subgroup * 4\n : (quad ? block_base + subgroup : block_base);'
if source.count(old_header)!=1:raise RuntimeError('B32 diamond schedule anchor changed')
source=source.replace(old_header,new_header,1);source=source.replace('(quad ? 4 : 1)','tile_copies');publish_tail=' if (diamond_first) {\n __syncthreads();\n __threadfence();\n __syncthreads();\n if (local_warp == 0 && lane == 0)\n atomicExch((int*)&ready0[block], epoch + 1);\n __syncthreads();\n }\n';tail_continue=' continue;\n }\n const int next_start'
if source.count(tail_continue)!=1:raise RuntimeError('B32 diamond tail anchor changed')
source=source.replace(tail_continue,publish_tail+' continue;\n }\n const int next_start',1);loop_end=' __syncthreads();\n }\n cluster.sync();\n }\n}';loop_end_new=' __syncthreads();\n if (diamond_first) {\n __threadfence();\n __syncthreads();\n if (local_warp == 0 && lane == 0)\n atomicExch((int*)&ready0[block], epoch + 1);\n __syncthreads();\n }\n }\n if (!diamond_first) cluster.sync();\n }\n}'
if source.count(loop_end)!=1:raise RuntimeError('B32 diamond convergence anchor changed')
return source.replace(loop_end,loop_end_new,1)
_E549_BAND32_N1024_CLUSTER4_NAME='e1073_band32_n1024_cluster4_frontier8_balanced';_E1760_TRIDIAG1024_PDL_NAME='e1760_tridiag1024_pdl_operator_builder';_E549_ACTIVE_SPARSE_WY_OPERATORS=32*33//2
@memo(maxsize=1)
def _e549_band32_n1024_cluster4_source():
source=_e184_n512_cluster2_source().replace('constexpr int n = 512;','constexpr int n = 1024;').replace('constexpr int max_blocks = 16;','constexpr int max_blocks = 32;').replace('__cluster_dims__(2, 1, 1)','__cluster_dims__(4, 1, 1)').replace('const int batch = blockIdx.x >> 1;','const int batch = blockIdx.x / 4;').replace('block += 2','block += 4').replace(_E184_N512_CLUSTER2_NAME,_E549_BAND32_N1024_CLUSTER4_NAME);marker=' // Only rank 0 owns the prefix row/column.\n'
if source.count(marker)!=1:raise RuntimeError('E549 reflector-chain/update boundary changed')
source=source.replace(marker,' cluster.sync();\n\n'+marker,1);source=source.replace('__shared__ float left_projection[storage_band];','__shared__ float left_projection[4][storage_band];').replace('__shared__ float right_projection[storage_band];','__shared__ float right_projection[4][storage_band];').replace('__shared__ float scalar;','__shared__ float scalar[4];\n __shared__ float diagonal_tile[4][32][32];\n __shared__ float upper_tile[4][32][32];\n __shared__ float lower_tile[4][32][32];\n __shared__ float chain_tile[2][32][32];');header=' for (int block = rank; block < block_count; block += 4) {'
if source.count(header)!=1:raise RuntimeError('E549 owned-block loop changed')
start=source.index(header);end=source.rindex(' cluster.sync();');segment=source[start:end];segment=segment.replace(header,' for (int block_wave = rank; block_wave < block_count; block_wave += (block_wave + 13 < block_count ? 16 : 4)) {\n const bool quad = block_wave + 13 < block_count;\n const int subgroup = quad ? (warp >> 3) : 0;\n const int local_warp = quad ? (warp & 7) : warp;\n const int local_warps = quad ? 8 : warps;\n const int block = block_wave + subgroup * 4;');segment=segment.replace('for (int row = warp; row < length; row += warps)','for (int row = local_warp; row < length; row += local_warps)').replace('for (int row = warp; row < next_length; row += warps)','for (int row = local_warp; row < next_length; row += local_warps)').replace('if (warp == 0)','if (local_warp == 0)').replace('right_projection[','right_projection[subgroup][').replace('left_projection[','left_projection[subgroup][').replace('if (lane == 0) scalar = projection;','if (lane == 0) scalar[subgroup] = projection;').replace('const float diagonal_scalar = scalar;','const float diagonal_scalar = scalar[subgroup];').replace('if (lane == 0) scalar = cross;','if (lane == 0) scalar[subgroup] = cross;').replace('const float cross_scalar = scalar;','const float cross_scalar = scalar[subgroup];');diagonal_marker=' // Diagonal tile H_k A_kk H_k.\n for (int row = local_warp; row < length; row += local_warps) {';diagonal_prefetch=' {\n #pragma unroll\n for (int q = 0; q < (quad ? 4 : 1); ++q) {\n const int row = local_warp + q * local_warps;\n if (row < length && lane < length) {\n const int dst = __cvta_generic_to_shared(\n &diagonal_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + row) * n\n + support_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n }\n asm volatile("cp.async.commit_group;");\n }\n\n // The neighboring cross tile is independent of the diagonal\n // update. Launch its copies now and consume them only after the\n // diagonal math, hiding their long-scoreboard latency.\n if (block + 1 < block_count) {\n const int prefetch_next_start = support_start + bandwidth;\n const int prefetch_next_length =\n min(bandwidth, n - prefetch_next_start);\n #pragma unroll\n for (int q = 0; q < (quad ? 4 : 1); ++q) {\n const int row = local_warp + q * local_warps;\n if (row < length && lane < prefetch_next_length) {\n const int dst = __cvta_generic_to_shared(\n &upper_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + row) * n\n + prefetch_next_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n if (row < prefetch_next_length && lane < length) {\n const int dst = __cvta_generic_to_shared(\n &lower_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(prefetch_next_start + row) * n\n + support_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n }\n asm volatile("cp.async.commit_group;");\n }\n if (block + 1 < block_count) {\n asm volatile("cp.async.wait_group 1;" ::: "memory");\n } else {\n asm volatile("cp.async.wait_group 0;" ::: "memory");\n }\n __syncwarp();\n\n // Diagonal tile H_k A_kk H_k.\n for (int row = local_warp; row < length; row += local_warps) {'
if segment.count(diagonal_marker)!=1:raise RuntimeError('E836 diagonal prefetch anchor changed')
segment=segment.replace(diagonal_marker,diagonal_prefetch,1);diagonal_projection=' dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j] * vectors[block][j];';diagonal_projection_staged=' const float tile_value =\n diagonal_tile[subgroup][row][j];\n dot += tile_value * vectors[block][j];'
if segment.count(diagonal_projection)!=1:raise RuntimeError('E836 diagonal projection anchor changed')
segment=segment.replace(diagonal_projection,diagonal_projection_staged,1);diagonal_update=' matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j]\n -= w_row * vectors[block][j] + u_row * w_j;';diagonal_update_staged=' const float tile_value =\n diagonal_tile[subgroup][row][j];\n matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j]\n = tile_value\n - (w_row * vectors[block][j] + u_row * w_j);'
if segment.count(diagonal_update)!=1:raise RuntimeError('E836 diagonal update anchor changed')
segment=segment.replace(diagonal_update,diagonal_update_staged,1);cross_marker=' // Both projections of the neighboring off-diagonal tile. Read\n // upper and lower copies row-wise so every transaction coalesces.\n for (int row = local_warp; row < length; row += local_warps) {';cross_prefetch=' asm volatile(\n "cp.async.wait_group 0;" ::: "memory");\n __syncwarp();\n\n // Both projections of the neighboring off-diagonal tile. Read\n // upper and lower copies row-wise so every transaction coalesces.\n for (int row = local_warp; row < length; row += local_warps) {'
if segment.count(cross_marker)!=1:raise RuntimeError('E836 cross prefetch anchor changed')
segment=segment.replace(cross_marker,cross_prefetch,1);upper_projection=' dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j] * vectors[block + 1][j];';upper_projection_staged=' const float tile_value =\n upper_tile[subgroup][row][j];\n dot += tile_value * vectors[block + 1][j];'
if segment.count(upper_projection)!=1:raise RuntimeError('E836 upper projection anchor changed')
segment=segment.replace(upper_projection,upper_projection_staged,1);lower_projection=' dot += matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j] * vectors[block][j];';lower_projection_staged=' const float tile_value =\n lower_tile[subgroup][row][j];\n dot += tile_value * vectors[block][j];'
if segment.count(lower_projection)!=1:raise RuntimeError('E836 lower projection anchor changed')
segment=segment.replace(lower_projection,lower_projection_staged,1);upper_update=' const float updated = matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j]\n - 2.0f * vectors[block][row] * left_projection[subgroup][j]';upper_update_staged=' const float updated =\n upper_tile[subgroup][row][j]\n - 2.0f * vectors[block][row] * left_projection[subgroup][j]'
if segment.count(upper_update)!=1:raise RuntimeError('E836 upper update anchor changed')
segment=segment.replace(upper_update,upper_update_staged,1);lower_update=' const float updated = matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j]\n - 2.0f * left_projection[subgroup][row] * vectors[block][j]';lower_update_staged=' const float updated =\n lower_tile[subgroup][row][j]\n - 2.0f * left_projection[subgroup][row] * vectors[block][j]'
if segment.count(lower_update)!=1:raise RuntimeError('E836 lower update anchor changed')
segment=segment.replace(lower_update,lower_update_staged,1);source=source[:start]+segment+source[end:];old_chain=' // The first column of H_k is all that is needed to create the next\n // bulge: x_{k+1} = A[S_{k+1},S_k] H_k e_0.\n for (int block = 1; block < block_count; ++block) {\n const int previous_start = first_support + (block - 1) * bandwidth;\n const int support_start = previous_start + bandwidth;\n const int length = min(bandwidth, n - support_start);\n for (int row = warp; row < length; row += warps) {\n float dot = 0.0f;\n for (int j = lane; j < bandwidth; j += 32) {\n const float h_column = (j == 0 ? 1.0f : 0.0f)\n - 2.0f * vectors[block - 1][j]\n * vectors[block - 1][0];\n dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + previous_start + j] * h_column;\n }\n dot = warp_sum(dot);\n if (lane == 0) vectors[block][row] = dot;\n }\n __syncthreads();\n if (warp == 0) {\n const int i0 = lane;\n const int i1 = lane + 32;\n float x0 = i0 < length ? vectors[block][i0] : 0.0f;\n float x1 = i1 < length ? vectors[block][i1] : 0.0f;\n float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);\n x0 *= inverse_norm;\n x1 *= inverse_norm;\n vectors[block][i0] = x0;\n vectors[block][i1] = x1;\n const long long destination = reflector_base\n + ((long long)output_column * max_blocks + block)\n * storage_band;\n reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;\n }\n __syncthreads();\n }\n';new_chain=' // The first column of H_k is all that is needed to create the next\n // bulge: x_{k+1} = A[S_{k+1},S_k] H_k e_0. Stage block 1, then\n // overlap every following tile with the current reflector reduction.\n if (block_count > 1) {\n const int support_start = first_support + bandwidth;\n const int length = min(bandwidth, n - support_start);\n {\n const int load_row = (warp << 1) + (lane >> 4);\n const int load_col = lane & 15;\n if (load_row < length) {\n const int dst0 = __cvta_generic_to_shared(\n &chain_tile[1][load_row][load_col]);\n const int dst1 = __cvta_generic_to_shared(\n &chain_tile[1][load_row][load_col + 16]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + load_row) * n\n + first_support + load_col;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst0), "l"(src));\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst1), "l"(src + 16));\n }\n }\n asm volatile("cp.async.commit_group;");\n }\n for (int block = 1; block < block_count; ++block) {\n const int previous_start = first_support + (block - 1) * bandwidth;\n const int support_start = previous_start + bandwidth;\n const int length = min(bandwidth, n - support_start);\n if (block + 1 < block_count) {\n const int next_previous = first_support + block * bandwidth;\n const int next_support = next_previous + bandwidth;\n const int next_length = min(bandwidth, n - next_support);\n {\n const int load_row = (warp << 1) + (lane >> 4);\n const int load_col = lane & 15;\n if (load_row < next_length) {\n const int dst0 = __cvta_generic_to_shared(\n &chain_tile[(block + 1) & 1][load_row][load_col]);\n const int dst1 = __cvta_generic_to_shared(\n &chain_tile[(block + 1) & 1][load_row][load_col + 16]);\n const float* src = matrices + matrix_base\n + (long long)(next_support + load_row) * n\n + next_previous + load_col;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst0), "l"(src));\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst1), "l"(src + 16));\n }\n }\n asm volatile("cp.async.commit_group;");\n asm volatile("cp.async.wait_group 1;" ::: "memory");\n } else {\n asm volatile("cp.async.wait_group 0;" ::: "memory");\n }\n __syncwarp();\n {\n const int half_lane = lane & 15;\n const int row = (warp << 1) + (lane >> 4);\n if (row < length) {\n const int j0 = half_lane;\n const int j1 = half_lane + 16;\n const float v0 = vectors[block - 1][0];\n const float h0 = (j0 == 0 ? 1.0f : 0.0f)\n - 2.0f * vectors[block - 1][j0] * v0;\n const float h1 = -2.0f * vectors[block - 1][j1] * v0;\n const float term0 = __fmul_rn(\n chain_tile[block & 1][row][j0], h0);\n const float term1 = __fmul_rn(\n chain_tile[block & 1][row][j1], h1);\n float dot = __fadd_rn(term0, term1);\n const unsigned mask = lane < 16\n ? 0x0000ffffu : 0xffff0000u;\n dot += __shfl_down_sync(mask, dot, 8, 16);\n dot += __shfl_down_sync(mask, dot, 4, 16);\n dot += __shfl_down_sync(mask, dot, 2, 16);\n dot += __shfl_down_sync(mask, dot, 1, 16);\n if (half_lane == 0) vectors[block][row] = dot;\n }\n }\n __syncthreads();\n if (warp == 0) {\n const int i0 = lane;\n const int i1 = lane + 32;\n float x0 = i0 < length ? vectors[block][i0] : 0.0f;\n float x1 = i1 < length ? vectors[block][i1] : 0.0f;\n float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);\n x0 *= inverse_norm;\n x1 *= inverse_norm;\n vectors[block][i0] = x0;\n vectors[block][i1] = x1;\n const long long destination = reflector_base\n + ((long long)output_column * max_blocks + block)\n * storage_band;\n reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;\n }\n __syncthreads();\n }\n'
if source.count(old_chain)!=1:raise RuntimeError('E843 reflector-chain anchor changed')
source=source.replace(old_chain,new_chain,1);source=_e1055_apply_b32_diamond(source);chain_shared=' __shared__ float chain_tile[2][32][32];'
if source.count(chain_shared)!=1:raise RuntimeError('B32 hoisted h-column shared anchor changed')
source=source.replace(chain_shared,chain_shared+'\n __shared__ float chain_h[32];',1);initial_vector=' vectors[0][i0] = x0;\n vectors[0][i1] = x1;';initial_h=initial_vector+'\n const float h_v0 = __shfl_sync(0xffffffffu, x0, 0);\n chain_h[i0] = (i0 == 0 ? 1.0f : 0.0f)\n - 2.0f * x0 * h_v0;'
if source.count(initial_vector)!=2:raise RuntimeError('B32 initial hoisted h-column anchor changed')
source=source.replace(initial_vector,initial_h);following_vector=' vectors[block][i0] = x0;\n vectors[block][i1] = x1;';following_h=following_vector+'\n const float h_v0 = __shfl_sync(0xffffffffu, x0, 0);\n chain_h[i0] = (i0 == 0 ? 1.0f : 0.0f)\n - 2.0f * x0 * h_v0;'
if source.count(following_vector)!=2:raise RuntimeError('B32 following hoisted h-column anchor changed')
source=source.replace(following_vector,following_h);repeated_h=' const float v0 = vectors[block - 1][0];\n const float h0 = (j0 == 0 ? 1.0f : 0.0f)\n - 2.0f * vectors[block - 1][j0] * v0;\n const float h1 = -2.0f * vectors[block - 1][j1] * v0;';hoisted_h=' const float h0 = chain_h[j0];\n const float h1 = chain_h[j1];'
if source.count(repeated_h)!=2:raise RuntimeError('B32 repeated h-column anchor changed')
source=source.replace(repeated_h,hoisted_h);chain_marker=' // Generate v_0 directly from the target band column.';half_barrier='if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");';cursor=0
for expected_barriers in(3,4):
chain_start=source.find(chain_marker,cursor);chain_end=source.find(' cluster.sync();',chain_start)
if chain_start<0 or chain_end<0:raise RuntimeError('B32 half-CTA chain section changed')
chain_section=source[chain_start:chain_end]
if chain_section.count('__syncthreads();')!=expected_barriers:raise RuntimeError('B32 half-CTA barrier count changed')
chain_section=chain_section.replace('__syncthreads();',half_barrier);source=source[:chain_start]+chain_section+source[chain_end:];cursor=chain_start+len(chain_section)
if source.count(half_barrier)!=7:raise RuntimeError('B32 half-CTA barrier rewrite changed')
return source
_E1208_B32_MAIN_NAME='e1208_b32_lower_main';_E1208_B32_TAIL_NAME='e1208_b32_production_tail';_E1208_B32_BOUNDARY=958
def _e1661_b32_prefix_overlap_source(source:str):
prefix_begin=source.index(' // Only rank 0 owns the prefix row/column.\n');prefix_end_marker=' const int full_groups = (block_count - 1) >> 2;';prefix_end=source.index(prefix_end_marker,prefix_begin);prefix=source[prefix_begin:prefix_end];old_owner=' if (rank == 0 && warp == 0) {'
if prefix.count(old_owner)!=1 or not prefix.endswith(' __syncthreads();\n\n'):raise RuntimeError('B32 prefix overlap anchor changed')
prefix=prefix.replace(' // Only rank 0 owns the prefix row/column.\n',' // Prefix update overlaps the following reflector producer.\n',1).replace(old_owner,' if (rank == 0 && warp == 16) {\n if (lane == 0) {\n while (atomicAdd((int*)&prefix_ready, 0) < epoch + 1)\n __nanosleep(64);\n }\n __syncwarp();',1);prefix=prefix[:-len(' __syncthreads();\n\n')];ready=' __shared__ volatile int ready[32];\n volatile int* ready0';ready_new=' __shared__ volatile int ready[32];\n __shared__ volatile int prefix_ready;\n volatile int* ready0';initialize=' if (tid < 32) ready[tid] = 0;\n cluster.sync();';initialize_new=' if (tid < 32) ready[tid] = 0;\n if (tid == 0) prefix_ready = 0;\n cluster.sync();'
if source.count(ready)!=1 or source.count(initialize)!=1:raise RuntimeError('B32 prefix readiness anchor changed')
source=source.replace(ready,ready_new,1).replace(initialize,initialize_new,1);marker=' // Generate v_0 directly from the target band column.\n';chain=' // The first column of H_k is all that is needed to create the next\n';publish_anchor=' chain_h[i0] = (i0 == 0 ? 1.0f : 0.0f)\n - 2.0f * x0 * h_v0;';publish=publish_anchor+'\n if (rank == 0 && lane == 0) {\n __threadfence_block();\n atomicExch((int*)&prefix_ready, epoch + 1);\n }';barrier=' if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");\n';cursor=0
for _ in range(2):
begin=source.find(marker,cursor);end=source.find(chain,begin)
if begin<0 or end<0:raise RuntimeError('B32 v0 overlap section changed')
section=source[begin:end]
if section.count(publish_anchor)!=1 or section.count(barrier)!=1:raise RuntimeError('B32 v0 publication anchor changed')
section=section.replace(publish_anchor,publish,1).replace(barrier,prefix+barrier,1);source=source[:begin]+section+source[end:];cursor=begin+len(section)+len(chain)
if source.find(marker,cursor)>=0:raise RuntimeError('B32 unexpected v0 overlap section')
prefix_begin=source.index(' // Only rank 0 owns the prefix row/column.\n');prefix_end=source.index(prefix_end_marker,prefix_begin);return source[:prefix_begin]+' // Prefix update completed concurrently in rank0 warp16.\n'+source[prefix_end:]
def _e1669_b32_dedicated_producer_quad_source(source:str):
declaration=' __shared__ volatile int ready[32];\n __shared__ volatile int prefix_ready;\n volatile int* ready0 = (volatile int*)cluster.map_shared_rank(\n (int*)ready, 0);\n if (tid < 32) ready[tid] = 0;\n if (tid == 0) prefix_ready = 0;\n cluster.sync();\n';replacement=' __shared__ volatile int ready[32];\n __shared__ volatile int produced[32];\n __shared__ volatile int prefix_ready;\n volatile int* ready0 = (volatile int*)cluster.map_shared_rank(\n (int*)ready, 0);\n volatile int* produced0 = (volatile int*)cluster.map_shared_rank(\n (int*)produced, 0);\n float* vectors0 = cluster.map_shared_rank(&vectors[0][0], 0);\n if (tid < 32) {\n ready[tid] = 0;\n produced[tid] = 0;\n }\n if (tid == 0) prefix_ready = 0;\n cluster.sync();\n'
if source.count(declaration)!=1:raise RuntimeError('E1669 shared publication anchor changed')
source=source.replace(declaration,replacement,1);begin=source.index(' } else {\n if (rank == 0) {');end_marker='\n // Prefix update completed concurrently in rank0 warp16.';end=source.index(end_marker,begin);producer=source[begin:end];barrier=' if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");\n';first_publish=barrier+' if (tid == 0) atomicExch((int*)&produced0[0], epoch + 1);\n'
if producer.count(barrier)!=4:raise RuntimeError(f"E1669 producer barrier count {producer.count(barrier)}")
producer=producer.replace(barrier,first_publish,1);tail=' }\n if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");\n }\n\n }\n cluster.sync();\n for (int index = tid; index < block_count * storage_band;\n index += blockDim.x) {\n const int block = index / storage_band;\n const int item = index - block * storage_band;\n vectors[block][item] = reflectors[\n reflector_base\n + ((long long)output_column * max_blocks + block)\n * storage_band + item];\n }\n __syncthreads();\n }\n';tail_new=' }\n if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");\n if (tid == 0)\n atomicExch((int*)&produced0[block], epoch + 1);\n }\n\n __syncthreads();\n\n }\n }\n'
if producer.count(tail)!=1:raise RuntimeError('E1669 producer tail changed')
producer=producer.replace(tail,tail_new,1);source=source[:begin]+producer+source[end:];phase_start=source.index(' for (int phase = 0; phase < 6; ++phase)');anchor=' const int support_start = first_support + block * bandwidth;\n const int length = min(bandwidth, n - support_start);\n\n {\n';location=source.find(anchor,phase_start)
if location<0 or source.find(anchor,location+1)>=0:raise RuntimeError('E1669 phase staging anchor changed')
staged=' const int support_start = first_support + block * bandwidth;\n const int length = min(bandwidth, n - support_start);\n\n if (diamond_second) {\n const bool load_next = block + 1 < block_count\n && (strided || !quad || subgroup == 3);\n const int required = load_next ? block + 1 : block;\n if (local_warp == 0 && lane == 0) {\n while (atomicAdd((int*)&produced0[required], 0)\n < epoch + 1)\n __nanosleep(64);\n }\n __syncthreads();\n for (int item = local_warp * 32 + lane;\n item < storage_band;\n item += local_warps * 32) {\n vectors[block][item] =\n vectors0[block * storage_band + item];\n if (load_next)\n vectors[block + 1][item] =\n vectors0[(block + 1) * storage_band + item];\n }\n __syncthreads();\n }\n\n {\n';source=source[:location]+source[location:].replace(anchor,staged,1);mapping_end=' if (!active) continue;\n const int subgroup =\n';mapping_new=' if (diamond_second) {\n active = false;\n quad = false;\n pair = false;\n strided = false;\n if (rank > 0) {\n const int full_quads = (block_count - 1) >> 2;\n const int quad_phases = (full_quads + 2) / 3;\n const int group = phase * 3 + rank - 1;\n if (phase < quad_phases && group < full_quads) {\n active = true;\n quad = true;\n block_base = group << 2;\n } else if (phase >= quad_phases) {\n const int item = (phase - quad_phases) * 3 + rank - 1;\n const int remainder = block_count - (full_quads << 2);\n active = item < remainder;\n block_base = (full_quads << 2) + item;\n }\n }\n }\n if (!active) continue;\n const int subgroup =\n'
if source.count(mapping_end)!=1:raise RuntimeError('E1669 schedule override anchor changed')
source=source.replace(mapping_end,mapping_new,1);waits=(' if (tid == 0) {\n while (atomicAdd((int*)&ready0[1], 0) < epoch)\n __nanosleep(64);\n }\n __syncthreads();',' if (tid == 0) {\n while (atomicAdd((int*)&ready0[1], 0) < epoch)\n __nanosleep(64);\n }\n if (warp == 0)\n asm volatile("fence.acquire.gpu;" ::: "memory");\n __syncthreads();'),(' if (tid == 0) {\n while (atomicAdd((int*)&ready0[dependency], 0) < epoch)\n __nanosleep(64);\n }\n if (warp < 16)',' if (tid == 0) {\n while (atomicAdd((int*)&ready0[dependency], 0) < epoch)\n __nanosleep(64);\n }\n if (warp < 16)\n asm volatile("fence.acquire.gpu;" ::: "memory");\n if (warp < 16)'),(' if (local_warp == 0 && lane == 0) {\n while (atomicAdd((int*)&produced0[required], 0)\n < epoch + 1)\n __nanosleep(64);\n }\n __syncthreads();',' if (local_warp == 0 && lane == 0) {\n while (atomicAdd((int*)&produced0[required], 0)\n < epoch + 1)\n __nanosleep(64);\n }\n if (local_warp < 2)\n asm volatile("fence.acquire.cluster;" ::: "memory");\n __syncthreads();')
for(old,new)in waits:
if source.count(old)!=1:raise RuntimeError('E1669 acquire-fence anchor changed')
source=source.replace(old,new,1)
return source
def _e1677_b32_selected_barrier_source(source:str):
transforms=(' }\n __syncthreads();\n\n if (block + 1 >= block_count) {',' }\n\n if (block + 1 >= block_count) {'),(' if (local_warp == 0 && lane == 0)\n atomicExch((int*)&ready0[block], epoch + 1);\n __syncthreads();\n }\n continue;',' if (local_warp == 0 && lane == 0)\n atomicExch((int*)&ready0[block], epoch + 1);\n }\n continue;'),(' if (local_warp == 0 && lane == 0)\n atomicExch((int*)&ready0[block], epoch + 1);\n __syncthreads();\n }\n }\n if (!diamond_first)',' if (local_warp == 0 && lane == 0)\n atomicExch((int*)&ready0[block], epoch + 1);\n }\n }\n if (!diamond_first)')
for(old,new)in transforms:
if source.count(old)!=1:raise RuntimeError('E1677 barrier-prune anchor changed')
source=source.replace(old,new,1)
return source
def _e1208_b32_main_source():
source=_e549_band32_n1024_cluster4_source().replace(_E549_BAND32_N1024_CLUSTER4_NAME,_E1208_B32_MAIN_NAME,1);source=source.replace(' __shared__ float diagonal_tile[4][32][32];',' __shared__ float diagonal_tile[4][32][33];',1).replace(' __shared__ float lower_tile[4][32][32];',' __shared__ float lower_tile[4][32][33];',1).replace(' __shared__ float upper_tile[4][32][32];\n','',1);old='for (int output_column = 0; output_column < n - 2; ++output_column)';source=source.replace(old,f"for (int output_column = 0; output_column < {_E1208_B32_BOUNDARY}; ++output_column)",1);predicate=' && block_count >= 3;';source=source.replace(predicate,predicate[:-1]+f"\n && output_column + 1 < {_E1208_B32_BOUNDARY};",1);upper_copy=' if (row < length && lane < prefetch_next_length) {\n const int dst = __cvta_generic_to_shared(\n &upper_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + row) * n\n + prefetch_next_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n';source=source.replace(upper_copy,'',1);first_wait=' if (block + 1 < block_count) {\n asm volatile("cp.async.wait_group 1;" ::: "memory");\n } else {\n asm volatile("cp.async.wait_group 0;" ::: "memory");\n }\n __syncwarp();\n\n // Diagonal tile H_k A_kk H_k.';first_publish=first_wait.replace(' __syncwarp();',' __syncthreads();');source=source.replace(first_wait,first_publish,1);cross_wait=' asm volatile(\n "cp.async.wait_group 0;" ::: "memory");\n __syncwarp();\n\n // Both projections of the neighboring off-diagonal tile.';cross_publish=cross_wait.replace(' __syncwarp();',' __syncthreads();')
if source.count(cross_wait)!=1:raise RuntimeError('E1208 lower cross publish anchor changed')
source=source.replace(cross_wait,cross_publish,1);diagonal='diagonal_tile[subgroup][row][j]';source=source.replace(diagonal,'diagonal_tile[subgroup][row > j ? row : j][row > j ? j : row]');prefix=' for (int j = lane; j < first_length; j += 32) {\n dot += matrices[matrix_base + (long long)prefix_row * n\n + first_support + j] * vectors[0][j];\n }\n dot = warp_sum(dot);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n for (int j = lane; j < first_length; j += 32) {\n const float updated = matrices[\n matrix_base + (long long)prefix_row * n\n + first_support + j] - 2.0f * dot * vectors[0][j];\n matrices[matrix_base + (long long)prefix_row * n\n + first_support + j] = updated;\n matrices[matrix_base + (long long)(first_support + j) * n\n + prefix_row] = updated;\n }';lower_prefix=' for (int j = lane; j < first_length; j += 32) {\n dot += matrices[matrix_base\n + (long long)(first_support + j) * n\n + prefix_row] * vectors[0][j];\n }\n dot = warp_sum(dot);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n for (int j = lane; j < first_length; j += 32) {\n const long long lower = matrix_base\n + (long long)(first_support + j) * n + prefix_row;\n matrices[lower] -= 2.0f * dot * vectors[0][j];\n }';source=source.replace(prefix,lower_prefix,1);source=source.replace('upper_tile[subgroup][row][j]','lower_tile[subgroup][j][row]',1);upper_update=' for (int row = local_warp; row < length; row += local_warps) {\n for (int j = lane; j < next_length; j += 32) {\n const float updated =\n upper_tile[subgroup][row][j]\n - 2.0f * vectors[block][row] * left_projection[subgroup][j]\n - 2.0f * right_projection[subgroup][row] * vectors[block + 1][j]\n + 4.0f * vectors[block][row] * cross_scalar\n * vectors[block + 1][j];\n matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j] = updated;\n }\n }\n';source=source.replace(upper_update,'',1);tail_upper=' matrices[matrix_base\n + (long long)(support_start + j) * n\n + tail_start] = updated;\n';source=source.replace(tail_upper,'',1)
if'upper_tile'in source or source.count(diagonal)!=0:raise RuntimeError('E1208 lower-main source rewrite changed')
source=_e1677_b32_selected_barrier_source(_e1669_b32_dedicated_producer_quad_source(_e1661_b32_prefix_overlap_source(source)));initial_old=' float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);';initial_new=' const float norm2 = warp_sum(x0 * x0 + x1 * x1);\n const float norm = lane == 0 ? sqrtf(norm2) : 0.0f;\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -norm : norm) : 0.0f;\n const float reflector_norm2 = lane == 0\n ? 2.0f * norm * (norm + fabsf(x0)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n float inverse_norm = lane == 0 && reflector_norm2 > 1.0e-40f\n ? rsqrtf(reflector_norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);';following_old=' float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);';following_new=' const float norm2 = warp_sum(x0 * x0 + x1 * x1);\n const float norm = lane == 0 ? sqrtf(norm2) : 0.0f;\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -norm : norm) : 0.0f;\n const float reflector_norm2 = lane == 0\n ? 2.0f * norm * (norm + fabsf(x0)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n float inverse_norm = lane == 0 && reflector_norm2 > 1.0e-40f\n ? rsqrtf(reflector_norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);'
if source.count(initial_old)!=2 or source.count(following_old)!=2:raise RuntimeError('E1712 B32 analytic-norm anchor changed')
return source.replace(initial_old,initial_new).replace(following_old,following_new)
_e1828_b32_main_base_source=_e1208_b32_main_source
def _e1828_b32_dual_warp_source(source:str):
pointers=' volatile int* produced0 = (volatile int*)cluster.map_shared_rank(\n (int*)produced, 0);\n float* vectors0 = cluster.map_shared_rank(&vectors[0][0], 0);\n';pointers_new=pointers+' volatile int* produced1 = (volatile int*)cluster.map_shared_rank(\n (int*)produced, 1);\n float* vectors1 = cluster.map_shared_rank(&vectors[0][0], 1);\n'
if source.count(pointers)!=1:raise RuntimeError('E1828 pointer anchor changed')
source=source.replace(pointers,pointers_new,1);vectors_decl=' __shared__ float vectors[max_blocks][storage_band];\n'
if source.count(vectors_decl)!=1:raise RuntimeError('E1828 helper-vector anchor changed')
source=source.replace(vectors_decl,vectors_decl+' __shared__ float helper_vectors[max_blocks][storage_band];\n',1);loop=f" for (int output_column = 0; output_column < {_E1208_B32_BOUNDARY}; ++output_column) {{\n const int remaining = n - output_column - 2;\n";loop_new=f""" for (int local_step = 0; ; ++local_step) {{
const int producer_column = rank < 2
? rank + 2 * local_step : -1;
int output_column = rank < 2 ? producer_column : local_step;
if ((rank < 2 && producer_column > {_E1208_B32_BOUNDARY})
|| (rank >= 2 && output_column >= {_E1208_B32_BOUNDARY})) break;
const bool producer_active =
rank < 2 && producer_column < {_E1208_B32_BOUNDARY};
bool consumer_active = rank >= 2;
volatile int* produced_owner = (output_column & 1)
? produced1 : produced0;
float* vectors_owner = (output_column & 1) ? vectors1 : vectors0;
int remaining = n - output_column - 2;
"""
if source.count(loop)!=1:raise RuntimeError('E1828 loop anchor changed')
source=source.replace(loop,loop_new,1);dimensions=' const int block_count = (remaining + bandwidth - 1) / bandwidth;\n const int first_support = output_column + 1;'
if source.count(dimensions)!=1:raise RuntimeError('E1828 dimension anchor changed')
source=source.replace(dimensions,' int block_count = (remaining + bandwidth - 1) / bandwidth;\n int first_support = output_column + 1;',1);predicates=f""" const bool diamond_first =
output_column < 1021
&& output_column + 1 < n - 2
&& block_count >= 3
&& output_column + 1 < {_E1208_B32_BOUNDARY};
const int previous_blocks =
(n - (output_column - 1) - 2 + bandwidth - 1) / bandwidth;
const bool diamond_second =
output_column > 0
&& output_column < 1022
&& previous_blocks >= 3;
const int epoch = output_column;
""";predicates_new=' bool diamond_first = true;\n int previous_blocks =\n (n - (output_column - 1) - 2 + bandwidth - 1) / bandwidth;\n bool diamond_second = true;\n int epoch = output_column;\n'
if source.count(predicates)!=1:raise RuntimeError('E1828 predicate anchor changed')
source=source.replace(predicates,predicates_new,1);dedicated=' } else {\n if (rank == 0) {'
if source.count(dedicated)!=1:raise RuntimeError('E1828 producer anchor changed')
source=source.replace(dedicated,' } else {\n if (producer_active) {',1);source=source.replace('if (rank == 0 && warp == 16)','if (rank < 2 && warp == 15)').replace('if (rank == 0 && lane == 0)','if (rank < 2 && lane == 0)').replace('if (rank == 0) {','if (rank < 2) {').replace('produced0[','produced_owner[').replace('vectors0[','vectors_owner[');producer_begin=source.index(' } else {\n if (producer_active) {');producer_end=source.index('\n // Prefix update completed concurrently in rank0 warp16.',producer_begin);producer=source[producer_begin:producer_end];producer=producer.replace(' __syncthreads();',' if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");').replace(' __syncthreads();',' if (warp < 16) asm volatile("bar.sync 1, 512;" ::: "memory");');source=source[:producer_begin]+producer+source[producer_end:];handoff=' // Prefix update completed concurrently in rank0 warp16.\n';handoff_new=' // Warps16--31 update the preceding column while warps0--15\n // produce the current odd/even column.\n if (rank < 2 && warp >= 16) {\n output_column = producer_column - 1;\n consumer_active = output_column >= 0;\n produced_owner = (output_column & 1) ? produced1 : produced0;\n vectors_owner = (output_column & 1) ? vectors1 : vectors0;\n remaining = n - output_column - 2;\n block_count = (remaining + bandwidth - 1) / bandwidth;\n first_support = output_column + 1;\n previous_blocks =\n (n - (output_column - 1) - 2 + bandwidth - 1) / bandwidth;\n diamond_first = true;\n diamond_second = true;\n epoch = output_column;\n }\n\n // Prefix update completed concurrently in the producing CTA.\n'
if source.count(handoff)!=1:raise RuntimeError('E1828 handoff anchor changed')
source=source.replace(handoff,handoff_new,1);mapping=' if (rank > 0) {\n const int full_quads = (block_count - 1) >> 2;\n const int quad_phases = (full_quads + 2) / 3;\n const int group = phase * 3 + rank - 1;\n if (phase < quad_phases && group < full_quads) {\n active = true;\n quad = true;\n block_base = group << 2;\n } else if (phase >= quad_phases) {\n const int item = (phase - quad_phases) * 3 + rank - 1;\n const int remainder = block_count - (full_quads << 2);\n active = item < remainder;\n block_base = (full_quads << 2) + item;\n }\n }\n';mapping_new=' if (consumer_active && (rank >= 2 || warp >= 16)) {\n const int full_quads = (block_count - 1) >> 2;\n const int consumer_slot = rank >= 2 ? rank - 1 : 0;\n const int quad_phases = (full_quads + 2) / 3;\n const int group = phase * 3 + consumer_slot;\n if (phase < quad_phases && group < full_quads) {\n active = true;\n quad = true;\n block_base = group << 2;\n } else if (phase >= quad_phases) {\n const int item =\n (phase - quad_phases) * 3 + consumer_slot;\n const int remainder = block_count - (full_quads << 2);\n active = item < remainder;\n block_base = (full_quads << 2) + item;\n }\n }\n'
if source.count(mapping)!=1:raise RuntimeError('E1828 consumer-map anchor changed')
source=source.replace(mapping,mapping_new,1);subgroup=' const int subgroup =\n quad ? (warp >> 3) : (pair ? (warp >> 4) : 0);\n const int local_warp =\n quad ? (warp & 7) : (pair ? (warp & 15) : warp);\n const int local_warps = quad ? 8 : (pair ? 16 : warps);\n const int tile_copies = quad ? 4 : (pair ? 2 : 1);\n';subgroup_new=' const bool helper = rank < 2;\n const bool helper_quad = helper && quad;\n const int subgroup = helper_quad\n ? ((warp - 16) >> 2)\n : (helper ? 0\n : (quad ? (warp >> 3) : (pair ? (warp >> 4) : 0)));\n const int local_warp = helper_quad\n ? ((warp - 16) & 3)\n : (helper ? warp - 16\n : (quad ? (warp & 7) : (pair ? (warp & 15) : warp)));\n const int local_warps = helper_quad\n ? 4 : (helper ? 16 : (quad ? 8 : (pair ? 16 : warps)));\n const int tile_copies = helper_quad\n ? 8 : (helper ? 2 : (quad ? 4 : (pair ? 2 : 1)));\n'
if source.count(subgroup)!=1:raise RuntimeError('E1828 subgroup anchor changed')
source=source.replace(subgroup,subgroup_new,1);phase_anchor=' for (int phase = 0; phase < 6; ++phase)';phase_pointer=' float (*consumer_vectors)[storage_band] =\n rank < 2 ? helper_vectors : vectors;\n for (int phase = 0; phase < 6; ++phase)'
if source.count(phase_anchor)!=1:raise RuntimeError('E1828 consumer-vector anchor changed')
source=source.replace(phase_anchor,phase_pointer,1);phase_begin=source.index(phase_pointer);phase_end=source.index(' if (!diamond_first) cluster.sync();',phase_begin);phase=source[phase_begin:phase_end].replace('vectors[','consumer_vectors[').replace(' __syncthreads();',' if (rank < 2)\n asm volatile("bar.sync 2, 512;" ::: "memory");\n else\n __syncthreads();');source=source[:phase_begin]+phase+source[phase_end:];finish=' if (!diamond_first) cluster.sync();\n }\n}\n'
if source.count(finish)!=1:raise RuntimeError('E1828 final anchor changed')
return source.replace(finish,' if (!diamond_first) cluster.sync();\n }\n cluster.sync();\n}\n',1)
def _e1208_b32_main_source():return _e1828_b32_dual_warp_source(_e1828_b32_main_base_source())
def _e1208_b32_tail_source():source=_e549_band32_n1024_cluster4_source().replace(_E549_BAND32_N1024_CLUSTER4_NAME,_E1208_B32_TAIL_NAME,1);old='for (int output_column = 0; output_column < n - 2; ++output_column)';source=source.replace(old,f"for (int output_column = {_E1208_B32_BOUNDARY}; output_column < n - 2; ++output_column)",1);source=source.replace(' output_column > 0',f" output_column > {_E1208_B32_BOUNDARY}",1);return _e1661_b32_prefix_overlap_source(source)
_e1831_b32_main_base_source=_e1208_b32_main_source;_e1831_b32_tail_base_source=_e1208_b32_tail_source
def _e1831_add_ready_argument(source:str):
old=' float* __restrict__ matrices,\n float* __restrict__ reflectors\n) {';new=' float* __restrict__ matrices,\n float* __restrict__ reflectors,\n int* __restrict__ matrix_ready\n) {'
if source.count(old)!=1:raise RuntimeError('E1831 signature anchor changed')
return source.replace(old,new,1)
def _e1208_b32_main_source():
source=_e1831_add_ready_argument(_e1831_b32_main_base_source());end=' cluster.sync();\n}\n';replacement=' __threadfence();\n cluster.sync();\n if (rank == 0 && tid == 0) {\n atomicExch(matrix_ready + batch, 1);\n }\n cluster.sync();\n if (tid == 0)\n asm volatile("griddepcontrol.launch_dependents;" ::: "memory");\n}\n'
if source.count(end)!=1:raise RuntimeError('E1831 main completion anchor changed')
return source.replace(end,replacement,1)
_e1950_b32_main_base_source=_e1208_b32_main_source
def _e2008_redundant_b32_scalar_source(source:str,sync:str,vectors:str):
diagonal_old=f""" {sync}
if (local_warp == 0) {{
float projection = 0.0f;
if (lane < length) {{
projection += {vectors}[block][lane]
* right_projection[subgroup][lane];
}}
if (lane + 32 < length) {{
projection += {vectors}[block][lane + 32]
* right_projection[subgroup][lane + 32];
}}
projection = warp_sum(projection);
if (lane == 0) scalar[subgroup] = projection;
}}
{sync}
const float diagonal_scalar = scalar[subgroup];""";diagonal_new=f""" {sync}
float diagonal_scalar = 0.0f;
if (lane < length) {{
diagonal_scalar += {vectors}[block][lane]
* right_projection[subgroup][lane];
}}
if (lane + 32 < length) {{
diagonal_scalar += {vectors}[block][lane + 32]
* right_projection[subgroup][lane + 32];
}}
diagonal_scalar = warp_sum(diagonal_scalar);
diagonal_scalar = __shfl_sync(
0xffffffffu, diagonal_scalar, 0);""";cross_old=f""" {sync}
if (local_warp == 0) {{
float cross = 0.0f;
if (lane < length) {{
cross += {vectors}[block][lane] * right_projection[subgroup][lane];
}}
if (lane + 32 < length) {{
cross += {vectors}[block][lane + 32]
* right_projection[subgroup][lane + 32];
}}
cross = warp_sum(cross);
if (lane == 0) scalar[subgroup] = cross;
}}
{sync}
const float cross_scalar = scalar[subgroup];""";cross_new=f""" {sync}
float cross_scalar = 0.0f;
if (lane < length) {{
cross_scalar += {vectors}[block][lane]
* right_projection[subgroup][lane];
}}
if (lane + 32 < length) {{
cross_scalar += {vectors}[block][lane + 32]
* right_projection[subgroup][lane + 32];
}}
cross_scalar = warp_sum(cross_scalar);
cross_scalar = __shfl_sync(0xffffffffu, cross_scalar, 0);"""
for(old,new,label)in((diagonal_old,diagonal_new,'diagonal'),(cross_old,cross_new,'cross')):
if source.count(old)!=1:raise RuntimeError(f"E2008 B32 {label} scalar anchor changed")
source=source.replace(old,new,1)
return source
def _e1208_b32_main_source():
source=_e1950_b32_main_base_source();marker='extern "C" __global__ __cluster_dims__(4, 1, 1)';helper='static __device__ __forceinline__ void e1950_phase_sync(\n bool helper, bool quad, int subgroup) {\n const int barrier = quad ? 3 + subgroup : (helper ? 2 : 0);\n const int arrivals = quad ? (helper ? 128 : 256)\n : (helper ? 512 : 1024);\n asm volatile("bar.sync %0, %1;"\n :: "r"(barrier), "r"(arrivals) : "memory");\n}\n\n'
if source.count(marker)!=1:raise RuntimeError('E1950 B32 kernel marker changed')
source=source.replace(marker,helper+marker,1);phase_begin=source.index(' for (int phase = 0; phase < 6; ++phase)');phase_end=source.index(' if (!diamond_first) cluster.sync();',phase_begin);phase=source[phase_begin:phase_end];shared_next='const bool load_next = block + 1 < block_count\n && (strided || !quad || subgroup == 3);';private_next='const bool load_next = block + 1 < block_count;'
if phase.count(shared_next)!=1:raise RuntimeError('E1950 next-vector sharing anchor changed')
phase=phase.replace(shared_next,private_next,1);barrier=' if (rank < 2)\n asm volatile("bar.sync 2, 512;" ::: "memory");\n else\n __syncthreads();'
if phase.count(barrier)!=13:raise RuntimeError('E1950 phase barrier count changed')
pieces=phase.split(barrier);sync=' e1950_phase_sync(rank < 2, quad, subgroup);';phase=pieces[0]
for(index,piece)in enumerate(pieces[1:]):
if index!=1:phase+=sync
phase+=piece
source=source[:phase_begin]+phase+source[phase_end:];source=_e2008_redundant_b32_scalar_source(source,'e1950_phase_sync(rank < 2, quad, subgroup);','consumer_vectors');acquire='asm volatile("fence.acquire.gpu;" ::: "memory");'
if source.count(acquire)!=2:raise RuntimeError('B32 in-cluster acquire anchor changed')
source=source.replace(acquire,'asm volatile("fence.acquire.cluster;" ::: "memory");');phase_begin=source.index(' for (int phase = 0; phase < 6; ++phase)');phase_end=source.index(' if (!diamond_first) cluster.sync();',phase_begin);phase=source[phase_begin:phase_end]
if phase.count('__threadfence();')!=2:raise RuntimeError('B32 in-cluster release anchor changed')
phase=phase.replace('__threadfence();','asm volatile("fence.release.cluster;" ::: "memory");');source=source[:phase_begin]+phase+source[phase_end:]
if source.count('__threadfence();')!=1:raise RuntimeError('B32 final GPU release anchor changed')
handoff=' for (int item = local_warp * 32 + lane;\n item < storage_band;\n item += local_warps * 32) {\n consumer_vectors[block][item] =\n vectors_owner[block * storage_band + item];\n if (load_next)\n consumer_vectors[block + 1][item] =\n vectors_owner[(block + 1) * storage_band + item];\n }\n ';handoff_four_warp=' if (local_warp < 4) {\n const int item = (local_warp & 1) * 32 + lane;\n if (local_warp < 2) {\n consumer_vectors[block][item] =\n vectors_owner[block * storage_band + item];\n } else if (load_next) {\n consumer_vectors[block + 1][item] =\n vectors_owner[(block + 1) * storage_band + item];\n }\n }\n '
if source.count(handoff)!=1:raise RuntimeError('B32 DSM handoff anchor changed')
source=source.replace(handoff,handoff_four_warp,1);phase_loop='for (int phase = 0; phase < 6; ++phase)'
if source.count(phase_loop)!=1:raise RuntimeError('B32 dead-phase anchor changed')
return source.replace(phase_loop,'for (int phase = 0; phase < 5; ++phase)',1)
def _e1208_b32_tail_source():
source=_e1831_add_ready_argument(_e1831_b32_tail_base_source());start=f" if (tid < 32) ready[tid] = 0;\n if (tid == 0) prefix_ready = 0;\n cluster.sync();\n\n for (int output_column = {_E1208_B32_BOUNDARY}; output_column < n - 2; ++output_column) {{";replacement=f''' if (tid < 32) ready[tid] = 0;
if (tid == 0) prefix_ready = 0;
if (rank == 0 && tid == 0) {{
while (atomicAdd(matrix_ready + batch, 0) == 0)
__nanosleep(64);
asm volatile("fence.acquire.gpu;" ::: "memory");
}}
cluster.sync();
constexpr int mirror_begin = {_E1208_B32_BOUNDARY};
constexpr int mirror_width = n - mirror_begin;
for (int item = rank * blockDim.x + tid;
item < mirror_width * mirror_width;
item += 4 * blockDim.x) {{
const int row = item / mirror_width;
const int column = item - row * mirror_width;
if (row > column)
matrices[matrix_base
+ (long long)(mirror_begin + column) * n
+ mirror_begin + row] = matrices[matrix_base
+ (long long)(mirror_begin + row) * n
+ mirror_begin + column];
}}
cluster.sync();
for (int output_column = {_E1208_B32_BOUNDARY}; output_column < n - 2; ++output_column) {{'''
if source.count(start)!=1:raise RuntimeError('E1831 tail startup anchor changed')
source=source.replace(start,replacement,1);return _e2008_redundant_b32_scalar_source(source,'__syncthreads();','vectors')
def _e2024_direct_b32_cluster_source(source:str):
helper='\ntemplate <typename T>\n__device__ __forceinline__ T* e2024_map_shared_rank(T* pointer, int rank) {\n unsigned long long remote;\n asm volatile("mapa.u64 %0, %1, %2;"\n : "=l"(remote)\n : "l"((unsigned long long)pointer), "r"(rank));\n return reinterpret_cast<T*>(remote);\n}\n';rewrites=('#include <cooperative_groups.h>\n',helper),('namespace cg = cooperative_groups;\n',''),(' cg::cluster_group cluster = cg::this_cluster();\n',''),(' const int rank = cluster.block_rank();',' const int rank = blockIdx.x & 3;')
for(old,new)in rewrites:
if source.count(old)!=1:raise RuntimeError('E2024 direct B32 cluster PTX anchor changed')
source=source.replace(old,new)
if source.count('cluster.map_shared_rank(')<1 or source.count('cluster.sync();')<1:raise RuntimeError('E2024 direct B32 DSM anchor changed')
source=source.replace('cluster.map_shared_rank(','e2024_map_shared_rank(');return source.replace('cluster.sync();','asm volatile("barrier.cluster.arrive.aligned;" ::: "memory");\n asm volatile("barrier.cluster.wait.aligned;" ::: "memory");')
@memo(maxsize=1)
def _e1208_b32_main_kernel():source=_e2024_direct_b32_cluster_source(_e1208_b32_main_source());return CUDAKernel(_fast_nvrtc_compile(source,_E1208_B32_MAIN_NAME),_E1208_B32_MAIN_NAME)
@memo(maxsize=1)
def _e1208_b32_tail_kernel():source=_e2024_direct_b32_cluster_source(_e1208_b32_tail_source());return CUDAKernel(_fast_nvrtc_compile(source,_E1208_B32_TAIL_NAME),_E1208_B32_TAIL_NAME)
@torch.no_grad()
def _e549_band32_to_tridiagonal_n1024(matrix:torch.Tensor):batch=matrix.shape[0];output=matrix;reflectors=torch.empty((batch,1024,32,64),device=matrix.device,dtype=torch.float32);matrix_ready=torch.zeros((batch,),device=matrix.device,dtype=torch.int32);_e1208_b32_main_kernel().launch((batch*4,1,1),(1024,1,1),(output,reflectors,matrix_ready));_e1208_b32_tail_kernel().launch((batch*4,1,1),(1024,1,1),(output,reflectors,matrix_ready));return output,reflectors
@memo(maxsize=1)
def _e1760_tridiag1024_pdl_kernel():
source=_TRIDIAG_TWISTED512_SOURCE.replace('constexpr int N = 512;','constexpr int N = 1024;').replace('__launch_bounds__(512, 1)','__launch_bounds__(256, 1)').replace('for (int step = 0; step < 30; ++step)','for (int step = 0; step < 42; ++step)').replace('tridiag_twisted512',_E1760_TRIDIAG1024_PDL_NAME);original=' const int batch = blockIdx.x;\n const int eigen = threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n diagonal[eigen] = matrix[mb + (long long)eigen * N + eigen];\n off[eigen] = eigen + 1 < N\n ? matrix[mb + (long long)(eigen + 1) * N + eigen] : 0.0f;\n __syncthreads();\n\n if (eigen == 0) {';split=' constexpr int SHARDS = 4;\n constexpr int EIGEN_PER_SHARD = 256;\n const int batch = blockIdx.x / SHARDS;\n if (threadIdx.x == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);\n const int shard = blockIdx.x - batch * SHARDS;\n const int eigen = shard * EIGEN_PER_SHARD + threadIdx.x;\n const long long mb = (long long)batch * N * N;\n __shared__ float diagonal[N];\n __shared__ float off[N];\n __shared__ float bounds[2];\n for (int item = threadIdx.x; item < N; item += blockDim.x) {\n diagonal[item] = matrix[mb + (long long)item * N + item];\n off[item] = item + 1 < N\n ? matrix[mb + (long long)(item + 1) * N + item] : 0.0f;\n }\n __syncthreads();\n\n if (threadIdx.x == 0) {'
if source.count(original)!=1:raise RuntimeError('E1760 tridiagonal split4 source anchor changed')
source=source.replace(original,split,1);source=source.replace(' __shared__ float bounds[2];\n',' __shared__ float bounds[2];\n __shared__ int coarse_counts[EIGEN_PER_SHARD];\n',1);old_bisection=' float lower = bounds[0];\n float upper = bounds[1];\n #pragma unroll 1\n for (int step = 0; step < 42; ++step) {';coarse_bisection=' const float global_lower = bounds[0];\n const float global_upper = bounds[1];\n const float coarse_step =\n (global_upper - global_lower) * (1.0f / 257.0f);\n const float coarse_shift =\n global_lower + (threadIdx.x + 1) * coarse_step;\n float coarse_pivot = diagonal[0] - coarse_shift;\n int coarse_count = coarse_pivot < 0.0f;\n #pragma unroll 1\n for (int row = 1; row < N; ++row) {\n if (fabsf(coarse_pivot) < 1.0e-12f)\n coarse_pivot = copysignf(\n 1.0e-12f,\n coarse_pivot == 0.0f ? -1.0f : coarse_pivot);\n coarse_pivot = diagonal[row] - coarse_shift\n - off[row - 1] * off[row - 1] / coarse_pivot;\n coarse_count += coarse_pivot < 0.0f;\n }\n coarse_counts[threadIdx.x] = coarse_count;\n __syncthreads();\n\n int low_probe = -1;\n int high_probe = EIGEN_PER_SHARD;\n #pragma unroll\n for (int level = 0; level < 9; ++level) {\n if (high_probe - low_probe > 1) {\n const int middle = (low_probe + high_probe) >> 1;\n if (coarse_counts[middle] <= eigen) low_probe = middle;\n else high_probe = middle;\n }\n }\n float lower = low_probe < 0 ? global_lower\n : global_lower + (low_probe + 1) * coarse_step;\n float upper = high_probe >= EIGEN_PER_SHARD ? global_upper\n : global_lower + (high_probe + 1) * coarse_step;\n #pragma unroll 1\n for (int step = 0; step < 31; ++step) {'
if source.count(old_bisection)!=1:raise RuntimeError('E2307 coarse Sturm anchor changed')
source=source.replace(old_bisection,coarse_bisection,1);return CUDAKernel(_fast_nvrtc_compile(source,_E1760_TRIDIAG1024_PDL_NAME),_E1760_TRIDIAG1024_PDL_NAME)
@triton.jit
def _e549_build_band32_sparse_wy1024(reflectors,operators):
batch=tl.program_id(0);packed_id=tl.program_id(1);root=tl.sqrt(tl.cast(8*packed_id+1,tl.float32));window_id=tl.cast((root-1.)*.5,tl.int32);spatial_block=packed_id-window_id*(window_id+1)//2;local_row=tl.arange(0,64);rank=tl.arange(0,32);high_column=1021-window_id*32;column=high_column-rank;active=column>=0;reflector_lane=local_row[:,None]-(31-rank[None,:]);support=column[None,:]+1+spatial_block*32;block_count=(1022-column+31)//32;reflector_mask=active[None,:]&(spatial_block<block_count[None,:])&(reflector_lane>=0)&(reflector_lane<32)&(support+reflector_lane<1024);address=reflectors+batch*1024*32*64+(column[None,:]*32+spatial_block)*64+reflector_lane;vectors=tl.load(address,mask=reflector_mask,other=.0);vectors16=vectors.to(tl.float16);gram=tl.dot(tl.trans(vectors16),vectors16,out_dtype=tl.float32);row=rank[:,None];col=rank[None,:];triangular=tl.where((row==col)&active[:,None],2.,.0).to(tl.float32)
for current_column in tl.static_range(1,32):gram_column=tl.sum(tl.where(col==current_column,gram,.0),axis=1);overlap=tl.sum(tl.where(col<current_column,triangular*gram_column[None,:],.0),axis=1);triangular=tl.where((col==current_column)&(row<current_column),(-2.*overlap)[:,None],triangular)
weighted=tl.dot(vectors16,tl.trans(triangular.to(tl.float16)),out_dtype=tl.float32);product=tl.dot(weighted.to(tl.float16),tl.trans(vectors16),out_dtype=tl.float32);operator_row=local_row[:,None];operator_col=local_row[None,:];operator=tl.where(operator_row==operator_col,1.,.0)-product;operator_base=((batch*32+window_id)*32+spatial_block)*64*64;tl.store(operators+operator_base+operator_row*64+operator_col,operator)
@triton.jit
def _e549_band32_sparse_wy1024_replay(q,operators,block_cols:tl.constexpr,signal_dependents:tl.constexpr):
if signal_dependents:tl_cuda.gdc_launch_dependents()
batch=tl.program_id(0);column_tile=tl.program_id(1);local_row=tl.arange(0,64);rank=tl.arange(0,32);columns=column_tile*block_cols+tl.arange(0,block_cols);column_mask=columns<1024;q_base=batch*1024*1024;operator_row=local_row[:,None];operator_col=local_row[None,:]
for window_id in tl.range(0,32):
high_column=1021-window_id*32;low_column=tl.maximum(0,high_column-31);block_limit=(1024-low_column-2+31)//32
for spatial_block in tl.range(0,block_limit):union_start=high_column-30+spatial_block*32;operator_base=((batch*32+window_id)*32+spatial_block)*64*64;operator=tl.load(operators+operator_base+operator_row*64+operator_col).to(tl.float16);rows=union_start+local_row;row_mask=(rows>=0)&(rows<1024);tile=tl.load(q+q_base+rows[:,None]*1024+columns[None,:],mask=row_mask[:,None]&column_mask[None,:],other=.0).to(tl.float16);tile=tl.dot(operator,tile,out_dtype=tl.float32);tl.store(q+q_base+rows[:,None]*1024+columns[None,:],tile,mask=row_mask[:,None]&column_mask[None,:])
@torch.no_grad()
def _e1760_tridiag_solve_build1024(matrix:torch.Tensor,reflectors:torch.Tensor):batch=matrix.shape[0];vectors=torch.empty_like(matrix);values=torch.empty((batch,1024),device=matrix.device);workspace=torch.empty_like(matrix);operators=torch.empty((batch,32,32,64,64),device=matrix.device,dtype=torch.float16);_e1760_tridiag1024_pdl_kernel().launch((batch*4,1,1),(256,1,1),(matrix,vectors,values,workspace));_e549_build_band32_sparse_wy1024[batch,_E549_ACTIVE_SPARSE_WY_OPERATORS](reflectors,operators,num_warps=2,num_stages=1,launch_pdl=True);return vectors,values,operators
@triton.jit
def _e3386_pack_cuppen_children(matrix,children,rho_output,sign_output):batch=tl.program_id(0);rows=tl.arange(0,512);source_base=batch*1024*1024;left_base=batch*512*512;right_base=(batch+tl.num_programs(0))*512*512;coupling=tl.load(matrix+source_base+512*1024+511);rho=tl.abs(coupling);left_diagonal=tl.load(matrix+source_base+rows*1024+rows);right_rows=rows+512;right_diagonal=tl.load(matrix+source_base+right_rows*1024+right_rows);left_diagonal=tl.where(rows==511,left_diagonal-rho,left_diagonal);right_diagonal=tl.where(rows==0,right_diagonal-rho,right_diagonal);tl.store(children+left_base+rows*512+rows,left_diagonal);tl.store(children+right_base+rows*512+rows,right_diagonal);off_rows=rows+1;valid=off_rows<512;left_off=tl.load(matrix+source_base+off_rows*1024+rows,mask=valid,other=.0);right_off=tl.load(matrix+source_base+(off_rows+512)*1024+rows+512,mask=valid,other=.0);tl.store(children+left_base+off_rows*512+rows,left_off,mask=valid);tl.store(children+right_base+off_rows*512+rows,right_off,mask=valid);tl.store(rho_output+batch,rho);tl.store(sign_output+batch,tl.where(coupling<.0,-1.,1.))
@triton.jit
def _e3386_pack_active_cuppen(source_poles,source_update,packed_poles,packed_update,packed_indices,active_counts):matrix=tl.program_id(0);offsets=tl.arange(0,1024);base=matrix*1024;poles=tl.load(source_poles+base+offsets);update=tl.load(source_update+base+offsets);scale=tl.maximum(tl.max(tl.abs(update),axis=0),1e-30);active=tl.abs(update)>1e-07*scale;active_position=tl.cumsum(active.to(tl.int32),axis=0)-1;count=tl.sum(active.to(tl.int32),axis=0);inactive_position=tl.cumsum((~active).to(tl.int32),axis=0)-1;position=tl.where(active,active_position,count+inactive_position);tl.store(packed_poles+base+position,poles);tl.store(packed_update+base+position,update);tl.store(packed_indices+base+position,offsets);tl.store(active_counts+matrix,count)
@triton.jit
def _e3386_cuppen_root_norm(diagonal,update,active_counts,output,tau_output,origin_output,inverse_norm_output):
root=tl.program_id(0);matrix=tl.program_id(1);count=tl.load(active_counts+matrix);safe_root=tl.minimum(root,count-1);offsets=tl.arange(0,1024);mask=offsets<count;base=matrix*1024;poles=tl.load(diagonal+base+offsets,mask=mask,other=.0);vector=tl.load(update+base+offsets,mask=mask,other=.0);weights=vector*vector;left_pole=tl.load(diagonal+base+safe_root);norm2=tl.sum(tl.where(mask,weights,.0),axis=0);final_derivative=norm2*.0
if safe_root+1==count:
origin_value=left_pole;lower=norm2*.0;upper=norm2*1.000001
for _ in range(40):middle=.5*(lower+upper);denominator=poles-origin_value-middle;secular=1.+tl.sum(tl.where(mask,weights/denominator,.0),axis=0);lower=tl.where(secular<.0,middle,lower);upper=tl.where(secular<.0,upper,middle)
tau=.5*(lower+upper);value=origin_value+tau;origin_index=safe_root;delta=poles-origin_value-tau;final_derivative=tl.sum(tl.where(mask,weights/(delta*delta),.0),axis=0)
else:
right_pole=tl.load(diagonal+base+safe_root+1);pole_gap=right_pole-left_pole;midpoint=.5*(left_pole+right_pole);denominator=poles-midpoint;fractions=tl.where(mask,weights/denominator,.0);full_mid=1.+tl.sum(fractions,axis=0);left_weight=tl.load(update+base+safe_root);left_weight*=left_weight;right_weight=tl.load(update+base+safe_root+1);right_weight*=right_weight;constant=full_mid-left_weight/(left_pole-midpoint)-right_weight/(right_pole-midpoint);coefficient_a=constant*pole_gap+left_weight+right_weight;coefficient_b=left_weight*pole_gap;discriminant=tl.sqrt(tl.abs(coefficient_a*coefficient_a-4.*coefficient_b*constant));tau_left=tl.where(coefficient_a>.0,2.*coefficient_b/(coefficient_a+discriminant),(coefficient_a-discriminant)/(2.*constant));coefficient_a_right=constant*pole_gap-left_weight-right_weight;coefficient_b_right=right_weight*pole_gap;discriminant_right=tl.sqrt(tl.abs(coefficient_a_right*coefficient_a_right+4.*coefficient_b_right*constant));tau_right=tl.where(coefficient_a_right<.0,2.*coefficient_b_right/(coefficient_a_right-discriminant_right),-(coefficient_a_right+discriminant_right)/(2.*constant));origin_left=full_mid>.0;origin_value=tl.where(origin_left,left_pole,right_pole);origin_index=tl.where(origin_left,safe_root,safe_root+1);tau=tl.where(origin_left,tau_left,tau_right);lower=left_pole-origin_value;upper=right_pole-origin_value;tau=tl.maximum(lower,tl.minimum(upper,tau));iterations=0;converged=False
while(iterations<20)&~converged:delta=poles-origin_value-tau;fractions=tl.where(mask,weights/delta,.0);secular=1.+tl.sum(fractions,axis=0);derivative=tl.sum(tl.where(mask,weights/(delta*delta),.0),axis=0);final_derivative=derivative;error_scale=1.+tl.sum(tl.abs(fractions),axis=0);newton_step=-secular/derivative;converged=(iterations>=1)&((tl.abs(secular)<=1.1920928955078125e-07*error_scale)|(tl.abs(newton_step)<=1.1920928955078125e-07*tl.maximum(tl.abs(origin_value+tau),pole_gap)));lower=tl.where(secular<.0,tau,lower);upper=tl.where(secular<.0,upper,tau);delta_left=left_pole-origin_value-tau;delta_right=right_pole-origin_value-tau;coefficient_c_left=secular-delta_right*derivative+pole_gap*left_weight/(delta_left*delta_left);coefficient_c_right=secular-delta_left*derivative-pole_gap*right_weight/(delta_right*delta_right);coefficient_c=tl.where(origin_left,coefficient_c_left,coefficient_c_right);coefficient_a=(delta_left+delta_right)*secular-delta_left*delta_right*derivative;coefficient_b=delta_left*delta_right*secular;discriminant=tl.sqrt(tl.abs(coefficient_a*coefficient_a-4.*coefficient_b*coefficient_c));rational_step=tl.where(coefficient_a<=.0,(coefficient_a-discriminant)/(2.*coefficient_c),2.*coefficient_b/(coefficient_a+discriminant));rational_step=tl.where((secular*rational_step<.0)&(rational_step==rational_step),rational_step,newton_step);candidate=tau+rational_step;safe=(candidate==candidate)&(candidate>lower)&(candidate<upper);next_value=tl.where(safe,candidate,.5*(lower+upper));tau=tl.where(converged,tau,next_value);iterations+=1
value=origin_value+tau
active=root<count;tl.store(output+base+root,value,mask=active);tl.store(tau_output+base+root,tau,mask=active);tl.store(origin_output+base+root,origin_index,mask=active);tl.store(inverse_norm_output+base+root,tl.rsqrt(final_derivative),mask=active)
@triton.jit
def _e3386_build_cuppen_maps(child_order,packed_indices,packed_poles,packed_update,roots,active_counts,half_poles,half_weights,half_active,combined_values,combined_original):matrix=tl.program_id(0);offsets=tl.arange(0,1024);base=matrix*1024;count=tl.load(active_counts+matrix);active=offsets<count;packed_index=tl.load(packed_indices+base+offsets);original=tl.load(child_order+base+packed_index);pole=tl.load(packed_poles+base+offsets);weight=tl.load(packed_update+base+offsets,mask=active,other=.0);root=tl.load(roots+base+offsets,mask=active,other=.0);tl.store(half_poles+base+original,tl.where(active,pole,.0));tl.store(half_weights+base+original,weight);tl.store(half_active+base+original,active.to(tl.int8));tl.store(combined_values+base+offsets,tl.where(active,root,pole));tl.store(combined_original+base+offsets,tl.where(active,-1,original))
@triton.jit
def _e3386_cuppen_consumer(left_basis,right_basis,packed_poles,tau,origin,inverse_norm,half_poles,half_weights,half_active,active_counts,final_order,combined_original,output,BLOCK_M:tl.constexpr,BLOCK_N:tl.constexpr):
batch=tl.program_id(0);row_tile=tl.program_id(1);column_tile=tl.program_id(2);half=row_tile//(512//BLOCK_M);local_rows=(row_tile-half*(512//BLOCK_M))*BLOCK_M+tl.arange(0,BLOCK_M);columns=column_tile*BLOCK_N+tl.arange(0,BLOCK_N);count=tl.load(active_counts+batch);combined=tl.load(final_order+batch*1024+columns);active_column=combined<count;safe_root=tl.where(active_column,combined,0);root_origin=tl.load(origin+batch*1024+safe_root);root_origin_pole=tl.load(packed_poles+batch*1024+root_origin);root_tau=tl.load(tau+batch*1024+safe_root);root_scale=tl.load(inverse_norm+batch*1024+safe_root);accumulator=tl.zeros((BLOCK_M,BLOCK_N),tl.float32);half_base=(batch*2+half)*512;basis_source=tl.where(half==0,left_basis,right_basis)
for start in tl.static_range(0,512,32):k=start+tl.arange(0,32);basis_tile=tl.load(basis_source+batch*512*512+local_rows[:,None]*512+k[None,:]);pole=tl.load(half_poles+half_base+k);weight=tl.load(half_weights+half_base+k);factor_active=tl.load(half_active+half_base+k)!=0;denominator=pole[:,None]-root_origin_pole[None,:]-root_tau[None,:];cauchy=weight[:,None]/tl.where(factor_active[:,None],denominator,1.);cauchy*=root_scale[None,:];cauchy=tl.where(factor_active[:,None]&active_column[None,:],cauchy,.0);accumulator+=tl.dot(basis_tile.to(tl.float16),cauchy.to(tl.float16),out_dtype=tl.float32)
inactive_original=tl.load(combined_original+batch*1024+combined);local_inactive=inactive_original-half*512;inactive_match=~active_column&(local_inactive>=0)&(local_inactive<512);accumulator+=tl.load(basis_source+batch*512*512+local_rows[:,None]*512+local_inactive[None,:],mask=inactive_match[None,:],other=.0);tl.store(output+batch*1024*1024+(half*512+local_rows)[:,None]*1024+columns[None,:],accumulator)
_E3386_REPAIR_NAME='e3386_inplace_adjacent_mgs512';_E3386_REPAIR_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(512,2)\nvoid e3386_inplace_adjacent_mgs512(float* q,const float* l,int batch,float threshold){\n constexpr int N=512,W=16;int matrix=blockIdx.x;if(matrix>=batch)return;\n int warp=threadIdx.x>>5,lane=threadIdx.x&31;long long base=(long long)matrix*N*N;\n const float* values=l+(long long)matrix*N;\n float span=fmaxf(values[N-1]-values[0],1.0e-30f);\n #pragma unroll\n for(int parity=0;parity<2;++parity){\n for(int column=1+parity+2*warp;column<N;column+=2*W){\n float gap=values[column]-values[column-1];\n if(gap<threshold*span){\n float dot=0.f;\n #pragma unroll\n for(int row=lane;row<N;row+=32)\n dot=fmaf(q[base+(long long)row*N+column-1],q[base+(long long)row*N+column],dot);\n #pragma unroll\n for(int offset=16;offset>0;offset>>=1)dot+=__shfl_down_sync(0xffffffffu,dot,offset);\n dot=__shfl_sync(0xffffffffu,dot,0);\n float inverse=rsqrtf(fmaxf(1.f-dot*dot,1.0e-12f));\n #pragma unroll\n for(int row=lane;row<N;row+=32){long long target=base+(long long)row*N+column;\n q[target]=(q[target]-dot*q[base+(long long)row*N+column-1])*inverse;}\n }\n }\n __syncthreads();\n }\n}\n'
@memo(maxsize=1)
def _e3386_repair_kernel():return CUDAKernel(_fast_nvrtc_compile(_E3386_REPAIR_SOURCE,_E3386_REPAIR_NAME),_E3386_REPAIR_NAME)
@torch.no_grad()
def _e3386_cuppen_solve_build1024(matrix:torch.Tensor,reflectors:torch.Tensor):batch=matrix.shape[0];child_input=torch.empty((2*batch,512,512),device=matrix.device);rho=torch.empty((batch,),device=matrix.device);sign=torch.empty_like(rho);_e3386_pack_cuppen_children[batch,](matrix,child_input,rho,sign,num_warps=8);child_vectors=torch.empty_like(child_input);child_values=torch.empty((2*batch,512),device=matrix.device);child_workspace=torch.empty_like(child_input);_e1782_mixed_tridiag24_pdl_kernel().launch((2*batch*_E1782_MIXED_TRI24_SHARDS,1,1),(_E1782_MIXED_TRI24_EIGEN_PER_SHARD,1,1),(child_input,child_vectors,child_values,child_workspace));operators=torch.empty((batch,32,32,64,64),device=matrix.device,dtype=torch.float16);_e549_build_band32_sparse_wy1024[batch,_E549_ACTIVE_SPARSE_WY_OPERATORS](reflectors,operators,num_warps=2,num_stages=1,launch_pdl=True);_e3386_repair_kernel().launch((2*batch,1,1),(512,1,1),(child_vectors,child_values,2*batch,1e-05));update=torch.cat((child_vectors[:batch,-1,:],sign[:,None]*child_vectors[batch:,0,:]),dim=1);update*=rho.sqrt()[:,None];poles,order=torch.cat((child_values[:batch],child_values[batch:]),dim=1).sort(1);update=update.gather(1,order);order=order.to(torch.int32);packed_poles=torch.empty_like(poles);packed_update=torch.empty_like(update);packed_indices=torch.empty_like(order);counts=torch.empty((batch,),device=matrix.device,dtype=torch.int32);_e3386_pack_active_cuppen[batch,](poles,update,packed_poles,packed_update,packed_indices,counts,num_warps=8);roots=torch.empty_like(packed_poles);tau=torch.empty_like(packed_poles);origin=torch.empty_like(packed_indices);inverse_norm=torch.empty_like(packed_poles);_e3386_cuppen_root_norm[1024,batch](packed_poles,packed_update,counts,roots,tau,origin,inverse_norm,num_warps=8);half_poles=torch.empty_like(packed_poles);half_weights=torch.empty_like(packed_update);half_active=torch.empty_like(packed_poles,dtype=torch.int8);combined_values=torch.empty_like(packed_poles);combined_original=torch.empty_like(packed_indices);_e3386_build_cuppen_maps[batch,](order,packed_indices,packed_poles,packed_update,roots,counts,half_poles,half_weights,half_active,combined_values,combined_original,num_warps=8);values,final_order=combined_values.sort(1);vectors=torch.empty_like(matrix);_e3386_cuppen_consumer[batch,8,32](child_vectors[:batch],child_vectors[batch:],packed_poles,tau,origin,inverse_norm,half_poles,half_weights,half_active,counts,final_order,combined_original,vectors,BLOCK_M=128,BLOCK_N=32,num_warps=8,num_stages=2);return vectors,values,operators
@torch.no_grad()
def _e1760_sparse_wy1024_replay(vectors:torch.Tensor,operators:torch.Tensor,prefix_launcher=None):
_e549_band32_sparse_wy1024_replay[vectors.shape[0],16](vectors,operators,block_cols=64,signal_dependents=prefix_launcher is not None,num_warps=4,num_stages=1)
if prefix_launcher is not None:prefix_launcher()
return vectors
@torch.no_grad()
def _e549_clustered1024_impl(matrix:torch.Tensor,*,strong_range:bool):
batch,n,_=matrix.shape;negative_rank=n//3;padded_rank=384;identity=torch.eye(n,device=matrix.device).expand(batch,-1,-1);projector=.5*(identity-matrix);leverage=projector.diagonal(dim1=-2,dim2=-1);indices=leverage.topk(negative_rank,dim=-1).indices;low=projector.gather(2,indices[:,None,:].expand(-1,n,-1)).contiguous();low=projector@low;low=cholesky_orthonormalize(low,passes=1,ridge=1e-05,final_ridge=1e-07,inverse_precision='highest',gram_precision='highest')
if strong_range:low=projector@low;low=cholesky_orthonormalize(low,passes=1,ridge=1e-06,final_ridge=1e-08,inverse_precision='highest',gram_precision='highest')
seed=torch.empty((batch,n,padded_rank),device=matrix.device,dtype=torch.float32);seed[:,:,:negative_rank]=low;seed[:,:,negative_rank:]=identity[:,:,negative_rank:padded_rank];vectors=compact_wy_orthogonalize_1024x384(None,seed,complete=True,newton_steps=1,newton_precision='highest',paired_replay=True);product=matrix@vectors;values=(vectors*product).sum(1);eigen_l1=(product-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);matrix_l1=matrix.abs().sum(dim=1).amax(dim=1).clamp_min(1e-30);scaled_eigen=eigen_l1/(1.1920928955078125e-07*n*matrix_l1);risk=~torch.isfinite(scaled_eigen)|(scaled_eigen>64.);values,order=values.sort(1);vectors=vectors.gather(2,order[:,None,:].expand(-1,n,-1));return vectors,values,risk
@torch.no_grad()
def _e549_clustered1024(matrix:torch.Tensor):
vectors,values,risk=_e549_clustered1024_impl(matrix,strong_range=False)
if bool(risk.any().item()):repaired_vectors,repaired_values,_=_e549_clustered1024_impl(matrix[risk].contiguous(),strong_range=True);vectors=vectors.clone();values=values.clone();vectors[risk]=repaired_vectors;values[risk]=repaired_values
return vectors,values
@triton.jit
def _e549_native_band32_flags_kernel(matrix,outside_nonzero,ring_nonzero,chunks_per_matrix:tl.constexpr,N:tl.constexpr,BLOCK:tl.constexpr):program=tl.program_id(0);batch=program//chunks_per_matrix;chunk=program-batch*chunks_per_matrix;element=chunk*BLOCK+tl.arange(0,BLOCK);valid=element<N*N;row=element//N;column=element-row*N;distance=tl.abs(row-column);value=tl.load(matrix+batch*N*N+element,mask=valid,other=.0);nonzero=value!=.0;outside_hit=tl.max(tl.where(valid&(distance>32)&nonzero,1,0),axis=0);ring_hit=tl.max(tl.where(valid&(distance==32)&nonzero,1,0),axis=0);tl.atomic_or(outside_nonzero+batch,outside_hit);tl.atomic_or(ring_nonzero+batch,ring_hit)
@torch.no_grad()
def _e549_native_band32_mask(matrix:torch.Tensor):batch=matrix.shape[0];outside=torch.zeros(batch,device=matrix.device,dtype=torch.int32);ring=torch.zeros_like(outside);block=4096;chunks=triton.cdiv(1024*1024,block);_e549_native_band32_flags_kernel[batch*chunks,](matrix,outside,ring,chunks_per_matrix=chunks,N=1024,BLOCK=block,num_warps=8,num_stages=1);return(outside==0)&(ring!=0)
@triton.jit
def _e1378_exact1024_gram_affine_risk_kernel(gram,risk,chunks_per_matrix:tl.constexpr,threshold:tl.constexpr,N:tl.constexpr,BLOCK:tl.constexpr):program=tl.program_id(0);batch=program//chunks_per_matrix;chunk=program-batch*chunks_per_matrix;linear=chunk*BLOCK+tl.arange(0,BLOCK);mask=linear<N*N;row=linear//N;column=linear-row*N;value=tl.load(gram+batch*N*N+linear,mask=mask,other=.0);hit=tl.max(tl.where(mask&(row!=column)&(tl.abs(value)>threshold),1,0),axis=0);tl.atomic_or(risk+batch,hit);value=-.5*value+tl.where(row==column,1.5,.0);tl.store(gram+batch*N*N+linear,value,mask=mask)
@torch.no_grad()
def _e1378_exact1024_gram_affine_risk(gram:torch.Tensor):batch,n,_=gram.shape;block=4096;chunks=triton.cdiv(n*n,block);risk=torch.zeros((batch,),device=gram.device,dtype=torch.int32);_e1378_exact1024_gram_affine_risk_kernel[batch*chunks,](gram,risk,chunks_per_matrix=chunks,threshold=.03,N=n,BLOCK=block,num_warps=8,num_stages=1);return risk
@torch.no_grad()
def _e1853_finalize_owned_exact1024(state,holder=None):
vectors,values,operators,dense_count,panels=state;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
if dense_count:dense_panel_backtransform(vectors[:dense_count],panels,precision='tf32')
vectors16=vectors.to(torch.float16);gram=torch.bmm(vectors16.mT,vectors16,out_dtype=torch.float32);risk=_e1378_exact1024_gram_affine_risk(gram);vectors=torch.bmm(vectors16,gram.to(torch.float16),out_dtype=torch.float32)
if bool(risk.any().item()):
repair=risk!=0;repaired=vectors[repair].contiguous()
for _ in range(3):repair_gram=repaired.mT@repaired;repair_gram.mul_(-.5);repair_gram.diagonal(dim1=-2,dim2=-1).add_(1.5);repaired=repaired@repair_gram
vectors[repair]=repaired
vectors=normalize_columns_(vectors)
finally:torch.set_float32_matmul_precision(previous)
if holder is not None:holder[0]=vectors;holder[1]=values;return holder
return vectors,values
@torch.no_grad()
def _e549_owned_exact1024(matrix:torch.Tensor,*,dense_count:int|None=None,prefix_launcher=None,defer_finalize:bool=False,cuppen:bool=False):
grouped=dense_count is not None
if grouped:
band=matrix
if dense_count:_,panels=dense_to_band32(band[:dense_count],precision='tf32',save_panels=True,symmetric_update=True,cluster_gram=True)
else:panels=[]
else:
native_band=_e549_native_band32_mask(matrix);dense_mask=~native_band
if not bool(native_band.any()):band,panels=dense_to_band32(matrix,precision='tf32',save_panels=True,symmetric_update=True,cluster_gram=True);dense_count=matrix.shape[0]
else:
dense_input=matrix[dense_mask].contiguous();band_input=matrix[native_band].contiguous()
if dense_input.shape[0]:reduced,panels=dense_to_band32(dense_input,precision='tf32',save_panels=True,symmetric_update=True,cluster_gram=True)
else:reduced,panels=dense_input,[]
band=torch.cat((reduced,band_input),dim=0);dense_count=dense_input.shape[0]
tridiagonal,reflectors=_e549_band32_to_tridiagonal_n1024(band)
if cuppen:vectors,values,operators=_e3386_cuppen_solve_build1024(tridiagonal,reflectors)
else:vectors,values,operators=_e1760_tridiag_solve_build1024(tridiagonal,reflectors)
previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:vectors=_e1760_sparse_wy1024_replay(vectors,operators,prefix_launcher=prefix_launcher)
finally:torch.set_float32_matmul_precision(previous)
state=vectors,values,operators,dense_count,panels
if defer_finalize:
if not grouped:raise ValueError('deferred exact1024 finalize requires grouped input')
holder=[vectors,values];return holder,state
vectors,values=_e1853_finalize_owned_exact1024(state);dense_vectors=vectors[:dense_count]
if grouped:return vectors,values
if not bool(native_band.any()):return dense_vectors,values
output_vectors=torch.empty_like(matrix);output_values=torch.empty(matrix.shape[:-1],device=matrix.device);output_vectors[dense_mask]=dense_vectors;output_values[dense_mask]=values[:dense_count];output_vectors[native_band]=vectors[dense_count:];output_values[native_band]=values[dense_count:];return output_vectors,output_values
@triton.jit
def _e992_mixed1024_route_scatter_kernel(q0,q1,q2,q3,l0,l1,l2,l3,order,counts,output_q,output_l,N:tl.constexpr,ELEMENTS:tl.constexpr,BLOCK:tl.constexpr):packed_matrix=tl.program_id(0);chunk=tl.program_id(1);c0=tl.load(counts);c1=tl.load(counts+1);c2=tl.load(counts+2);s1=c0;s2=c0+c1;s3=s2+c2;route=(packed_matrix>=s1).to(tl.int32)+(packed_matrix>=s2).to(tl.int32)+(packed_matrix>=s3).to(tl.int32);start=tl.where(route==0,0,tl.where(route==1,s1,tl.where(route==2,s2,s3)));local_matrix=packed_matrix-start;destination=tl.load(order+packed_matrix).to(tl.int64);linear=chunk*BLOCK+tl.arange(0,BLOCK);valid=linear<ELEMENTS;source=local_matrix*ELEMENTS+linear;value=tl.load(q0+source,mask=valid&(route==0),other=.0);value+=tl.load(q1+source,mask=valid&(route==1),other=.0);value+=tl.load(q2+source,mask=valid&(route==2),other=.0);value+=tl.load(q3+source,mask=valid&(route==3),other=.0);tl.store(output_q+destination*ELEMENTS+linear,value,mask=valid);value_mask=(chunk==0)&(linear<N);value_source=local_matrix*N+linear;eigen=tl.load(l0+value_source,mask=value_mask&(route==0),other=.0);eigen+=tl.load(l1+value_source,mask=value_mask&(route==1),other=.0);eigen+=tl.load(l2+value_source,mask=value_mask&(route==2),other=.0);eigen+=tl.load(l3+value_source,mask=value_mask&(route==3),other=.0);tl.store(output_l+destination*N+linear,eigen,mask=value_mask)
@torch.no_grad()
def _e992_mixed1024_route_scatter(outputs:tuple[output_t,...],order:torch.Tensor,counts:torch.Tensor):batch=order.shape[0];n=outputs[0][0].shape[-1];vectors=torch.empty((batch,n,n),device=order.device,dtype=outputs[0][0].dtype);values=torch.empty((batch,n),device=order.device,dtype=outputs[0][1].dtype);block=4096;_e992_mixed1024_route_scatter_kernel[batch,triton.cdiv(n*n,block)](*(output[0]for output in outputs),*(output[1]for output in outputs),order,counts,vectors,values,N=n,ELEMENTS=n*n,BLOCK=block,num_warps=8,num_stages=1);return vectors,values
@torch.no_grad()
def _mixed_1024_repeated_partitioned(data:torch.Tensor,repeated:torch.Tensor,stats:torch.Tensor,route_metadata):
batch=data.shape[0];refined_route,counts_device,counts,dense_exact_count,extreme_exact_count=route_metadata;_,order=torch.sort(refined_route);packed=data.index_select(0,order);packed_stats=stats.index_select(0,order);starts=[0]
for count in counts:starts.append(starts[-1]+count)
groups=tuple(packed[starts[index]:starts[index+1]]for index in range(4));dense_input,repeated_input,clustered_input,exact_input=groups;normal_exact_count=counts[3]-extreme_exact_count;normal_exact_input=exact_input[:normal_exact_count];extreme_exact_input=exact_input[normal_exact_count:];n=data.shape[-1];empty_values=data.new_empty((0,n));root_lanczos_stats=None;prefix_launched=False;exact_finalize_state=None
if counts[0]and counts[3]:
median=torch.empty((counts[0],),device=data.device);lower=torch.empty_like(median);upper=torch.empty_like(median);dense_flags=torch.empty((counts[0],),device=data.device,dtype=torch.bool);root_lanczos_stats=median,lower,upper
def launch_dense_lanczos():nonlocal prefix_launched;_e189_n1024_lanczos3_kernel().launch_pdl((counts[0],1,1),(512,1,1),(dense_input,median,lower,upper,dense_flags));prefix_launched=True
else:launch_dense_lanczos=None
if normal_exact_count:
if counts[0]:exact_output,exact_finalize_state=_e549_owned_exact1024(normal_exact_input,dense_count=dense_exact_count,prefix_launcher=launch_dense_lanczos,defer_finalize=True,cuppen=True)
else:exact_output=_e549_owned_exact1024(normal_exact_input,dense_count=dense_exact_count,prefix_launcher=launch_dense_lanczos,cuppen=True)
else:exact_output=normal_exact_input,empty_values
if root_lanczos_stats is not None and not prefix_launched:root_lanczos_stats=None
if counts[0]:
try:dense_output=_dense_1024_specialized(dense_input,low_precision_sign=True,matrix_scale=packed_stats[:counts[0],4],root_lanczos_stats=root_lanczos_stats)
except torch.linalg.LinAlgError:dense_output=_dense_1024_specialized(dense_input,low_precision_sign=False,matrix_scale=packed_stats[:counts[0],4],certify_output=False)
else:dense_output=dense_input,empty_values
if exact_finalize_state is not None:_e1853_finalize_owned_exact1024(exact_finalize_state,exact_output)
if counts[1]:
previous=torch.get_float32_matmul_precision()
try:repeated_output=mixed_repeated_krylov_eigh(repeated_input,seed_width=112);torch.set_float32_matmul_precision(previous)
except torch.linalg.LinAlgError:torch.set_float32_matmul_precision(previous);repeated_output=_e549_owned_exact1024(repeated_input,dense_count=counts[1])
else:repeated_output=repeated_input,empty_values
if counts[2]:
try:clustered_output=_e549_clustered1024(clustered_input)
except torch.linalg.LinAlgError:clustered_output=_e549_owned_exact1024(clustered_input,dense_count=counts[2])
else:clustered_output=clustered_input,empty_values
if extreme_exact_count:
extreme_values,extreme_vectors=torch.linalg.eigh(extreme_exact_input);extreme_vectors=extreme_vectors.contiguous()
if normal_exact_count:exact_output=torch.cat((exact_output[0],extreme_vectors),dim=0),torch.cat((exact_output[1],extreme_values),dim=0)
else:exact_output=extreme_vectors,extreme_values
outputs=dense_output,repeated_output,clustered_output,exact_output;return _e992_mixed1024_route_scatter(outputs,order,counts_device)
_N160_POTRF_SHARED_BYTES=(160*161//2+1)*4
@memo(maxsize=None)
def _e083_n160_cholesky_kernel(name:str):
source=_N160_PACKED_FULL_CHOLESKY_SOURCE if name=='potrf160_packed_shared_full_output'else _N160_FULL_CHOLESKY_SOURCE
if name=='potrf160_packed_shared_full_output':source=source.replace('__launch_bounds__(256, 1)\nvoid potrf160_packed_shared_full_output(','__launch_bounds__(256, 4)\nvoid potrf160_packed_shared_full_output(')
source=_fast_only_cuda_kernel(source,name);image=_fast_nvrtc_compile(source,name);return CUDAKernel(image,name)
@torch.no_grad()
def _e083_potrf160(gram:torch.Tensor,ridge:float):lower=torch.empty_like(gram);_e083_n160_cholesky_kernel('potrf160_packed_shared_full_output').launch((gram.shape[0],1,1),(256,1,1),(gram,lower,gram.shape[0],float(ridge)),shared_mem=_N160_POTRF_SHARED_BYTES);return lower
_E2409_POTRF160_WMMA_NAME='e2409_potrf160_wmma_tf32x3'
@memo(maxsize=1)
def _e2409_potrf160_wmma_kernel():source=_e202_wmma_potrf_source(160,_E2409_POTRF160_WMMA_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E2409_POTRF160_WMMA_NAME),_E2409_POTRF160_WMMA_NAME)
@torch.no_grad()
def _e2409_potrf160_wmma(gram:torch.Tensor,ridge:float):lower=torch.empty_like(gram);_e2409_potrf160_wmma_kernel().launch((gram.shape[0],1,1),(320,1,1),(gram,lower,gram.shape[0],float(ridge)),shared_mem=(160*160+1)*4);return lower
@torch.no_grad()
def _e083_trsm160(matrix:torch.Tensor,lower:torch.Tensor):output=torch.empty_like(matrix);_e083_n160_cholesky_kernel('right_trsm160_block16_rows32').launch((matrix.shape[0],(matrix.shape[1]+31)//32,1),(256,1,1),(matrix,lower,output,matrix.shape[0],matrix.shape[1]),shared_mem=160*33*4);return output
_E106_TRSM160_HALF2_NAME='e106_right_trsm160_half2_panel16';_E106_TRSM160_HALF2_SOURCE='\n#include <cuda_runtime.h>\n#include <cuda_fp16.h>\n\nextern "C" __global__ __launch_bounds__(256, 8)\nvoid e106_right_trsm160_half2_panel16(\n const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch,\n int rows) {\n constexpr int N = 160;\n constexpr int ROWS = 32;\n constexpr int PAIRS = 16;\n constexpr int PANEL = 16;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int row_base = blockIdx.y * ROWS;\n if (matrix_id >= batch) return;\n extern __shared__ unsigned char storage[];\n __half* rhs = reinterpret_cast<__half*>(storage);\n __half* factor16 = rhs + N * ROWS;\n const float* source = matrix + (long long)matrix_id * rows * N;\n const float* factor = lower + (long long)matrix_id * N * N;\n float* destination = output + (long long)matrix_id * rows * N;\n\n for (int index = tid; index < N * ROWS; index += blockDim.x) {\n const int local_row = index / N;\n const int column = index - local_row * N;\n const int row = row_base + local_row;\n rhs[column * ROWS + local_row] = row < rows\n ? __float2half_rn(source[(long long)row * N + column])\n : __float2half_rn(0.0f);\n }\n __syncthreads();\n\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += PANEL) {\n const int end = min(panel + PANEL, N);\n const int factor_count = (N - panel) * PANEL;\n for (int index = tid; index < factor_count; index += blockDim.x) {\n const int column = panel + index / PANEL;\n const int k_offset = index & 15;\n const int k = panel + k_offset;\n factor16[column * PANEL + k_offset] = k <= column\n ? __float2half_rn(factor[column * N + k])\n : __float2half_rn(0.0f);\n }\n __syncthreads();\n\n if (tid < PAIRS) {\n const int pair = tid;\n #pragma unroll\n for (int offset = 0; offset < PANEL; ++offset) {\n const int column = panel + offset;\n __half2 value = *reinterpret_cast<__half2*>(\n rhs + column * ROWS + 2 * pair);\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n if (k_offset >= offset) break;\n const __half factor_value = factor16[\n column * PANEL + k_offset];\n const __half2 coefficient = __halves2half2(\n __hneg(factor_value), __hneg(factor_value));\n const __half2 solved = *reinterpret_cast<__half2*>(\n rhs + (panel + k_offset) * ROWS + 2 * pair);\n value = __hfma2(coefficient, solved, value);\n }\n const float inverse = 1.0f / __half2float(\n factor16[column * PANEL + offset]);\n value = __floats2half2_rn(\n __low2float(value) * inverse,\n __high2float(value) * inverse);\n *reinterpret_cast<__half2*>(\n rhs + column * ROWS + 2 * pair) = value;\n }\n }\n __syncthreads();\n\n const int remaining = (N - end) * PAIRS;\n for (int index = tid; index < remaining; index += blockDim.x) {\n const int column = end + index / PAIRS;\n const int pair = index - (index / PAIRS) * PAIRS;\n __half2 value = *reinterpret_cast<__half2*>(\n rhs + column * ROWS + 2 * pair);\n #pragma unroll\n for (int k_offset = 0; k_offset < PANEL; ++k_offset) {\n const __half factor_value = factor16[\n column * PANEL + k_offset];\n const __half2 coefficient = __halves2half2(\n __hneg(factor_value), __hneg(factor_value));\n const __half2 solved = *reinterpret_cast<__half2*>(\n rhs + (panel + k_offset) * ROWS + 2 * pair);\n value = __hfma2(coefficient, solved, value);\n }\n *reinterpret_cast<__half2*>(\n rhs + column * ROWS + 2 * pair) = value;\n }\n __syncthreads();\n }\n\n for (int index = tid; index < N * ROWS; index += blockDim.x) {\n const int local_row = index / N;\n const int column = index - local_row * N;\n const int row = row_base + local_row;\n if (row < rows)\n destination[(long long)row * N + column] =\n __half2float(rhs[column * ROWS + local_row]);\n }\n}\n'
@memo(maxsize=1)
def _e106_trsm160_half2_kernel():return CUDAKernel(_fast_nvrtc_compile(_E106_TRSM160_HALF2_SOURCE,_E106_TRSM160_HALF2_NAME),_E106_TRSM160_HALF2_NAME)
@torch.no_grad()
def _e106_trsm160_half2(matrix:torch.Tensor,lower:torch.Tensor):batch,rows,n=matrix.shape;output=torch.empty_like(matrix);_e106_trsm160_half2_kernel().launch((batch,(rows+31)//32,1),(256,1,1),(matrix,lower,output,batch,rows),shared_mem=(n*32+n*16)*2);return output
_E102_REDUCE160_NAME='e102_reduce160_singlecta';_E102_SOLVE160_NAME='e102_solve160_cluster';_E107_SOLVE160_NAME='e107_solve160_tridiagonal_cluster';_E1008_REDUCE160_BLOCK2_NAME='e1008_reduce160_block2_t384';_E1008_REDUCE160_BLOCK2_SHARED=(160*161//2+2*2*160+160+32+2*2)*4
def _e102_solve160_source():preamble_end=_N160_EIGH_SOURCE.index('extern "C" __global__');solve_start=_N160_EIGH_SOURCE.index(' // Parallel Sturm bisection followed');preamble=_N160_EIGH_SOURCE[:preamble_end];tail=_N160_EIGH_SOURCE[solve_start:];header=f'''extern "C" __global__
__cluster_dims__(2, 1, 1)
__launch_bounds__(384, 4)
void {_E102_SOLVE160_NAME}(
const float* __restrict__ saved_reflectors,
const float* __restrict__ diagonal_input,
const float* __restrict__ off_diagonal_input,
float* __restrict__ output_q,
float* __restrict__ output_l,
unsigned long long* __restrict__ timers) {{
cg::cluster_group cluster = cg::this_cluster();
const int rank = cluster.block_rank();
const int matrix_id = blockIdx.x >> 1;
const int tid = threadIdx.x;
const int first_row = rank * HALF;
const float* saved_matrix = saved_reflectors + (long long)matrix_id * N * N;
extern __shared__ __align__(16) unsigned char storage[];
float* q = reinterpret_cast<float*>(storage);
float* v = q + TRI;
float* w = v + N;
float* partial = w + N;
float* scratch = partial + N;
float* diagonal = scratch + 16;
float* off_diagonal = diagonal + N;
float* partial0 = cluster.map_shared_rank(partial, 0);
if (tid < N) {{
diagonal[tid] = diagonal_input[matrix_id * N + tid];
off_diagonal[tid] = off_diagonal_input[matrix_id * N + tid];
}}
__syncthreads();
if (rank == 0 && tid == 0) {{
const unsigned long long now = clock64();
timers[matrix_id * 5] = now;
timers[matrix_id * 5 + 1] = now;
}}
''';return preamble+header+tail
def _e107_solve160_source():source=_e102_solve160_source().replace(_E102_SOLVE160_NAME,_E107_SOLVE160_NAME);wy_start=source.index(' // Apply normalized reflectors in reverse blocks of WY_BLOCK.');store_start=source.index(' // Sturm bisection emits ascending eigenvalues.',wy_start);return source[:wy_start]+source[store_start:]
@memo(maxsize=1)
def _e107_solve160_kernel():
source=_e107_solve160_source().replace('step < 24','step < 22');source=_parallelize_n176_adjacent_mgs(source)
if source.count('local_eigen += 16')!=1:raise RuntimeError('n160 parallel MGS stride template changed')
source=source.replace('local_eigen += 16','local_eigen += 24',1);anchor=' const int tid = threadIdx.x;\n const int first_row = rank * HALF;';wait=' const int tid = threadIdx.x;\n if (rank == 0 && tid == 0) {\n volatile unsigned long long* ready = timers + matrix_id * 5 + 4;\n while (*ready == 0ULL) __nanosleep(64);\n }\n cluster.sync();\n const int first_row = rank * HALF;'
if source.count(anchor)!=1:raise RuntimeError('n160 solve ready template changed')
source=source.replace(anchor,wait,1);source=_e2337_cluster_coarse_sturm(source,steps=22,probes=128,fine_steps=15,levels=8);return CUDAKernel(_fast_nvrtc_compile(source,_E107_SOLVE160_NAME),_E107_SOLVE160_NAME)
@memo(maxsize=1)
def _e1008_reduce160_block2_kernel():
source=_E185_N176_BLOCK4_REDUCE_SOURCE.replace('constexpr int N=176,TRI=N*(N+1)/2,B=4;','constexpr int N=160,TRI=N*(N+1)/2,B=2;').replace(_E185_N176_BLOCK4_REDUCE_NAME,_E1008_REDUCE160_BLOCK2_NAME).replace('__launch_bounds__(768,1)','__launch_bounds__(320,1)').replace('total=tid<24?scratch[tid]:0.f','total=tid<10?scratch[tid]:0.f').replace('constexpr int G=4;','constexpr int G=2;').replace('constexpr int PAIRS=16,ROWS=48;','constexpr int PAIRS=16,ROWS=20;');entry=' int mid=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;'
if source.count(entry)!=1:raise RuntimeError('n160 reducer entry template changed')
source=source.replace(entry,entry+'\n if(tid==0) asm volatile("griddepcontrol.launch_dependents;":::);',1);tail=' timers[mid*5+1]=clock64();\n }\n}';ready=' timers[mid*5+1]=clock64();\n }\n __threadfence();\n __syncthreads();\n if(tid==0) timers[mid*5+4]=1ULL;\n}'
if source.count(tail)!=1:raise RuntimeError('n160 reducer tail template changed')
source=source.replace(tail,ready,1);return CUDAKernel(_fast_nvrtc_compile(source,_E1008_REDUCE160_BLOCK2_NAME),_E1008_REDUCE160_BLOCK2_NAME)
@torch.no_grad()
def _e107_eigh160(data:torch.Tensor):batch,n,_=data.shape;vectors=torch.empty_like(data);values=torch.empty(data.shape[:-1],device=data.device);saved=torch.empty_like(data);timers=torch.zeros((batch,5),device=data.device,dtype=torch.int64);diagonal=torch.empty_like(values);off_diagonal=torch.empty_like(values);_e1008_reduce160_block2_kernel().launch((batch,1,1),(320,1,1),(data,saved,diagonal,off_diagonal,timers),shared_mem=_E1008_REDUCE160_BLOCK2_SHARED);_e107_solve160_kernel().launch_pdl((batch*2,1,1),(384,1,1),(saved,diagonal,off_diagonal,vectors,values,timers),shared_mem=56*1024);return _tensor_wy_backtransform_saved_(vectors,saved,use_tf32=True),values
def _e083_checkpoint_orthogonalizer160(start_checkpoint:int=0):
schedules=(3e-06,),(3e-07,),(1e-08,),(1e-08,);state={'checkpoint':start_checkpoint}
def orthogonalize(matrix:torch.Tensor):
checkpoint=min(state['checkpoint'],3);state['checkpoint']+=1;result=matrix;solve=_e106_trsm160_half2 if checkpoint<2 else _e083_trsm160
for ridge in schedules[checkpoint]:previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');gram=result.mT@result;torch.set_float32_matmul_precision(previous);lower=_e2409_potrf160_wmma(gram,ridge)if checkpoint>=2 else _e083_potrf160(gram,ridge);result=solve(result,lower)
return result
return orthogonalize
_E188_N320_LANCZOS3_NAME='e188_n320_fused_lanczos3'
@memo(maxsize=1)
def _e188_n320_lanczos3_kernel():source=_E096_LANCZOS3_SOURCE.replace(_E096_LANCZOS3_NAME,_E188_N320_LANCZOS3_NAME).replace('constexpr int N = 352;','constexpr int N = 320;').replace('0.05330017908890261f','0.05590169943749474f');return CUDAKernel(_fast_nvrtc_compile(source,_E188_N320_LANCZOS3_NAME),_E188_N320_LANCZOS3_NAME)
@torch.no_grad()
def _e188_fused_n320_lanczos3(matrix:torch.Tensor):batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);dense_flags=torch.empty((batch,),device=matrix.device,dtype=torch.bool);_e188_n320_lanczos3_kernel().launch((batch,1,1),(512,1,1),(matrix,median,lower,upper,dense_flags));return median,lower,upper
_E189_N1024_LANCZOS3_NAME='e189_n1024_fused_lanczos3'
@memo(maxsize=1)
def _e189_n1024_lanczos3_kernel():source=_E096_LANCZOS3_SOURCE.replace(_E096_LANCZOS3_NAME,_E189_N1024_LANCZOS3_NAME).replace('constexpr int N = 352;','constexpr int N = 1024;').replace('0.05330017908890261f','0.03125f');return CUDAKernel(_fast_nvrtc_compile(source,_E189_N1024_LANCZOS3_NAME),_E189_N1024_LANCZOS3_NAME)
@torch.no_grad()
def _e189_fused_n1024_lanczos3(matrix:torch.Tensor,*,steps:int=3,probes:int=1):
if matrix.shape[-1]!=1024 or steps!=3:return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);dense_flags=torch.empty((batch,),device=matrix.device,dtype=torch.bool);_e189_n1024_lanczos3_kernel().launch((batch,1,1),(512,1,1),(matrix,median,lower,upper,dense_flags));return median,lower,upper
_E192_N512_LANCZOS5_NAME='e942_n512_fused_lanczos5_sturm'
@memo(maxsize=1)
def _e192_n512_lanczos5_kernel():
source=_E096_LANCZOS3_SOURCE.replace(_E096_LANCZOS3_NAME,_E192_N512_LANCZOS5_NAME).replace('constexpr int N = 352;','constexpr int N = 512;').replace('0.05330017908890261f','0.04419417382415922f').replace('__shared__ float diagonal[3];','__shared__ float diagonal[5];').replace('__shared__ float off[2];','__shared__ float off[4];').replace('step < 3','step < 5').replace('step < 2 && tid == 0','step < 4 && tid == 0');marker='extern "C" __global__ __launch_bounds__(512, 1)\n';helpers='\nstatic __device__ __forceinline__ int e942_sturm_count(\n const float* diagonal, const float* off, double shift) {\n double pivot = (double)diagonal[0] - shift;\n if (fabs(pivot) < 1.0e-30) pivot = -1.0e-30;\n int count = pivot < 0.0;\n #pragma unroll\n for (int index = 1; index < 5; ++index) {\n const double coupling = (double)off[index - 1];\n pivot = (double)diagonal[index] - shift\n - coupling * coupling / pivot;\n if (fabs(pivot) < 1.0e-30) pivot = -1.0e-30;\n count += pivot < 0.0;\n }\n return count;\n}\n\nstatic __device__ __forceinline__ double e942_node(\n const float* diagonal, const float* off, int order,\n double global_lower, double global_upper) {\n double lower = global_lower;\n double upper = global_upper;\n #pragma unroll 1\n for (int iteration = 0; iteration < 48; ++iteration) {\n const double middle = 0.5 * (lower + upper);\n if (e942_sturm_count(diagonal, off, middle) <= order)\n lower = middle;\n else\n upper = middle;\n }\n return 0.5 * (lower + upper);\n}\n\n'
if source.count(marker)!=1:raise RuntimeError('E942 launch marker changed')
source=source.replace(marker,helpers+marker,1);analytic=source.index(' if (tid == 0) {\n const float a = diagonal[0];');source=source[:analytic]+' if (tid == 0) {\n double global_lower = 1.0e300;\n double global_upper = -1.0e300;\n #pragma unroll\n for (int index = 0; index < 5; ++index) {\n const double radius = (index > 0 ? fabs((double)off[index - 1]) : 0.0)\n + (index < 4 ? fabs((double)off[index]) : 0.0);\n global_lower = fmin(global_lower, (double)diagonal[index] - radius);\n global_upper = fmax(global_upper, (double)diagonal[index] + radius);\n }\n const double padding = fmax(\n 1.0e-12, (global_upper - global_lower) * 1.0e-7);\n global_lower -= padding;\n global_upper += padding;\n\n double nodes[5];\n double weights[5];\n double weight_sum = 0.0;\n #pragma unroll\n for (int order = 0; order < 5; ++order) {\n const double node = e942_node(\n diagonal, off, order, global_lower, global_upper);\n nodes[order] = node;\n double previous = 0.0;\n double component = 1.0;\n double norm = 1.0;\n #pragma unroll\n for (int row = 0; row < 4; ++row) {\n double coupling = (double)off[row];\n if (fabs(coupling) < 1.0e-30)\n coupling = copysign(1.0e-30, coupling == 0.0 ? 1.0 : coupling);\n const double next = ((node - (double)diagonal[row]) * component\n - (row > 0 ? (double)off[row - 1] * previous : 0.0))\n / coupling;\n previous = component;\n component = next;\n norm = fma(component, component, norm);\n }\n weights[order] = 1.0 / norm;\n weight_sum += weights[order];\n }\n\n double cumulative = 0.0;\n double selected = nodes[4];\n #pragma unroll\n for (int index = 0; index < 5; ++index) {\n cumulative += weights[index] / weight_sum;\n if (cumulative >= 0.5) {\n selected = nodes[index];\n break;\n }\n }\n median[batch] = (float)selected;\n lower[batch] = (float)nodes[0];\n upper[batch] = (float)nodes[4];\n }\n}\n';return CUDAKernel(_fast_nvrtc_compile(source,_E192_N512_LANCZOS5_NAME),_E192_N512_LANCZOS5_NAME)
@torch.no_grad()
def _e192_fused_n512_lanczos5(matrix:torch.Tensor,*,steps:int=5,probes:int=1):
if matrix.shape[-1]!=512 or steps!=5:return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);dense_flags=torch.empty((batch,),device=matrix.device,dtype=torch.bool);_e192_n512_lanczos5_kernel().launch((batch,1,1),(512,1,1),(matrix,median,lower,upper,dense_flags));return median,lower,upper
_E924_N512_LANCZOS3_NAME='e924_n512_fused_lanczos3'
@memo(maxsize=1)
def _e924_n512_lanczos3_kernel():source=_E096_LANCZOS3_SOURCE.replace(_E096_LANCZOS3_NAME,_E924_N512_LANCZOS3_NAME).replace('constexpr int N = 352;','constexpr int N = 512;').replace('0.05330017908890261f','0.04419417382415922f');return CUDAKernel(_fast_nvrtc_compile(source,_E924_N512_LANCZOS3_NAME),_E924_N512_LANCZOS3_NAME)
@torch.no_grad()
def _e924_fused_n512_lanczos3(matrix:torch.Tensor,*,steps:int=3,probes:int=2):
if matrix.shape[-1]!=512 or steps!=3:return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);dense_flags=torch.empty((batch,),device=matrix.device,dtype=torch.bool);_e924_n512_lanczos3_kernel().launch((batch,1,1),(512,1,1),(matrix,median,lower,upper,dense_flags));return median,lower,upper
_E931_N2048_LANCZOS3_NAME='e931_n2048_cluster8_lanczos3';_E931_N2048_LANCZOS3_SOURCE='\n#include <cuda_runtime.h>\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;\n\nstatic __device__ __forceinline__ float e931_warp_sum(float value) {\n value += __shfl_down_sync(0xffffffffu, value, 16);\n value += __shfl_down_sync(0xffffffffu, value, 8);\n value += __shfl_down_sync(0xffffffffu, value, 4);\n value += __shfl_down_sync(0xffffffffu, value, 2);\n value += __shfl_down_sync(0xffffffffu, value, 1);\n return value;\n}\n\nextern "C" __global__\n__cluster_dims__(8, 1, 1)\n__launch_bounds__(256, 1)\nvoid e931_n2048_cluster8_lanczos3(\n const float* __restrict__ matrix,\n float* __restrict__ median,\n float* __restrict__ lower,\n float* __restrict__ upper) {\n constexpr int N = 2048;\n constexpr int CLUSTER = 8;\n constexpr int SLICE = N / CLUSTER;\n constexpr int WARPS = 8;\n cg::cluster_group cluster = cg::this_cluster();\n const int rank = cluster.block_rank();\n const int batch = blockIdx.x / CLUSTER;\n const int tid = threadIdx.x;\n const int lane = tid & 31;\n const int warp = tid >> 5;\n const long long matrix_base = (long long)batch * N * N;\n\n __shared__ float previous[SLICE];\n __shared__ float current[SLICE];\n __shared__ float product[SLICE];\n __shared__ float warp_partial[WARPS];\n __shared__ float cta_reduction[1];\n __shared__ float root_scalar[1];\n __shared__ float root_diagonal[3];\n __shared__ float root_off[2];\n\n previous[tid] = 0.0f;\n current[tid] = 0.02209708691207961f;\n cluster.sync();\n float beta = 0.0f;\n\n #pragma unroll\n for (int step = 0; step < 3; ++step) {\n for (int local_row = warp; local_row < SLICE; local_row += WARPS) {\n const int row = rank * SLICE + local_row;\n float value = 0.0f;\n #pragma unroll\n for (int source_rank = 0; source_rank < CLUSTER; ++source_rank) {\n float* remote = cluster.map_shared_rank(current, source_rank);\n const int column_base = source_rank * SLICE;\n for (int local_column = lane; local_column < SLICE; local_column += 32) {\n const float a = matrix[\n matrix_base + (long long)row * N + column_base + local_column];\n value = fmaf(a, remote[local_column], value);\n }\n }\n value = e931_warp_sum(value);\n if (lane == 0)\n product[local_row] = value - beta * previous[local_row];\n }\n __syncthreads();\n\n float partial = current[tid] * product[tid];\n partial = e931_warp_sum(partial);\n if (lane == 0) warp_partial[warp] = partial;\n __syncthreads();\n if (warp == 0) {\n float value = lane < WARPS ? warp_partial[lane] : 0.0f;\n value = e931_warp_sum(value);\n if (lane == 0) cta_reduction[0] = value;\n }\n cluster.sync();\n if (rank == 0 && tid == 0) {\n float value = 0.0f;\n #pragma unroll\n for (int source_rank = 0; source_rank < CLUSTER; ++source_rank)\n value += cluster.map_shared_rank(cta_reduction, source_rank)[0];\n root_scalar[0] = value;\n root_diagonal[step] = value;\n }\n cluster.sync();\n const float alpha = cluster.map_shared_rank(root_scalar, 0)[0];\n\n const float residual = product[tid] - alpha * current[tid];\n product[tid] = residual;\n partial = e931_warp_sum(residual * residual);\n if (lane == 0) warp_partial[warp] = partial;\n __syncthreads();\n if (warp == 0) {\n float value = lane < WARPS ? warp_partial[lane] : 0.0f;\n value = e931_warp_sum(value);\n if (lane == 0) cta_reduction[0] = value;\n }\n cluster.sync();\n if (rank == 0 && tid == 0) {\n float value = 0.0f;\n #pragma unroll\n for (int source_rank = 0; source_rank < CLUSTER; ++source_rank)\n value += cluster.map_shared_rank(cta_reduction, source_rank)[0];\n value = sqrtf(fmaxf(value, 1.0e-40f));\n root_scalar[0] = value;\n if (step < 2) root_off[step] = value;\n }\n cluster.sync();\n const float next_beta = cluster.map_shared_rank(root_scalar, 0)[0];\n previous[tid] = current[tid];\n current[tid] = residual / next_beta;\n beta = next_beta;\n cluster.sync();\n }\n\n if (rank == 0 && tid == 0) {\n const float a = root_diagonal[0];\n const float d = root_diagonal[1];\n const float f = root_diagonal[2];\n const float b = root_off[0];\n const float e = root_off[1];\n const float q = (a + d + f) * (1.0f / 3.0f);\n const float aa = a - q;\n const float dd = d - q;\n const float ff = f - q;\n const float p = sqrtf(fmaxf(\n (aa * aa + dd * dd + ff * ff + 2.0f * (b * b + e * e))\n * (1.0f / 6.0f), 1.0e-30f));\n const float ia = aa / p;\n const float id = dd / p;\n const float iff = ff / p;\n const float ib = b / p;\n const float ie = e / p;\n const float determinant = ia * (id * iff - ie * ie) - ib * ib * iff;\n const float r = fminf(1.0f, fmaxf(-1.0f, 0.5f * determinant));\n const float phi = acosf(r) * (1.0f / 3.0f);\n const float maximum = q + 2.0f * p * cosf(phi);\n const float minimum = q + 2.0f * p * cosf(phi + 2.0943951023931953f);\n lower[batch] = minimum;\n upper[batch] = maximum;\n median[batch] = 3.0f * q - minimum - maximum;\n }\n}\n'
@memo(maxsize=1)
def _e931_n2048_lanczos3_kernel():return CUDAKernel(_fast_nvrtc_compile(_E931_N2048_LANCZOS3_SOURCE,_E931_N2048_LANCZOS3_NAME),_E931_N2048_LANCZOS3_NAME)
@torch.no_grad()
def _e931_fused_n2048_lanczos3(matrix:torch.Tensor,*,steps:int=3,probes:int=2):
if matrix.shape[-1]!=2048 or steps!=3 or probes!=2:return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);_e931_n2048_lanczos3_kernel().launch((batch*8,1,1),(256,1,1),(matrix,median,lower,upper));return median,lower,upper
_E940_N1024_LANCZOS3_NAME='e940_n1024_cluster4_fused_two_probe_lanczos3';_E940_N1024_LANCZOS3_SOURCE='\n#include <cuda_runtime.h>\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;\n\nstatic __device__ __forceinline__ float e940_warp_sum(float value) {\n value += __shfl_down_sync(0xffffffffu, value, 16);\n value += __shfl_down_sync(0xffffffffu, value, 8);\n value += __shfl_down_sync(0xffffffffu, value, 4);\n value += __shfl_down_sync(0xffffffffu, value, 2);\n value += __shfl_down_sync(0xffffffffu, value, 1);\n return value;\n}\n\nextern "C" __global__\n__cluster_dims__(4, 1, 1)\n__launch_bounds__(256, 1)\nvoid e940_n1024_cluster4_fused_two_probe_lanczos3(\n const float* __restrict__ matrix,\n float* __restrict__ median_out,\n float* __restrict__ lower_out,\n float* __restrict__ upper_out) {\n constexpr int N = 1024;\n constexpr int CLUSTER = 4;\n constexpr int SLICE = N / CLUSTER;\n constexpr int WARPS = 8;\n cg::cluster_group cluster = cg::this_cluster();\n const int rank = cluster.block_rank();\n const int batch = blockIdx.x / CLUSTER;\n const int tid = threadIdx.x;\n const int lane = tid & 31;\n const int warp = tid >> 5;\n const long long matrix_base = (long long)batch * N * N;\n\n __shared__ float previous[2][SLICE];\n __shared__ float current[2][SLICE];\n __shared__ float product[2][SLICE];\n __shared__ float warp_partial[2][WARPS];\n __shared__ float cta_reduction[2];\n __shared__ float root_scalar[2];\n __shared__ float root_diagonal[2][3];\n __shared__ float root_off[2][2];\n\n const int global_index = rank * SLICE + tid;\n previous[0][tid] = 0.0f;\n previous[1][tid] = 0.0f;\n current[0][tid] = 0.03125f;\n current[1][tid] = (global_index & 1) ? -0.03125f : 0.03125f;\n cluster.sync();\n float beta0 = 0.0f;\n float beta1 = 0.0f;\n\n #pragma unroll\n for (int step = 0; step < 3; ++step) {\n for (int local_row = warp; local_row < SLICE; local_row += WARPS) {\n const int row = rank * SLICE + local_row;\n float value0 = 0.0f;\n float value1 = 0.0f;\n #pragma unroll\n for (int source_rank = 0; source_rank < CLUSTER; ++source_rank) {\n float* remote0 = cluster.map_shared_rank(current[0], source_rank);\n float* remote1 = cluster.map_shared_rank(current[1], source_rank);\n const int column_base = source_rank * SLICE;\n for (int local_column = lane; local_column < SLICE; local_column += 32) {\n const float a = matrix[\n matrix_base + (long long)row * N + column_base + local_column];\n value0 = fmaf(a, remote0[local_column], value0);\n value1 = fmaf(a, remote1[local_column], value1);\n }\n }\n value0 = e940_warp_sum(value0);\n value1 = e940_warp_sum(value1);\n if (lane == 0) {\n product[0][local_row] = value0 - beta0 * previous[0][local_row];\n product[1][local_row] = value1 - beta1 * previous[1][local_row];\n }\n }\n __syncthreads();\n\n #pragma unroll\n for (int probe = 0; probe < 2; ++probe) {\n float partial = current[probe][tid] * product[probe][tid];\n partial = e940_warp_sum(partial);\n if (lane == 0) warp_partial[probe][warp] = partial;\n __syncthreads();\n if (warp == 0) {\n float value = lane < WARPS ? warp_partial[probe][lane] : 0.0f;\n value = e940_warp_sum(value);\n if (lane == 0) cta_reduction[probe] = value;\n }\n __syncthreads();\n }\n cluster.sync();\n if (rank == 0 && tid < 2) {\n const int probe = tid;\n float value = 0.0f;\n #pragma unroll\n for (int source_rank = 0; source_rank < CLUSTER; ++source_rank)\n value += cluster.map_shared_rank(cta_reduction, source_rank)[probe];\n root_scalar[probe] = value;\n root_diagonal[probe][step] = value;\n }\n cluster.sync();\n const float alpha0 = cluster.map_shared_rank(root_scalar, 0)[0];\n const float alpha1 = cluster.map_shared_rank(root_scalar, 0)[1];\n\n const float residual0 = product[0][tid] - alpha0 * current[0][tid];\n const float residual1 = product[1][tid] - alpha1 * current[1][tid];\n product[0][tid] = residual0;\n product[1][tid] = residual1;\n float norm0 = e940_warp_sum(residual0 * residual0);\n float norm1 = e940_warp_sum(residual1 * residual1);\n if (lane == 0) {\n warp_partial[0][warp] = norm0;\n warp_partial[1][warp] = norm1;\n }\n __syncthreads();\n if (warp == 0) {\n float value0 = lane < WARPS ? warp_partial[0][lane] : 0.0f;\n float value1 = lane < WARPS ? warp_partial[1][lane] : 0.0f;\n value0 = e940_warp_sum(value0);\n value1 = e940_warp_sum(value1);\n if (lane == 0) {\n cta_reduction[0] = value0;\n cta_reduction[1] = value1;\n }\n }\n cluster.sync();\n if (rank == 0 && tid < 2) {\n const int probe = tid;\n float value = 0.0f;\n #pragma unroll\n for (int source_rank = 0; source_rank < CLUSTER; ++source_rank)\n value += cluster.map_shared_rank(cta_reduction, source_rank)[probe];\n value = sqrtf(fmaxf(value, 1.0e-40f));\n root_scalar[probe] = value;\n if (step < 2) root_off[probe][step] = value;\n }\n cluster.sync();\n const float next_beta0 = cluster.map_shared_rank(root_scalar, 0)[0];\n const float next_beta1 = cluster.map_shared_rank(root_scalar, 0)[1];\n previous[0][tid] = current[0][tid];\n previous[1][tid] = current[1][tid];\n current[0][tid] = residual0 / next_beta0;\n current[1][tid] = residual1 / next_beta1;\n beta0 = next_beta0;\n beta1 = next_beta1;\n cluster.sync();\n }\n\n if (rank == 0 && tid == 0) {\n float nodes[6];\n float weights[6];\n #pragma unroll\n for (int probe = 0; probe < 2; ++probe) {\n const float a = root_diagonal[probe][0];\n const float d = root_diagonal[probe][1];\n const float f = root_diagonal[probe][2];\n const float b = root_off[probe][0];\n const float ee = root_off[probe][1];\n const float q = (a + d + f) * (1.0f / 3.0f);\n const float aa = a - q;\n const float dd = d - q;\n const float ff = f - q;\n const float p = sqrtf(fmaxf(\n (aa * aa + dd * dd + ff * ff + 2.0f * (b * b + ee * ee))\n * (1.0f / 6.0f), 1.0e-30f));\n const float ia = aa / p;\n const float id = dd / p;\n const float iff = ff / p;\n const float ib = b / p;\n const float ie = ee / p;\n const float determinant = ia * (id * iff - ie * ie) - ib * ib * iff;\n const float r = fminf(1.0f, fmaxf(-1.0f, 0.5f * determinant));\n const float phi = acosf(r) * (1.0f / 3.0f);\n const float maximum = q + 2.0f * p * cosf(phi);\n const float minimum = q + 2.0f * p * cosf(phi + 2.0943951023931953f);\n const float middle = 3.0f * q - minimum - maximum;\n const int base = 3 * probe;\n nodes[base + 0] = minimum;\n nodes[base + 1] = middle;\n nodes[base + 2] = maximum;\n float local_weights[3];\n #pragma unroll\n for (int index = 0; index < 3; ++index) {\n const float lambda = nodes[base + index];\n const float other0 = nodes[base + ((index + 1) % 3)];\n const float other1 = nodes[base + ((index + 2) % 3)];\n const float numerator = (lambda - d) * (lambda - f) - ee * ee;\n const float denominator = (lambda - other0) * (lambda - other1);\n local_weights[index] = fmaxf(numerator / denominator, 0.0f);\n }\n const float inverse_weight = 0.5f / fmaxf(\n local_weights[0] + local_weights[1] + local_weights[2], 1.0e-20f);\n weights[base + 0] = local_weights[0] * inverse_weight;\n weights[base + 1] = local_weights[1] * inverse_weight;\n weights[base + 2] = local_weights[2] * inverse_weight;\n }\n #pragma unroll\n for (int index = 1; index < 6; ++index) {\n const float node = nodes[index];\n const float weight = weights[index];\n int position = index;\n while (position > 0 && nodes[position - 1] > node) {\n nodes[position] = nodes[position - 1];\n weights[position] = weights[position - 1];\n --position;\n }\n nodes[position] = node;\n weights[position] = weight;\n }\n float cumulative = 0.0f;\n float median = nodes[5];\n #pragma unroll\n for (int index = 0; index < 6; ++index) {\n cumulative += weights[index];\n if (cumulative >= 0.5f) {\n median = nodes[index];\n break;\n }\n }\n median_out[batch] = median;\n lower_out[batch] = nodes[0];\n upper_out[batch] = nodes[5];\n }\n}\n'
@memo(maxsize=1)
def _e940_n1024_lanczos3_kernel():return CUDAKernel(_fast_nvrtc_compile(_E940_N1024_LANCZOS3_SOURCE,_E940_N1024_LANCZOS3_NAME),_E940_N1024_LANCZOS3_NAME)
@torch.no_grad()
def _e940_fused_n1024_lanczos3(matrix:torch.Tensor,*,steps:int=3,probes:int=2):
if matrix.shape[-1]!=1024 or steps!=3 or probes!=2:return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median);_e940_n1024_lanczos3_kernel().launch((batch*4,1,1),(256,1,1),(matrix,median,lower,upper));return median,lower,upper
_E1052_N768_LANCZOS3_NAME='e1052_n768_cluster3_fused_two_probe_lanczos3';_E1052_N384_LANCZOS3_NAME='e1052_n384_cluster2_fused_two_probe_lanczos3'
def _e1052_nearrank_lanczos_source(n:int,cluster:int,threads:int,warps:int,name:str,inverse_sqrt:str):
source=_E940_N1024_LANCZOS3_SOURCE;replacements=(_E940_N1024_LANCZOS3_NAME,name),('__cluster_dims__(4, 1, 1)',f"__cluster_dims__({cluster}, 1, 1)"),('__launch_bounds__(256, 1)',f"__launch_bounds__({threads}, 1)"),('constexpr int N = 1024;',f"constexpr int N = {n};"),('constexpr int CLUSTER = 4;',f"constexpr int CLUSTER = {cluster};"),('constexpr int WARPS = 8;',f"constexpr int WARPS = {warps};"),('0.03125f',inverse_sqrt)
for(old,new)in replacements:
if old not in source:raise RuntimeError(f"nearrank Lanczos source anchor changed: {old}")
source=source.replace(old,new)
return source
@memo(maxsize=1)
def _e1052_n768_lanczos3_kernel():source=_e1052_nearrank_lanczos_source(768,3,256,8,_E1052_N768_LANCZOS3_NAME,'0.03608439182435161f');return CUDAKernel(_fast_nvrtc_compile(source,_E1052_N768_LANCZOS3_NAME),_E1052_N768_LANCZOS3_NAME)
@memo(maxsize=1)
def _e1052_n384_lanczos3_kernel():source=_e1052_nearrank_lanczos_source(384,2,192,6,_E1052_N384_LANCZOS3_NAME,'0.05103103630798288f');return CUDAKernel(_fast_nvrtc_compile(source,_E1052_N384_LANCZOS3_NAME),_E1052_N384_LANCZOS3_NAME)
@torch.no_grad()
def _e1052_fused_nearrank_lanczos3(matrix:torch.Tensor,*,steps:int=3,probes:int=2):
n=matrix.shape[-1]
if steps!=3 or probes!=2 or n not in(384,768):return stochastic_lanczos_stats(matrix,steps=steps,probes=probes)
batch=matrix.shape[0];median=torch.empty((batch,),device=matrix.device);lower=torch.empty_like(median);upper=torch.empty_like(median)
if n==768:_e1052_n768_lanczos3_kernel().launch((batch*3,1,1),(256,1,1),(matrix,median,lower,upper))
else:_e1052_n384_lanczos3_kernel().launch((batch*2,1,1),(192,1,1),(matrix,median,lower,upper))
return median,lower,upper
@torch.no_grad()
def _quintic_sign_corrected_step(sign:torch.Tensor,*,slope:float=2.3):cubic_coefficient=2.5-2.*slope;quintic_coefficient=slope-1.5;square=torch.bmm(sign,sign);polynomial=torch.baddbmm(square,square,square,beta=cubic_coefficient,alpha=quintic_coefficient);polynomial.diagonal(dim1=-2,dim2=-1).add_(slope);return torch.bmm(sign,polynomial)
_E945_N8_JACOBI_NAME='e945_n8_warp_jacobi_s4';_E945_N8_JACOBI_SOURCE='\n#include <cuda_runtime.h>\n\nextern "C" __global__ __launch_bounds__(32, 8)\nvoid e945_n8_warp_jacobi_s4(\n const float* __restrict__ input,\n float* __restrict__ vectors,\n float* __restrict__ values,\n int batch) {\n constexpr int N = 8;\n constexpr int PAIRS = 4;\n constexpr int SWEEPS = 4;\n const int matrix = (int)blockIdx.x;\n const int tid = (int)threadIdx.x;\n if (matrix >= batch) return;\n const long long base = (long long)matrix * N * N;\n __shared__ double a[N * N];\n __shared__ double q[N * N];\n __shared__ double cosine[PAIRS];\n __shared__ double sine[PAIRS];\n __shared__ int order[N];\n __shared__ int next_order[N];\n __shared__ int sorted[N];\n\n for (int index = tid; index < N * N; index += 32) {\n const int row = index / N;\n const int column = index - row * N;\n // Match torch.linalg.eigh\'s default UPLO=\'L\' representation.\n a[index] = (double)(row >= column\n ? input[base + index]\n : input[base + (long long)column * N + row]);\n q[index] = row == column ? 1.0 : 0.0;\n }\n if (tid < N) order[tid] = tid;\n __syncwarp();\n\n #pragma unroll\n for (int sweep = 0; sweep < SWEEPS; ++sweep) {\n #pragma unroll\n for (int round = 0; round < N - 1; ++round) {\n if (tid < PAIRS) {\n const int p = order[tid];\n const int r = order[N - 1 - tid];\n const double app = a[p * N + p];\n const double arr = a[r * N + r];\n const double apr = 0.5 * (a[p * N + r] + a[r * N + p]);\n double c = 1.0;\n double s = 0.0;\n const double threshold = 1.0e-13\n * sqrt(fmax(fabs(app * arr), 1.0e-300));\n if (fabs(apr) > threshold) {\n const double tau = (arr - app) / (2.0 * apr);\n const double t = copysign(\n 1.0 / (fabs(tau) + sqrt(1.0 + tau * tau)), tau);\n c = 1.0 / sqrt(1.0 + t * t);\n s = t * c;\n }\n cosine[tid] = c;\n sine[tid] = s;\n }\n __syncwarp();\n\n const int row = tid / PAIRS;\n const int pair = tid - row * PAIRS;\n const int p = order[pair];\n const int r = order[N - 1 - pair];\n const double c = cosine[pair];\n const double s = sine[pair];\n double x = a[row * N + p];\n double y = a[row * N + r];\n a[row * N + p] = fma(-s, y, c * x);\n a[row * N + r] = fma(s, x, c * y);\n __syncwarp();\n\n const int column = row;\n x = a[p * N + column];\n y = a[r * N + column];\n a[p * N + column] = fma(-s, y, c * x);\n a[r * N + column] = fma(s, x, c * y);\n __syncwarp();\n\n x = q[row * N + p];\n y = q[row * N + r];\n q[row * N + p] = fma(-s, y, c * x);\n q[row * N + r] = fma(s, x, c * y);\n __syncwarp();\n\n if (tid < N) {\n if (tid == 0) next_order[tid] = order[tid];\n else if (tid == 1) next_order[tid] = order[N - 1];\n else next_order[tid] = order[tid - 1];\n }\n __syncwarp();\n if (tid < N) order[tid] = next_order[tid];\n __syncwarp();\n }\n }\n\n if (tid == 0) {\n #pragma unroll\n for (int index = 0; index < N; ++index) sorted[index] = index;\n #pragma unroll\n for (int index = 1; index < N; ++index) {\n const int item = sorted[index];\n const double value = a[item * N + item];\n int position = index;\n while (position > 0\n && a[sorted[position - 1] * N + sorted[position - 1]] > value) {\n sorted[position] = sorted[position - 1];\n --position;\n }\n sorted[position] = item;\n }\n }\n __syncwarp();\n for (int index = tid; index < N * N; index += 32) {\n const int row = index / N;\n const int column = index - row * N;\n vectors[base + index] = (float)q[row * N + sorted[column]];\n }\n if (tid < N)\n values[matrix * N + tid] = (float)a[sorted[tid] * N + sorted[tid]];\n}\n'
@memo(maxsize=1)
def _e945_n8_jacobi_kernel():return CUDAKernel(_fast_nvrtc_compile(_E945_N8_JACOBI_SOURCE,_E945_N8_JACOBI_NAME),_E945_N8_JACOBI_NAME)
@torch.no_grad()
def _e945_n8_eigh(matrix:torch.Tensor):vectors=torch.empty_like(matrix);values=torch.empty(matrix.shape[:-1],device=matrix.device);_e945_n8_jacobi_kernel().launch((matrix.shape[0],1,1),(32,1,1),(matrix,vectors,values,matrix.shape[0]));return vectors,values
@triton.jit
def _e2695_nonic_tail_kernel(square,fourth,tail,total:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);mask=offsets<total;square_value=tl.load(square+offsets,mask=mask,other=.0).to(tl.float32);fourth_value=tl.load(fourth+offsets,mask=mask,other=.0).to(tl.float32);scaled=(-10.7625*square_value).to(tl.float16).to(tl.float32);tl.store(tail+offsets,scaled+2.6125*fourth_value,mask=mask)
@triton.jit
def _e2695_nonic_finish_kernel(polynomial,square,fourth,total:tl.constexpr,n:tl.constexpr,block:tl.constexpr):offsets=tl.program_id(0)*block+tl.arange(0,block);mask=offsets<total;element=offsets%(n*n);row=element//n;column=element-row*n;value=tl.load(polynomial+offsets,mask=mask,other=.0).to(tl.float32);square_value=tl.load(square+offsets,mask=mask,other=.0).to(tl.float32);fourth_value=tl.load(fourth+offsets,mask=mask,other=.0).to(tl.float32);value=(value-12.6375*square_value).to(tl.float16).to(tl.float32);value=(value+16.9875*fourth_value).to(tl.float16).to(tl.float32);value=tl.where(row==column,(value+4.8).to(tl.float16).to(tl.float32),value);tl.store(polynomial+offsets,value,mask=mask)
@torch.no_grad()
def _e2695_nonic_sign_step320(sign:torch.Tensor):square=torch.bmm(sign,sign);fourth=torch.bmm(square,square);tail=torch.empty_like(square);total=square.numel();_e2695_nonic_tail_kernel[triton.cdiv(total,4096),](square,fourth,tail,total=total,block=4096,num_warps=8,num_stages=1);polynomial=torch.bmm(fourth,tail);_e2695_nonic_finish_kernel[triton.cdiv(total,4096),](polynomial,square,fourth,total=total,n=sign.shape[-1],block=4096,num_warps=8,num_stages=1);return torch.bmm(sign,polynomial)
@torch.no_grad()
def _e2695_nonic_precompile():sign=torch.empty((1,320,320),device='cuda',dtype=torch.float32);_e2695_nonic_sign_step320(sign)
@torch.no_grad()
def _e083_split320_impl(projected:torch.Tensor,*,skip_first_checkpoint:bool):
global _E083_DEFERRED_INNER_FAILED;batch,n,_=projected.shape;rank=160;trace_center=projected.diagonal(dim1=-2,dim2=-1).mean(dim=-1);lanczos_median,lower,upper=_e188_fused_n320_lanczos3(projected);spectral_scale=torch.maximum(lower.abs(),upper.abs()).clamp_min(1e-20);center=torch.where((lower>=-.01*spectral_scale)|(upper<=.01*spectral_scale),lanczos_median,trace_center);radius=torch.maximum((lower-center).abs(),(upper-center).abs());radius=radius.clamp_min(1e-20)*1.2;sign=_shift_scale_symmetric_half(projected,center,radius)
for _ in range(2):square=torch.bmm(sign,sign);sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5)
sign=_e2695_nonic_sign_step320(sign);square=torch.bmm(sign,sign);sign=torch.baddbmm(sign,sign,square,beta=1.5,alpha=-.5);low_projector=_e483_half_sign_projector(sign);rows=torch.arange(n,device=projected.device)[:,None]+.5;columns=torch.arange(n,device=projected.device)[None];dct=torch.cos(3.141592653589793*rows*columns/n)*.07905694150420949;dct[:,0]*=.7071067811865476;low=dct[:,:rank];low_operator=low_projector@low_projector;low_operator=low_operator@low_operator;orthogonalize=_e083_checkpoint_orthogonalizer160(start_checkpoint=1 if skip_first_checkpoint else 0)
for checkpoint in range(3):
low=low_operator@low
if checkpoint!=0 or not skip_first_checkpoint:low=orthogonalize(low)
high=_rankdef_householder_complement(low,wmma160=True);low_product=projected@low;high_product=projected@high;children=torch.empty((2*batch,rank,rank),device=projected.device,dtype=projected.dtype);torch.bmm(low.mT,low_product,out=children[:batch]);torch.bmm(high.mT,high_product,out=children[batch:]);cross_block=low.mT@high_product;child_vectors,child_values=_e107_eigh160(children);low_child=child_vectors[:batch];high_child=child_vectors[batch:];low_vectors=low@low_child;high_vectors=high@high_child;low_values=child_values[:batch];high_values=child_values[batch:];cross=low_child.mT@cross_block@high_child;denominator=high_values[:,None,:]-low_values[:,:,None];signed_floor=torch.where(denominator>=.0,torch.full_like(denominator,1e-08),torch.full_like(denominator,-1e-08));denominator=torch.where(denominator.abs()>=1e-08,denominator,signed_floor);correction=(cross/denominator).clamp(min=-.05,max=.05);original_low,original_high=low_vectors,high_vectors;low_vectors=original_low-original_high@correction.mT;high_vectors=original_high+original_low@correction;vectors=torch.cat((low_vectors,high_vectors),dim=-1);values=torch.cat((low_values,high_values),dim=-1);begin,end=rank-4,rank+4;window=vectors[:,:,begin:end];rotation,local_values=_e945_n8_eigh(window.mT@projected@window);vectors[:,:,begin:end]=window@rotation;values[:,begin:end]=local_values;torch.set_float32_matmul_precision('high');gram=vectors.mT@vectors;vectors=torch.baddbmm(vectors,vectors,gram,beta=1.5,alpha=-.5);torch.set_float32_matmul_precision('high');values=(vectors*(projected@vectors)).sum(dim=1);values,order=values.sort(dim=-1);vectors=_e160_gather_columns320(vectors,order)
if skip_first_checkpoint:_E083_DEFERRED_INNER_FAILED=~_e083_finite_rows(vectors,values)
else:_E083_DEFERRED_INNER_FAILED=torch.zeros(batch,device=vectors.device,dtype=torch.bool)
return vectors,values
@torch.no_grad()
def _e083_split320(projected:torch.Tensor):return _e083_split320_impl(projected,skip_first_checkpoint=True)
@triton.jit
def _e159_dense512_merge_values(top_values,values,negative_counts):batch=tl.program_id(0);columns=tl.arange(0,512);top=tl.load(top_values+batch*320+columns,mask=columns<320,other=.0);negative=tl.sum(tl.where(columns<320,top<.0,0).to(tl.int32),axis=0);tl.store(negative_counts+batch,negative);from_negative=columns<negative;from_tail=(columns>=negative)&(columns<negative+192);source_column=tl.where(from_negative,columns,columns-192);source_value=tl.load(top_values+batch*320+source_column,mask=~from_tail&(source_column>=0)&(source_column<320),other=.0);tl.store(values+batch*512+columns,source_value)
@triton.jit
def _e159_dense512_merge_vectors(top,tail,negative_counts,output,total,tail_batch_stride,tail_row_stride,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;matrix=offsets//(512*512);element=offsets-matrix*512*512;row=element//512;column=element-row*512;negative=tl.load(negative_counts+matrix,mask=mask,other=0);from_negative=column<negative;from_tail=(column>=negative)&(column<negative+192);top_column=tl.where(from_negative,column,column-192);tail_column=column-negative;top_value=tl.load(top+matrix*512*320+row*320+top_column,mask=mask&~from_tail&(top_column>=0)&(top_column<320),other=.0);tail_value=tl.load(tail+matrix*tail_batch_stride+row*tail_row_stride+tail_column,mask=mask&from_tail,other=.0);tl.store(output+offsets,tl.where(from_tail,tail_value,top_value),mask=mask)
@triton.jit
def _e3260_dense512_merge_vectors_ordered(top,tail,top_order,negative_counts,output,total,tail_batch_stride,tail_row_stride,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;matrix=offsets//(512*512);element=offsets-matrix*512*512;row=element//512;column=element-row*512;negative=tl.load(negative_counts+matrix,mask=mask,other=0);from_negative=column<negative;from_tail=(column>=negative)&(column<negative+192);sorted_top_column=tl.where(from_negative,column,column-192);valid_top=mask&~from_tail&(sorted_top_column>=0)&(sorted_top_column<320);source_top_column=tl.load(top_order+matrix*320+sorted_top_column,mask=valid_top,other=0);tail_column=column-negative;top_value=tl.load(top+matrix*512*320+row*320+source_top_column,mask=valid_top,other=.0);tail_value=tl.load(tail+matrix*tail_batch_stride+row*tail_row_stride+tail_column,mask=mask&from_tail,other=.0);tl.store(output+offsets,tl.where(from_tail,tail_value,top_value),mask=mask)
@triton.jit
def _e160_gather_columns_kernel(source,order,output,total,N:tl.constexpr,BLOCK:tl.constexpr):offsets=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total;matrix=offsets//(N*N);element=offsets-matrix*N*N;row=element//N;column=element-row*N;source_column=tl.load(order+matrix*N+column,mask=mask,other=0);value=tl.load(source+matrix*N*N+row*N+source_column,mask=mask,other=.0);tl.store(output+offsets,value,mask=mask)
@torch.no_grad()
def _e160_gather_columns320(vectors:torch.Tensor,order:torch.Tensor):return _e160_gather_columns(vectors,order)
@torch.no_grad()
def _e160_gather_columns(vectors:torch.Tensor,order:torch.Tensor):
if not vectors.is_contiguous()or vectors.shape[-1]!=vectors.shape[-2]:return vectors.gather(2,order[:,None,:].expand(-1,vectors.shape[-2],-1))
output=torch.empty_like(vectors);total=output.numel();_e160_gather_columns_kernel[triton.cdiv(total,4096),](vectors,order,output,total,N=vectors.shape[-1],BLOCK=4096,num_warps=8,num_stages=1);return output
@torch.no_grad()
def _e159_dense512_merge(top:torch.Tensor,tail:torch.Tensor,top_values:torch.Tensor):batch=top.shape[0];values=torch.empty((batch,512),device=top.device);negative_counts=torch.empty((batch,),device=top.device,dtype=torch.int32);_e159_dense512_merge_values[batch,](top_values,values,negative_counts,num_warps=8,num_stages=1);vectors=torch.empty((batch,512,512),device=top.device,dtype=top.dtype);total=vectors.numel();_e159_dense512_merge_vectors[triton.cdiv(total,1024),](top,tail,negative_counts,vectors,total,tail.stride(0),tail.stride(1),BLOCK=1024,num_warps=8,num_stages=1);return vectors,values
@torch.no_grad()
def _e3260_dense512_merge_ordered(top:torch.Tensor,tail:torch.Tensor,top_values:torch.Tensor,top_order:torch.Tensor):batch=top.shape[0];values=torch.empty((batch,512),device=top.device);negative_counts=torch.empty((batch,),device=top.device,dtype=torch.int32);_e159_dense512_merge_values[batch,](top_values,values,negative_counts,num_warps=8,num_stages=1);vectors=torch.empty((batch,512,512),device=top.device,dtype=top.dtype);total=vectors.numel();_e3260_dense512_merge_vectors_ordered[triton.cdiv(total,1024),](top,tail,top_order,negative_counts,vectors,total,tail.stride(0),tail.stride(1),BLOCK=1024,num_warps=8,num_stages=1);return vectors,values
@triton.jit
def _e2539_dense512_pregram_risk_kernel(gram,risk,threshold:tl.constexpr,N:tl.constexpr,BLOCK:tl.constexpr):matrix=tl.program_id(0);offsets=tl.arange(0,BLOCK);mask=offsets<N;values=tl.load(gram+matrix*N*N+offsets*N+offsets,mask=mask,other=1.);finite=(values<=3.402823466e38)&(values>=-3.402823466e38);bad=tl.max(tl.where(mask&((tl.abs(values-1.)>threshold)|~finite),1,0),axis=0);tl.store(risk+matrix,bad!=0)
def _e2539_dense512_pregram_risk(gram:torch.Tensor,threshold:float=.02):risk=torch.empty((gram.shape[0],),device=gram.device,dtype=torch.bool);_e2539_dense512_pregram_risk_kernel[gram.shape[0],](gram,risk,threshold=threshold,N=512,BLOCK=512,num_warps=8,num_stages=1);return risk
_E2760_DENSE512_NEWTON_GUARD_NAME='e2760_dense512_fused_newton_guard';_E2760_DENSE512_NEWTON_GUARD_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(256, 2)\nvoid e2760_dense512_fused_newton_guard(\n float* __restrict__ gram,\n int* __restrict__ risk,\n int batch, float threshold) {\n constexpr int N = 512;\n const int matrix = blockIdx.x;\n const int tile = blockIdx.y;\n const int tid = threadIdx.x;\n const int lane = tid & 31;\n const int warp = tid >> 5;\n if (matrix >= batch) return;\n float* source = gram + (long long)matrix * N * N;\n #pragma unroll\n for (int group = 0; group < 4; ++group) {\n const int row = tile * 32 + warp + group * 8;\n float total = 0.0f;\n int finite = 1;\n #pragma unroll\n for (int column = lane; column < N; column += 32) {\n const long long offset = (long long)row * N + column;\n const float value = source[offset];\n finite &= isfinite(value);\n total += fabsf(value - (row == column ? 1.0f : 0.0f));\n source[offset] = -0.5f * value\n + (row == column ? 1.5f : 0.0f);\n }\n #pragma unroll\n for (int delta = 16; delta; delta >>= 1) {\n total += __shfl_down_sync(0xffffffffu, total, delta);\n finite &= __shfl_down_sync(0xffffffffu, finite, delta);\n }\n if (lane == 0 && (!finite || total > threshold))\n atomicExch(risk + matrix, 1);\n }\n}\n'
@memo(maxsize=1)
def _e2760_dense512_newton_guard_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2760_DENSE512_NEWTON_GUARD_SOURCE,_E2760_DENSE512_NEWTON_GUARD_NAME),_E2760_DENSE512_NEWTON_GUARD_NAME)
@torch.no_grad()
def _e2760_dense512_newton_guard_(gram:torch.Tensor):batch=gram.shape[0];risk=torch.zeros((batch,),device=gram.device,dtype=torch.int32);_e2760_dense512_newton_guard_kernel().launch((batch,16,1),(256,1,1),(gram,risk,batch,.1));return risk!=0
_E2860_DENSE512_STAGE_SELECT_PACK_NAME='e2860_dense512_stage_select_pack';_E2860_DENSE512_STAGE_SELECT_PACK_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(256, 2)\nvoid e2860_dense512_stage_select_pack(\n const float* __restrict__ projected,\n const float* __restrict__ cross,\n const float* __restrict__ values,\n long long* __restrict__ indices,\n float* __restrict__ local,\n int current,\n int batch) {\n constexpr int MAX_N = 256;\n constexpr int KEEP = 64;\n const int matrix = blockIdx.x;\n const int tid = threadIdx.x;\n if (matrix >= batch) return;\n const int warp = tid >> 5;\n const int lane = tid & 31;\n __shared__ float energies[MAX_N];\n __shared__ int order[MAX_N];\n #pragma unroll 1\n for (int wave = 0; wave < 32; ++wave) {\n const int row = wave * 8 + warp;\n float energy = 0.0f;\n if (row < current) {\n const long long base = (long long)matrix * current * KEEP\n + (long long)row * KEEP;\n energy = cross[base + lane] * cross[base + lane]\n + cross[base + lane + 32] * cross[base + lane + 32];\n }\n #pragma unroll\n for (int delta = 16; delta > 0; delta >>= 1)\n energy += __shfl_down_sync(0xffffffffu, energy, delta);\n if (lane == 0) {\n energies[row] = row < current\n ? energy : -__int_as_float(0x7f800000);\n order[row] = row;\n }\n }\n __syncthreads();\n for (int width = 2; width <= MAX_N; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n const int peer = tid ^ stride;\n if (peer > tid) {\n const float left_value = energies[tid];\n const float right_value = energies[peer];\n const int left_index = order[tid];\n const int right_index = order[peer];\n const bool ascending = (tid & width) == 0;\n const bool greater = left_value > right_value\n || (left_value == right_value && left_index > right_index);\n const bool swap = ascending ? greater : !greater;\n if (swap) {\n energies[tid] = right_value;\n energies[peer] = left_value;\n order[tid] = right_index;\n order[peer] = left_index;\n }\n }\n __syncthreads();\n }\n }\n if (tid < KEEP)\n indices[(long long)matrix * KEEP + tid] = order[MAX_N - 1 - tid];\n __syncthreads();\n const long long local_base = (long long)matrix * 128 * 128;\n const long long cross_base = (long long)matrix * current * KEEP;\n const long long projected_base = (long long)matrix * 320 * 320;\n for (int linear = tid; linear < 128 * 128; linear += blockDim.x) {\n const int row = linear >> 7;\n const int column = linear & 127;\n float value = 0.0f;\n if (row < KEEP && column < KEEP) {\n if (row == column) {\n const int source = order[MAX_N - 1 - row];\n value = values[(long long)matrix * current + source];\n }\n } else if (row < KEEP) {\n const int source = order[MAX_N - 1 - row];\n value = cross[cross_base + (long long)source * KEEP + column - KEEP];\n } else if (column < KEEP) {\n const int source = order[MAX_N - 1 - column];\n value = cross[cross_base + (long long)source * KEEP + row - KEEP];\n } else {\n value = projected[projected_base\n + (long long)(current + row - KEEP) * 320\n + current + column - KEEP];\n }\n local[local_base + linear] = value;\n }\n}\n'
@memo(maxsize=1)
def _e2860_dense512_stage_select_pack_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2860_DENSE512_STAGE_SELECT_PACK_SOURCE,_E2860_DENSE512_STAGE_SELECT_PACK_NAME),_E2860_DENSE512_STAGE_SELECT_PACK_NAME)
@torch.no_grad()
def _e2860_dense512_stage_select_pack(projected:torch.Tensor,cross:torch.Tensor,values:torch.Tensor,current:int):batch=projected.shape[0];indices=torch.empty((batch,64),device=projected.device,dtype=torch.int64);local=torch.empty((batch,128,128),device=projected.device);_e2860_dense512_stage_select_pack_kernel().launch((batch,1,1),(256,1,1),(projected,cross,values,indices,local,current,batch));return indices,local
@triton.jit
def _e2796_dense512_cross_energy_kernel(cross,energy,current:tl.constexpr,BLOCK:tl.constexpr):row=tl.program_id(0);batch=tl.program_id(1);columns=tl.arange(0,BLOCK);values=tl.load(cross+batch*current*64+row*64+columns,mask=columns<64,other=.0);tl.store(energy+batch*current+row,tl.sum(values*values,axis=0))
@triton.jit
def _e2796_dense512_pack_local128_kernel(projected,cross,values,indices,local,current:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);offsets=tl.program_id(1)*BLOCK+tl.arange(0,BLOCK);mask=offsets<128*128;row=offsets//128;column=offsets-row*128;upper_row=tl.load(indices+batch*64+row,mask=mask&(row<64),other=0);upper_column=tl.load(indices+batch*64+column,mask=mask&(column<64),other=0);diagonal=tl.load(values+batch*current+upper_row,mask=mask&(row<64)&(column<64)&(row==column),other=.0);upper=tl.load(cross+batch*current*64+upper_row*64+column-64,mask=mask&(row<64)&(column>=64),other=.0);lower=tl.load(cross+batch*current*64+upper_column*64+row-64,mask=mask&(row>=64)&(column<64),other=.0);bottom=tl.load(projected+batch*320*320+(current+row-64)*320+current+column-64,mask=mask&(row>=64)&(column>=64),other=.0);value=tl.where(row<64,tl.where(column<64,diagonal,upper),tl.where(column<64,lower,bottom));tl.store(local+batch*128*128+offsets,value,mask=mask)
@triton.jit
def _e2796_dense512_expand_base_kernel(old,repaired_top,local_rotation,output,current:tl.constexpr,total:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);offsets=tl.program_id(1)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total*total;row=offsets//total;column=offsets-row*total;old_value=tl.load(old+batch*current*current+row*current+column,mask=mask&(row<current)&(column<current),other=.0);new_column=column-current;top_new=tl.load(repaired_top+batch*current*128+row*128+64+new_column,mask=mask&(row<current)&(column>=current),other=.0);new_row=row-current;bottom_new=tl.load(local_rotation+batch*128*128+(64+new_row)*128+64+new_column,mask=mask&(row>=current)&(column>=current),other=.0);value=tl.where(column<current,old_value,tl.where(row<current,top_new,bottom_new));tl.store(output+batch*total*total+offsets,value,mask=mask)
@triton.jit
def _e2796_dense512_expand_selected_kernel(indices,repaired_top,local_rotation,output,current:tl.constexpr,total:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);offsets=tl.program_id(1)*BLOCK+tl.arange(0,BLOCK);mask=offsets<total*64;row=offsets//64;position=offsets-row*64;column=tl.load(indices+batch*64+position,mask=mask,other=0);top=tl.load(repaired_top+batch*current*128+row*128+position,mask=mask&(row<current),other=.0);new_row=row-current;bottom=tl.load(local_rotation+batch*128*128+(64+new_row)*128+position,mask=mask&(row>=current),other=.0);value=tl.where(row<current,top,bottom);tl.store(output+batch*total*total+row*total+column,value,mask=mask)
@triton.jit
def _e2796_dense512_values_base_kernel(old,local_values,output,current:tl.constexpr,total:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);columns=tl.arange(0,BLOCK);old_value=tl.load(old+batch*current+columns,mask=columns<current,other=.0);new_value=tl.load(local_values+batch*128+64+columns-current,mask=(columns>=current)&(columns<total),other=.0);tl.store(output+batch*total+columns,tl.where(columns<current,old_value,new_value),mask=columns<total)
@triton.jit
def _e2796_dense512_values_selected_kernel(indices,local_values,output,total:tl.constexpr,BLOCK:tl.constexpr):batch=tl.program_id(0);positions=tl.arange(0,BLOCK);mask=positions<64;columns=tl.load(indices+batch*64+positions,mask=mask,other=0);value=tl.load(local_values+batch*128+positions,mask=mask,other=.0);tl.store(output+batch*total+columns,value,mask=mask)
def _e2796_dense512_expand_stage(rotation:torch.Tensor,values:torch.Tensor,indices:torch.Tensor,local_rotation:torch.Tensor,local_values:torch.Tensor,current:int):batch=rotation.shape[0];total=current+64;selected_rotation=rotation.gather(2,indices[:,None,:].expand(-1,current,-1));repaired_top=selected_rotation@local_rotation[:,:64,:];expanded=torch.empty((batch,total,total),device=rotation.device);block=256;_e2796_dense512_expand_base_kernel[batch,triton.cdiv(total*total,block)](rotation,repaired_top,local_rotation,expanded,current=current,total=total,BLOCK=block,num_warps=8,num_stages=1);_e2796_dense512_expand_selected_kernel[batch,triton.cdiv(total*64,block)](indices,repaired_top,local_rotation,expanded,current=current,total=total,BLOCK=block,num_warps=8,num_stages=1);expanded_values=torch.empty((batch,total),device=rotation.device);_e2796_dense512_values_base_kernel[batch,](values,local_values,expanded_values,current=current,total=total,BLOCK=512,num_warps=8,num_stages=1);_e2796_dense512_values_selected_kernel[batch,](indices,local_values,expanded_values,total=total,BLOCK=64,num_warps=2,num_stages=1);return expanded,expanded_values
@triton.jit
def _e2796_dense512_rayleigh_energy_kernel(rayleigh,energy,BLOCK:tl.constexpr):row=tl.program_id(0);batch=tl.program_id(1);columns=tl.arange(0,BLOCK);mask=columns<320;value=tl.load(rayleigh+batch*320*320+row*320+columns,mask=mask,other=.0);value=tl.where(columns==row,.0,value);tl.store(energy+batch*320+row,tl.sum(value*value,axis=0))
@triton.jit
def _e2796_dense512_pack_rayleigh128_kernel(rayleigh,indices,local,BLOCK:tl.constexpr):batch=tl.program_id(0);offsets=tl.program_id(1)*BLOCK+tl.arange(0,BLOCK);mask=offsets<128*128;row=offsets//128;column=offsets-row*128;source_row=tl.load(indices+batch*128+row,mask=mask,other=0);source_column=tl.load(indices+batch*128+column,mask=mask,other=0);forward=tl.load(rayleigh+batch*320*320+source_row*320+source_column,mask=mask,other=.0);reverse=tl.load(rayleigh+batch*320*320+source_column*320+source_row,mask=mask,other=.0);tl.store(local+batch*128*128+offsets,.5*(forward+reverse),mask=mask)
@triton.jit
def _e3263_dense512_fused_tail_update(rotation,indices,local_rotation,output,BLOCK_ROWS:tl.constexpr,BLOCK_STATE:tl.constexpr):batch=tl.program_id(0);row_base=tl.program_id(1)*BLOCK_ROWS;rows=row_base+tl.arange(0,BLOCK_ROWS)[:,None];state_columns=tl.arange(0,BLOCK_STATE)[None,:];state=tl.load(rotation+batch*320*320+rows*320+state_columns,mask=(rows<320)&(state_columns<320),other=.0);tl.store(output+batch*320*320+rows*320+state_columns,state,mask=(rows<320)&(state_columns<320));tl.debug_barrier();k=tl.arange(0,128);positions=tl.arange(0,128);selected_columns=tl.load(indices+batch*128+k);selected=tl.load(rotation+batch*320*320+rows*320+selected_columns[None,:],mask=rows<320,other=.0);factor=tl.load(local_rotation+batch*128*128+k[:,None]*128+positions[None,:]);transformed=tl.dot(selected,factor,input_precision='tf32',out_dtype=tl.float32);destinations=tl.load(indices+batch*128+positions);tl.store(output+batch*320*320+rows*320+destinations[None,:],transformed,mask=rows<320)
@triton.jit
def _e3200_dense512_indexed_factor_left_apply(x,q,indices,output,current:tl.constexpr,factor_current:tl.constexpr,initial:tl.constexpr,X_BATCH_STRIDE:tl.constexpr,X_ROW_STRIDE:tl.constexpr,BLOCK_N:tl.constexpr):
batch=tl.program_id(0);column_base=tl.program_id(1)*BLOCK_N;out_row=tl.arange(0,128);k=tl.arange(0,128);columns=column_base+tl.arange(0,BLOCK_N)
if initial:out_coordinates=out_row;k_coordinates=k
else:out_pick=tl.load(indices+batch*64+out_row,mask=out_row<64,other=0);k_pick=tl.load(indices+batch*64+k,mask=k<64,other=0);out_coordinates=tl.where(out_row<64,out_pick,factor_current+out_row-64);k_coordinates=tl.where(k<64,k_pick,factor_current+k-64)
q_tile=tl.load(q+batch*128*128+k[:,None]*128+out_row[None,:]);x_tile=tl.load(x+batch*X_BATCH_STRIDE+k_coordinates[:,None]*X_ROW_STRIDE+columns[None,:],mask=columns[None,:]<64,other=.0);result=tl.dot(tl.trans(q_tile.to(tl.float16)),x_tile.to(tl.float16));tl.store(output+batch*current*64+out_coordinates[:,None]*64+columns[None,:],result,mask=columns[None,:]<64)
if initial:bottom_row=tl.arange(0,128)+128;bottom=tl.load(x+batch*X_BATCH_STRIDE+bottom_row[:,None]*X_ROW_STRIDE+columns[None,:],mask=(bottom_row[:,None]<current)&(columns[None,:]<64),other=.0);tl.store(output+batch*current*64+bottom_row[:,None]*64+columns[None,:],bottom,mask=(bottom_row[:,None]<current)&(columns[None,:]<64))
@torch.no_grad()
def _e3200_dense512_factor_cross(projected:torch.Tensor,initial_rotation:torch.Tensor,factors:list,current:int):
batch=projected.shape[0];raw=projected[:,:current,current:current+64];transformed=torch.empty((batch,current,64),device=projected.device);_e3200_dense512_indexed_factor_left_apply[batch,1](raw,initial_rotation,raw,transformed,current=current,factor_current=0,initial=True,X_BATCH_STRIDE=320*320,X_ROW_STRIDE=320,BLOCK_N=64,num_warps=8,num_stages=1)
for(factor_current,indices,local_rotation,_)in factors:_e3200_dense512_indexed_factor_left_apply[batch,1](transformed,local_rotation,indices,transformed,current=current,factor_current=factor_current,initial=False,X_BATCH_STRIDE=current*64,X_ROW_STRIDE=64,BLOCK_N=64,num_warps=8,num_stages=1)
return transformed
@triton.jit
def _e3200_dense512_fused_factor_owner(initial,indices1,factor1,indices2,factor2,indices3,factor3,output,BLOCK_ROWS:tl.constexpr,BLOCK_STATE:tl.constexpr):batch=tl.program_id(0);row_base=tl.program_id(1)*BLOCK_ROWS;local_row=tl.arange(0,BLOCK_ROWS)[:,None];row=row_base+local_row;state_column=tl.arange(0,BLOCK_STATE)[None,:];initial_value=tl.load(initial+batch*128*128+row*128+state_column,mask=(row<128)&(state_column<128),other=.0);identity=(row==state_column)&(row>=128)&(state_column<320);state=initial_value+identity.to(tl.float32);state_pointer=output+batch*320*320+row*320+state_column;tl.store(state_pointer,state,mask=state_column<320);tl.debug_barrier();position=tl.arange(0,128)[None,:];k=tl.arange(0,64);pick1=tl.load(indices1+batch*64+k);selected1=tl.load(output+batch*320*320+row*320+pick1[None,:]).to(tl.float16);top1=tl.load(factor1+batch*128*128+k[:,None]*128+position).to(tl.float16);product1=tl.dot(selected1,top1,out_dtype=tl.float32);direct1=tl.load(factor1+batch*128*128+(64+local_row)*128+position,mask=row_base==128,other=.0);transformed1=tl.where(row_base<128,product1,direct1);destination1=tl.where(position<64,tl.load(indices1+batch*64+position,mask=position<64,other=0),128+position-64);tl.store(output+batch*320*320+row*320+destination1,transformed1,mask=row_base<=128);tl.debug_barrier();pick2=tl.load(indices2+batch*64+k);selected2=tl.load(output+batch*320*320+row*320+pick2[None,:]).to(tl.float16);top2=tl.load(factor2+batch*128*128+k[:,None]*128+position).to(tl.float16);product2=tl.dot(selected2,top2,out_dtype=tl.float32);direct2=tl.load(factor2+batch*128*128+(64+local_row)*128+position,mask=row_base==192,other=.0);transformed2=tl.where(row_base<192,product2,direct2);destination2=tl.where(position<64,tl.load(indices2+batch*64+position,mask=position<64,other=0),192+position-64);tl.store(output+batch*320*320+row*320+destination2,transformed2,mask=row_base<=192);tl.debug_barrier();pick3=tl.load(indices3+batch*64+k);selected3=tl.load(output+batch*320*320+row*320+pick3[None,:]).to(tl.float16);top3=tl.load(factor3+batch*128*128+k[:,None]*128+position).to(tl.float16);product3=tl.dot(selected3,top3,out_dtype=tl.float32);direct3=tl.load(factor3+batch*128*128+(64+local_row)*128+position,mask=row_base==256,other=.0);transformed3=tl.where(row_base<256,product3,direct3);destination3=tl.where(position<64,tl.load(indices3+batch*64+position,mask=position<64,other=0),256+position-64);tl.store(output+batch*320*320+row*320+destination3,transformed3,mask=row_base<=256)
@torch.no_grad()
def _e3200_dense512_materialize(initial_rotation:torch.Tensor,factors:list):batch=initial_rotation.shape[0];output=torch.empty((batch,320,320),device=initial_rotation.device);(_,indices1,factor1,_),(_,indices2,factor2,_),(_,indices3,factor3,_)=factors;_e3200_dense512_fused_factor_owner[batch,5](initial_rotation,indices1,factor1,indices2,factor2,indices3,factor3,output,BLOCK_ROWS=64,BLOCK_STATE=512,num_warps=4,num_stages=1);return output
@torch.no_grad()
def _e3200_dense512_update_values(values:torch.Tensor,indices:torch.Tensor,local_values:torch.Tensor,current:int):batch=values.shape[0];total=current+64;result=torch.empty((batch,total),device=values.device);_e2796_dense512_values_base_kernel[batch,](values,local_values,result,current=current,total=total,BLOCK=512,num_warps=8,num_stages=1);_e2796_dense512_values_selected_kernel[batch,](indices,local_values,result,total=total,BLOCK=64,num_warps=2,num_stages=1);return result
@torch.no_grad()
def _e2760_dense512_projected_eigh(projected:torch.Tensor):
batch=projected.shape[0];initial_rotation,values=_e196_lapack_n128_leaf_eigh(projected[:,:128,:128].contiguous(),fuse_mgs=True);factors=[]
for current in(128,192,256):cross=_e3200_dense512_factor_cross(projected,initial_rotation,factors,current);indices,local=_e2860_dense512_stage_select_pack(projected,cross,values,current);local_rotation,local_values=_e196_lapack_n128_leaf_eigh(local,fuse_mgs=current==128,compact_t_tf32=True);factors.append((current,indices,local_rotation,local_values));values=_e3200_dense512_update_values(values,indices,local_values,current)
rotation=_e3200_dense512_materialize(initial_rotation,factors);rayleigh=rotation.mT@projected@rotation;diagonal=rayleigh.diagonal(dim1=1,dim2=2);energy=torch.empty((batch,320),device=projected.device);_e2796_dense512_rayleigh_energy_kernel[320,batch](rayleigh,energy,BLOCK=512,num_warps=8,num_stages=1);indices=energy.topk(128,dim=1).indices;local=torch.empty((batch,128,128),device=projected.device);_e2796_dense512_pack_rayleigh128_kernel[batch,triton.cdiv(128*128,256)](rayleigh,indices,local,BLOCK=256,num_warps=8,num_stages=1);local_rotation,local_values=_e196_lapack_n128_leaf_eigh(local,fuse_mgs=True,compact_t_tf32=True);updated=torch.empty_like(rotation);_e3263_dense512_fused_tail_update[batch,20](rotation,indices,local_rotation,updated,BLOCK_ROWS=16,BLOCK_STATE=512,num_warps=4,num_stages=1);values=diagonal.clone();values.scatter_(1,indices,local_values);values,order=values.sort(dim=1);return updated,values,order
@torch.no_grad()
def _e083_dense512_legacy_impl(data:torch.Tensor,*,safe_split:bool):
batch,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:
basis=data@data[:,:,:320];basis=normalize_columns_(basis);complete=mixed_qr_active(basis.contiguous(),complete=True,module=None);top=complete[:,:,:320];tail=complete[:,:,320:];projected=top.mT@data@top
if safe_split:rotation,top_values=_e083_split320_impl(projected,skip_first_checkpoint=False)
else:rotation,top_values=_e083_split320(projected)
top=top@rotation;vectors,values=_e159_dense512_merge(top,tail,top_values);gram=vectors.mT@vectors;norm_risk=_e2539_dense512_pregram_risk(gram);gram.mul_(-.5);gram.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors=vectors@gram;return vectors,values,norm_risk
finally:torch.set_float32_matmul_precision(previous)
@torch.no_grad()
def _e083_dense512_impl(data:torch.Tensor,*,safe_split:bool):
global _E083_DEFERRED_INNER_FAILED
if safe_split:return _e083_dense512_legacy_impl(data,safe_split=True)
batch=data.shape[0];previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('high')
try:basis=normalize_columns_(data@data[:,:,:320]);complete=mixed_qr_active(basis.contiguous(),complete=True,module=None);active_basis=complete[:,:,:320];tail=complete[:,:,320:];projected=active_basis.mT@data@active_basis;rotation,active_values,active_order=_e2760_dense512_projected_eigh(projected);active=active_basis@rotation;vectors,values=_e3260_dense512_merge_ordered(active,tail,active_values,active_order);gram=vectors.mT@vectors;risk=_e2760_dense512_newton_guard_(gram);vectors=vectors@gram;_E083_DEFERRED_INNER_FAILED=torch.zeros(batch,device=data.device,dtype=torch.bool);return vectors,values,risk
finally:torch.set_float32_matmul_precision(previous)
@triton.jit
def _e083_nonfinite_row_flags_kernel(vectors,values,flags,N:tl.constexpr,BLOCK:tl.constexpr):matrix=tl.program_id(0);chunk=tl.program_id(1);offsets=chunk*BLOCK+tl.arange(0,BLOCK);q_elements=N*N;total=N*N+N;q_mask=offsets<q_elements;l_mask=(offsets>=q_elements)&(offsets<total);q=tl.load(vectors+matrix*q_elements+offsets,mask=q_mask,other=.0);l=tl.load(values+matrix*N+offsets-q_elements,mask=l_mask,other=.0);value=tl.where(q_mask,q,l);valid=q_mask|l_mask;finite=(value<=3.402823466e38)&(value>=-3.402823466e38);bad=tl.max(tl.where(valid&~finite,1,0),axis=0);tl.atomic_or(flags+matrix,bad)
_E2225_TCGEN_TRSM256_NAME='e2225_trsm256_tcgen_raw';_E2225_TCGEN_TRSM256_SOURCE='\n#include <cuda_runtime.h>\n\n__device__ __forceinline__ unsigned smem_address(const void* pointer) {\n return (unsigned)__cvta_generic_to_shared(pointer);\n}\n__device__ __forceinline__ void tc_alloc(unsigned* destination) {\n const unsigned columns = 128;\n asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"\n : : "r"(smem_address(destination)), "r"(columns) : "memory");\n}\n__device__ __forceinline__ void tc_relinquish() {\n asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;"\n : : : "memory");\n}\n__device__ __forceinline__ void tc_dealloc(unsigned address) {\n const unsigned columns = 128;\n asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"\n : : "r"(address), "r"(columns) : "memory");\n}\n__device__ __forceinline__ void barrier_init(unsigned long long* barrier) {\n const unsigned count = 1;\n asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;"\n : : "r"(smem_address(barrier)), "r"(count) : "memory");\n}\n__device__ __forceinline__ bool barrier_wait(\n unsigned long long* barrier, unsigned phase) {\n unsigned complete;\n asm volatile("{ .reg .pred p;"\n "mbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;"\n "selp.b32 %0, 1, 0, p; }"\n : "=r"(complete)\n : "r"(smem_address(barrier)), "r"(phase) : "memory");\n return complete != 0;\n}\n__device__ __forceinline__ void barrier_invalidate(\n unsigned long long* barrier) {\n asm volatile("mbarrier.inval.shared::cta.b64 [%0];"\n : : "r"(smem_address(barrier)) : "memory");\n}\n__device__ __forceinline__ void tc_mma(\n unsigned destination, unsigned long long a, unsigned long long b,\n unsigned descriptor, bool accumulate) {\n const unsigned zero = 0, enabled = accumulate ? 1u : 0u;\n asm volatile(\n "{ .reg .pred p; setp.ne.b32 p, %8, 0;"\n "tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3,"\n "{%4, %5, %6, %7}, p; }"\n : : "r"(destination), "l"(a), "l"(b), "r"(descriptor),\n "r"(zero), "r"(zero), "r"(zero), "r"(zero), "r"(enabled)\n : "memory");\n}\n__device__ __forceinline__ void tc_commit(unsigned long long* barrier) {\n asm volatile(\n "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"\n : : "r"(smem_address(barrier)) : "memory");\n}\n__device__ __forceinline__ void tc_load8(\n unsigned (&output)[8], unsigned address) {\n asm volatile(\n "tcgen05.ld.sync.aligned.16x32bx2.x8.b32 "\n "{%0,%1,%2,%3,%4,%5,%6,%7}, [%8], 8;"\n : "=r"(output[0]), "=r"(output[1]), "=r"(output[2]), "=r"(output[3]),\n "=r"(output[4]), "=r"(output[5]), "=r"(output[6]), "=r"(output[7])\n : "r"(address) : "memory");\n}\n__device__ __forceinline__ void tc_wait_load() {\n asm volatile("tcgen05.wait::ld.sync.aligned;" : : : "memory");\n}\n\n__device__ __forceinline__ int swizzle64_index(int index) {\n return index ^ ((index >> 3) & 12);\n}\n\n__device__ __forceinline__ unsigned long long smem_desc(const void* pointer) {\n const unsigned address = (unsigned)__cvta_generic_to_shared(pointer);\n return 0x8000402000000000ull | ((unsigned long long)(address & 0x3ffff) >> 4);\n}\n\nextern "C" __global__ __launch_bounds__(128, 1)\nvoid e2225_trsm256_tcgen_raw(const float* __restrict__ matrix,\n const float* __restrict__ lower,\n float* __restrict__ output,\n int batch,\n int source_rows,\n int source_columns,\n int row_offset,\n int column_offset,\n int rows,\n int destination_columns,\n int destination_column_offset) {\n const int matrix_id = (int)blockIdx.x;\n const int row_base = (int)blockIdx.y * 64;\n if (matrix_id >= batch || row_base >= rows) return;\n __shared__ __align__(1024) unsigned char storage[7168];\n float* sa = reinterpret_cast<float*>(storage);\n float* sb = reinterpret_cast<float*>(storage + 4096);\n unsigned* tmem_pointer = reinterpret_cast<unsigned*>(storage + 5120);\n unsigned long long* done = reinterpret_cast<unsigned long long*>(storage + 5128);\n const int tid = (int)threadIdx.x;\n const int lane = tid & 31;\n const int local_row = (tid >> 5) * 16 + (lane & 15);\n const int row = row_base + local_row;\n const int column_half = (lane >> 4) * 8;\n const float* source = matrix +\n (long long)matrix_id * source_rows * source_columns;\n const float* factor = lower + (long long)matrix_id * 256 * 256;\n float* destination = output +\n (long long)matrix_id * rows * destination_columns;\n\n if (tid < 32) tc_alloc(tmem_pointer);\n __syncthreads();\n const unsigned tmem = *tmem_pointer;\n if (tid < 32) tc_relinquish();\n if (tid == 0) barrier_init(done);\n __syncthreads();\n unsigned phase = 0;\n\n #pragma unroll 1\n for (int panel = 0; panel < 256; panel += 16) {\n float rhs[8] = {0.0f, 0.0f, 0.0f, 0.0f,\n 0.0f, 0.0f, 0.0f, 0.0f};\n if (row < rows) {\n const float* rhs_pointer = source +\n (long long)(row_offset + row) * source_columns +\n column_offset + panel + column_half;\n const float4 first = *reinterpret_cast<const float4*>(rhs_pointer);\n const float4 second = *reinterpret_cast<const float4*>(rhs_pointer + 4);\n rhs[0] = first.x; rhs[1] = first.y;\n rhs[2] = first.z; rhs[3] = first.w;\n rhs[4] = second.x; rhs[5] = second.y;\n rhs[6] = second.z; rhs[7] = second.w;\n }\n\n bool first_product = true;\n #pragma unroll 1\n for (int history = 0; history < panel; history += 16) {\n if (tid < 64) {\n const int local_m = tid;\n const int global_row = row_base + local_m;\n const int canonical_base =\n (local_m & 7) * 16 + (local_m >> 3) * 128;\n #pragma unroll\n for (int k = 0; k < 16; k += 4) {\n float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);\n if (global_row < rows) {\n const float* pointer = destination +\n (long long)global_row * destination_columns +\n destination_column_offset + history + k;\n value = *reinterpret_cast<const float4*>(pointer);\n }\n *reinterpret_cast<float4*>(\n sa + swizzle64_index(canonical_base + k)) = value;\n }\n }\n {\n const int n = tid >> 3;\n const int k = (tid & 7) * 2;\n const int canonical = (n & 7) * 16 + (n >> 3) * 128 + k;\n const float2 value = *reinterpret_cast<const float2*>(\n factor + (panel + n) * 256 + history + k);\n *reinterpret_cast<float2*>(\n sb + swizzle64_index(canonical)) = value;\n }\n __syncthreads();\n asm volatile("fence.proxy.async.shared::cta;" : : : "memory");\n if (tid == 0) {\n const unsigned idesc = 0x04040910u;\n const unsigned long long adesc = smem_desc(sa);\n const unsigned long long bdesc = smem_desc(sb);\n tc_mma(tmem, adesc, bdesc, idesc, !first_product);\n tc_mma(tmem, adesc + 2, bdesc + 2, idesc, true);\n tc_commit(done);\n while (!barrier_wait(done, phase)) {}\n }\n __syncthreads();\n phase ^= 1;\n first_product = false;\n }\n\n if (panel != 0) {\n unsigned product[8];\n const unsigned warp_tmem = tmem + ((unsigned)(tid >> 5) << 21);\n tc_load8(product, warp_tmem);\n tc_wait_load();\n #pragma unroll\n for (int i = 0; i < 8; ++i) rhs[i] -= __uint_as_float(product[i]);\n }\n\n {\n const int diagonal_row = tid >> 3;\n const int diagonal_column = (tid & 7) * 2;\n const float2 value = *reinterpret_cast<const float2*>(\n factor + (panel + diagonal_row) * 256 + panel + diagonal_column);\n *reinterpret_cast<float2*>(\n sb + diagonal_row * 16 + diagonal_column) = value;\n }\n __syncthreads();\n\n if (lane < 16) {\n #pragma unroll\n for (int j = 0; j < 8; ++j) {\n float value = rhs[j];\n #pragma unroll\n for (int k = 0; k < j; ++k) value -= sb[j * 16 + k] * rhs[k];\n rhs[j] = value / sb[j * 16 + j];\n }\n }\n float first_half[8];\n #pragma unroll\n for (int k = 0; k < 8; ++k)\n first_half[k] = __shfl_sync(0xffffffffu, rhs[k], lane & 15);\n if (lane >= 16) {\n #pragma unroll\n for (int j = 0; j < 8; ++j) {\n const int diagonal_row = 8 + j;\n float value = rhs[j];\n #pragma unroll\n for (int k = 0; k < 8; ++k)\n value -= sb[diagonal_row * 16 + k] * first_half[k];\n #pragma unroll\n for (int k = 0; k < j; ++k)\n value -= sb[diagonal_row * 16 + 8 + k] * rhs[k];\n rhs[j] = value / sb[diagonal_row * 16 + diagonal_row];\n }\n }\n if (row < rows) {\n float* pointer = destination + (long long)row * destination_columns +\n destination_column_offset + panel + column_half;\n *reinterpret_cast<float4*>(pointer) =\n make_float4(rhs[0], rhs[1], rhs[2], rhs[3]);\n *reinterpret_cast<float4*>(pointer + 4) =\n make_float4(rhs[4], rhs[5], rhs[6], rhs[7]);\n }\n __syncthreads();\n }\n\n barrier_invalidate(done);\n __syncthreads();\n if (tid < 32) tc_dealloc(tmem);\n}\n'
@memo(maxsize=1)
def _e2225_tcgen_trsm256_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2225_TCGEN_TRSM256_SOURCE,_E2225_TCGEN_TRSM256_NAME),_E2225_TCGEN_TRSM256_NAME)
@torch.no_grad()
def _e2225_tcgen_trsm256(matrix:torch.Tensor,lower:torch.Tensor,*,source_rows:int,source_columns:int,row_offset:int,column_offset:int,rows:int,destination:torch.Tensor|None=None,destination_column_offset:int=0):
if destination is None:destination=torch.empty((matrix.shape[0],rows,256),device=matrix.device,dtype=matrix.dtype)
destination_columns=destination.shape[-1];_e2225_tcgen_trsm256_kernel().launch((matrix.shape[0],(rows+63)//64,1),(128,1,1),(matrix,lower,destination,matrix.shape[0],source_rows,source_columns,row_offset,column_offset,rows,destination_columns,destination_column_offset));return destination[:,:,destination_column_offset:destination_column_offset+256]
def _e2369_tcgen_trsm_source(n:int,name:str,*,implicit_identity:bool=False):
source=_E2225_TCGEN_TRSM256_SOURCE.replace(_E2225_TCGEN_TRSM256_NAME,name,1).replace('256',str(n))
if implicit_identity:
old=' float rhs[8] = {0.0f, 0.0f, 0.0f, 0.0f,\n 0.0f, 0.0f, 0.0f, 0.0f};\n if (row < rows) {\n const float* rhs_pointer = source +\n (long long)(row_offset + row) * source_columns +\n column_offset + panel + column_half;\n const float4 first = *reinterpret_cast<const float4*>(rhs_pointer);\n const float4 second = *reinterpret_cast<const float4*>(rhs_pointer + 4);\n rhs[0] = first.x; rhs[1] = first.y;\n rhs[2] = first.z; rhs[3] = first.w;\n rhs[4] = second.x; rhs[5] = second.y;\n rhs[6] = second.z; rhs[7] = second.w;\n }';new=' float rhs[8];\n #pragma unroll\n for (int i = 0; i < 8; ++i)\n rhs[i] = row < rows && row == panel + column_half + i ? 1.0f : 0.0f;'
if source.count(old)!=1:raise RuntimeError('E2369 implicit-identity tcgen anchor changed')
source=source.replace(old,new,1)
return source
@memo(maxsize=None)
def _e2369_tcgen_trsm_kernel(n:int,implicit_identity:bool):suffix='inverse'if implicit_identity else'trsm';name=f"e2369_{suffix}{n}_tcgen_raw";source=_e2369_tcgen_trsm_source(n,name,implicit_identity=implicit_identity);return CUDAKernel(_fast_nvrtc_compile(source,name),name)
def _e2369_tcgen_trsm160_kernel():return _e2369_tcgen_trsm_kernel(160,False)
def _e2373_tcgen_trsm176_kernel():return _e2369_tcgen_trsm_kernel(176,False)
def _e2369_tcgen_inverse128_kernel():return _e2369_tcgen_trsm_kernel(128,True)
def _e2369_tcgen_inverse192_kernel():return _e2369_tcgen_trsm_kernel(192,True)
@torch.no_grad()
def _e2369_tcgen_trsm160(matrix:torch.Tensor,lower:torch.Tensor):
batch,rows,n=matrix.shape
if n!=160:raise ValueError('E2369 tcgen TRSM expects n=160')
output=torch.empty_like(matrix);_e2369_tcgen_trsm160_kernel().launch((batch,(rows+63)//64,1),(128,1,1),(matrix,lower,output,batch,rows,n,0,0,rows,n,0));return output
@torch.no_grad()
def _e2373_tcgen_trsm176(matrix:torch.Tensor,lower:torch.Tensor):
batch,rows,n=matrix.shape
if n!=176:raise ValueError('E2373 tcgen TRSM expects n=176')
output=torch.empty_like(matrix);_e2373_tcgen_trsm176_kernel().launch((batch,(rows+63)//64,1),(128,1,1),(matrix,lower,output,batch,rows,n,0,0,rows,n,0));return output
@torch.no_grad()
def _e2369_tcgen_inverse_lt(lower:torch.Tensor):
batch,n,_=lower.shape
if n==128:implementation=_e2369_tcgen_inverse128_kernel()
elif n==192:implementation=_e2369_tcgen_inverse192_kernel()
else:raise ValueError('E2369 tcgen inverse expects n=128 or n=192')
output=torch.empty_like(lower);implementation.launch((batch,(n+63)//64,1),(128,1,1),(lower,lower,output,batch,n,n,0,0,n,n,0));return output
_E3180_POTRF_X1_NAME='e3180_potrf256_x1_panel_producer';_E3180_POTRF_X3_NAME='e3180_potrf256_x3_panel_producer';_E3180_TCGEN_NAME='e3180_tcgen_panel_consumer';_E3180_L10_NAME='e3180_l10_panel_consumer_producer';_E3180_TENSOR_NAME='e3180_tensor_panel_consumer'
def _e3180_potrf_source(x3:bool):
source=_E084_POTRF256_SOURCE if x3 else _dense2048_potrf256_x1_source();old='potrf256_strided_tf32x3_w19'if x3 else _DENSE2048_POTRF256_X1_NAME;name=_E3180_POTRF_X3_NAME if x3 else _E3180_POTRF_X1_NAME;source=source.replace(old,name,1);signature=' float ridge) {'
if source.count(signature)!=1:raise RuntimeError('E3180 POTRF signature anchor changed')
source=source.replace(signature,' float ridge, int* __restrict__ ready) {',1);start=' const float shift = ridge * (*diagonal_mean);\n\n'
if source.count(start)!=1:raise RuntimeError('E3180 POTRF start anchor changed')
source=source.replace(start,start+' if (tid == 0)\n asm volatile("griddepcontrol.launch_dependents;" ::: "memory");\n\n',1);anchor=' __syncthreads();\n\n const int trailing_tiles = (N - end) / PANEL;';publish=' __syncthreads();\n\n float* panel_destination = lower + (long long)matrix_id * NN;\n for (int index = tid; index < N * PANEL; index += blockDim.x) {\n const int row = index / PANEL;\n const int column = panel + index - row * PANEL;\n if (row >= column)\n panel_destination[row * N + column] = factor[pidx(row, column)];\n }\n __threadfence();\n __syncthreads();\n if (tid == 0) atomicExch(ready + matrix_id, end / PANEL);\n\n const int trailing_tiles = (N - end) / PANEL;'
if source.count(anchor)!=1:raise RuntimeError('E3180 POTRF publish anchor changed')
return source.replace(anchor,publish,1)
def _e3180_tcgen_source():
source=_E2225_TCGEN_TRSM256_SOURCE.replace(_E2225_TCGEN_TRSM256_NAME,_E3180_TCGEN_NAME,1);signature=' int destination_column_offset) {'
if source.count(signature)!=1:raise RuntimeError('E3180 TCGEN signature anchor changed')
source=source.replace(signature,' int destination_column_offset,\n const int* __restrict__ ready) {',1);loop=' for (int panel = 0; panel < 256; panel += 16) {';wait=' for (int panel = 0; panel < 256; panel += 16) {\n if (tid == 0) {\n const int epoch = panel / 16 + 1;\n while (atomicAdd((int*)ready + matrix_id, 0) < epoch)\n __nanosleep(64);\n asm volatile("fence.acquire.gpu;" ::: "memory");\n }\n __syncthreads();'
if source.count(loop)!=1:raise RuntimeError('E3180 TCGEN loop anchor changed')
return source.replace(loop,wait,1)
def _e3180_l10_source():
source=_E084_SIMT_TRSM256_SOURCE.replace('right_trsm256_strided_panel32_rows32',_E3180_L10_NAME,1);signature=' int batch, int source_rows, int source_columns, int row_offset, int column_offset, int rows)\n{'
if source.count(signature)!=1:raise RuntimeError('E3180 L10 signature anchor changed')
source=source.replace(signature,' int batch, int source_rows, int source_columns, int row_offset, int column_offset, int rows,\n const int* __restrict__ ready)\n{',1);gate=' if (matrix_id >= batch) return;\n'
if source.count(gate)!=1:raise RuntimeError('E3180 L10 gate anchor changed')
source=source.replace(gate,gate+' if (tid == 0)\n asm volatile("griddepcontrol.launch_dependents;" ::: "memory");\n',1);loop=' for (int panel = 0; panel < N; panel += PANEL) {';wait=' for (int panel = 0; panel < N; panel += PANEL) {\n if (tid == 0) {\n const int epoch = (panel + PANEL) / 16;\n while (atomicAdd((int*)ready + matrix_id, 0) < epoch)\n __nanosleep(64);\n asm volatile("fence.acquire.gpu;" ::: "memory");\n }\n __syncthreads();'
if source.count(loop)!=1:raise RuntimeError('E3180 L10 loop anchor changed')
return source.replace(loop,wait,1)
def _e3180_tensor_source():
source=_E084_TENSOR_TRSM256_SOURCE.replace('right_trsm256_wmma_tf32x1_w4',_E3180_TENSOR_NAME,1);signature=' int destination_columns, int destination_column_offset) {'
if source.count(signature)!=1:raise RuntimeError('E3180 tensor signature anchor changed')
source=source.replace(signature,' int destination_columns, int destination_column_offset,\n const int* __restrict__ ready) {',1);loop=' for (int panel = 0; panel < N; panel += PANEL) {';wait=' for (int panel = 0; panel < N; panel += PANEL) {\n if (tid == 0) {\n const int epoch = (panel + PANEL) / 16;\n while (atomicAdd((int*)ready + matrix_id, 0) < epoch)\n __nanosleep(64);\n asm volatile("fence.acquire.gpu;" ::: "memory");\n }\n __syncthreads();'
if source.count(loop)!=1:raise RuntimeError('E3180 tensor loop anchor changed')
return source.replace(loop,wait,1)
@memo(maxsize=1)
def _e3180_potrf_x1_kernel():return CUDAKernel(_fast_nvrtc_compile(_e3180_potrf_source(False),_E3180_POTRF_X1_NAME),_E3180_POTRF_X1_NAME)
@memo(maxsize=1)
def _e3180_potrf_x3_kernel():return CUDAKernel(_fast_nvrtc_compile(_e3180_potrf_source(True),_E3180_POTRF_X3_NAME),_E3180_POTRF_X3_NAME)
@memo(maxsize=1)
def _e3180_tcgen_kernel():return CUDAKernel(_fast_nvrtc_compile(_e3180_tcgen_source(),_E3180_TCGEN_NAME),_E3180_TCGEN_NAME)
@memo(maxsize=1)
def _e3180_l10_kernel():return CUDAKernel(_fast_nvrtc_compile(_e3180_l10_source(),_E3180_L10_NAME),_E3180_L10_NAME)
@memo(maxsize=1)
def _e3180_tensor_kernel():return CUDAKernel(_fast_nvrtc_compile(_e3180_tensor_source(),_E3180_TENSOR_NAME),_E3180_TENSOR_NAME)
@torch.no_grad()
def _e3180_start_potrf(gram:torch.Tensor,scale:torch.Tensor,ridge:float,*,x3:bool,leading_dimension:int=256,row_offset:int=0,column_offset:int=0):batch=gram.shape[0];lower=torch.empty((batch,256,256),device=gram.device,dtype=gram.dtype);ready=torch.zeros((batch,),device=gram.device,dtype=torch.int32);kernel=_e3180_potrf_x3_kernel()if x3 else _e3180_potrf_x1_kernel();threads=608 if x3 else _DENSE2048_POTRF256_X1_WARPS*32;shared=_E084_POTRF256_SHARED if x3 else _DENSE2048_POTRF256_X1_SHARED;kernel.launch((batch,1,1),(threads,1,1),(gram,scale,lower,batch,leading_dimension,row_offset,column_offset,float(ridge),ready),shared_mem=shared);return lower,ready
@torch.no_grad()
def _e3180_launch_solve(matrix:torch.Tensor,lower:torch.Tensor,output:torch.Tensor,ready:torch.Tensor,*,destination_column_offset:int,tcgen:bool):
batch,rows,columns=matrix.shape;destination_columns=output.shape[-1]
if tcgen:_e3180_tcgen_kernel().launch_pdl((batch,(rows+63)//64,1),(128,1,1),(matrix,lower,output,batch,rows,columns,0,0,rows,destination_columns,destination_column_offset,ready))
else:_e3180_tensor_kernel().launch_pdl((batch,(rows+31)//32,1),(128,1,1),(matrix,lower,output,batch,rows,columns,0,0,rows,destination_columns,destination_column_offset,ready),shared_mem=_E084_TENSOR_TRSM256_SHARED)
return output[:,:,destination_column_offset:destination_column_offset+256]
@torch.no_grad()
def _e3180_panel_cqr256_half(matrix:torch.Tensor,*,x3:bool,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,**_):
result=matrix
for pass_index in range(passes):gram=_dense2048_half_gram(result);scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);lower,ready=_e3180_start_potrf(gram,scale,pass_ridge,x3=x3);output=torch.empty_like(result);_e3180_launch_solve(result,lower,output,ready,destination_column_offset=0,tcgen=True);result=output
return result
@torch.no_grad()
def _e3180_panel_cqr512_half(matrix:torch.Tensor,*,solve_call_base:int,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,**_):
result=matrix
for pass_index in range(passes):gram=_dense2048_half_gram(result);scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);batch,rows,_=result.shape;output=torch.empty((batch,rows,512),device=result.device,dtype=result.dtype);l00,ready00=_e3180_start_potrf(gram,scale,pass_ridge,x3=False,leading_dimension=512);l10=torch.empty((batch,256,256),device=result.device,dtype=result.dtype);_e3180_l10_kernel().launch_pdl((batch,8,1),(256,1,1),(gram,l00,l10,batch,512,512,256,0,256,ready00),shared_mem=_E084_SIMT_TRSM256_SHARED);y0=_e3180_launch_solve(result,l00,output,ready00,destination_column_offset=0,tcgen=solve_call_base<3);schur=torch.baddbmm(gram[:,256:,256:],l10,l10.mT,beta=1.,alpha=-1.);residual=torch.baddbmm(result[:,:,256:],y0,l10.mT,beta=1.,alpha=-1.);l11,ready11=_e3180_start_potrf(schur,scale,pass_ridge,x3=False);_e3180_launch_solve(residual,l11,output,ready11,destination_column_offset=256,tcgen=solve_call_base+1<3);result=output;solve_call_base+=2
return result
@torch.no_grad()
def _e083_finite_rows(vectors:torch.Tensor,values:torch.Tensor):batch,n,_=vectors.shape;block=8192;chunks=triton.cdiv(n*n+n,block);flags=torch.zeros(batch,device=vectors.device,dtype=torch.int32);_e083_nonfinite_row_flags_kernel[batch,chunks](vectors,values,flags,N=n,BLOCK=block,num_warps=8,num_stages=1);return flags==0
_E084_POTRF256_SHARED=(256*257//2+8+19*5*256)*4;_DENSE2048_POTRF256_X1_WARPS=31;_DENSE2048_POTRF256_X1_NAME='dense2048_potrf256_tf32x1_w31';_DENSE2048_POTRF256_X1_SHARED=(256*257//2+8+_DENSE2048_POTRF256_X1_WARPS*3*256)*4
def _dense2048_potrf256_x1_source():
source=_E084_POTRF256_SOURCE.replace('potrf256_strided_tf32x3_w19',_DENSE2048_POTRF256_X1_NAME,1);source=source.replace('constexpr int WARPS = 19;',f"constexpr int WARPS = {_DENSE2048_POTRF256_X1_WARPS};",1).replace('__launch_bounds__(608, 1)',f"__launch_bounds__({_DENSE2048_POTRF256_X1_WARPS*32}, 1)",1);source=source.replace(' float* a_hi = stage + (warp * 5 + 0) * TILE_FLOATS;\n float* a_lo = stage + (warp * 5 + 1) * TILE_FLOATS;\n float* b_hi = stage + (warp * 5 + 2) * TILE_FLOATS;\n float* b_lo = stage + (warp * 5 + 3) * TILE_FLOATS;\n float* c_tile = stage + (warp * 5 + 4) * TILE_FLOATS;',' float* a_hi = stage + (warp * 3 + 0) * TILE_FLOATS;\n float* b_hi = stage + (warp * 3 + 1) * TILE_FLOATS;\n float* c_tile = stage + (warp * 3 + 2) * TILE_FLOATS;',1)
for line in(' a_lo[index] = wmma::__float_to_tf32(av - ah);\n',' b_lo[index] = -wmma::__float_to_tf32(bv - bh);\n',' wmma::load_matrix_sync(al, a_lo + k0, 16);\n',' wmma::load_matrix_sync(bl, b_lo + k0 * 16, 16);\n',' wmma::mma_sync(c, al, bh, c);\n',' wmma::mma_sync(c, ah, bl, c);\n'):source=source.replace(line,'',1)
return source.replace('wmma::precision::tf32, wmma::row_major> ah, al;','wmma::precision::tf32, wmma::row_major> ah;',1).replace('wmma::precision::tf32, wmma::row_major> bh, bl;','wmma::precision::tf32, wmma::row_major> bh;',1)
@memo(maxsize=1)
def _dense2048_potrf256_x1_kernel():return CUDAKernel(_fast_nvrtc_compile(_dense2048_potrf256_x1_source(),_DENSE2048_POTRF256_X1_NAME),_DENSE2048_POTRF256_X1_NAME)
@torch.no_grad()
def _dense2048_potrf256_x1(matrix:torch.Tensor,diagonal_scale:torch.Tensor,*,leading_dimension:int,row_offset:int,column_offset:int,ridge:float):output=torch.empty((matrix.shape[0],256,256),device=matrix.device,dtype=matrix.dtype);_dense2048_potrf256_x1_kernel().launch((matrix.shape[0],1,1),(_DENSE2048_POTRF256_X1_WARPS*32,1,1),(matrix,diagonal_scale,output,matrix.shape[0],leading_dimension,row_offset,column_offset,float(ridge)),shared_mem=_DENSE2048_POTRF256_X1_SHARED);return output
_E084_SIMT_TRSM256_SHARED=256*32*4;_E084_TENSOR_TRSM256_SHARED=(256*32+4*3*256)*4
@memo(maxsize=None)
def _e084_cholesky_kernel(name:str):
if name=='potrf256_strided_tf32x3_w19':source=_E084_POTRF256_SOURCE
elif name=='right_trsm256_strided_panel32_rows32':source=_E084_SIMT_TRSM256_SOURCE
elif name=='right_trsm256_wmma_tf32x1_w4':source=_E084_TENSOR_TRSM256_SOURCE
else:raise ValueError(name)
return CUDAKernel(_fast_nvrtc_compile(source,name),name)
@torch.no_grad()
def _e084_potrf256_strided(matrix:torch.Tensor,diagonal_scale:torch.Tensor,*,leading_dimension:int,row_offset:int,column_offset:int,ridge:float):output=torch.empty((matrix.shape[0],256,256),device=matrix.device,dtype=matrix.dtype);_e084_cholesky_kernel('potrf256_strided_tf32x3_w19').launch((matrix.shape[0],1,1),(608,1,1),(matrix,diagonal_scale,output,matrix.shape[0],leading_dimension,row_offset,column_offset,float(ridge)),shared_mem=_E084_POTRF256_SHARED);return output
@torch.no_grad()
def _e084_trsm256_simt_strided(matrix:torch.Tensor,lower:torch.Tensor,*,source_rows:int,source_columns:int,row_offset:int,column_offset:int,rows:int):output=torch.empty((matrix.shape[0],rows,256),device=matrix.device,dtype=matrix.dtype);_e084_cholesky_kernel('right_trsm256_strided_panel32_rows32').launch((matrix.shape[0],(rows+31)//32,1),(256,1,1),(matrix,lower,output,matrix.shape[0],source_rows,source_columns,row_offset,column_offset,rows),shared_mem=_E084_SIMT_TRSM256_SHARED);return output
@torch.no_grad()
def _e084_trsm256_tensor_strided(matrix:torch.Tensor,lower:torch.Tensor,*,source_rows:int,source_columns:int,row_offset:int,column_offset:int,rows:int,destination:torch.Tensor|None=None,destination_column_offset:int=0):
if destination is None:destination=torch.empty((matrix.shape[0],rows,256),device=matrix.device,dtype=matrix.dtype)
destination_columns=destination.shape[-1];_e084_cholesky_kernel('right_trsm256_wmma_tf32x1_w4').launch((matrix.shape[0],(rows+31)//32,1),(128,1,1),(matrix,lower,destination,matrix.shape[0],source_rows,source_columns,row_offset,column_offset,rows,destination_columns,destination_column_offset),shared_mem=_E084_TENSOR_TRSM256_SHARED);return destination[:,:,destination_column_offset:destination_column_offset+256]
@torch.no_grad()
def _e084_factor512(gram:torch.Tensor,ridge:float,potrf_fn=None):
if potrf_fn is None:potrf_fn=_e084_potrf256_strided
scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);l00=potrf_fn(gram,scale,leading_dimension=512,row_offset=0,column_offset=0,ridge=ridge);l10=_e084_trsm256_simt_strided(gram,l00,source_rows=512,source_columns=512,row_offset=256,column_offset=0,rows=256);schur=torch.baddbmm(gram[:,256:,256:],l10,l10.mT,beta=1.,alpha=-1.);l11=potrf_fn(schur,scale,leading_dimension=256,row_offset=0,column_offset=0,ridge=ridge);return l00,l10,l11
@torch.no_grad()
def _e084_solve512(matrix:torch.Tensor,factors,trsm_fn=None):
if trsm_fn is None:trsm_fn=_e084_trsm256_tensor_strided
l00,l10,l11=factors;rows=matrix.shape[1];output=torch.empty((matrix.shape[0],rows,512),device=matrix.device,dtype=matrix.dtype);y0=trsm_fn(matrix,l00,source_rows=rows,source_columns=512,row_offset=0,column_offset=0,rows=rows,destination=output,destination_column_offset=0);residual=torch.baddbmm(matrix[:,:,256:],y0,l10.mT,beta=1.,alpha=-1.);trsm_fn(residual,l11,source_rows=rows,source_columns=256,row_offset=0,column_offset=0,rows=rows,destination=output,destination_column_offset=256);return output
_E1074_L10_PDL_NAME='e1074_right_trsm256_l10_pdl_producer'
@memo(maxsize=1)
def _e1074_l10_pdl_kernel():
source=_E084_SIMT_TRSM256_SOURCE.replace('right_trsm256_strided_panel32_rows32',_E1074_L10_PDL_NAME);anchor=' if (matrix_id >= batch) return;\n'
if source.count(anchor)!=1:raise RuntimeError('CQR512 PDL producer anchor changed')
source=source.replace(anchor,anchor+' if (tid == 0) asm volatile("griddepcontrol.launch_dependents;":::);\n',1);return CUDAKernel(_fast_nvrtc_compile(source,_E1074_L10_PDL_NAME),_E1074_L10_PDL_NAME)
@torch.no_grad()
def _e1074_cholesky_orthonormalize512_pdl(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,inverse_precision:str='highest',gram_precision:str='highest',trsm_fn=None):
if trsm_fn is None:trsm_fn=_e084_trsm256_tensor_strided
result=matrix
for pass_index in range(passes):previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=result.mT@result;torch.set_float32_matmul_precision(previous);scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);l00=_e084_potrf256_strided(gram,scale,leading_dimension=512,row_offset=0,column_offset=0,ridge=pass_ridge);batch,rows,_=result.shape;l10=torch.empty((batch,256,256),device=result.device,dtype=result.dtype);output=torch.empty((batch,rows,512),device=result.device,dtype=result.dtype);_e1074_l10_pdl_kernel().launch((batch,8,1),(256,1,1),(gram,l00,l10,batch,512,512,256,0,256),shared_mem=_E084_SIMT_TRSM256_SHARED);_e084_cholesky_kernel('right_trsm256_wmma_tf32x1_w4').launch_pdl((batch,(rows+31)//32,1),(128,1,1),(result,l00,output,batch,rows,512,0,0,rows,512,0),shared_mem=_E084_TENSOR_TRSM256_SHARED);y0=output[:,:,:256];schur=torch.baddbmm(gram[:,256:,256:],l10,l10.mT,beta=1.,alpha=-1.);l11=_e084_potrf256_strided(schur,scale,leading_dimension=256,row_offset=0,column_offset=0,ridge=pass_ridge);residual=torch.baddbmm(result[:,:,256:],y0,l10.mT,beta=1.,alpha=-1.);trsm_fn(residual,l11,source_rows=rows,source_columns=256,row_offset=0,column_offset=0,rows=rows,destination=output,destination_column_offset=256);result=output
return result
@torch.no_grad()
def _e084_cholesky_orthonormalize512(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,inverse_precision:str='highest',gram_precision:str='highest',potrf_fn=None,trsm_fn=None):
result=matrix
for pass_index in range(passes):previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=result.mT@result;torch.set_float32_matmul_precision(previous);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);result=_e084_solve512(result,_e084_factor512(gram,pass_ridge,potrf_fn=potrf_fn),trsm_fn=trsm_fn)
return result
@torch.no_grad()
def _e084_cholesky_orthonormalize256(matrix:torch.Tensor,*,passes:int=2,ridge:float=1e-05,final_ridge:float=1e-08,inverse_precision:str='highest',gram_precision:str='highest',potrf_fn=None,trsm_fn=None):
if trsm_fn is None:trsm_fn=_e084_trsm256_tensor_strided
result=matrix
for pass_index in range(passes):previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision(gram_precision);gram=result.mT@result;torch.set_float32_matmul_precision(previous);scale=gram.diagonal(dim1=-2,dim2=-1).mean(dim=-1);pass_ridge=ridge if pass_index==0 else min(ridge,final_ridge);factor=_e084_potrf256_strided if potrf_fn is None else potrf_fn;lower=factor(gram,scale,leading_dimension=256,row_offset=0,column_offset=0,ridge=pass_ridge);result=trsm_fn(result,lower,source_rows=result.shape[1],source_columns=256,row_offset=0,column_offset=0,rows=result.shape[1])
return result
_E150_N352_CLUSTER2_NAME='e150_band32_n352_cluster2'
def _e150_n352_cluster2_source():
source=_OWNED_N256_COLUMN_BASE_SOURCE.replace('#include <cuda_runtime.h>','#include <cuda_runtime.h>\n#include <cooperative_groups.h>\nnamespace cg = cooperative_groups;').replace('constexpr int n = 512;','constexpr int n = 352;').replace('constexpr int bandwidth = 63;','constexpr int bandwidth = 32;').replace('constexpr int max_blocks = 9;','constexpr int max_blocks = 11;').replace('extern "C" __global__ void column_tiled_band_to_tridiagonal_512(',f'extern "C" __global__ __cluster_dims__(2, 1, 1) void {_E150_N352_CLUSTER2_NAME}(').replace(' const int batch = blockIdx.x;',' cg::cluster_group cluster = cg::this_cluster();\n const int rank = cluster.block_rank();\n const int batch = blockIdx.x >> 1;').replace(' reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;',' if (rank == 0) {\n reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;\n }',1).replace(' reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;',' if (rank == 0) {\n reflectors[destination + i0] = x0;\n reflectors[destination + i1] = x1;\n }',1).replace(' // Transform the one structurally nonzero prefix row against S_0.\n if (warp == 0) {',' // Only rank 0 owns the prefix row/column.\n if (rank == 0 && warp == 0) {',1).replace(' for (int block = 0; block < block_count; ++block) {',' for (int block = rank; block < block_count; block += 2) {',1);end=' __syncthreads();\n }\n }\n}\n';replacement=' __syncthreads();\n }\n cluster.sync();\n }\n}\n'
if not source.endswith(end):raise RuntimeError('n352 cluster2 source tail changed')
return source[:-len(end)]+replacement
_E184_N512_CLUSTER2_NAME='e184_band32_n512_cluster2'
def _e184_n512_cluster2_source():return _e150_n352_cluster2_source().replace('constexpr int n = 352;','constexpr int n = 512;').replace('constexpr int max_blocks = 11;','constexpr int max_blocks = 16;').replace(_E150_N352_CLUSTER2_NAME,_E184_N512_CLUSTER2_NAME)
@memo(maxsize=1)
def _e184_n512_cluster2_kernel():return CUDAKernel(_fast_nvrtc_compile(_e184_n512_cluster2_source(),_E184_N512_CLUSTER2_NAME),_E184_N512_CLUSTER2_NAME)
_E214_N512_CLUSTER8_NAME='e214_band32_n512_cluster8'
def _e214_async_stage_segment(segment:str):
length=' const int length = min(bandwidth, n - support_start);\n';prefetch=length+'\n const int tile_copies = paired ? 4 : 2;\n {\n #pragma unroll\n for (int q = 0; q < tile_copies; ++q) {\n const int row = local_warp + q * local_warps;\n if (row < length && lane < length) {\n const int dst = __cvta_generic_to_shared(\n &diagonal_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + row) * n\n + support_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n }\n asm volatile("cp.async.commit_group;");\n }\n if (block + 1 < block_count) {\n const int prefetch_next_start = support_start + bandwidth;\n const int prefetch_next_length =\n min(bandwidth, n - prefetch_next_start);\n #pragma unroll\n for (int q = 0; q < tile_copies; ++q) {\n const int row = local_warp + q * local_warps;\n if (row < length && lane < prefetch_next_length) {\n const int dst = __cvta_generic_to_shared(\n &upper_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + row) * n\n + prefetch_next_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n if (row < prefetch_next_length && lane < length) {\n const int dst = __cvta_generic_to_shared(\n &lower_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(prefetch_next_start + row) * n\n + support_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n }\n asm volatile("cp.async.commit_group;");\n }\n if (block + 1 < block_count)\n asm volatile("cp.async.wait_group 1;" ::: "memory");\n else\n asm volatile("cp.async.wait_group 0;" ::: "memory");\n __syncwarp();\n'
if segment.count(length)!=1:raise RuntimeError('n512 async length anchor changed')
segment=segment.replace(length,prefetch,1);old=' dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j] * vectors[block][j];';new=' const float tile_value =\n diagonal_tile[subgroup][row][j];\n dot += tile_value * vectors[block][j];'
if segment.count(old)!=1:raise RuntimeError('n512 async diagonal projection changed')
segment=segment.replace(old,new,1);old=' matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j]\n -= w_row * vectors[block][j] + u_row * w_j;';new=' const float tile_value =\n diagonal_tile[subgroup][row][j];\n matrices[matrix_base\n + (long long)(support_start + row) * n\n + support_start + j]\n = tile_value\n - (w_row * vectors[block][j] + u_row * w_j);'
if segment.count(old)!=1:raise RuntimeError('n512 async diagonal update changed')
segment=segment.replace(old,new,1);marker=' // Both projections of the neighboring off-diagonal tile.'
if segment.count(marker)!=1:raise RuntimeError('n512 async cross anchor changed')
segment=segment.replace(marker,' asm volatile("cp.async.wait_group 0;" ::: "memory");\n __syncwarp();\n\n'+marker,1);old=' dot += matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j] * vectors[block + 1][j];';new=' const float tile_value =\n upper_tile[subgroup][row][j];\n dot += tile_value * vectors[block + 1][j];'
if segment.count(old)!=1:raise RuntimeError('n512 async upper projection changed')
segment=segment.replace(old,new,1);old=' dot += matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j] * vectors[block][j];';new=' const float tile_value =\n lower_tile[subgroup][row][j];\n dot += tile_value * vectors[block][j];'
if segment.count(old)!=1:raise RuntimeError('n512 async lower projection changed')
segment=segment.replace(old,new,1);old='const float updated = matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j]'
if segment.count(old)!=1:raise RuntimeError('n512 async upper update changed')
segment=segment.replace(old,'const float updated = upper_tile[subgroup][row][j]',1);old='const float updated = matrices[matrix_base\n + (long long)(next_start + row) * n\n + support_start + j]'
if segment.count(old)!=1:raise RuntimeError('n512 async lower update changed')
return segment.replace(old,'const float updated = lower_tile[subgroup][row][j]',1)
def _e214_n512_cluster8_source():
source=_e184_n512_cluster2_source().replace('__cluster_dims__(2, 1, 1)','__cluster_dims__(8, 1, 1)').replace('const int batch = blockIdx.x >> 1;','const int batch = blockIdx.x >> 3;').replace('block += 2','block += 8').replace(_E184_N512_CLUSTER2_NAME,_E214_N512_CLUSTER8_NAME);source=source.replace('__shared__ float left_projection[storage_band];','__shared__ float left_projection[2][storage_band];').replace('__shared__ float right_projection[storage_band];','__shared__ float right_projection[2][storage_band];').replace('__shared__ float scalar;','__shared__ float scalar[2];\n __shared__ float diagonal_tile[2][32][32];\n __shared__ float upper_tile[2][32][32];\n __shared__ float lower_tile[2][32][32];');header=' for (int block = rank; block < block_count; block += 8) {';start=source.index(header);end=source.rindex(' cluster.sync();');segment=source[start:end];segment=segment.replace(header,' for (int block_wave = rank; block_wave < block_count; block_wave += (block_wave + 9 < block_count ? 16 : 8)) {\n const bool paired = block_wave + 9 < block_count;\n const int subgroup = paired && warp >= 8;\n const int local_warp = paired ? (warp & 7) : warp;\n const int local_warps = paired ? 8 : warps;\n const int block = block_wave + subgroup * 8;');segment=segment.replace('for (int row = warp; row < length; row += warps)','for (int row = local_warp; row < length; row += local_warps)').replace('for (int row = warp; row < next_length; row += warps)','for (int row = local_warp; row < next_length; row += local_warps)').replace('if (warp == 0)','if (local_warp == 0)').replace('right_projection[','right_projection[subgroup][').replace('left_projection[','left_projection[subgroup][').replace('if (lane == 0) scalar = projection;','if (lane == 0) scalar[subgroup] = projection;').replace('const float diagonal_scalar = scalar;','const float diagonal_scalar = scalar[subgroup];').replace('if (lane == 0) scalar = cross;','if (lane == 0) scalar[subgroup] = cross;').replace('const float cross_scalar = scalar;','const float cross_scalar = scalar[subgroup];');segment=_e214_async_stage_segment(segment);source=source[:start]+segment+source[end:];initial_old=' float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);';initial_new=' const float norm2 = warp_sum(x0 * x0 + x1 * x1);\n const float norm = lane == 0 ? sqrtf(norm2) : 0.0f;\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -norm : norm) : 0.0f;\n const float reflector_norm2 = lane == 0\n ? 2.0f * norm * (norm + fabsf(x0)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n float inverse_norm = lane == 0 && reflector_norm2 > 1.0e-40f\n ? rsqrtf(reflector_norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);';following_old=' float norm2 = warp_sum(x0 * x0 + x1 * x1);\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -sqrtf(norm2) : sqrtf(norm2)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n norm2 = warp_sum(x0 * x0 + x1 * x1);\n float inverse_norm = lane == 0 && norm2 > 1.0e-40f\n ? rsqrtf(norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);';following_new=' const float norm2 = warp_sum(x0 * x0 + x1 * x1);\n const float norm = lane == 0 ? sqrtf(norm2) : 0.0f;\n float alpha = lane == 0\n ? (x0 >= 0.0f ? -norm : norm) : 0.0f;\n const float reflector_norm2 = lane == 0\n ? 2.0f * norm * (norm + fabsf(x0)) : 0.0f;\n alpha = __shfl_sync(0xffffffffu, alpha, 0);\n if (lane == 0) x0 -= alpha;\n float inverse_norm = lane == 0 && reflector_norm2 > 1.0e-40f\n ? rsqrtf(reflector_norm2) : 0.0f;\n inverse_norm = __shfl_sync(0xffffffffu, inverse_norm, 0);'
if source.count(initial_old)!=1 or source.count(following_old)!=1:raise RuntimeError('E1717 n512 analytic-norm anchor changed')
return source.replace(initial_old,initial_new).replace(following_old,following_new)
@memo(maxsize=1)
def _e214_n512_cluster8_kernel():return CUDAKernel(_fast_nvrtc_compile(_e214_n512_cluster8_source(),_E214_N512_CLUSTER8_NAME),_E214_N512_CLUSTER8_NAME)
@torch.no_grad()
def _e214_band32_to_tridiagonal_n512(matrix:torch.Tensor):output=matrix.clone();reflectors=torch.empty((matrix.shape[0],512,16,64),device=matrix.device,dtype=torch.float32);_e214_n512_cluster8_kernel().launch((matrix.shape[0]*8,1,1),(512,1,1),(output,reflectors));return output,reflectors
_E1214_N512_LOWER_NAME='e1214_n512_lower_authoritative'
def _e1214_n512_lower_source():
source=_e214_n512_cluster8_source().replace(_E214_N512_CLUSTER8_NAME,_E1214_N512_LOWER_NAME,1);source=source.replace(' __shared__ float diagonal_tile[2][32][32];',' __shared__ float diagonal_tile[2][32][33];',1).replace(' __shared__ float lower_tile[2][32][32];',' __shared__ float lower_tile[2][32][33];',1).replace(' __shared__ float upper_tile[2][32][32];\n','',1);upper_copy=' if (row < length && lane < prefetch_next_length) {\n const int dst = __cvta_generic_to_shared(\n &upper_tile[subgroup][row][lane]);\n const float* src = matrices + matrix_base\n + (long long)(support_start + row) * n\n + prefetch_next_start + lane;\n asm volatile(\n "cp.async.ca.shared.global [%0], [%1], 4;"\n :: "r"(dst), "l"(src));\n }\n';source=source.replace(upper_copy,'',1);first_wait=' if (block + 1 < block_count)\n asm volatile("cp.async.wait_group 1;" ::: "memory");\n else\n asm volatile("cp.async.wait_group 0;" ::: "memory");\n __syncwarp();\n\n // Diagonal tile H_k A_kk H_k.';source=source.replace(first_wait,first_wait.replace(' __syncwarp();',' __syncthreads();'),1);cross_wait=' asm volatile("cp.async.wait_group 0;" ::: "memory");\n __syncwarp();\n\n // Both projections of the neighboring off-diagonal tile.';source=source.replace(cross_wait,cross_wait.replace(' __syncwarp();',' __syncthreads();'),1);diagonal='diagonal_tile[subgroup][row][j]';source=source.replace(diagonal,'diagonal_tile[subgroup][row > j ? row : j][row > j ? j : row]');prefix=' for (int j = lane; j < first_length; j += 32) {\n dot += matrices[matrix_base + (long long)prefix_row * n\n + first_support + j] * vectors[0][j];\n }\n dot = warp_sum(dot);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n for (int j = lane; j < first_length; j += 32) {\n const float updated = matrices[\n matrix_base + (long long)prefix_row * n\n + first_support + j] - 2.0f * dot * vectors[0][j];\n matrices[matrix_base + (long long)prefix_row * n\n + first_support + j] = updated;\n matrices[matrix_base + (long long)(first_support + j) * n\n + prefix_row] = updated;\n }';lower_prefix=' for (int j = lane; j < first_length; j += 32) {\n dot += matrices[matrix_base\n + (long long)(first_support + j) * n\n + prefix_row] * vectors[0][j];\n }\n dot = warp_sum(dot);\n dot = __shfl_sync(0xffffffffu, dot, 0);\n for (int j = lane; j < first_length; j += 32) {\n const long long lower = matrix_base\n + (long long)(first_support + j) * n + prefix_row;\n matrices[lower] -= 2.0f * dot * vectors[0][j];\n }';source=source.replace(prefix,lower_prefix,1);source=source.replace('upper_tile[subgroup][row][j]','lower_tile[subgroup][j][row]',1);upper_update=' for (int row = local_warp; row < length; row += local_warps) {\n for (int j = lane; j < next_length; j += 32) {\n const float updated = upper_tile[subgroup][row][j]\n - 2.0f * vectors[block][row] * left_projection[subgroup][j]\n - 2.0f * right_projection[subgroup][row] * vectors[block + 1][j]\n + 4.0f * vectors[block][row] * cross_scalar\n * vectors[block + 1][j];\n matrices[matrix_base\n + (long long)(support_start + row) * n\n + next_start + j] = updated;\n }\n }\n';source=source.replace(upper_update,'',1);tail_upper=' matrices[matrix_base\n + (long long)(support_start + j) * n\n + tail_start] = updated;\n';source=source.replace(tail_upper,'',1)
if'upper_tile'in source or source.count(diagonal):raise RuntimeError('E1214 lower n512 source rewrite changed')
cluster_rewrites=('#include <cooperative_groups.h>\n',''),('namespace cg = cooperative_groups;\n',''),(' cg::cluster_group cluster = cg::this_cluster();\n',''),(' const int rank = cluster.block_rank();',' const int rank = blockIdx.x & 7;'),(' cluster.sync();',' asm volatile("barrier.cluster.arrive.aligned;" ::: "memory");\n asm volatile("barrier.cluster.wait.aligned;" ::: "memory");')
for(old,new)in cluster_rewrites:
if source.count(old)!=1:raise RuntimeError('E2018 direct cluster PTX anchor changed')
source=source.replace(old,new,1)
return source
@memo(maxsize=1)
def _e1214_n512_lower_kernel():return CUDAKernel(_fast_nvrtc_compile(_e1214_n512_lower_source(),_E1214_N512_LOWER_NAME),_E1214_N512_LOWER_NAME)
_E2349_N512_LOWER_EARLY_PREFIX_NAME='e2349_n512_lower_early_prefix'
def _e2349_n512_lower_early_prefix_source():
source=_e1214_n512_lower_source().replace(_E1214_N512_LOWER_NAME,_E2349_N512_LOWER_EARLY_PREFIX_NAME,1);anchor=' const int tid = threadIdx.x;\n';replacement=anchor+' // The mixed-route prefixes do not read this native-band chase.\n'+' // Publish PDL completion at residency so they can run in its shadow.\n'+' if (tid == 0)\n'+' asm volatile("griddepcontrol.launch_dependents;" ::: "memory");\n'
if source.count(anchor)!=1:raise RuntimeError('E2349 n512 B32 entry anchor changed')
return source.replace(anchor,replacement,1)
@memo(maxsize=1)
def _e2349_n512_lower_early_prefix_kernel():return CUDAKernel(_fast_nvrtc_compile(_e2349_n512_lower_early_prefix_source(),_E2349_N512_LOWER_EARLY_PREFIX_NAME),_E2349_N512_LOWER_EARLY_PREFIX_NAME)
@torch.no_grad()
def _e1214_band32_to_tridiagonal_n512(matrix:torch.Tensor,*,signal_dependents:bool=False):output=matrix.clone();reflectors=torch.empty((matrix.shape[0],512,16,64),device=matrix.device,dtype=torch.float32);kernel=_e2349_n512_lower_early_prefix_kernel()if signal_dependents else _e1214_n512_lower_kernel();kernel.launch((matrix.shape[0]*8,1,1),(512,1,1),(output,reflectors));return output,reflectors
@torch.no_grad()
def _mixed512_safe_native_band(data:torch.Tensor):tridiagonal=data.clone();reflectors=torch.zeros((data.shape[0],512,16,64),device=data.device,dtype=torch.float32);_e184_n512_cluster2_kernel().launch((data.shape[0]*2,1,1),(512,1,1),(tridiagonal,reflectors));vectors=torch.empty_like(data);values=torch.empty((data.shape[0],512),device=data.device,dtype=torch.float32);workspace=torch.empty_like(data);_e187_mixed_tridiag24_kernel().launch((data.shape[0]*_E187_MIXED_TRI24_SHARDS,1,1),(_E187_MIXED_TRI24_EIGEN_PER_SHARD,1,1),(tridiagonal,vectors,values,workspace));vectors=_band16_repair_twisted512(vectors,values,gap=.0001);vectors=_e243_band32_sparse_wy_replay(vectors,reflectors);return vectors,values
@triton.jit
def _e094_n1024_row_stats(data,row_energy,diagonal,row_l1,total_rows,N:tl.constexpr,BLOCK_ROWS:tl.constexpr,BLOCK_N:tl.constexpr):group=tl.program_id(0);first_row=group*BLOCK_ROWS;linear=tl.arange(0,BLOCK_ROWS*BLOCK_N);local_row=linear//BLOCK_N;column=linear%BLOCK_N;global_row=first_row+local_row;matrix_row=global_row%N;mask=(global_row<total_rows)&(column<N);values=tl.load(data+global_row*N+column,mask=mask,other=.0);values=tl.reshape(values,(BLOCK_ROWS,BLOCK_N));columns=tl.reshape(column,(BLOCK_ROWS,BLOCK_N));matrix_rows=tl.reshape(matrix_row,(BLOCK_ROWS,BLOCK_N));energies=tl.maximum(tl.sum(values*values,axis=1),1e-30);l1=tl.sum(tl.abs(values),axis=1);diag_values=tl.sum(tl.where(columns==matrix_rows,values,.0),axis=1);rows=first_row+tl.arange(0,BLOCK_ROWS);row_mask=rows<total_rows;tl.store(row_energy+rows,energies,mask=row_mask);tl.store(diagonal+rows,diag_values,mask=row_mask);tl.store(row_l1+rows,l1,mask=row_mask)
@triton.jit
def _e094_n1024_matrix_stats(row_energy,diagonal,row_l1,stats,N:tl.constexpr,BLOCK_N:tl.constexpr):matrix=tl.program_id(0);columns=tl.arange(0,BLOCK_N);rows=tl.load(row_energy+matrix*N+columns);diag=tl.load(diagonal+matrix*N+columns);l1=tl.load(row_l1+matrix*N+columns);tl.store(stats+matrix*5,tl.sum(diag,axis=0));tl.store(stats+matrix*5+1,tl.sum(rows,axis=0));tl.store(stats+matrix*5+2,tl.min(rows,axis=0));tl.store(stats+matrix*5+3,tl.max(rows,axis=0));tl.store(stats+matrix*5+4,tl.max(l1,axis=0))
@triton.jit
def _e2440_n1024_row_stats_native(data,row_energy,diagonal,row_l1,row_outside,row_ring,total_rows,N:tl.constexpr,BLOCK_ROWS:tl.constexpr,BLOCK_N:tl.constexpr):group=tl.program_id(0);first_row=group*BLOCK_ROWS;linear=tl.arange(0,BLOCK_ROWS*BLOCK_N);local_row=linear//BLOCK_N;column=linear%BLOCK_N;global_row=first_row+local_row;matrix_row=global_row%N;mask=(global_row<total_rows)&(column<N);values=tl.load(data+global_row*N+column,mask=mask,other=.0);values=tl.reshape(values,(BLOCK_ROWS,BLOCK_N));columns=tl.reshape(column,(BLOCK_ROWS,BLOCK_N));matrix_rows=tl.reshape(matrix_row,(BLOCK_ROWS,BLOCK_N));energies=tl.maximum(tl.sum(values*values,axis=1),1e-30);l1=tl.sum(tl.abs(values),axis=1);diag_values=tl.sum(tl.where(columns==matrix_rows,values,.0),axis=1);distance=tl.abs(columns-matrix_rows);nonzero=values!=.0;outside=tl.max(tl.where((distance>32)&nonzero,1,0),axis=1);ring=tl.max(tl.where((distance==32)&nonzero,1,0),axis=1);rows=first_row+tl.arange(0,BLOCK_ROWS);row_mask=rows<total_rows;tl.store(row_energy+rows,energies,mask=row_mask);tl.store(diagonal+rows,diag_values,mask=row_mask);tl.store(row_l1+rows,l1,mask=row_mask);tl.store(row_outside+rows,outside,mask=row_mask);tl.store(row_ring+rows,ring,mask=row_mask)
@triton.jit
def _e2440_n1024_matrix_stats_native(row_energy,diagonal,row_l1,row_outside,row_ring,stats,native,N:tl.constexpr,BLOCK_N:tl.constexpr):matrix=tl.program_id(0);columns=tl.arange(0,BLOCK_N);rows=tl.load(row_energy+matrix*N+columns);diag=tl.load(diagonal+matrix*N+columns);l1=tl.load(row_l1+matrix*N+columns);outside=tl.load(row_outside+matrix*N+columns);ring=tl.load(row_ring+matrix*N+columns);tl.store(stats+matrix*5,tl.sum(diag,axis=0));tl.store(stats+matrix*5+1,tl.sum(rows,axis=0));tl.store(stats+matrix*5+2,tl.min(rows,axis=0));tl.store(stats+matrix*5+3,tl.max(rows,axis=0));tl.store(stats+matrix*5+4,tl.max(l1,axis=0));outside_any=tl.max(outside,axis=0);ring_any=tl.max(ring,axis=0);tl.store(native+matrix,(outside_any==0)&(ring_any!=0))
@triton.jit
def _e279_n1024_fast_geometric_gate(data,row_energy,stats,output,batch,BLOCK_BATCH:tl.constexpr,N:tl.constexpr,BLOCK_N:tl.constexpr):matrix=tl.arange(0,BLOCK_BATCH);columns=tl.arange(0,BLOCK_N);mask=matrix<batch;rows=tl.load(row_energy+matrix[:,None]*N+columns[None,:],mask=mask[:,None]&(columns[None,:]<N),other=.0);frob2=tl.load(stats+matrix*5+1,mask=mask,other=.0);row4=tl.sum(rows*rows,axis=1);denominator=.5*((N+2.)*row4-frob2*frob2);participation=frob2*frob2/tl.maximum(denominator,1e-30);sample=tl.arange(0,8);sample_rows=sample*64;sample_columns=N-1-sample*64;addresses=data+matrix[:,None]*N*N+sample_rows[None,:]*N+sample_columns[None,:];far=tl.min(tl.abs(tl.load(addresses,mask=mask[:,None],other=.0)),axis=1);scale=tl.sqrt(tl.maximum(frob2,.0))/N;dense=far>1e-06*scale;finite=(frob2==frob2)&(denominator==denominator);magnitude_safe=finite&(frob2>1e-20)&(frob2<1e20)&(denominator>1e-20);geometric=magnitude_safe&dense&(participation>45.)&(participation<9e1);all_safe=tl.min(tl.where(mask,magnitude_safe,1),axis=0);all_geometric=tl.min(tl.where(mask,geometric,1),axis=0);state=tl.where(all_safe!=0,all_geometric,-1);tl.store(output,state.to(tl.int32))
@triton.jit
def _e2440_n1024_route_and_subgroups(data,stats,native,metadata,repeated_output,refined_route,batch,BLOCK_BATCH:tl.constexpr):
matrix=tl.arange(0,BLOCK_BATCH);mask=matrix<batch;trace=tl.load(stats+matrix*5,mask=mask,other=.0);frob2=tl.load(stats+matrix*5+1,mask=mask,other=1.);minimum=tl.load(stats+matrix*5+2,mask=mask,other=1.);maximum=tl.load(stats+matrix*5+3,mask=mask,other=1.);dynamic=maximum/tl.maximum(minimum,1e-30);ratio=trace/tl.sqrt(tl.maximum(frob2,1e-30));effective=trace*trace/tl.maximum(frob2,1e-30);base=matrix*1024*1024;probe0=tl.abs(tl.load(data+base+1,mask=mask,other=.0));probe1=tl.abs(tl.load(data+base+512*1024+513,mask=mask,other=.0));probe2=tl.abs(tl.load(data+base+1022*1024+1023,mask=mask,other=.0));probe3=tl.abs(tl.load(data+base+1023,mask=mask,other=.0));off=tl.maximum(tl.maximum(probe0,probe1),tl.maximum(probe2,probe3));diagonal=tl.max(tl.where(mask,off>.0,0),axis=0)==0;geometric=tl.min(tl.where(mask,(dynamic<1e1)&(tl.abs(ratio)<5.),1),axis=0);nearrank=tl.min(tl.where(mask,(dynamic<3.)&(ratio>2e1),1),axis=0);dense_all=tl.min(tl.where(mask,(dynamic>1e3)&(dynamic<1e6)&(tl.abs(ratio)<.5),1),axis=0);repeated=(dynamic<1.8)&(tl.abs(ratio)<.0001)&(frob2>1e-20)&(frob2<1e20);tl.store(repeated_output+matrix,repeated,mask=mask);repeated_any=tl.max(tl.where(mask,repeated,0),axis=0);effective_min=tl.min(tl.where(mask,effective,float('inf')),axis=0);effective_max=tl.max(tl.where(mask,effective,-float('inf')),axis=0);ratio_min=tl.min(tl.where(mask,ratio,float('inf')),axis=0);ratio_max=tl.max(tl.where(mask,ratio,-float('inf')),axis=0);heterogeneous=(effective_max-effective_min>8e1)|(ratio_max-ratio_min>8.);flags=diagonal.to(tl.int32)+2*geometric.to(tl.int32)+4*nearrank.to(tl.int32)+8*dense_all.to(tl.int32)+16*repeated_any.to(tl.int32)+32*heterogeneous.to(tl.int32);tl.store(metadata,flags);dense=(dynamic>1e3)&(dynamic<1e6)&(tl.abs(ratio)<.5);clustered=(dynamic<1.01)&(ratio>1.)&~repeated;route=tl.where(dense,0,tl.where(clustered,2,tl.where(repeated,1,3))).to(tl.int32);native_band=tl.load(native+matrix,mask=mask,other=0).to(tl.int1);exact=route==3;extreme=exact&(maximum>1e11*minimum);refined=tl.where(extreme,5,tl.where(exact&native_band,4,route));tl.store(refined_route+matrix,refined,mask=mask)
for index in tl.static_range(0,4):count=tl.sum(tl.where(mask&(route==index),1,0),axis=0);tl.store(metadata+2+index,count)
dense_exact_count=tl.sum(tl.where(mask&exact&~native_band&~extreme,1,0),axis=0);extreme_count=tl.sum(tl.where(mask&extreme,1,0),axis=0);tl.store(metadata+6,dense_exact_count);tl.store(metadata+7,extreme_count)
@torch.no_grad()
def _e094_fused_n1024_route(data:torch.Tensor):batch,n,_=data.shape;row_energy=torch.empty((batch,n),device=data.device);diagonal=torch.empty_like(row_energy);row_l1=torch.empty_like(row_energy);row_outside=torch.empty((batch,n),device=data.device,dtype=torch.int8);row_ring=torch.empty_like(row_outside);stats=torch.empty((batch,5),device=data.device);native=torch.empty((batch,),device=data.device,dtype=torch.bool);total_rows=batch*n;_e2440_n1024_row_stats_native[triton.cdiv(total_rows,8),](data,row_energy,diagonal,row_l1,row_outside,row_ring,total_rows,N=n,BLOCK_ROWS=8,BLOCK_N=1024,num_warps=8,num_stages=1);_e2440_n1024_matrix_stats_native[batch,](row_energy,diagonal,row_l1,row_outside,row_ring,stats,native,N=n,BLOCK_N=1024,num_warps=8,num_stages=1);metadata=torch.empty((8,),device=data.device,dtype=torch.int32);repeated=torch.empty((batch,),device=data.device,dtype=torch.bool);refined_route=torch.empty((batch,),device=data.device,dtype=torch.int32);_e2440_n1024_route_and_subgroups[1,](data,stats,native,metadata,repeated,refined_route,batch,BLOCK_BATCH=triton.next_power_of_2(batch),num_warps=4,num_stages=1);_e279_n1024_fast_geometric_gate[1,](data,row_energy,stats,metadata[1:2],batch,BLOCK_BATCH=triton.next_power_of_2(batch),N=n,BLOCK_N=1024,num_warps=8,num_stages=1);host=[int(value)for value in metadata.cpu().tolist()];mixed_metadata=refined_route,metadata[2:6],host[2:6],host[6],host[7];return host[0],repeated,host[1],stats,mixed_metadata
@torch.no_grad()
def _e153_fused_n512_stats(data:torch.Tensor):batch,n,_=data.shape;row_energy=torch.empty((batch,n),device=data.device);diagonal=torch.empty_like(row_energy);row_l1=torch.empty_like(row_energy);stats=torch.empty((batch,5),device=data.device);total_rows=batch*n;_e094_n1024_row_stats[triton.cdiv(total_rows,16),](data,row_energy,diagonal,row_l1,total_rows,N=n,BLOCK_ROWS=16,BLOCK_N=512,num_warps=8,num_stages=1);_e094_n1024_matrix_stats[batch,](row_energy,diagonal,row_l1,stats,N=n,BLOCK_N=512,num_warps=8,num_stages=1);return stats
@triton.jit
def _e154_n512_route_flags(data,stats,output,batch,BLOCK_BATCH:tl.constexpr):matrix=tl.arange(0,BLOCK_BATCH);mask=matrix<batch;trace=tl.load(stats+matrix*5,mask=mask,other=.0);frob2=tl.load(stats+matrix*5+1,mask=mask,other=1.);minimum=tl.load(stats+matrix*5+2,mask=mask,other=1.);maximum=tl.load(stats+matrix*5+3,mask=mask,other=1.);dynamic=maximum/tl.maximum(minimum,1e-30);ratio=trace/tl.sqrt(tl.maximum(frob2,1e-30));effective=trace*trace/tl.maximum(frob2,1e-30);base=matrix*512*512;probe0=tl.abs(tl.load(data+base+1,mask=mask,other=.0));probe1=tl.abs(tl.load(data+base+256*512+257,mask=mask,other=.0));probe2=tl.abs(tl.load(data+base+510*512+511,mask=mask,other=.0));probe3=tl.abs(tl.load(data+base+511,mask=mask,other=.0));off=tl.maximum(tl.maximum(probe0,probe1),tl.maximum(probe2,probe3));diagonal=tl.max(tl.where(mask,off>.0,0),axis=0)==0;effective_min=tl.min(tl.where(mask,effective,float('inf')),axis=0);effective_max=tl.max(tl.where(mask,effective,-float('inf')),axis=0);ratio_min=tl.min(tl.where(mask,ratio,float('inf')),axis=0);ratio_max=tl.max(tl.where(mask,ratio,-float('inf')),axis=0);heterogeneous=(effective_max-effective_min>8e1)|(ratio_max-ratio_min>8.);dense_profile=(dynamic>1e3)&(dynamic<1e6)&(tl.abs(ratio)<.9);dense=tl.min(tl.where(mask,dense_profile,1),axis=0);has_dense_profile=tl.max(tl.where(mask,dense_profile,0),axis=0);rankdef=(ratio_min>1e1)&(effective_min>2e2);lapack_even=tl.min(tl.where(mask,dynamic<2.,1),axis=0)&(ratio_max-ratio_min>1.);cluster_candidate=tl.min(tl.where(mask,dynamic<1.01,1),axis=0);negative_rank=.5*(512.-ratio*22.627416997969522);cluster_rank=tl.min(tl.where(mask,tl.abs(negative_rank-17e1)<.25,1),axis=0);flags=diagonal.to(tl.int32)+2*heterogeneous.to(tl.int32)+4*dense.to(tl.int32)+8*rankdef.to(tl.int32)+16*lapack_even.to(tl.int32)+32*cluster_candidate.to(tl.int32)+64*cluster_rank.to(tl.int32)+128*has_dense_profile.to(tl.int32);tl.store(output,flags)
@torch.no_grad()
def _e154_fused_n512_route(data:torch.Tensor):stats=_e153_fused_n512_stats(data);flags=torch.empty((),device=data.device,dtype=torch.int32);_e154_n512_route_flags[1,](data,stats,flags,data.shape[0],BLOCK_BATCH=triton.next_power_of_2(data.shape[0]),num_warps=4,num_stages=1);return int(flags.item()),stats
@triton.jit
def _e243_build_band32_sparse_wy_kernel(reflectors,operators):
batch=tl.program_id(0);packed_id=tl.program_id(1);root=tl.sqrt(tl.cast(8*packed_id+1,tl.float32));window_id=tl.cast((root-1.)*.5,tl.int32);spatial_block=packed_id-window_id*(window_id+1)//2;local_row=tl.arange(0,64);rank=tl.arange(0,32);high_column=509-window_id*32;column=high_column-rank;active=column>=0;reflector_lane=local_row[:,None]-(31-rank[None,:]);support=column[None,:]+1+spatial_block*32;reflector_mask=active[None,:]&(reflector_lane>=0)&(reflector_lane<32)&(support+reflector_lane<512);address=reflectors+batch*512*16*64+(column[None,:]*16+spatial_block)*64+reflector_lane;vectors=tl.load(address,mask=reflector_mask,other=.0);gram=tl.dot(tl.trans(vectors),vectors,input_precision='ieee');row=rank[:,None];col=rank[None,:];triangular=tl.where((row==col)&active[:,None],2.,.0).to(tl.float32)
for current_column in tl.static_range(1,32):gram_column=tl.sum(tl.where(col==current_column,gram,.0),axis=1);overlap=tl.sum(tl.where(col<current_column,triangular*gram_column[None,:],.0),axis=1);triangular=tl.where((col==current_column)&(row<current_column),(-2.*overlap)[:,None],triangular)
weighted=tl.dot(vectors,tl.trans(triangular),input_precision='tf32x3');product=tl.dot(weighted,tl.trans(vectors),input_precision='tf32x3');operator_row=local_row[:,None];operator_col=local_row[None,:];operator=tl.where(operator_row==operator_col,1.,.0)-product;operator_base=(batch*136+packed_id)*64*64;tl.store(operators+operator_base+operator_row*64+operator_col,operator)
@triton.jit
def _e243_band32_sparse_wy_replay_kernel(q,operators,block_cols:tl.constexpr,signal_dependents:tl.constexpr):
if signal_dependents:tl_cuda.gdc_launch_dependents()
batch=tl.program_id(0);column_tile=tl.program_id(1);local_row=tl.arange(0,64);columns=column_tile*block_cols+tl.arange(0,block_cols);column_mask=columns<512;q_base=batch*512*512;operator_row=local_row[:,None];operator_col=local_row[None,:]
for window_id in tl.range(0,16):
high_column=509-window_id*32;low_column=tl.maximum(0,high_column-31);block_limit=(512-low_column-2+31)//32;first_operator=window_id*(window_id+1)//2
for spatial_block in tl.range(0,block_limit):union_start=high_column-30+spatial_block*32;operator=tl.load(operators+(batch*136+first_operator+spatial_block)*64*64+operator_row*64+operator_col);rows=union_start+local_row;row_mask=(rows>=0)&(rows<512);address=q+q_base+rows[:,None]*512+columns[None,:];tile=tl.load(address,mask=row_mask[:,None]&column_mask[None,:],other=.0);operator16=operator.to(tl.float16);tile16=tile.to(tl.float16);transformed=tl.dot(operator16,tile16,out_dtype=tl.float32);tile_residual=(tile-tile16.to(tl.float32)).to(tl.float16);transformed+=tl.dot(operator16,tile_residual,out_dtype=tl.float32);operator_residual=(operator-operator16.to(tl.float32)).to(tl.float16);transformed+=tl.dot(operator_residual,tile16,out_dtype=tl.float32);tl.store(address,transformed,mask=row_mask[:,None]&column_mask[None,:])
@torch.no_grad()
def _e243_band32_sparse_wy_replay(vectors:torch.Tensor,reflectors:torch.Tensor):batch=vectors.shape[0];operators=torch.empty((batch,136,64,64),device=vectors.device,dtype=torch.float32);_e243_build_band32_sparse_wy_kernel[batch,136](reflectors,operators,num_warps=2,num_stages=1);output=vectors.clone();_e243_band32_sparse_wy_replay_kernel[batch,8](output,operators,block_cols=64,signal_dependents=False,num_warps=4,num_stages=1);return output
_E095_POTRF128_NAME='e095_potrf128_packed_shared_full_output';_E095_TRSM128_NAME='e095_right_trsm128_block16_rows32';_RANKDEF_TRSM128_PARENT512_NAME='e1593_right_trsm128_parent_ld512';_E095_N128_TRI=128*129//2;_E095_POTRF128_SHARED_BYTES=(128*128+1)*4;_RANKDEF_POTRF96_NAME='rankdef_potrf96_packed_shared_full_output';_RANKDEF_TRSM96_NAME='rankdef_right_trsm96_block16_rows32';_RANKDEF_N96_TRI=96*97//2;_RANKDEF_POTRF192_NAME='rankdef_potrf192_packed_shared_full_output';_RANKDEF_ROOT_WMMA_POTRF192_NAME='e2386_rankdef_root_potrf192_wmma_tf32x3';_RANKDEF_TRSM192_NAME='rankdef_right_trsm192_block16_rows32';_RANKDEF_INVERSE_LT192_NAME='e1589_inverse_lt192_implicit_identity';_RANKDEF_N192_TRI=192*193//2
@memo(maxsize=1)
def _rankdef_potrf96_kernel():source=_N160_PACKED_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 96;').replace('potrf160_packed_shared_full_output',_RANKDEF_POTRF96_NAME).replace('right_trsm160_packed_block16_rows32','rankdef_unused_trsm96');source=_fast_only_cuda_kernel(source,_RANKDEF_POTRF96_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_POTRF96_NAME),_RANKDEF_POTRF96_NAME)
@memo(maxsize=1)
@memo(maxsize=1)
def _rankdef_trsm96_kernel():source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 96;').replace('constexpr int PANEL = 16;','constexpr int PANEL = 32;').replace('right_trsm160_block16_rows32',_RANKDEF_TRSM96_NAME);source=_fast_only_cuda_kernel(source,_RANKDEF_TRSM96_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_TRSM96_NAME),_RANKDEF_TRSM96_NAME)
@memo(maxsize=1)
def _rankdef_potrf192_kernel():source=_N160_PACKED_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 192;').replace('potrf160_packed_shared_full_output',_RANKDEF_POTRF192_NAME).replace('right_trsm160_packed_block16_rows32','rankdef_unused_trsm192');source=_fast_only_cuda_kernel(source,_RANKDEF_POTRF192_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_POTRF192_NAME),_RANKDEF_POTRF192_NAME)
@memo(maxsize=1)
def _rankdef_root_wmma_potrf192_kernel():source=_e202_wmma_potrf_source(192,_RANKDEF_ROOT_WMMA_POTRF192_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_ROOT_WMMA_POTRF192_NAME),_RANKDEF_ROOT_WMMA_POTRF192_NAME)
@memo(maxsize=1)
@memo(maxsize=1)
def _rankdef_trsm192_kernel():source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 192;').replace('right_trsm160_block16_rows32',_RANKDEF_TRSM192_NAME);source=_fast_only_cuda_kernel(source,_RANKDEF_TRSM192_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_TRSM192_NAME),_RANKDEF_TRSM192_NAME)
@memo(maxsize=1)
@memo(maxsize=1)
def _rankdef_inverse_lt192_kernel():
source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 192;').replace('right_trsm160_block16_rows32',_RANKDEF_INVERSE_LT192_NAME);old=' const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows ? source[(long long)row * N + column] : 0.0f;';new=' const int row = row_base + local_row;\n rhs[column * PITCH + local_row] =\n row < rows && row == column ? 1.0f : 0.0f;'
if source.count(old)!=1:raise RuntimeError('inverse-LT192 identity anchor changed')
source=source.replace(old,new,1);source=_fast_only_cuda_kernel(source,_RANKDEF_INVERSE_LT192_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_INVERSE_LT192_NAME),_RANKDEF_INVERSE_LT192_NAME)
@torch.no_grad()
def _rankdef_inverse_lt192(lower:torch.Tensor):return _e2369_tcgen_inverse_lt(lower)
@memo(maxsize=1)
def _e095_potrf128_kernel():source=_e202_wmma_potrf_source(128,_E095_POTRF128_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E095_POTRF128_NAME),_E095_POTRF128_NAME)
_E1994_POTRF_INVERSE128_NAME='e1994_potrf_inverse128_packed_shared';_E1994_POTRF_INVERSE128_SHARED_BYTES=(128*129//2+1+128*129)*4
def _e1994_potrf_inverse128_source():
source=_N160_PACKED_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 128;',1).replace('potrf160_packed_shared_full_output',_E1994_POTRF_INVERSE128_NAME,1).replace('right_trsm160_packed_block16_rows32','e1994_unused_trsm128',1);old=' float* destination = packed_lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n destination[index] = row >= column ? factor[pidx(row, column)] : 0.0f;\n }\n}';new=' constexpr int PITCH = N + 1;\n constexpr int IPANEL = 32;\n float* rhs = factor + TRI + 1;\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n rhs[column * PITCH + row] = row == column ? 1.0f : 0.0f;\n }\n __syncthreads();\n #pragma unroll 1\n for (int panel = 0; panel < N; panel += IPANEL) {\n const int end = panel + IPANEL;\n if (tid < N) {\n const int local_row = tid;\n #pragma unroll\n for (int offset = 0; offset < IPANEL; ++offset) {\n const int column = panel + offset;\n float value = rhs[column * PITCH + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < IPANEL; ++k_offset) {\n const int k = panel + k_offset;\n if (k >= column) break;\n value = fmaf(-factor[pidx(column, k)],\n rhs[k * PITCH + local_row], value);\n }\n rhs[column * PITCH + local_row] =\n value / factor[pidx(column, column)];\n }\n }\n __syncthreads();\n const int remaining = (N - end) * N;\n for (int index = tid; index < remaining; index += blockDim.x) {\n const int column = end + index / N;\n const int local_row = index - (index / N) * N;\n float value = rhs[column * PITCH + local_row];\n #pragma unroll\n for (int k_offset = 0; k_offset < IPANEL; ++k_offset) {\n const int k = panel + k_offset;\n value = fmaf(-factor[pidx(column, k)],\n rhs[k * PITCH + local_row], value);\n }\n rhs[column * PITCH + local_row] = value;\n }\n __syncthreads();\n }\n float* destination = packed_lower + (long long)matrix_id * NN;\n for (int index = tid; index < NN; index += blockDim.x) {\n const int row = index / N;\n const int column = index - row * N;\n destination[index] = rhs[column * PITCH + row];\n }\n}'
if source.count(old)!=1:raise RuntimeError('E1994 fused POTRF/inverse128 anchor changed')
return source.replace(old,new,1)
@memo(maxsize=1)
def _e1994_potrf_inverse128_kernel():source=_fast_only_cuda_kernel(_e1994_potrf_inverse128_source(),_E1994_POTRF_INVERSE128_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E1994_POTRF_INVERSE128_NAME),_E1994_POTRF_INVERSE128_NAME)
@memo(maxsize=1)
@memo(maxsize=1)
def _e095_trsm128_kernel():source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 128;').replace('constexpr int PANEL = 16;','constexpr int PANEL = 32;').replace('right_trsm160_block16_rows32',_E095_TRSM128_NAME);source=_fast_only_cuda_kernel(source,_E095_TRSM128_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_E095_TRSM128_NAME),_E095_TRSM128_NAME)
@memo(maxsize=1)
@memo(maxsize=1)
def _rankdef_trsm128_parent512_kernel():source=_N160_FULL_CHOLESKY_SOURCE.replace('constexpr int N = 160;','constexpr int N = 128;').replace('constexpr int PANEL = 16;','constexpr int PANEL = 32;').replace('right_trsm160_block16_rows32',_RANKDEF_TRSM128_PARENT512_NAME);source=source.replace('const float* source = matrix + (long long)matrix_id * rows * N;','const float* source = matrix + (long long)matrix_id * 512 * 512;',1).replace('row < rows ? source[(long long)row * N + column] : 0.0f;','row < rows ? source[(long long)row * 512 + column] : 0.0f;',1);source=_fast_only_cuda_kernel(source,_RANKDEF_TRSM128_PARENT512_NAME);return CUDAKernel(_fast_nvrtc_compile(source,_RANKDEF_TRSM128_PARENT512_NAME),_RANKDEF_TRSM128_PARENT512_NAME)
@triton.jit
def _e095_rademacher128_kernel(output,total:tl.constexpr):offsets=tl.program_id(0)*256+tl.arange(0,256);within=offsets%(512*128);row=(within//128).to(tl.uint32);column=(within%128).to(tl.uint32);value=row*2654435761;value=value^column*2246822519^1;value=value^value>>16;value=value*2146121005;value=value^value>>15;value=value*2221713035;value=value^value>>16;sign=(value>>31&1).to(tl.float32)*2.-1.;tl.store(output+offsets,sign*.04419417382415922,mask=offsets<total)
@torch.no_grad()
def _e095_cqr128(matrix:torch.Tensor,ridge:float):gram=matrix.mT@matrix;lower=torch.empty_like(gram);_e095_potrf128_kernel().launch((matrix.shape[0],1,1),(256,1,1),(gram,lower,matrix.shape[0],float(ridge)),shared_mem=_E095_POTRF128_SHARED_BYTES);output=torch.empty_like(matrix);_e095_trsm128_kernel().launch((matrix.shape[0],16,1),(256,1,1),(matrix,lower,output,matrix.shape[0],512),shared_mem=128*33*4);return output
@torch.no_grad()
def _e095_repair_cluster128(tridiagonal:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor):
batch=vectors.shape[0];rest=vectors[:,:,128:];basis=torch.empty((batch,512,128),device=vectors.device,dtype=torch.float32);total=basis.numel();_e095_rademacher128_kernel[triton.cdiv(total,256),](basis,total=total,num_warps=8,num_stages=1);basis=basis-rest@(rest.mT@basis);basis=_e095_cqr128(basis,.0001);basis=_e095_cqr128(basis,.0)
for _ in range(2):basis=basis-rest@(rest.mT@basis);basis=_e095_cqr128(basis,.0)
repaired_values=(basis*(tridiagonal@basis)).sum(dim=1);scale=values.abs().amax(dim=1).clamp_min(1e-30);clustered=(values[:,127]-values[:,0]).abs()<=.0001*scale;first_vectors=torch.where(clustered[:,None,None],basis,vectors[:,:,:128]);first_values=torch.where(clustered[:,None],repaired_values,values[:,:128]);output_vectors=torch.cat((first_vectors,rest),dim=2);output_values=torch.cat((first_values,values[:,128:]),dim=1);output_values,order=output_values.sort(dim=1);output_vectors=output_vectors.gather(2,order[:,None,:].expand(-1,512,-1));return output_vectors,output_values
_E2167_DENSE_CERTIFICATE_RESIDUAL_NAME='e2557_dense512_certificate_residual_m256n16';_E2167_DENSE_CERTIFICATE_FINISH_NAME='e2167_dense512_certificate_adjacent_finish32';_E2691_DENSE_CERTIFICATE_RESIDUAL_NAME='e2694_dense512_certificate_residual_m64n64k32'
def _e2167_dense_certificate_residual_source():
text=_e2133_mixed_certificate_source();entry=text.index('extern "C" __global__ __launch_bounds__(256,2)');finish=text.index('extern "C" __global__ __launch_bounds__(256,2)',entry+1);body=text[entry:finish];body=body.replace(_E2133_MIXED_CERTIFICATE_B64_NAME,_E2167_DENSE_CERTIFICATE_RESIDUAL_NAME,1);body=body.replace('__launch_bounds__(256,2)','__launch_bounds__(512,2)',1);body=body.replace('BM=32,BN=64','BM=256,BN=16',1);body=body.replace('const int wm=w/4,wn=w-wm*4;','const int wm=w,wn=0;',1);body=body.replace('owner=(r>>4)*4+(j>>4)','owner=(r>>4)',1);reduction=' atomicAdd(sums+(long long)m*128+pos0+t,total);\n }\n}';augmented=' atomicAdd(sums+(long long)m*128+pos0+t,total);\n float norm=0.f,adjacent=0.f;\n int c=indices[pos0+t];\n int next=indices[(pos0+t+1)%count];\n #pragma unroll\n for(int r=0;r<BM;++r){\n float value=q[base+(long long)(row0+r)*N+c];\n float neighbor=q[base+(long long)(row0+r)*N+next];\n norm=fmaf(value,value,norm);\n adjacent=fmaf(value,neighbor,adjacent);\n }\n atomicAdd(sums+(long long)m*128+count+pos0+t,norm);\n atomicAdd(sums+(long long)m*128+2*count+pos0+t,adjacent);\n }\n}'
if body.count(reduction)!=1:raise RuntimeError('E2167 dense certificate source changed')
body=body.replace(reduction,augmented,1).replace('m*128','m*256');return text[:text.index('extern "C"')]+body
def _e2691_dense_certificate_residual_source():
source=_e2167_dense_certificate_residual_source();replacements=(_E2167_DENSE_CERTIFICATE_RESIDUAL_NAME,_E2691_DENSE_CERTIFICATE_RESIDUAL_NAME),('__launch_bounds__(512,2)','__launch_bounds__(512,2)'),('BM=256,BN=16,BK=64','BM=64,BN=64,BK=32'),('const int wm=w,wn=0;','const int wm=w/4,wn=w-wm*4;'),('owner=(r>>4)','owner=(r>>4)*4+(j>>4)')
for(old,new)in replacements:
if source.count(old)!=1:raise RuntimeError('E2691 dense certificate source changed')
source=source.replace(old,new,1)
return source
_E2167_DENSE_CERTIFICATE_FINISH_SOURCE=f'''
#include <cuda_runtime.h>
__device__ __forceinline__ float e2167_max(float x){{
#pragma unroll
for(int o=16;o;o>>=1)x=fmaxf(x,__shfl_down_sync(0xffffffffu,x,o));
return x;
}}
extern "C" __global__ __launch_bounds__(128,2)
void {_E2167_DENSE_CERTIFICATE_FINISH_NAME}(
const float* __restrict__ values,const float* __restrict__ scale,
const float* __restrict__ sums,const bool* __restrict__ deferred,
bool* __restrict__ risk,int batch,int count,float margin){{
constexpr int N=512;
int m=blockIdx.x,t=threadIdx.x,w=t>>5,lane=t&31;
if(m>=batch)return;
const float inf=__int_as_float(0x7f800000);
float eigen=t<count?sums[(long long)m*256+t]:0.f;
float norm=t<count?fabsf(sums[(long long)m*256+count+t]-1.f):0.f;
float adjacent=t<count?fabsf(sums[(long long)m*256+2*count+t]):0.f;
eigen=isfinite(eigen)?eigen:inf;
norm=isfinite(norm)?norm:inf;
adjacent=isfinite(adjacent)?adjacent:inf;
eigen=e2167_max(eigen);norm=e2167_max(norm);adjacent=e2167_max(adjacent);
float ordering=-1.e30f,value_scale=1.f;
for(int c=t;c<N;c+=128){{
float value=values[(long long)m*N+c];
if(!isfinite(value)){{ordering=inf;value_scale=inf;}}
else{{
value_scale=fmaxf(value_scale,fabsf(value));
if(c+1<N){{
float next=values[(long long)m*N+c+1];
ordering=!isfinite(next)?inf:fmaxf(ordering,value-next);
}}
}}
}}
ordering=e2167_max(ordering);value_scale=e2167_max(value_scale);
__shared__ float es[4],ns[4],as[4],os[4],vs[4];
if(lane==0){{es[w]=eigen;ns[w]=norm;as[w]=adjacent;os[w]=ordering;vs[w]=value_scale;}}
__syncthreads();
if(w==0){{
eigen=e2167_max(lane<4?es[lane]:0.f);
norm=e2167_max(lane<4?ns[lane]:0.f);
adjacent=e2167_max(lane<4?as[lane]:0.f);
ordering=e2167_max(lane<4?os[lane]:-1.e30f);
value_scale=e2167_max(lane<4?vs[lane]:1.f);
if(lane==0){{
float e=eigen/fmaxf(scale[m],1.e-30f),o=ordering/value_scale;
risk[m]=deferred[m]||!isfinite(e)||!isfinite(o)||!isfinite(norm)||!isfinite(adjacent)
||e>margin*200.f*N*1.1920928955078125e-7f
||norm>0.002f||adjacent>0.002f
||o>0.975f*100.f*N*1.1920928955078125e-7f;
}}
}}
}}
'''
@memo(maxsize=1)
def _e2167_dense_certificate_residual_kernel():return CUDAKernel(_fast_nvrtc_compile(_e2691_dense_certificate_residual_source(),_E2691_DENSE_CERTIFICATE_RESIDUAL_NAME),_E2691_DENSE_CERTIFICATE_RESIDUAL_NAME)
@memo(maxsize=1)
def _e2167_dense_certificate_finish_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2167_DENSE_CERTIFICATE_FINISH_SOURCE,_E2167_DENSE_CERTIFICATE_FINISH_NAME),_E2167_DENSE_CERTIFICATE_FINISH_NAME)
@memo(maxsize=4)
def _e2167_dense_certificate_indices(device):return torch.tensor(tuple(range(150,162))+tuple(range(350,374))+tuple(range(494,512)),device=device,dtype=torch.int32)
_E2545_SORT512_NAME='e2545_dense512_repair_sort';_E2545_SORT512_SOURCE='\n#include <cuda_runtime.h>\nextern "C" __global__ __launch_bounds__(256, 4)\nvoid e2545_dense512_repair_sort(\n const float* __restrict__ input,\n float* __restrict__ output,\n int* __restrict__ order,\n int batch) {\n const int matrix = blockIdx.x;\n const int tid = threadIdx.x;\n if (matrix >= batch) return;\n __shared__ float values[512];\n __shared__ int sources[512];\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int index = tid + half * 256;\n values[index] = input[(long long)matrix * 512 + index];\n sources[index] = index;\n }\n __syncthreads();\n for (int width = 2; width <= 512; width <<= 1) {\n for (int stride = width >> 1; stride > 0; stride >>= 1) {\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int left = tid + half * 256;\n const int right = left ^ stride;\n if (right > left) {\n const float left_value = values[left];\n const float right_value = values[right];\n const int left_source = sources[left];\n const int right_source = sources[right];\n const bool ascending = (left & width) == 0;\n const bool greater = left_value > right_value\n || (left_value == right_value && left_source > right_source);\n if (ascending == greater) {\n values[left] = right_value;\n values[right] = left_value;\n sources[left] = right_source;\n sources[right] = left_source;\n }\n }\n }\n __syncthreads();\n }\n }\n #pragma unroll\n for (int half = 0; half < 2; ++half) {\n const int index = tid + half * 256;\n output[(long long)matrix * 512 + index] = values[index];\n order[(long long)matrix * 512 + index] = sources[index];\n }\n}\n'
@memo(maxsize=1)
def _e2545_sort512_kernel():return CUDAKernel(_fast_nvrtc_compile(_E2545_SORT512_SOURCE,_E2545_SORT512_NAME),_E2545_SORT512_NAME)
@torch.no_grad()
def _e2545_sort512(values):batch=values.shape[0];output=torch.empty_like(values);order=torch.empty(values.shape,device=values.device,dtype=torch.int32);_e2545_sort512_kernel().launch((batch,1,1),(256,1,1),(values,output,order,batch));return output,order
@torch.no_grad()
def _e2564_dense512_exact_certificate_risk(data,vectors,values,matrix_scale,candidate_risk):
'Recheck rare FP16-certificate positives with the checker residual.';rows=torch.nonzero(candidate_risk,as_tuple=False).flatten();local_data=data.index_select(0,rows).contiguous();indices=_e2167_dense_certificate_indices(data.device).long();local_vectors=vectors.index_select(0,rows).index_select(2,indices).contiguous();local_values=values.index_select(0,rows).index_select(1,indices).contiguous();local_scale=matrix_scale.index_select(0,rows).clamp_min(1e-30);previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:residual=local_data@local_vectors
finally:torch.set_float32_matmul_precision(previous)
residual=residual-local_vectors*local_values[:,None,:];eigen=residual.abs().sum(dim=1).amax(dim=1)/local_scale;norm=(local_vectors.square().sum(dim=1)-1.).abs().amax(dim=1);adjacent=(local_vectors*torch.roll(local_vectors,shifts=-1,dims=2)).sum(dim=1).abs().amax(dim=1);value_scale=values.index_select(0,rows).abs().amax(dim=1).clamp_min(1.);ordering=(values.index_select(0,rows)[:,:-1]-values.index_select(0,rows)[:,1:]).amax(dim=1)/value_scale;n=data.shape[-1];eps=torch.finfo(torch.float32).eps;local_risk=~torch.isfinite(eigen)|~torch.isfinite(norm)|~torch.isfinite(adjacent)|~torch.isfinite(ordering)|(eigen>.99*2e2*float(n)*eps)|(norm>.002)|(adjacent>.002)|(ordering>.975*1e2*float(n)*eps);exact_risk=torch.zeros_like(candidate_risk);exact_risk[rows]=local_risk;return exact_risk
@torch.no_grad()
def _e2167_certify_dense512(data,vectors,values,matrix_scale,norm_risk_override=None):
def certificate(local_data,local_vectors,local_values,local_scale,deferred,norm_risk_override=None):local_batch=local_data.shape[0];indices=_e2167_dense_certificate_indices(local_data.device);count=indices.numel();sums=torch.zeros((local_batch,256),device=local_data.device);local_risk=torch.empty((local_batch,),device=local_data.device,dtype=torch.bool);_e2167_dense_certificate_residual_kernel().launch((local_batch,8,triton.cdiv(count,64)),(512,1,1),(local_data,local_vectors,local_values,local_scale,indices,sums,local_batch,count),shared_mem=40960);_e2167_dense_certificate_finish_kernel().launch((local_batch,1,1),(128,1,1),(local_values,local_scale,sums,deferred,local_risk,local_batch,count,.95));norm_risk=_e2073_column_norm_risk(local_vectors,.002)if norm_risk_override is None else norm_risk_override;local_risk|=norm_risk;return local_risk,norm_risk,sums
def wide_retry(local_data,local_vectors,local_values,local_scale):
for(begin,end,leaf)in((336,512,_e202_n352_child176_eigh),(96,224,_e196_lapack_n128_leaf_eigh)):window=local_vectors[:,:,begin:end];projected=window.mT@local_data@window;projected=.5*(projected+projected.mT);rotation,repaired_values=leaf(projected);rotation_gram=rotation.mT@rotation;rotation_gram.mul_(-.5);rotation_gram.diagonal(dim1=-2,dim2=-1).add_(1.5);rotation=rotation@rotation_gram;local_vectors[:,:,begin:end]=window@rotation;local_values[:,begin:end]=repaired_values
local_values,order=_e2545_sort512(local_values);local_vectors=_e160_gather_columns(local_vectors,order);deferred=torch.zeros(local_data.shape[0],device=data.device,dtype=torch.bool);retry_risk,_,_=certificate(local_data,local_vectors,local_values,local_scale,deferred);return local_vectors,local_values,retry_risk
matrix_scale=matrix_scale.contiguous();risk,norm_risk,sums=certificate(data,vectors,values,matrix_scale,_E083_DEFERRED_INNER_FAILED,norm_risk_override)
if not bool(risk.any().item()):return vectors,values
direct=risk&(norm_risk|_E083_DEFERRED_INNER_FAILED);candidates=risk&~direct
if bool(candidates.any().item()):
risk=direct|_e2564_dense512_exact_certificate_risk(data,vectors,values,matrix_scale,candidates)
if not bool(risk.any().item()):return vectors,values
vectors=vectors.clone();values=values.clone();direct=risk&(norm_risk|_E083_DEFERRED_INNER_FAILED)
if bool(direct.any().item()):
direct_data=data[direct].contiguous();direct_scale=matrix_scale[direct].contiguous()
try:
safe_vectors,safe_values,safe_norm_risk=_e083_dense512_impl(direct_data,safe_split=True);safe_deferred=torch.zeros(direct_data.shape[0],device=data.device,dtype=torch.bool);safe_risk,_,_=certificate(direct_data,safe_vectors,safe_values,direct_scale,safe_deferred,safe_norm_risk)
if bool(safe_risk.any().item()):exact_values,exact_vectors=torch.linalg.eigh(direct_data[safe_risk].contiguous());safe_vectors=safe_vectors.clone();safe_values=safe_values.clone();safe_vectors[safe_risk]=exact_vectors;safe_values[safe_risk]=exact_values
except torch.linalg.LinAlgError:safe_values,safe_vectors=torch.linalg.eigh(direct_data)
vectors[direct]=safe_vectors;values[direct]=safe_values
repairable=risk&~direct
if bool(repairable.any().item()):
local_data=data[repairable].contiguous();local_vectors=vectors[repairable].contiguous();local_values=values[repairable].contiguous();local_scale=matrix_scale[repairable].contiguous();retry_vectors,retry_values,retry_risk=wide_retry(local_data,local_vectors,local_values,local_scale)
if bool(retry_risk.any().item()):exact_values,exact_vectors=torch.linalg.eigh(local_data[retry_risk].contiguous());retry_vectors=retry_vectors.clone();retry_values=retry_values.clone();retry_vectors[retry_risk]=exact_vectors;retry_values[retry_risk]=exact_values
vectors[repairable]=retry_vectors;values[repairable]=retry_values
return vectors,values
@torch.no_grad()
def _dense512_owned_staged_repair(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor):
'Repair staged rank loss or ordinary pre-Newton risk before exact EVD.';batch,n,_=data.shape;previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
vectors=vectors.clone();values=values.clone();column_norms=vectors.square().sum(dim=1);missing_columns=column_norms.argmin(dim=1);missing=column_norms.gather(1,missing_columns[:,None])[:,0]<.25
if bool(missing.any().item()):
local=vectors[missing];row_norms=local.square().sum(dim=2);seed_rows=row_norms.argmin(dim=1);complement=torch.zeros((local.shape[0],n,1),device=data.device,dtype=data.dtype);complement.scatter_(1,seed_rows[:,None,None],1.)
for _ in range(3):complement-=local@(local.mT@complement)
complement*=complement.square().sum(dim=1,keepdim=True).rsqrt();local_columns=missing_columns[missing];local.scatter_(2,local_columns[:,None,None].expand(-1,n,1),complement);products=data[missing]@complement;rayleigh=(complement*products).sum(dim=1)[:,0];local_values=values[missing];local_values.scatter_(1,local_columns[:,None],rayleigh[:,None]);local_values,order=local_values.sort(dim=1);local=local.gather(2,order[:,None,:].expand(-1,n,-1));vectors[missing]=local;values[missing]=local_values
regular=~missing
if bool(regular.any().item()):local=vectors[regular];gram=local.mT@local;transform=-.5*gram;transform.diagonal(dim1=-2,dim2=-1).add_(1.5);vectors[regular]=local@transform
products=data@vectors;gram=vectors.mT@vectors
finally:torch.set_float32_matmul_precision(previous)
residual=(products-vectors*values[:,None,:]).abs().sum(dim=1).amax(dim=1);eye=torch.eye(n,device=data.device,dtype=data.dtype);orthogonality=(gram-eye).abs().sum(dim=1).amax(dim=1);ordering=(values[:,:-1]-values[:,1:]).amax(dim=1);value_scale=values.abs().amax(dim=1).clamp_min(1.);eps=torch.finfo(torch.float32).eps;nonfinite=~torch.isfinite(residual)|~torch.isfinite(orthogonality)|~torch.isfinite(ordering);residual_risk=residual>.95*2e2*n*eps*matrix_scale;orthogonality_risk=orthogonality>.95*1e2*n*eps;ordering_risk=ordering>.95*1e2*n*eps*value_scale;hard=nonfinite|residual_risk|orthogonality_risk|ordering_risk
if _certificate_reason_ledger is not None:_certificate_reason_record('dense512_owned_staged_repair',_certificate_reason_bits(hard,(_CERT_REASON_NONFINITE,nonfinite),(_CERT_REASON_EIGEN_RESIDUAL,residual_risk),(_CERT_REASON_ORTHOGONALITY,orthogonality_risk),(_CERT_REASON_ORDERING,ordering_risk)),hard)
if bool(hard.any().item()):exact_values,exact_vectors=torch.linalg.eigh(data[hard].contiguous());vectors[hard]=exact_vectors;values[hard]=exact_values
return vectors,values
@torch.no_grad()
def _e083_dense512(data:torch.Tensor,matrix_scale:torch.Tensor):
vectors,values,risk=_e083_dense512_impl(data,safe_split=False)
if bool(risk.any().item()):local_data=data[risk].contiguous();local_scale=matrix_scale[risk].contiguous();repair_vectors,repair_values=_dense512_owned_staged_repair(local_data,vectors[risk].contiguous(),values[risk].contiguous(),local_scale);vectors=vectors.clone();values=values.clone();vectors[risk]=repair_vectors;values[risk]=repair_values
return vectors,values
_fast_precompile_ready=False;_e1342_band_probe_kernel=_e1342_band_probe_kernels
def _fast_precompile_tasks():names='_small_jacobi_kernel _e1630_reduce176_kernel _e1630_n176_jacobi_kernel _e096_lanczos3_kernel _e202_n352_potrf176_kernel _e185_reduce176_block4_kernel _small_n24_jacobi_kernel _dense_active320_panel_kernels _n512_gau_active_image _n352_gau_t_image _e188_n320_lanczos3_kernel _e106_trsm160_half2_kernel _e1008_reduce160_block2_kernel _e107_solve160_kernel _e945_n8_jacobi_kernel _e189_n1024_lanczos3_kernel _e1074_l10_pdl_kernel _n1024_rhh_left_kernels _n1024_rhh_right_kernels _n1024_rhh_t_kernel _dense1024_qr_tail128_kernel _e192_n512_lanczos5_kernel _lapack_active256_panel_kernels _active352_t64_kernel _e1521_reduce256_kernel _e1773_dense1024_solve256_pdl_kernel _e204_rankdef_n96_reduce_kernel _e1778_n96_solve_pdl_kernel _e1284_n96_mgs_kernel _e2073_column_norm_kernel _e931_n2048_lanczos3_kernel _e940_n1024_lanczos3_kernel _dense2048_potrf256_x1_kernel _e924_n512_lanczos3_kernel _e208_reduce256_kernel _e1789_dense2048_solve256_pdl_kernel _e095_potrf128_kernel _rankdef_trsm128_parent512_kernel _e959_lapack_inverse_lt128_kernel _e095_trsm128_kernel _rankdef_potrf192_kernel _rankdef_inverse_lt192_kernel _rankdef_potrf96_kernel _rankdef_trsm96_kernel _e951_n16_jacobi_kernel _e951_n32_jacobi_kernel _e1529_clustered_potrf170_kernel _e1529_clustered_trsm170_kernel _clustered_compact_qr_kernels _clustered_guard8_kernel _e1052_n768_lanczos3_kernel _e1186_qr768_left_kernels _e1186_qr768_right_kernels _e1052_n384_lanczos3_kernel _e842_reduce192_kernel _e1771_solve192_pdl_kernel _e1214_n512_lower_kernel _e1782_mixed_tridiag24_pdl_kernel _band16_repair512_kernel _e1342_band_probe_kernel _e147_reduce320_kernel _e1769_solve320_pdl_kernel _e1994_potrf_inverse128_kernel _rankdef_trsm192_kernel _e1344_psd_probe_kernel _e208_solve256_kernel _repair_kernel _e2115_mixed_certificate_kernels _e834_entry_sharded_gram_t32_kernel _e1208_b32_main_kernel _e1208_b32_tail_kernel _e1760_tridiag1024_pdl_kernel _lapack_child_active128_panel_kernels _t32_kernel _e196_lapack_n128_reduce_kernel _e1776_n128_solve_pdl_kernel _e2984_n128_solve_mgs_kernel _e1282_n128_mgs_kernel _e2860_dense512_stage_select_pack_kernel _lapack_rr96_select_pack_kernel _small_n96_kernel _lapack_rr96_rotate_order_kernel _lapack_rr96_rotate_order_apply_kernel'.split();names=['_e2545_sort512_kernel','_e2225_tcgen_trsm256_kernel','_e2369_tcgen_trsm160_kernel','_e2373_tcgen_trsm176_kernel','_e2369_tcgen_inverse128_kernel','_e2369_tcgen_inverse192_kernel','_rankdef_root_wmma_potrf192_kernel','_e2387_clustered_potrf170_kernel','_e2409_potrf160_wmma_kernel','_e2676_n352_certificate_kernels','_e2992_n352_norm_guard_kernel','_e2681_nearrank_certificate_kernel']+names;tasks=[(globals()[name],())for name in names];tasks.extend(((_e1776_n128_solve_pdl_kernel,(8,)),(_e1762_solve176_pdl_kernel,(23,)),(_e1762_solve176_pdl_kernel,(20,)),(_n176_choleskyqr_kernel,('right_trsm176_block16_rows16',)),(_e083_n160_cholesky_kernel,('potrf160_packed_shared_full_output',)),(_e083_n160_cholesky_kernel,('right_trsm160_block16_rows32',)),(_e084_cholesky_kernel,('potrf256_strided_tf32x3_w19',)),(_e084_cholesky_kernel,('right_trsm256_wmma_tf32x1_w4',)),(_e084_cholesky_kernel,('right_trsm256_strided_panel32_rows32',)),(_rankdef_small_active_panel_kernels,(384,192)),(_rankdef_small_active_panel_kernels,(192,96)),(_panel_kernels,(1024,)),(_e2695_nonic_precompile,())));return tasks
_e2167_base_precompile_tasks=_fast_precompile_tasks
def _fast_precompile_tasks():tasks=_e2167_base_precompile_tasks();tasks.extend(((_e2167_dense_certificate_residual_kernel,()),(_e2167_dense_certificate_finish_kernel,())));return tasks
_E2175_GEOMETRIC_RESIDUAL_NAME='e2175_geometric_fused_residual';_E2175_GEOMETRIC_FINISH_NAME='e2175_geometric_fused_finish';_E2175_GEOMETRIC_SOURCE='\n#include <cuda_runtime.h>\n#include <cuda_fp16.h>\n#include <mma.h>\nusing namespace nvcuda;\n\n__device__ __forceinline__ float e2175_tf32(float value) {\n unsigned int bits = __float_as_uint(value);\n const unsigned int exponent = bits & 0x7f800000u;\n if (exponent != 0x7f800000u)\n bits = (bits + 0x00000fffu + ((bits >> 13) & 1u)) & 0xffffe000u;\n return __uint_as_float(bits);\n}\n\n__device__ __forceinline__ float e2175_warp_max(float value) {\n #pragma unroll\n for (int offset = 16; offset; offset >>= 1)\n value = fmaxf(value, __shfl_down_sync(0xffffffffu, value, offset));\n return value;\n}\n\nextern "C" __global__ __launch_bounds__(256, 2)\nvoid e2175_geometric_fused_residual(\n const float* __restrict__ matrix,\n const float* __restrict__ vectors,\n const float* __restrict__ values,\n const float* __restrict__ scale,\n int scale_stride,\n float* __restrict__ sums,\n int batch) {\n constexpr int N = 1024, BM = 128, BN = 16, BK = 64;\n const int matrix_id = blockIdx.x;\n const int row0 = blockIdx.y * BM;\n const int tid = threadIdx.x;\n const int warp = tid >> 5;\n if (matrix_id >= batch) return;\n\n extern __shared__ __align__(1024) unsigned char raw[];\n __half* matrix_tile = reinterpret_cast<__half*>(raw);\n __half* vector_tile = matrix_tile + BM * BK;\n float* output_tiles = reinterpret_cast<float*>(vector_tile + BK * BN);\n float* gram_tile = output_tiles + 8 * 16 * 16;\n const long long matrix_base = (long long)matrix_id * N * N;\n const float matrix_scale = fmaxf(scale[(long long)matrix_id * scale_stride], 1.e-30f);\n const float inverse_scale = 1.f / matrix_scale;\n\n __shared__ int selected_indices[16];\n __shared__ int magnitude_indices[4];\n if (tid == 0) {\n // Values are sorted. Two ten-step lower bounds replace the original\n // comparison/reduction/cat/topk PyTorch chain.\n int lo = 0, hi = N;\n while (lo < hi) {\n const int middle = (lo + hi) >> 1;\n if (values[(long long)matrix_id * N + middle] < 0.f) lo = middle + 1;\n else hi = middle;\n }\n const int negative_end = lo;\n lo = 0; hi = N;\n while (lo < hi) {\n const int middle = (lo + hi) >> 1;\n if (values[(long long)matrix_id * N + middle] <= 0.f) lo = middle + 1;\n else hi = middle;\n }\n const int positive_begin = lo;\n #pragma unroll\n for (int index = 0; index < 6; ++index) {\n selected_indices[index] = max(0, min(N - 1, negative_end + index - 3));\n selected_indices[6 + index] = max(0, min(N - 1, positive_begin + index - 3));\n }\n\n const int nonzero_count = N - (positive_begin - negative_end);\n const float edge0 = fabsf(values[(long long)matrix_id * N]);\n const float edge1 = fabsf(values[(long long)matrix_id * N + N - 1]);\n const float threshold = fmaxf(edge0, edge1)\n * exp2f(-0.02150537634408597f * (0.5f * nonzero_count - 0.5f));\n lo = 0; hi = negative_end;\n while (lo < hi) {\n const int middle = (lo + hi) >> 1;\n if (values[(long long)matrix_id * N + middle] < -threshold) lo = middle + 1;\n else hi = middle;\n }\n const int negative_crossing = lo;\n lo = positive_begin; hi = N;\n while (lo < hi) {\n const int middle = (lo + hi) >> 1;\n if (values[(long long)matrix_id * N + middle] < threshold) lo = middle + 1;\n else hi = middle;\n }\n const int positive_crossing = lo;\n int candidates[12];\n int candidate_count = 0;\n #pragma unroll\n for (int delta = -3; delta <= 2; ++delta) {\n if (negative_end > 0) {\n const int candidate = max(0, min(negative_end - 1, negative_crossing + delta));\n bool unique = true;\n #pragma unroll\n for (int prior = 0; prior < candidate_count; ++prior)\n unique &= candidates[prior] != candidate;\n if (unique) candidates[candidate_count++] = candidate;\n }\n if (positive_begin < N) {\n const int candidate = max(positive_begin, min(N - 1, positive_crossing + delta));\n bool unique = true;\n #pragma unroll\n for (int prior = 0; prior < candidate_count; ++prior)\n unique &= candidates[prior] != candidate;\n if (unique) candidates[candidate_count++] = candidate;\n }\n }\n #pragma unroll\n for (int output = 0; output < 4; ++output) {\n float best_distance = 1.e30f;\n int best_position = -1;\n #pragma unroll\n for (int position = 0; position < 12; ++position) {\n if (position < candidate_count && candidates[position] >= 0) {\n const int candidate = candidates[position];\n const float distance = fabsf(\n fabsf(values[(long long)matrix_id * N + candidate]) - threshold);\n if (distance < best_distance) {\n best_distance = distance;\n best_position = position;\n }\n }\n }\n const int chosen = best_position >= 0 ? candidates[best_position] : output;\n magnitude_indices[output] = chosen;\n selected_indices[12 + output] = chosen;\n if (best_position >= 0) candidates[best_position] = -1;\n }\n }\n __syncthreads();\n\n wmma::fragment<wmma::accumulator, 16, 16, 16, float> residual;\n wmma::fragment<wmma::accumulator, 16, 16, 16, float> gram;\n wmma::fill_fragment(residual, 0.f);\n wmma::fill_fragment(gram, 0.f);\n\n #pragma unroll\n for (int start = 0; start < N; start += BK) {\n for (int offset = tid; offset < BM * BK; offset += blockDim.x) {\n const int row = offset / BK;\n const int inner = offset - row * BK;\n matrix_tile[offset] = __float2half_rn(\n matrix[matrix_base + (long long)(row0 + row) * N + start + inner]\n * inverse_scale);\n }\n for (int offset = tid; offset < BK * BN; offset += blockDim.x) {\n const int inner = offset / BN;\n const int position = offset - inner * BN;\n const int column = selected_indices[position];\n vector_tile[offset] = __float2half_rn(\n vectors[matrix_base + (long long)(start + inner) * N + column]);\n }\n __syncthreads();\n\n #pragma unroll\n for (int inner = 0; inner < BK; inner += 16) {\n wmma::fragment<wmma::matrix_a, 16, 16, 16,\n __half, wmma::row_major> matrix_fragment;\n wmma::fragment<wmma::matrix_b, 16, 16, 16,\n __half, wmma::row_major> vector_fragment;\n wmma::load_matrix_sync(\n matrix_fragment, matrix_tile + warp * 16 * BK + inner, BK);\n wmma::load_matrix_sync(\n vector_fragment, vector_tile + inner * BN, BN);\n wmma::mma_sync(residual, matrix_fragment, vector_fragment, residual);\n\n if (warp == 0) {\n wmma::fragment<wmma::matrix_a, 16, 16, 16,\n __half, wmma::col_major> transpose_fragment;\n wmma::load_matrix_sync(\n transpose_fragment, vector_tile + inner * BN, BN);\n wmma::mma_sync(gram, transpose_fragment, vector_fragment, gram);\n }\n }\n __syncthreads();\n }\n\n wmma::store_matrix_sync(\n output_tiles + warp * 16 * 16, residual, 16, wmma::mem_row_major);\n if (warp == 0)\n wmma::store_matrix_sync(gram_tile, gram, 16, wmma::mem_row_major);\n __syncthreads();\n\n if (tid < BN) {\n float total = 0.f;\n const int column = selected_indices[tid];\n const float lambda = values[(long long)matrix_id * N + column] * inverse_scale;\n #pragma unroll\n for (int row = 0; row < BM; ++row) {\n const int owner = row >> 4;\n const float product = output_tiles[\n owner * 256 + (row & 15) * 16 + tid];\n total += fabsf(product\n - vectors[matrix_base + (long long)(row0 + row) * N + column] * lambda);\n }\n sums[(long long)matrix_id * 416 + blockIdx.y * BN + tid] = total;\n }\n\n if (row0 == 0 && tid < BN) {\n float total = 0.f;\n #pragma unroll\n for (int row = 0; row < BN; ++row)\n total += fabsf(gram_tile[row * BN + tid]\n - (row == tid ? 1.f : 0.f));\n sums[(long long)matrix_id * 416 + 400 + tid] = total;\n }\n\n // Preserve the strict FP32 34-column norm guard that caught geometric\n // underflow failures. Four row tiles contribute through atomics and the\n // following finish kernel applies the threshold after all contributions.\n if (tid < 34) {\n long long x0 = magnitude_indices[0];\n long long x1 = magnitude_indices[1];\n long long x2 = magnitude_indices[2];\n long long x3 = magnitude_indices[3];\n #define E2175_SWAP(a, b) do { if ((a) > (b)) { long long tmp = (a); (a) = (b); (b) = tmp; } } while (0)\n E2175_SWAP(x0, x1); E2175_SWAP(x2, x3); E2175_SWAP(x0, x2);\n E2175_SWAP(x1, x3); E2175_SWAP(x1, x2);\n const int center = tid < 17 ? int((x0 + x1) >> 1) : int((x2 + x3) >> 1);\n const int column = max(0, min(N - 1, center + (tid % 17) - 8));\n float norm = 0.f;\n #pragma unroll 4\n for (int row = row0; row < row0 + BM; ++row) {\n const float value = vectors[matrix_base + (long long)row * N + column];\n norm = fmaf(value, value, norm);\n }\n sums[(long long)matrix_id * 416 + 128 + blockIdx.y * 34 + tid] = norm;\n }\n}\n\nextern "C" __global__ __launch_bounds__(256, 2)\nvoid e2175_geometric_fused_finish(\n const float* __restrict__ sums,\n const float* __restrict__ values,\n bool* __restrict__ risk,\n float* __restrict__ scores,\n int batch, float margin) {\n constexpr int N = 1024;\n const int matrix_id = blockIdx.x;\n const int tid = threadIdx.x;\n const int warp = tid >> 5;\n const int lane = tid & 31;\n if (matrix_id >= batch) return;\n\n float eigen = 0.f;\n float norm_error = 0.f;\n if (tid < 16) {\n #pragma unroll\n for (int tile = 0; tile < 8; ++tile)\n eigen += sums[(long long)matrix_id * 416 + tile * 16 + tid];\n }\n float orthogonal = tid < 16\n ? sums[(long long)matrix_id * 416 + 400 + tid] : 0.f;\n if (tid < 34) {\n float norm = 0.f;\n #pragma unroll\n for (int tile = 0; tile < 8; ++tile)\n norm += sums[(long long)matrix_id * 416 + 128 + tile * 34 + tid];\n norm_error = fabsf(norm - 1.f);\n }\n float ordering = -1.e30f;\n float value_scale = 1.f;\n for (int column = tid; column < N; column += blockDim.x) {\n const float value = values[(long long)matrix_id * N + column];\n value_scale = fmaxf(value_scale, fabsf(value));\n if (column + 1 < N)\n ordering = fmaxf(\n ordering, value - values[(long long)matrix_id * N + column + 1]);\n }\n eigen = e2175_warp_max(eigen);\n orthogonal = e2175_warp_max(orthogonal);\n norm_error = e2175_warp_max(norm_error);\n ordering = e2175_warp_max(ordering);\n value_scale = e2175_warp_max(value_scale);\n\n __shared__ float eigen_warp[8], orthogonal_warp[8], norm_warp[8];\n __shared__ float ordering_warp[8], scale_warp[8];\n if (lane == 0) {\n eigen_warp[warp] = eigen;\n orthogonal_warp[warp] = orthogonal;\n norm_warp[warp] = norm_error;\n ordering_warp[warp] = ordering;\n scale_warp[warp] = value_scale;\n }\n __syncthreads();\n if (warp == 0) {\n eigen = e2175_warp_max(lane < 8 ? eigen_warp[lane] : 0.f);\n orthogonal = e2175_warp_max(lane < 8 ? orthogonal_warp[lane] : 0.f);\n norm_error = e2175_warp_max(lane < 8 ? norm_warp[lane] : 0.f);\n ordering = e2175_warp_max(lane < 8 ? ordering_warp[lane] : -1.e30f);\n value_scale = e2175_warp_max(lane < 8 ? scale_warp[lane] : 1.f);\n if (lane == 0) {\n const float ordering_score = ordering / value_scale;\n const bool finite = isfinite(eigen) && isfinite(orthogonal)\n && isfinite(norm_error) && isfinite(ordering_score);\n risk[matrix_id] = !finite\n || eigen > margin * 200.f * N * 1.1920928955078125e-7f\n || orthogonal > 0.002f\n || norm_error > 0.0015f\n || ordering_score > 0.975f * 100.f * N * 1.1920928955078125e-7f;\n scores[(long long)matrix_id * 4] = eigen;\n scores[(long long)matrix_id * 4 + 1] = orthogonal;\n scores[(long long)matrix_id * 4 + 2] = norm_error;\n scores[(long long)matrix_id * 4 + 3] = ordering_score;\n }\n }\n}\n'
@memo(maxsize=1)
def _e2175_geometric_certificate_kernels():image=_fast_nvrtc_compile(_E2175_GEOMETRIC_SOURCE,_E2175_GEOMETRIC_RESIDUAL_NAME);return CUDAKernel(image,_E2175_GEOMETRIC_RESIDUAL_NAME),CUDAKernel(image,_E2175_GEOMETRIC_FINISH_NAME)
@torch.no_grad()
def _e2175_certify_geometric_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor|None=None):
batch=data.shape[0]
if matrix_scale is None:matrix_scale=data.abs().sum(dim=1).amax(dim=1)
sums=torch.empty((batch,416),device=data.device,dtype=torch.float32);risk=torch.empty((batch,),device=data.device,dtype=torch.bool);scores=torch.empty((batch,4),device=data.device,dtype=torch.float32);residual,finish=_e2175_geometric_certificate_kernels();residual.launch((batch,8,1),(256,1,1),(data,vectors,values,matrix_scale,matrix_scale.stride(0),sums,batch),shared_mem=27648);finish.launch((batch,1,1),(256,1,1),(sums,values,risk,scores,batch,.9));norm_risk=_e2073_column_norm_risk(vectors,.0015);probe_orthogonal_risk=_e2657_orthogonality_probe_risk(vectors,count=4,threshold=.000355);risk|=probe_orthogonal_risk;norm_only=norm_risk&~risk
if _certificate_reason_ledger is not None:eps=torch.finfo(torch.float32).eps;finite_risk=~torch.isfinite(scores).all(dim=1);eigen_risk=scores[:,0]>.9*2e2*1024.*eps;kernel_orthogonal_risk=scores[:,1]>.002;kernel_norm_risk=scores[:,2]>.0015;ordering_risk=scores[:,3]>.975*1e2*1024.*eps;classified=finite_risk|eigen_risk|kernel_orthogonal_risk|kernel_norm_risk|ordering_risk|probe_orthogonal_risk;unknown_risk=risk&~classified;_certificate_reason_record('geometric1024_output',_certificate_reason_bits(risk,(_CERT_REASON_NONFINITE,finite_risk),(_CERT_REASON_EIGEN_RESIDUAL,eigen_risk),(_CERT_REASON_ORTHOGONALITY,kernel_orthogonal_risk|probe_orthogonal_risk),(_CERT_REASON_NORM,kernel_norm_risk|norm_only),(_CERT_REASON_ORDERING,ordering_risk),(_CERT_REASON_ROUTE,unknown_risk)),risk)
risk_flag,norm_flag=[bool(value)for value in torch.stack((risk.any(),norm_only.any())).cpu().tolist()]
if norm_flag:local_vectors=vectors[norm_only].contiguous();previous=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest');gram=local_vectors.mT@local_vectors;local_vectors=torch.baddbmm(local_vectors,local_vectors,gram,beta=1.5,alpha=-.5);torch.set_float32_matmul_precision(previous);vectors=vectors.clone();vectors[norm_only]=local_vectors
if risk_flag:exact_values,exact_vectors=torch.linalg.eigh(data[risk].contiguous());vectors=vectors.clone();values=values.clone();vectors[risk]=exact_vectors;values[risk]=exact_values
return vectors,values
@torch.no_grad()
def _certify_geometric_output(data:torch.Tensor,vectors:torch.Tensor,values:torch.Tensor,matrix_scale:torch.Tensor|None=None):return _e2175_certify_geometric_output(data,vectors,values,matrix_scale)
_e2175_base_precompile_tasks=_fast_precompile_tasks
def _fast_precompile_tasks():tasks=_e2175_base_precompile_tasks();tasks.append((_e2175_geometric_certificate_kernels,()));tasks.append((_e2760_dense512_newton_guard_kernel,()));return tasks
_e3180_base_precompile_tasks=_fast_precompile_tasks
def _fast_precompile_tasks():tasks=_e3180_base_precompile_tasks();tasks.extend(((_e3180_potrf_x1_kernel,()),(_e3180_potrf_x3_kernel,()),(_e3180_tcgen_kernel,()),(_e3180_l10_kernel,()),(_e3180_tensor_kernel,())));return tasks
def _fast_precompile_worker(device_index:int,function,args):torch.cuda.set_device(device_index);return function(*args)
def _fast_prepare_cuda_kernels(device_index:int):
global _fast_precompile_ready
if _fast_precompile_ready:return
torch.cuda.set_device(device_index)
with ThreadPoolExecutor(max_workers=8)as pool:
futures=[pool.submit(_fast_precompile_worker,device_index,function,args)for(function,args)in _fast_precompile_tasks()]
for future in futures:
try:future.result()
except Exception:pass
_fast_precompile_ready=True
def custom_kernel(data:input_t):
device_index=data.device.index
if device_index is None:device_index=torch.cuda.current_device()
if not _fast_precompile_ready:_fast_prepare_cuda_kernels(device_index)
n=data.shape[-1]
if n==32:return _small_jacobi_eigh(data)
if n==176:return _e185_guarded_n176_eigh(data)
if n==352:
lanczos_stats,dense_flags,matrix_scale=_e133_n352_probe(data)
try:return _e202_guarded_small_n352_eigh(data,lanczos_stats=lanczos_stats,matrix_scale=matrix_scale,route_flags=dense_flags)
except torch.linalg.LinAlgError:pass
values,vectors=torch.linalg.eigh(data);return vectors,values
if n==2048:
if _is_large_dense_cond1(data):return _guarded_large_eigh(data)
values,vectors=torch.linalg.eigh(data);return vectors,values
if n==1024:
batch,n,_=data.shape;route_flags,repeated_profile,fast_geometric,route_stats,mixed_route_metadata=_e094_fused_n1024_route(data)
if route_flags&1:values,order=data.diagonal(dim1=-2,dim2=-1).sort(dim=-1);eye=torch.eye(n,device=data.device,dtype=torch.float32);vectors=eye.expand(batch,-1,-1).gather(2,order[:,None,:].expand(-1,n,-1));return vectors,values
if route_flags&2:
geometric=fast_geometric==1 or fast_geometric<0 and is_lapack_geometric_1024(data)
if geometric:vectors,values=_lowrank_geometric352_eigh(data);return _certify_geometric_output(data,vectors,values,matrix_scale=route_stats[:,4])
if route_flags&4:
previous_precision=torch.get_float32_matmul_precision()
try:
try:vectors,values=_nearrank1024_eigh(data,low_precision_projector=False,low_precision_sign=True)
except torch.linalg.LinAlgError:exact_values,exact_vectors=torch.linalg.eigh(data);return exact_vectors,exact_values
vectors,values=_certify_eigh_output(data,vectors,values,((1,4,-4,4),),adaptive_gap_count=8,eigen_margin=.82,residual_probe_count=0,residual_probe_threshold=.0,matrix_scale=route_stats[:,4]);return vectors,values
finally:torch.set_float32_matmul_precision(previous_precision)
if route_flags&8:
try:vectors,values=_dense_1024_specialized(data,low_precision_sign=True,matrix_scale=route_stats[:,4],leaf_eigh_fn=_e507_dense1024_leaf_eigh_skip_newton,child_complete_qr=True,nonic_sign=True)
except torch.linalg.LinAlgError:vectors,values=_dense_1024_specialized(data,low_precision_sign=False,matrix_scale=route_stats[:,4],leaf_eigh_fn=_e507_dense1024_leaf_eigh_skip_newton,child_complete_qr=True)
return vectors,values
if route_flags&(16|32):vectors,values=_mixed_1024_repeated_partitioned(data,repeated_profile,route_stats,mixed_route_metadata);return _e2657_certify_n1024_output(data,vectors,values,route_stats[:,4],residual_probe=False,orthogonality_probe=True,selected_orthogonality=True)
values,vectors=torch.linalg.eigh(data);return vectors,values
if n!=512:values,vectors=torch.linalg.eigh(data);return vectors,values
batch,n,_=data.shape;route_flags,stats=_e154_fused_n512_route(data)
if route_flags&1:eye=torch.eye(n,device=data.device,dtype=torch.float32);values,order=data.diagonal(dim1=-2,dim2=-1).sort(dim=-1);vectors=eye.expand(batch,-1,-1).gather(2,order[:,None,:].expand(-1,n,-1));return vectors,values
if route_flags&(2|32|64)==32|64:vectors,values=_clustered512_eigh_qr(data,fast_active=True,cqr_fn=_e1529_clustered_cqr170);return vectors,values
heterogeneous=bool(route_flags&2)
if heterogeneous:
previous_precision=torch.get_float32_matmul_precision();torch.set_float32_matmul_precision('highest')
try:
if not route_flags&128:values,vectors=torch.linalg.eigh(data);return vectors,values
return _mixed512_eigh_scheduled(data,stats)
finally:torch.set_float32_matmul_precision(previous_precision)
if route_flags&4:return _e083_dense512(data,stats[:,4])
rankdef_like=bool(route_flags&8)
if rankdef_like:vectors,values=_rankdef_eigh_512_guarded(data);return _e204_certify_rankdef512_output(data,vectors,values,matrix_scale=stats[:,4])
lapack_even_like=bool(route_flags&16)
if lapack_even_like:
try:vectors,values=_lapack_even_guarded_fast(data);return _certify_eigh_output(data,vectors,values,((0,1,24,64),(1,4,-8,9),(3,4,-32,9)),matrix_scale=stats[:,4])
except torch.linalg.LinAlgError:values,vectors=torch.linalg.eigh(data);return vectors,values
if not route_flags&32:values,vectors=torch.linalg.eigh(data);return vectors,values
trace=stats[:,0];frobenius_squared=stats[:,1].clamp_min(1e-30);alpha=frobenius_squared/float(n);scale=alpha.clamp_min(1e-30).sqrt();negative_rank=.5*(n-trace/scale);clustered=(negative_rank-17e1).abs()<.25
if bool(clustered.all().item()):vectors,values=_clustered512_eigh_qr(data,fast_active=True,cqr_fn=_e1529_clustered_cqr170);return vectors,values
if bool(clustered.any().item()):values,vectors=torch.linalg.eigh(data);return vectors,values
values,vectors=torch.linalg.eigh(data);return vectors,values
__all__=['custom_kernel']
N =116
_rn ="r116"
_sn ="s116"
_rs =(N *(N +1 )//2 +2 *4 *N +N +32 +2 *4 )*4
_pr =False
@memo (maxsize =1 )
def _r116 ():
source =(_E185_N176_BLOCK4_REDUCE_SOURCE .replace ("constexpr int N=176,TRI=N*(N+1)/2,B=4;","constexpr int N=116,TRI=N*(N+1)/2,B=4;").replace (_E185_N176_BLOCK4_REDUCE_NAME ,_rn ).replace ("__launch_bounds__(768,1)","__launch_bounds__(128,1)").replace ("total=tid<24?scratch[tid]:0.f","total=tid<4?scratch[tid]:0.f").replace ("constexpr int G=4;","constexpr int G=1;").replace ("constexpr int PAIRS=16,ROWS=48;","constexpr int PAIRS=4,ROWS=32;"))
entry =" int mid=blockIdx.x,tid=threadIdx.x,warp=tid>>5,lane=tid&31;"
source =source .replace (" if(tid==0)timers[mid*5]=clock64();","",1 ).replace (entry ,entry +'\n if(tid==0){\n timers[mid]=0ULL;\n __threadfence();\n asm volatile("griddepcontrol.launch_dependents;":::);\n }',1 )
tail =" timers[mid*5+1]=clock64();\n }\n}"
source =source .replace (tail ," }\n __threadfence();\n __syncthreads();\n if(tid==0) timers[mid]=1ULL;\n}",1 )
return CUDAKernel (_fast_nvrtc_compile (source ,_rn ),_rn )
@memo (maxsize =1 )
def _s116 ():
source =_e1282_n128_solve_source (8 )
for old ,new in (("constexpr int N = 128;","constexpr int N = 116;"),("__launch_bounds__(128, 8)","__launch_bounds__(116, 8)"),(_E1282_N128_SOLVE_NAME ,_sn )):
if source .count (old )!=1 :raise RuntimeError ("n116 solve anchor changed")
source =source .replace (old ,new ,1 )
source =_e2333_small_leaf_coarse_sturm (source ,116 ,116 .bit_length ())
ready =" diagonal[eigen] = diagonal_input[(long long)batch * N + eigen];\n off[eigen] = off_input[(long long)batch * N + eigen];\n off_squared[eigen] = off[eigen] * off[eigen];\n __syncthreads();"
source =source .replace (ready ,ready +'\n if (eigen == 0)\n asm volatile("griddepcontrol.launch_dependents;" :::);',1 )
tail =" for (int row = 0; row < N; ++row)\n vectors[mb + (long long)row * N + eigen] *= inverse;\n}"
fused =''' for (int row = 0; row < N; ++row)
vectors[mb + (long long)row * N + eigen] *= inverse;
__syncthreads();
const int warp = eigen >> 5;
const int lane = eigen & 31;
#pragma unroll
for (int parity = 0; parity < 2; ++parity) {
if (warp < 3) {
for (int column = 1 + parity + 2 * warp;
column < N; column += 6) {
const float gap = values[(long long)batch * N + column]
- values[(long long)batch * N + column - 1];
if (gap < 0.005f) {
float dot = 0.0f;
#pragma unroll
for (int row = lane; row < N; row += 32)
dot = fmaf(vectors[mb + (long long)row * N + column - 1], vectors[mb + (long long)row * N + column], dot);
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
dot += __shfl_down_sync(0xffffffffu, dot, offset);
dot = __shfl_sync(0xffffffffu, dot, 0);
const float repair_inverse = rsqrtf(fmaxf(1.0f - dot * dot, 1.0e-12f));
#pragma unroll
for (int row = lane; row < N; row += 32)
vectors[mb + (long long)row * N + column] = (vectors[mb + (long long)row * N + column] - dot * vectors[mb + (long long)row * N + column - 1]) * repair_inverse;
}
}
}
__syncthreads();
}
}'''
if source .count (tail )!=1 :raise RuntimeError ("n116 solve tail changed")
return CUDAKernel (_fast_nvrtc_compile (source .replace (tail ,fused ,1 ),_sn ),_sn )
@torch .no_grad ()
def _n116 (data :torch .Tensor ):
batch =data .shape [0 ]
saved =torch .empty_like (data )
diagonal =torch .empty ((batch ,N ),device =data .device ,dtype =torch .float32 )
off =torch .empty_like (diagonal )
vectors =torch .empty_like (data )
values =torch .empty_like (diagonal )
workspace =torch .empty_like (data )
timers =torch .empty ((batch ,),device =data .device ,dtype =torch .int64 )
_r116 ().launch ((batch ,1 ,1 ),(128 ,1 ,1 ),(data ,saved ,diagonal ,off ,timers ),shared_mem =_rs ,)
_s116 ().launch_pdl ((batch ,1 ,1 ),(N ,1 ,1 ),(diagonal ,off ,vectors ,values ,workspace ,timers ),)
panel_total =(N -2 +31 )//32
triangular =torch .empty ((batch ,panel_total ,32 ,32 ),device =data .device ,dtype =data .dtype )
_householder_compact_t16_kernel [batch ,panel_total ](saved ,triangular ,N ,panel_total ,panel_width =32 ,block_rows =32 ,use_tf32 =True ,num_warps =1 ,num_stages =2 ,launch_pdl =True ,)
return _e1762_tensor_wy_apply_saved_ (vectors ,saved ,triangular ,block_cols =128 ),values
@triton .jit
def _q116 (rayleigh ,score ,BLOCK :tl .constexpr ):
row =tl .program_id (0 )
batch =tl .program_id (1 )
columns =tl .arange (0 ,BLOCK )
mask =columns <320
value =tl .load (rayleigh +batch *320 *320 +row *320 +columns ,mask =mask ,other =0.0 ,)
value =tl .where (columns ==row ,0.0 ,value )
absolute =tl .abs (value )
combined =tl .sqrt (tl .sum (value *value ,axis =0 ))+1.5 *tl .max (absolute ,axis =0 )
tl .store (score +batch *320 +row ,combined )
@triton .jit
def _p116 (rayleigh ,indices ,local ,BLOCK :tl .constexpr ):
batch =tl .program_id (0 )
offsets =tl .program_id (1 )*BLOCK +tl .arange (0 ,BLOCK )
mask =offsets <116 *116
row =offsets //116
column =offsets -row *116
source_row =tl .load (indices +batch *116 +row ,mask =mask ,other =0 )
source_column =tl .load (indices +batch *116 +column ,mask =mask ,other =0 )
forward =tl .load (rayleigh +batch *320 *320 +source_row *320 +source_column ,mask =mask ,other =0.0 ,)
reverse =tl .load (rayleigh +batch *320 *320 +source_column *320 +source_row ,mask =mask ,other =0.0 ,)
tl .store (local +batch *116 *116 +offsets ,0.5 *(forward +reverse ),mask =mask ,)
@triton .jit
def _f116 (rotation ,indices ,local_rotation ,output ,BLOCK_ROWS :tl .constexpr ,BLOCK_STATE :tl .constexpr ,FACTOR_PITCH :tl .constexpr ):
batch =tl .program_id (0 )
row_base =tl .program_id (1 )*BLOCK_ROWS
rows =row_base +tl .arange (0 ,BLOCK_ROWS )[:,None ]
state_columns =tl .arange (0 ,BLOCK_STATE )[None ,:]
state =tl .load (rotation +batch *320 *320 +rows *320 +state_columns ,mask =(rows <320 )&(state_columns <320 ),other =0.0 ,)
tl .store (output +batch *320 *320 +rows *320 +state_columns ,state ,mask =(rows <320 )&(state_columns <320 ),)
tl .debug_barrier ()
k =tl .arange (0 ,128 )
positions =tl .arange (0 ,128 )
active_k =k <116
active_positions =positions <116
selected_columns =tl .load (indices +batch *116 +k ,mask =active_k ,other =0 )
selected =tl .load (rotation +batch *320 *320 +rows *320 +selected_columns [None ,:],mask =(rows <320 )&active_k [None ,:],other =0.0 ,)
factor =tl .load (local_rotation +batch *FACTOR_PITCH *FACTOR_PITCH +k [:,None ]*FACTOR_PITCH +positions [None ,:],mask =active_k [:,None ]&active_positions [None ,:],other =0.0 ,)
transformed =tl .dot (selected ,factor ,input_precision ="tf32",out_dtype =tl .float32 )
destinations =tl .load (indices +batch *116 +positions ,mask =active_positions ,other =0 )
tl .store (output +batch *320 *320 +rows *320 +destinations [None ,:],transformed ,mask =(rows <320 )&active_positions [None ,:],)
@torch .no_grad ()
def _d116 (projected :torch .Tensor ):
batch =projected .shape [0 ]
initial_rotation ,values =_e196_lapack_n128_leaf_eigh (projected [:,:128 ,:128 ].contiguous (),fuse_mgs =True )
factors =[]
for current in (128 ,192 ,256 ):
cross =_e3200_dense512_factor_cross (projected ,initial_rotation ,factors ,current )
indices ,local =_e2860_dense512_stage_select_pack (projected ,cross ,values ,current )
local_rotation ,local_values =_e196_lapack_n128_leaf_eigh (local ,fuse_mgs =current ==128 ,compact_t_tf32 =True ,)
factors .append ((current ,indices ,local_rotation ,local_values ))
values =_e3200_dense512_update_values (values ,indices ,local_values ,current )
rotation =_e3200_dense512_materialize (initial_rotation ,factors )
rayleigh =rotation .mT @projected @rotation
diagonal =rayleigh .diagonal (dim1 =1 ,dim2 =2 )
score =torch .empty ((batch ,320 ),device =projected .device )
_q116 [320 ,batch ](rayleigh ,score ,BLOCK =512 ,num_warps =8 ,num_stages =1 )
indices =score .topk (N ,dim =1 ).indices .contiguous ()
local =torch .empty ((batch ,N ,N ),device =projected .device )
_p116 [batch ,triton .cdiv (N *N ,256 )](rayleigh ,indices ,local ,BLOCK =256 ,num_warps =8 ,num_stages =1 ,)
local_rotation ,local_values =_n116 (local )
factor_pitch =116
if _pr :
padded_rotation =torch .empty ((batch ,128 ,128 ),device =projected .device )
_pad_rotation116 [batch ,triton .cdiv (128 *128 ,256 )](local_rotation ,padded_rotation ,BLOCK =256 ,num_warps =8 ,num_stages =1 ,)
local_rotation =padded_rotation
factor_pitch =128
updated =torch .empty_like (rotation )
_f116 [batch ,20 ](rotation ,indices ,local_rotation ,updated ,BLOCK_ROWS =16 ,BLOCK_STATE =512 ,FACTOR_PITCH =factor_pitch ,num_warps =4 ,num_stages =1 ,)
values =diagonal .clone ()
values .scatter_ (1 ,indices ,local_values )
values ,order =values .sort (dim =1 )
return updated ,values ,order
_e2760_dense512_projected_eigh =_d116
scrolls · 3967 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