Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
6.92ms
#8 of 286
2026-07-14

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