Skip to content
KernelIndex
Search⌘K

submission 877172

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

cand_1152_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-877172?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
18.4ms
#31 of 286
2026-07-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d6624e3d7b8d655381ddf7e66bf17e4c0696f9400ca0bd0064f794bf7fbb970f
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-26

Techniques

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

cluster__global__ void __cluster_dims__(C, 1, 1)
mma"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
shared-memoryextern __shared__ float smem[];
vector-width = float4float4 v = *reinterpret_cast<const float4*>(&Hb[(size_t)(j0 + r) * n + (j0 + c)]);

Kernel source

cand_1152_submission.py4552 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200

import base64
import ctypes
import lzma

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

_TRSM170_CUBIN_XZ = b"""
/Td6WFoAAATm1rRGAgAhARYAAAB0L+Wj4bHvYhldAD+RRYRoPYmv2/G7y/KieK7m/PvfO/psG7IR
q5HCJJ7JRyAUR5ceivar+NbA+7NNX/dbsL4uFJ946ckTOq54VocX98aSN9ak98UmLtb2pkyJFDjI
QwjDe2SHAQJuoPzzdblOXSz86mClGh1N6wXo8R4eurDRNpq4yN4dIVxjDQO1GdPvpwGUzJMha8i2
PXzWQc6isLZdcTqxtt5DiLTXBkDrzyV2Wu7F0Zq/Xkt3LajqST1/ciokU52VQJoxTqSP9JIkuNCL
y9XUq7DKJn7iBrN1si5xLnIHfGyVUMCGRYuiigh7BStimzRFiED5rYGMERgwSwamNnTRE1ua1dDT
XIdT9NCuNzFeE1vPfRWHwZTYuBcQzuZSD7zA/zTJEQsj1bInch0dyNlsGbD5+z31PNlLcd7ix8si
+ibna+57TQUMMRsrfF5Yhf9H8xgm+w1fUq4aad7Ze+m8UAFt08t53dYYHoRXslY+GpZ/kVkq8/FH
q8f+UHTsPpp+3rfUo2akMfNS5YCXHMpa8iUs/E3UlorJRA4HcUxBC/TZlUGW6l2gsMqFlktKiKqo
8de4uroHMOAm727QQPkuE0Wuco73xnaV4cBO8yr56rkJfWKeNzCWEByXMeBL8t/WatbV2rQqMqCK
SJZFEUbXCZyCWHjF9aKNCm8d+RlwEdfUm8CSRGhjr6Z8xypx7YWvuOWnEenFPtkZNHnQMR03jnBy
zLTfPXP/p1DqH2t2ijXEEjmJVDG5FL8uMiZey8l3dkljYQFeFj5rQE/78kAgNLKax8bA6I3j/NbN
BovlUw5HcqqUr4G3iU9gVdsHHV53VxMCeazp49cQw9yPVn3pinOvSqtdTun0M9+uZATI7QK1+y3x
xibhfDa5ihiAuFVasxpWKUuC7hhuXGnK9HQ8f/FJmAkcK9v9CpOfVq+cSHK1nTAutxctzIFzucCb
9Yt1n7vJ7JVy+syjr2evdCyFsQsSeORafLWBTBs7pBjJV89Gtwa7DspJ3Zi0188GC0HzUtPLK0pY
91fS0QBiRE1dfS0BgLEERD2wk3NgLmTXF/2OCWA2N+LzDYQjO/2tc4p5c1xx+ph7pbHNuuMT68vP
3DlX6jw96+Q+n6mUfmwdoXvrsiN/RFD3yaEWLv5uQsAU0t1b0LHbh768jQY3IFhDON4vxnqvGvnG
Vt7D/P8FtUO8tNWulY4I9TkLdSlxdSZ1RsB9CJrcIA5HYCw6xHTKjrzvEGloTkuan3CTIXwWKfso
ohG8SyspgMIqVJL5Z6tHDlfiNExwSElzusi0tYV8ZOXDwiMiSMFAXSAcR/tQdYWdkq79x9c9VvfT
l+bXOJWAFMGPsU1B3I1EZJBCAiviWHalIi7u8tHcs7TujkJXFNlt+xQQMwH6v/Q1lMiO2rDvPW5Q
X+NFWqjJEtk54iOqiAz4wJ/TutBxNzChKlgx+V/mJWlD46J0VbClfzP3bCdFVuRtwYtjUk2/xgbv
0bRQVtyyz9adop8w3rXEzgjID1cRpzkI+jlSovTQ/T900ksz8dU2GmJQlAescgJaigBGnsTtOtKd
IfIOzHQPro+KBNn25u+vBP5alUhOxjlbqY7xoa0Qu+S48CfBHttIfOJWgMNKnYn8PN7qEEmi+d4I
zEgbu19cptow85i9bVGaW16ELmGoO/+0ta1lHfs0MXRfCfVgbeQp47dvD4mgW6ajAnKQHxGlEDN3
CSgrWr3gmLIwNSbpA4xkp+X7eJ+dJGUEgUlsBjcSVq9/B4xTHFcK6MQO9DjS/GcrqlC18El+Rlxk
6M7XCpfKlmZHex1JjKH7BoXE91Jqt4EMbqmBUFiqH7b0Sy5FAOf5ab7E1nZG4aD/k4pJKdQmxsZD
OEo3pSPzd3Clsreo6SErVZxkiGLb2ThvBj8b+FqGTD7FzJrHUhwsipnUQ6Bdm9fmWAW3jy2hVoe2
Q+mNUgx67JM6BwvPG/VZycuNJXpYerqgFbG4Yb5JFmb649wAGQFEbnU6fQkpJOFd+L/7fGUVdYtZ
V4HkWwvyNF5m0Qc9TASCPEVm4klYDw3TK7lsirNqfFLkFiJQvgKF78bj1UmgT+pEgpkSlzMEcmOt
chiwL6jKttg3SlzCspg8DWRizskmO5ytaN3C5KvEOe+Dmd8gpQw2quvMiuGVP6/BhUsmUvIxvsol
bcba1tAqbF3KsL9PeCdhbenBM8r1KYy2ksuWQOl88RUbwsfHLY7Wyxfv/0hYmnRNVEz+RbeyNKDW
FZrZbebpjHXzSieR6kQSssfG+wj+nN1ivfX/30gldXWesMba7gRqLwhDEopb5h1/r3XLzlKI2DQN
5INolzSzG6EzwFQ6KBMrTlN0D42fkwMJYGdqoK4LQUHb7S62pHTTRAxmJrDropQHONKKDeWKjJ99
tdDu4XzjVpyUB0U3cExD6jplDuDjUkdIKq6hOQmwLsxq2ulNWrR7GXp4dFGhd7NKUvo32ur3+LzC
cMx8eS1DdkCzrpx3300unfixTATDWMmrYW/czSap0yQcM377XKlZ68sIPq8oh+rqd4MW/wOb1MTt
ZSJ+uq/RiSwoGsBA6uyrUFA2DHF/dWf6GtoCXk0LemPz5vbVHlDpJrXh4f3mfv94z+/Hzz97iDbX
3t/iSEhx44rJQpNIVXJeNVLrD6Nng7PxNLybqSV3HEybv+llYGiCnTmNwMa2VoFE0FbWdsMmlCA1
6Eg/Y2vOlq8eQAtn/dAK27NVFgdGkMoLH5+0z1cBkb8ImMgRlregSSsNkgnHK0j3XpjC67bTix+s
WiYbzCog82KPB0CjAlfeJyrodrAtQb1Cks/z132e/mSmo98MIHo7FqUwYjX2XHoSnfaSO8bZOv3w
P5L7+wXYnsapFPm16PCntDVjV8BiNvzNAJ/x0R36szNs6GXKVKyQdPTjj2vlYIi+B09Jy2ttstLv
fK/3jXIJPt1ch9jF11i6FbmVquLIoEVOfskBFbY/eEcReeI8xg7sHZU8kHJDXpLK8I/cEo3TXKHk
hkbqwjiMJKquQ1ozy56fcZJiEO9HBKejRm3chrggJuqoEIBK4oZa7KeaV91+38zuFzq7cJVVeQOP
Dm5gz8sZm/ZESlg1JptK1xSqsEC64JCCqwGAVfS3ZDOrabuPaFfk2ES+g1XrP+nNJZP1K/U2hKQi
SXeCiB7XJTEX1tDlE13K41kYZjAwn8bvVOsqS0Y96TiEu1j3GTOok9i7ggk4ACsgzzM56GC570yj
da2YzDVGpShMKxc8XzM8EqWxFooaK3bwbaaWP13prbDRcsmnOo2tArJE42W2JR4WrxAfJlOpNvUI
Mk75dKlp/SWR7wN9bNL8aQ0p747NaXPsejScxMGmPR31QhxNVFB1vpDAax3kISF2FXXSXAr5oLhT
DVa/Pu1Mz3M3C4r+AIsecO7rmHMaGKuxGxgWBYWpRsA10Zo5K3UwwgAhRIiiuY8HxKvRk22ROH33
dbybKijudlX1x6yo2srRAvmU+Oa4CzKFvmFVX0f7363/kov1Ng2WOcsGtQFUdaYr0YLsGXar2sY7
HVPGwWkpFzleW4f5nNWQkxEz+BEwxaN6qXCqHUIYkOIq+qGBpD3drwxyTsHcJf9Y7kXwODCa3krs
6PEpu8qwtjoVJkHpV5GMcNETmS8TDTKhKe/ZkqM+Hu5IofO+kHcJObrj/W8TAE1gZLK8baItmnmU
7bBmYYSJbZFWUEhZ3fPGOCjXhkmWZJGsi8PFX0ZoAvna6RNKPGw8kD3eRRYLJpEHuho5LIc56Zxs
x5VyhqRdmO1BXaM6gwLVu5pG67nrnRncXNiUAB2YhYbfwMDAIG859HSWc3A4ha6ThbC9pSQBzIq4
eiboo4cWt/5fBDEJ5jZuyiwNwS3PA22kmQRsNHTj6k1+1aA0hAZOYf3/+JeO44Q8MKTpal6E8ac1
FTS07RHeOQA8VoLYE6JUdQruHbxX4GrFs+ugQHV6CwNUyEJe5SEJkI2Ov0mBRqVIs0p1X8rz3t+u
yx25zX6GzqRe2NwFQR32BXn4RxlA8WIh/ghxAnYgb+G+wPUnKlcsi/WlBnQwqZCp4hZzhF6uP0w6
zUrxaqWvKU/sPzJSq8+azCtgPHtidJPlNF3Dvwv/chYzD6CRhcRdpv4BsfcINhOce4VRE0cTrdwk
G9JJuOTk6vLfavDCkiQZfBGbr2H8zuagYYBxNfy712ILFekt3jauIvZr+bv9NQIP3Out06gAlRTD
vLBLOPlzUO5/YeBdZNB5ZvM08o0rV8hksOFzxnGeQtbew5UlEBsXLn5Gd8XoxEObkA+/0s2zoGRx
E4k0290Udj7RwVqBvZEEkgU47GFlewvH2RhPMSq2fQe973acUOg/PHGo67YdiVTaQAmcpuPpmKyX
OvGti2mg2qA4iPo1In7PZZAGXlB5MBXOZgwRDB7O1rrN4DBZzUslFHf04aC7m0Vg1YY4x8VYsRkf
EDEsv/xCcGuc/mshuvOvmsyL6ILRq/AYqMTqO4H2XbtkVyx1XC6t+WuAVPOEr7n4U0etLCRhWn2j
vJmsOWviI8gJ8Vm4BHs091Ub1iT/mg192yu/qDi85h55wDFmkwNNB6DEogqJ+POfXswkT1M7jz4w
qKTEHwGDRBx1LV74SSE2Wb0x2zhYes3lnqu1taHIe56YLxrxvt1vdr6mbB5nSShIKTrF7rthV9bI
KC/cs2pyUxDDQZwC42sojzvq8jzX8aZncZXMTBeHkl/BAwD1zu0eUt83B5wxHduZyH9CI8K9H7bx
OvVLkjX6tZctAldf+4Ie+vNDFGrmO/yJZo1W/Rm0rQlN1F9O+Sat3h0lnp+s08Mbi/dLa8/RQxru
cMzTtVhaIvkE1tqECVj7mjdFTtqJt28ZYTrpUxeAX0K6l8SZqlzNIpU65s5iLJwUTKO2wCvMgv8z
8jZjBh8HZ0efNJnVPjOoimSl1AR02DnaRj3SZkXG4X4puwhripyfGitlnV6ohllEavUFMzL+QtVD
5vdVMFwZu2JK/+FWbQpfw1HL7rtwveTwwljbNmrQ7ayFtiC3wZPEn+PHi2htJJkd4tsoR9E0Lpf5
9FHVEg0zpNSNmgno9vBvOEVj1OxDFNM42hRzrjZClQD4jsE33YS7XOlFVIXEcpOq2h39fd/SBz+9
IF0NaEMgi5LTwkKWWi4UleD0jc5A1iYnz9tflN4K+eIKESSYqkWxQpNTYzyQewI1BrJe08j8RY2y
0tHpULYqSDetxtj7AIlp/p/VLGvCHxHH5D/VahYF66t0Oojn/iQMV9t0XDx9/HDSZx4yWIoKC1P5
jE2fqbkx22NnkY73z/2Vmwy3WTzSCKplPxZjSYcr74kH1DGyM5xdK2QTDLvFJ0DxtALE9Ot/SMx3
Hk09tiF2giMCkN4ucjdBBCY2yLf3xlnWRDj8F8sxqzRHtWU+QVzy9gYl4Cj88BEZiXsNlhBzxPT9
NPodxrqsDjM4yDsLAkSeSpQs2uWye3S2olnQVdW/p54T9gJIRNnrIEUGCCqJJrAB19ofl4+Ri0sM
0b4crKeNs/ZBA6e0bd3/F2h1Qtw4Xg7qPOGw8ZpImJVbQj3OxlEBKVqPK5FKjHIq1aPy1GEcl0vz
RhFDtHyLDvFUqOwre0RAa4+8yJtOP6a5KMmHpBtIAhaTL+E7/ADUtm5Enzfwg8Hnh4Ycu8Z5E+kp
ELPiRcs7HI0cGTQy1s8RJh5/W4neBUydm19rQXkQsogeHQKKTHiMlkmps0Oht5trAW/EbeyV+XIH
Z+5Z0tQfT+2/4YAzC04GWK3XwfBo1kbSEPfDR6nDlMBTH/wxnQHPD8PHPIOm6f01rzkrjXnEomKT
edR3hnE0vm/18K58w95DqA+RBYPFuTfyotfnoVvQhZ9KFLv073LN2VMXHIBaOqF679FhpdMyRjDz
cfkKv/bvOK144abeYxlEvpRQxq2tiQlScWTbOr5e/wA6tQJHAKXqKo9gtryqowFgcGdBWnv+3gx1
XrcAO0kPRwXUaBmpAtfigctScsLSXQUlxw9laIRHOXNAFwTooB+kSJV3dGkURS2TBHGHm/vb3bZ3
op2eGyFcHjUbUxJ6uue1RfZzGbAydsjdcHgSxG6DkqK9Ki+NdPlMcKwndG0/lAFRXoAUmYQ33rri
Wt+bJcpRvGLxjt4pVAGYPHkKF7jDs8Pv5ObqEWpti82zzPxl7ekJ7ax6GiuZG+wZhv6Nch58nuxr
36t8iNTts+0wFr9Q/L0xczE801vOMctPVWqNy3ct/M5X6sNjWZ9u5X0OzHMXwbzgfDtfuAyz3vCU
YXb7pfNC3Ymv6RiIQlAUxx+PxeDWG+l/SPArQkKA+IYKovWUcdv1fXTJhVansq9qKs+bE4V8ybsb
pX8SOxaVDzi0WYbYp6yVnmqamXKYtoFmqny2qRRCcbNI+XStLB62NNcQroJtUJaAU/VdG74+Vq8P
vS6Zr6MUOiVDOzA4BGn+V0CLD+qFzNaZOj+HVRBWzs2w2Vf9xIqiXfL66GpKCVbdJKFIkR42kgKa
HR07DfVMC9wFiPw3OtRgEluxqrP84iZsaXiwIzTGKVOnTRaRh2jlCaDetiypWia7Jc41zjg2y3XU
faD6j+S2ERJxgNx3PEr2Fe3C0BhVG6xfTlSS/na4A1G218xWV0ah6c7WidPVcqjz8hxnQ+HxWjzC
YIIMYPUBbaD8UfpFhZ5rXymz2XP/DxM98iWTp/50ba0e5ED+0B6a1iBt/A/l7FU4CgZTYoIoy/u2
xhyd4bjZxbxASZoUw/YLp5T4K7rgfajUbSr8Fmj8M1reHaNoaqaXT5TBzUJCc7JCoQXPzecMVmPn
hxaVp0A1KhVDwj8Ws8EwDykpc6F+PwbLlCGuNPsJleo1tsduaK+op3zeQpj+mieuiR54z0cth4ws
IEGhCcplwjg672E07wxSwjQ1eZtPdrmjdKfLEl1ujqFNPu6GvUsFp4LgmtA0C66RuD6nuVK/oayS
7CCwZ8iYMfwN0oFEvnDgR6ocC4JZlKtnpf4zu8Zn/2ffroMU3iZYlql96/AVl1qbpVT4KWTfHmAD
WpW1V6mkg/I9HRY4iC2nwVRnLGzbC0mJNmGkomDHBeuJpSTSHNnXMfdgPLRu03voPDWcFi8BIyHi
+J0XYpjSmukGUYApajH8yOzpzifBvs+ji6/YLpgDCSPKGHHkncBIPaA5v/UF6B1JlbJZ9J8bCjTb
PeqwX6KKM1lzllxZg69S9pvB46BzAqS4e4P9BQgbegh3T//j7dKHn3wpwt0PmsBeMda9EcFfsSxX
uupJLtnikxGFCmol+GvvPtbvtR1U0h8JDCzFJew49p4ShV1i3pkLw97QeJ20utd70t+qzA1Fpvoq
hLYx+tu+HLmN2aDhFNjA8s+XRZGs0ChRc/mT842B5TQSr1iRAZU1P2XjldZ827ZnOHhI47w0emsS
2YLo2FSM+UkHlqTunGAIfupIx577f2bihysbq8kXVvN2wegbWjFFd2hxskxSSpuCV4GXK5n0PIxA
z5WldGG3MQWOKVzueEgElsMfdsyDGaZXV6jDlEMgGzeh+DjYKxFrYVxkd0Nyh2mTutqvDjcVzZBc
gubOQJzAGDvJ0f1XTX8hTWXKYuyFF9aZv5JvIfOXWjKI8dg06p1viNIy6kmX1h+rC1lupuJjWXC/
jbp0auRztvnqpx9wnulrhDoYFBYoVRFSXHVm7CEGhsH/gDCPQsUWk8TQwvDTmfkVBmB7KH6TBQDe
UDuhikx4l9RE9651gvaP3+4aA0NHi50zcUlc2jYURR+O7lANP79CMPEZ/5NwFtAzzm2jcFkxbC00
n3/kVdRKahuW/p6Rx4dnCIwGYSIxcy2X3ciih0zBsVAe6y50gBskVK2c1+FfiX2l8exWIV+ywlFK
s18RDesMuKsMVYsqCiu268r4RR1CGp31f3Iju+vT/TyHNA2yLlSAbUO+5g96ejrgNjAJMNFQXAYo
7EohxwffUonl7aGudy7xF7GphdiHlQJpFpu8UhHqRAhc6u1cEH4ezWQo3++JbkcXkOjXqa9YhepV
31QE4gGB8JrN+E/15MC/D8whtJG0QkPpVBYrDanFKDVv//3cCqzkxAiTmXAr4quk++2JY7DX4Vk0
HfUBnTW4kepOevmyFrzwEaw4DN94djTm/fneI+TvWOwYzodLcM9CJ26kNmLNyf3CvcKkny1xwBzm
T+OfX0u9Z2F22WLXPAHVS/8X7Sjt2KTyO2b1sUnMdY74DU3SmhCMi35igNmEG0eIClPyrrs8WEDD
dq+Kz67aHYbbPCwhBHqc75a9c3FwwOnlUFUKlANDGbtT1jnVuVQd/fmXhQRDU3zAjlwYBFv4sgYe
HdHzCpxax3LxnfR89OZBD5m78elxDJCYkb1OZwk2IlxDZSpjpLxz/8VXhuqQuB4JrIBzKUVf14P+
71zgMo6XadxOERhVI+RfxOhHCPxTqlaFIauMoBbX10xPpe2pipk7f1CgcAY1LbU/BsYX/2iGvh0I
Byygln1R2SI9SrUsc57/zp8RUpqc52HZZIkUZBAJuNzh5Mxvx7Y241RxehtTsuVy/dJiQqRezvaM
rOZu5FlyZBQQnhBHUxv7mUB6l39zknUtUc9DCfRjXo2RaC7ZDnAA/mNbiYYvxz21oVTBGuBxp1Jg
+4CCSnLldvvCJTrYGSPnfnOKrrPm+Vdajtjmqab/oLFO94U26jRZxCSWEHBdeAK5Tk1rGcRuJRQE
/930qaC8QmPgkBZUxb9M/+y1GxGPo9OcsEUKiL8NKbpn7hq4nAjfiDu2+VFmcfGOlKu5VBjXLekk
J8KFgc8FtKkWColcngN1dQrQ5+dQ/GhJJojMBOaBgp0obRlABv7cH80hvDEFUeczweXUlLh64TeJ
6v/RtZg27TTroK91yniNHwU7yVx7TzH5xzD6rihau54Ls87Xi4WSuJmQ2dBoub1gWyJfog6HtkFv
On6cwprXMw5JgnfZV9szPHMl7NIbmI2anSVTQMx5kFG8nY4rmJ5f7I65lTema8lTD+qgh4PknDOf
jZW0ElCoJW8fCzm1fk5sDhJAcfcOoSNAZQ+qpsnjx5m/Uj6naBcmLZga5UsQ9UCL2iDDiPGRw6qu
r0sbAtM01EHoUa3bLtNSpc/HhCUgm3VuG4KPcK1xhBc4Q49XS3phNJCG2SOlCJ1C1G0l40qbRVHY
MmtTUmuMQ9DzMLiI9OR+jrv0F20j5prnLpRsZwwDNfMeFzqPgrsI9d+pgGvRNwX18T8DVnYhbkM0
aG5w99JMR8Rxc91TyXu7iKCLRc6f935qKTBAiicagZi8OoKq/sZdK1YUELXfpcK1DPZshN2c8Dqn
O9fY3aJgZo7fOhD01OA+OJVrTZdGXCsuzWLcStWHz+fMuaiSYu1CB/lx02iPZ3otCTuh2XEfqijT
Hmv26eLR+CVSosiKuHQRxvenk3Bnm8aLCljFCe7nXu4UE665GO/oy3IRyhqt7LqmGEhsifBXLB1x
Eh0qXZ0mEI98rIlxtDUpDb59tX1dl22OBMYflhnOhQJyflU/g/HW5JYGG6sEp3tFKFfk7vLxkth4
Idu6JQfQDwewkpge2EWFlSf0tAy+p9eO63JgQ235Up5gAwa/4QN61dXj4gS2zqiQUaEYLxkYLsOb
bxzHdSuwH4J9hnZLUewGzcRpG68Hq6SZey11TdJe/PQOOBKghZ8kL06IIVTxhOnxJ4RHNAWJmiP4
ILvOnLzkQy9DkS+MVa552LhJTaJuRpOmMPH8HpgPeHO8iAAksagBxFVOEnFa4Rca4+e3OMCu5IQB
ldkCi15BFJfFDdnjRvkr3yMD5ybW1q2beX+S9Et9J0/0A2+Q1EBj/qQpVjJB/GWEdb3ZOIrolJme
Vw4p3KDtP8jql0/qgB3Lqbmo0cwqP19iZW9qNRUK9XrG07A76mgjmsUnHQUl4vw7dUAbzVd95Xbq
RK4HtWu+6RezPi23m1dWL5ciV2xJmolXKU4mXxzLWVNNSBST64podhzy06Rz1OjuIv+g7412yEDG
laGzrgTIPWEwkUR3foLLn7Ihhp006ea50I4DmUnWt4XB5t+bOm0A53LpyZ9npkkIIVys4K0lHGwr
AlErzOCwK7jQis2jPPEgFFjzsk4295WRTGeQtepx/+MQ4F8SHZ991BzCoNTgvkzr7Ro3Bv0oKMgD
6e54fBt0r/T6C/cdKUYbxcyeYSdCXImOfzY3qH65lr5eLfre7MhzCo1JLVVw+yUZBGtMSepPypE8
gtF9ROdH6KizyGlr8yFEVZwHLBxl5H5xsfKumHCeOCeh+bpfs7ainEFPIq/U+wnyETG1StnddO4I
Bnl15+phsGnLLCYttb3vJdazzxARPuDgKhsCUXyjLTvRsZt7mtCKR4tIdQ270gkGiPpD2u8hroc/
zcnleBKnmjkmstm1zpQOd0rx7Lb7cIxCihyQ1vK1pE32VlVBWc5pilYUI6ScZxfQ4y29M/z2tMB6
X6IrHL/iFfg3C8rjrRTJprFPEzBXEudp45p2sc/Q/v8GDBb3ElGjc+X2QMVMm3Gskh+Ey3ECWicj
cP8uhpEmNM0TiDjJthF0vRi8S5oO2Y3dyKt3oM/wETpIByvtHU08t0sn5a7rgNIjYdkXCJLhcJGx
uM70bo/0lRTp/eLkf5vxtykxsWV9YNT3RbdDjMVn/6QqZ0P5skaQXvoCs8pJLUBPsdq8jmO8JjT0
zZtkJXd+dUKEcDcALQh/VNGsc8Gm1JSa7m6FCFA9d6dY2CvzV9LZayHKUx+ce+Zw0Twp//2haDBL
Lbsjooj1j0THmbZQNWw1uH6cC++FeBx3Ok8nhyVxSy74T5QL8wpjkAP8ZXDzPYh48lVxesVQutUM
9WXiWrNwmfzJAbV8XpfTtHIlfV1GB9cZfGXgr9ADtWbSLv5S9CrTbAWnjvrmO3V9Z0RD+I38I55L
A5F5r1D4BYcSjj0Y24d4e78YjJUPeXxnPuzvRjoLKuDJQZVNsDFVJJv4RSQ4nKeLMpfASy2mr0Et
56PE1l1a3VNkfrepWyekw4uvquDCXS07XsnjsfnNT1Uqp8eCrvs+claEAmNeFeQMtrL93u5HTOew
4FtXhE1C1l+GSCThK3rFSxmG1lJxI3gkjdJqHwcOGFjnb76Y7gtsICYpqu2wS+TqeMKc4QxFyrmq
L3n2gY4k+KF+93j/0KMhu0ZMRNAzdpbtuJqQAs/PrOX4XJX4fcYrvAHHnhN8cvMx3eEOsFdW8Ml6
Y+VETU6LdqnnjXTFl6kHHctZBkUb5xERLxRQFOPc7wRzfu9QtQR6fRFP6ojQ97D9uQfSM6HoyjMj
l6FIxtS1DWfZgFEDNRPBX5rSbce3B1lbBz0MZ27Op2sNoi8r1p0Iep+XqvN+JTnPbZFnoa2FxNWx
juzpJB/OYTw5wpHJwPsz97iinzJjWDbk00cP00A2YdAmxQxyV7nA75DGQuPvHd2DPUv804Zau6as
Ka48hOldRYSQFXXzdqqbs98+Cguf8Wp/SKQjoHW0N5nPZTjx6/cxork0eYTRSAMQWzEubvPtk7mA
1YZ4HEhIbNs0i/k1ilhpRx2HL9E3cr12CU3uo3+f7ri/0wt7ZugIs5C+zYU+Mw3x6x2vHf3r8uzr
wHhnHtoEKOnyCxHF5ZP2x5ryyWTXhVvB+gHruGF7rAjA0DjauO1fqSshutELAqxcR4ORykMzAybq
4t+nvSTqYY9RN2irRJvU12Sz6FqGYjKmx3fxHxCCWwP717ZMM9F5DxB3xZf5r0FCo0G1y5n/Ja8w
FJq3hYGYl1PJeSyeFNyHb/95QQpF5ZC/3gIneWUecnDvIGh5ZHdGHnhfANk7SftDgvs1gvs/OeHf
+ziLzmSWIzDuqyNX3JIkJ7k7shXARHvoycMBfKLoBSZqPXuvu2IC5kEDAskVUDYmKpn73N6AM/1t
AgP7HeuVwpud74QbxFI6yKUIOnvaIj/wg01vdExjuXy0FBZB9HZfH/lclgAJ8Ph2xoVsFXrp3m4r
KAldFx30yj3z/sH45B9k4bq7fDY0fJ8LJKsjOh9gsfOsTv7jgGzYcbniYLqfAYsEVCcOpa7aliHb
UIZSur150iV5lnQ/yDV1J+FdcWA2/yRVytJPwhqXoiXlfJS+4kcCE6hhgjygBeq0XwhcQ2mbanZr
DWbH2WxTC29OeldnT6BPScOWVxVoh8MegWYSYvLBersYY4mohM+jF3vdlkp+/GaHguDwzGItCt2x
2Tk/LHj+Abtsw/Tbe9jAeYDdn8FicwBsMwLC9eYGvSYkPDYguYTKRgpfWJSs0NYzRXMBAKeHuJDH
jZMcLA4zvUZOXVCXCLxCv//U97n+9ba1KQLVvcMAgBgROEPQljViPJu5fIKFrAw9JKsIKdS08xke
6jaJslcczX4SFXA2TIZsZs8QDCt7vjojnh/ukvLJqWO5/iyi/t+C57dZPDFRUOIii/KUkv3r9inx
B3++MVxhV5/CHFA4AzcmtxPdIjREvK8DsuHonFyUeut2NXZvrFvR90rDhnzGLFMEQz0hFlXNjTEn
2OlT8ghAKlYf1vkKtuKfoIN6YnwTlq2mdFRRGuues5M36s7NQZ+sWWwjS5K4k6eGM1JnXTRmxp+1
UYUi1xvT6/ZMM5ZWYVbLSEFiN8mtvo/a4g1A7rXBPvhQoxIMsThdfw6YJNqPghBVG+oBqPuoS9Jn
0a4OmdPrQj5FcZ1bpaCVRgHHJeCrlrGg9Y0fw8qDoZ5dXRRBGycovULYspnOmUU9s1L0kOPbkzvu
PRHP03NTqLqGeAJCcsIWAdiYsj6n8XSr9dDinh0ZWVS+DSZ60EyCznfTEZpS5EWHcHtyAmErOHZ2
hnuWXW4CRcR66orlXMFl/R0RgJsL2hj8+3U44fL78TP0coHM78FYVfg7DMw3eHky6Nl34tHjCqHt
AK47mxAGnDcTFRFVzv/98YAIhyiBrVnWrEC9MJWV0FyUZreDx562FrbvhqaACPRT21fKI74r0F9I
XgYwlLvxxEMbzwzyJugc2++SC3/TXHfALqgJudewHgNEnrdXs/pPVIcadd77UNDSjgOb/JS85BSn
FnOwd5cxK7Vxe9r0s/+V2HoB78jpVZNSINyAjPbjxXS/zgCoInOi14FE+mfksBEs9dHv5ocFrbTj
rh6pyMeIoqzCOqtfRq9+Oj+kX+r0c665RW6a8D/B0i1/GEkjg32qTu9JHWeSEFbrPBe/E5tMT+6h
vVb3H/D4MP1SzoLLbOAmrvW8xpVk9sCJ/dcRk+urL9umZhROoBFKdSlrxRPNztiwGCyVYLorItdu
qsaS5f/g6txbk2U/Ow1YGesZDa99r4BVY6cgXxOcB7aaO/3T9mEgzF7C3WmMtQw1xo3q+Hv8yhbD
HzCWgW0+7SPcKYSM3/hWOBtCnGVBJA1pUh/CUmWWJwmKlP9ag4KdG4vxLp9oZXI+KBgE5fr9JlMG
QoN97FcuZgOf4I+0X59wZwKq6sl/bVoPVi66EL0b4AOH6eF7Bz92WnJo9SPNyZIYm+DVHiUoIINK
BISiKieceejvytEvhfZ0JF6+W+bKXZEMacHa5jgIgep+/D++u/wYodPxuqNJtWJwCHzW0AWbpa4B
MelHvFOLiidYQneMJSo4iVdPZ5+wDPDgV17bvo3H2hU7ZNrsy6X+Al2G2mGoVCzt/LEMMEVe2mz5
TAqs12x1AScb86+uGXL474shcLEkx6MieOD5ZDkBRl5ARSO+UmWrog7CxFTaphiJ4sTlWc3nJtpO
MMXlGmXXJHLZiStgjMw6QNLk56Rweb8aG+oX/u9rhDsQuu81sA7hWSD8VXaSoxnJRcVxVDE6Dwr6
9A2QLe1kEN/sGoPGZkFyZ8EmtaXwQl2JfAusuJ7o2qDQeHWSBiiYhM2+b3hwPO6vnlkdD37NcNGo
6UvHnMqRlEoR4tSCVB+OLFtOvuWHC69Jh4/Yjh+G0KUkPhCs0b7j1z0rAWTSOqHt+QVQirwy8xi5
dUQXdi5l3E01URz0YCLbN8htcRAOWh18PpodvQP3xggjYs5kQGVQsuO6ORBhX775xo3tjD7eEwWq
n37mj2ZvqS5srusBfSsH471Q5JR6F9lpsIJiRDWk5eusq+umvqNkuN7Yje8ruFuC+NqsCEehYlSQ
L+2VJNhyD0TwHypdoGYY7LQAw6Tc1+AWKi7/F48ym452o5hEEt63FrmyHOp4okPQLSILBFYbIv92
LuVdBaYxIKpbOwL+F0oyfaEYIrNC108CH5u0rXtsdJ+niXfnWivzLj0wKfR8Jwy1/d2xnzawLf8i
rxDZcn49QEJqEQEs5RSRDMWPLU4u0e3td3i+/wygdNLzVxaOSPIGQ6nMvk2jeuo70z5IWvUNukc4
Y7S3G9LoF497c8Ph7ktTXd0A7ibLbgN/mNm/4B6G+ZPjqn/bsGX3+aCciijY5DIlGNiGtkbLnL7Q
Jm+NB2I47h3E3d2baDWgp0WLczbFsAVu6duO8xKxxLswnItBdf/z+tH0vclbZ+IanBN/JnXg3lc2
r0AvLOT/nmScxzr6WPOQ0VaJMwB/i611aj1BKDnBe4KpZYBhiGfJQDLnNGJlqzEb9fk2R8NkQz0m
HzYRb6kiD2fDthK1EBdqmWtX66NoRKOG2MCKDRbpIZJWebliRFfz/1HBaShFOQqVPpAgXbiPiNz0
hyPKChgAMec2Uheidf3Kdpk6EYxFVl5J06hNwnssIgQMmCXFcYbiq8Jc7V9cj4hAp8Lxl41Y8sXk
DGkh2TOtWGIYO1YT7ki+jLhtYjBbnZBpWyFFBPP2EhAGOoGrABtUCLYg7xpc7t2oDfTbxjLwz/qK
00JX4Kq/XvYz8qUok+xonw2gc4AKOk8l2yxQSjWl3NnVuwUrbx9gJa8PuL0a4Wkw4lCnHxiW3oPB
ieUEtM2eh4ohCcA1/KSWrPC/nbI6slRX29jvMR5VOH1oyGdkbKIoIY+RABX7JjN5bH7bwUVALk4g
ISvyqqY7k51CE/5qd3/iisnTGlRfR/qbJRPAee0/3F2cEzpO1JEaw1kAMu8olokqWZmnyrJdFtmZ
WOk/kgmZ76AfHR6y5zqH9Cbb7BfPn0+yFDeE4+VYAl8ppnRhgiyL4PCTkMsgowIMPELq5XncIrJg
whk5IxVsdNODjWziYf2CDE/u890Xo8uCn+47JSAhlHJHoP4YARremuDI2VDJ6/TxqCk1zMWmieMP
0yZjEPbCKvhSl3fXm8BNI+o5TL7vJogiNtwlPrN/S56q7fHoQ3HbfshBOfsr1kK7ec/lJ17hWu7x
qI2PFjmHQIAQWGUrGtm603FDLsW577TFALPh9DSJZP/k50jmwg049ZYPXEx60xBiCYoIIYsoGid2
JdZJEe3gP13HncA9AUO3F8nNCMpfUIWyEmEAY132W5rBeOoa7UTUmkPximQuEaAsnQUYrvU2mIfS
Kt8oQ+GiTU7NA0fGM+G71dWwYBqJIZETfWEfPWGfRLy4wIUjg/h2w5GrjBUBvOuRmNOcpOu7E2uw
jTHooswq34QNfvoRUVB2nIXvDpxMG+cYmYMyqtXFNu7UdVx54wUIvHrcX6DjG5jKoXCVNZcZh5Ju
IE1InOZdg0uqmHet7bV+VhRy6E5kmfcRSC5Af6RfW3o5a7RBe3qpEjvlY7qqKvgrpeX3YLH5rxpH
CAOtxzzXDFU7D4byfiK8VhlS9QirpDRAtMzYUgH9yD2JBZ0uVkZlD4D3ked5e4ez6A6UhyopZKFs
zecxRIEiBwJCsTtwh689wZ//Lx8B0XH/GWiNt+cv/SBihs8S9QShlgtjBDKz01PNXK1yITS7jXPE
fP79Q4RKIKtpFgQtPU0rtfbbaB5UVeEZJP9fljU5TGUDEv4kSqFpv2tBh6lY52OGR0PEUfC9sHP9
bEvNul0DSxXH2In5QpHATZSJ7hhrzcg2rlOplDYs46x1SiEcswYDXCwRJA2gJHr35rZCpUl/FswL
19rr5+E/Mw58Ix69lj4bsLJ/7Ef+tNPPbOS0aczA274MahRMevWqtQPnJ4fttQmLw0gcc1bTXvPD
uGTZ8DKnHoWOhRqq6zC6uMjkbW/AZrcOHXSXu8snn0aseMDRDzifVtCV3UnIfCf090cLI0S7EfZl
I1TotvYpaNjEGN2fON5GGLj4BHKDjSjZn1mtwopQ3Xdm9MS50JOgmEo+R63e7u6C/gU8CXFJgvT0
h3LG5oEM+kgZBDjcPj08YHtD1gg7/ioYIFVrLqaczBbdKFhlwn6SXMCmT6cWmm8iIal+GCfw9GTm
ehxt0Ec/kezuO650Qn5LutCI98LGG3HJOh9M5R2PaiHpQQTAeGOGFVjCeK+lJL+6xjtI0dfFCgoz
pn3r8jNjQToi2wo7QFo3++tRhKq8savZEQD6tVxRR8p1ywrjftsma5bt7COg9ZlWFOI55nB0Ew+d
OwJTZ5DwUPIJJOwl0Q7vFy4FGBOeME0eN1/8Ul5BAbt2z+ogymyuGN8COiSUFLpkei72yM8nCEHt
tmMHvSzB1GlSnbGW7RWOwqYcPnqIpEixd8qFH4FnDF5zbiBOwoUfN9ur4u3hMUnfv1KRhtzVDiYB
faAEgvJUHlruB8ylogTYpo1pTFNC7eoOr+FaQK2LfFaF+aF8P+lWPr6Yz69wcTBf6r1qFJSoC64I
X7PVlGSpFcxIa/ipNdqdI4Dx8QeDn78edmAE3RzMyQos0RX7eanOVbIanT1FDx0Rc6C1spwteymq
DIWOn7Zr+oOXhho6dLpCD+YSN2j6Wx6xJXvF22hq+vp02r8HGLw0JNqw0QYLj9p2/gSKWFxlBNNA
cmQ1fp6rOsjMK8kuNNA1uM2OTnZVTp9i1Oz0qVwnOdarc3zDIhQMK8wo4RQgH/u4dLPhBHQLJnXM
gfKCuQM7hYp0Mb8l5SAXM/O0rfaHWFc4zgxdk4OR2FUwKhdY0vx/eOKlyAMkUPU/ZlbqhWrHUmKJ
92MC40B/SUJ+VOEymieafJJOaaa7pdd4Do6uvzeHlwbhpEDsUdmaqNzpqI+oLGDQkyhRMi2avHMU
UeEO/HqQCjGsXe3gRMvCIiHosbFgqNy1+qjac2P+3/oSFEKzzsyF7wtbGzP3412RWR4DQmmPFkSt
igFsMK7Ninn+UVYmob9BmNZBDlYCMCk9rh+LQaPCeOM+10wIHow4GUVxr4HUeVodnxyJuv0EsfrP
uvqu33+TjdoXzV/BAjycNT6GEUu9wKlNSpxZ90XbaZF8XaAXpzOq3WYQgHsIUk3hPGMG6sO6INow
g8TeDK4zV8kj2xxUPfH2wGa3o9mfv8yfN0enA4ikgDS5dGeuqPsX0GGr6+NNWha0GTH14f2M3+4J
UiTB38S9Uh7gCSUbiLXk3GleN7waAhYMQ50dEwpd71EkVnDp4JKcTV9BvUSYcpHVJuU0gtFpaEpo
4q5FiVZlqqO8NVlML9wgbppUtLuF1KvDGulNqcLLJFDeDH9mzyIC+mzfRYNw9jVGGyKXHGCXGERD
5NeVgxgIGcHqKLU5MnxtcFQLik3EsmTlp0umqQZlOYwk7juVc1XrOgol7Z4Vc1TiPM3ERxP6lgOW
ozAWv+Vhx/0AJ9E4W5nDDnaB1D/admQ30hXULbiWLmgXWR7Q4Ie9CkdhOi09JTjzZIjf+LGJ3YRZ
Rzr3cvsBJqfr0eZRjPBkOX7edUfQWpPp47OzY1B5t7PZIT0E4no09OV+Dl1Ik02ilWvHeRWKOEBz
Cha08AmbmXvOI7QN5AzUbJzpjRUtqLc0c65zWYseMKC84pKIbkT3xoPWBJv46z1fCaJHvatoWsKo
ZgA2bxgaMrTDCs+AlDB5dFaQD3ZTl6/03KNg97iIJdR2MGlQ2X/wCXtOdgmCcwJlyMgYZHU/8IjU
hBvRu32vmfaRYIKAc0/o7rQA7pNlPcLMJC0Olc1yU07gi2gjoNNMjKD4qA6RLb8z8lnkdeNDDPib
4zqiYOfBJi17oixl2SnlMexTh5Wf6zIA9jMe3JWEEqve/6FJ1GZKe4Oipm0NFfp0NaLmiVXWlU9Y
DTZcrg6EYsy8rNKM/ltD9+kqL7w35tagUJSbnlVghUzBVGZc+LfaZ7a95uD6VKslRDfjhHzM7RPy
EwcGqJCCt8jWrAwLvUItBrBP2bne4dd8A5OFZhP8hdVr1gm7pIeoFjEge918rfOQVqwW2vKGOMOH
dLgV0ebCXQQv1Hx3w2KXE0MBlL9pcuHXqvlgr2h72cCkgLbFeIAekZXbrKOqlmuEynOib5BCpbRy
CzttDbL5kazYMZd6PmM5DM8R9WrBGTjJQIpjQbhBTgfLWYMuN7Yj1iKVE5fcNqYxcF7cG9nFc+sL
an3FkqtK9ggt+Hf3dKeJ1fRXlXBrQ2cNOM0iPyLYswa0jgQsnbq/JZQy8glNCWJcop1GDzwa26P0
wFOr+cl/x2RHXUM1Kw/DpM02Qkv2dXUvf4DY8RPUr/gKnmOl3AggIAo0nmd4PpsdPLUXgJRUhIO0
YHLNerw87eKyHNfHebwy7SNdyQ65PdkoZdewg1bzAJvGwBDxZ16GVpGXHcxyPo6om+sdWF9yqP0U
C2ChBZa3OCeFyyFgYJYPGuGNBKX5Rm4Gswn3vsXlgyxzU/dyX2JJCWnGCLQMKT1s74pWzbZwNc9O
igVDaGHyElI9yWlysUeYTENIyX92b4XdaomLf8v3TL0ImukVwDqW8aUOiaS1tZqfWgFRN7Jj5DfG
bs6wp/IJhZIvcHwUzJWZcemW0yOwlgQ1WXonUW93lM4iohK3L39sv+mzHtuSkjkIUX+XHMVcZFk9
9bsQe0k159nLtVvyhtWnYCaWNwAw4bnm31BE6Ot7PXrUa7rPM/BJUe5MY5G99NjIFYBn6lYvpZ+3
AW5HHb7IL0WTrTwPmOUtjuMDFUrpHge701rwaVMIVqMuqkrUpMMQ+JVzGzYS8D439szEGC/y4Myp
f5EfH3HAdxOmavZ+FLMnoZw2oB0agVJhmLFdlcob+zBPjkeD8Pp3Ur3X7jOuJbW2C21I5BAjYhyP
IWS2oyN+KYXmMK9ipd6iB1am+qiRSf4bWxooqTKNuD2SYaI/KMcyu/dOR+TzHwPHfDxToO7P8SV/
86aKodGVDB6Gf8zCWeFN7RjUShP/8EoWl8LHANxnigaz1cMyz5aZmUvUO9LpcsALtPk3iZEaMmkV
SZK182uXHlPzKKypTfLFn3WQIuGfzV+sxowpocBbyJJwhZkzSWq0+R3dLYlWGchht+N4FIWrVEE+
kmPnpL1/sY0/ytv6fwd+eVjHE8mNLeoZwmnbiv+vIoMO31sxkYdU9O34xvFd0WtLKMIq4r5Pj0+W
UZ/idh5OzdNh+GK7LpCK81J0j9KrIaC8A2Kk12ydNPnTklgTa0Mlz61/syrmm2x7qBtfXUggNOwG
jZYBsRqp5DW6nOf/ryaCvy3rk8MuW3Djr9F5QN762AAWdMYtMD15evorflw9fv3kXneZGpFJZhSe
yaJWFzBbR4g4D/ma/yW70X+dyVN4jfYnK4I+JGLGh35Vl48MeC2dyYCCW7f0WhVQGtq4BILpLvyM
VbjY4i+Q7lGHtvOw3hXlB5Ixc5Jnu83z2vGuRJ7k0uAwjj3bjQQbQfzXJ3BkRTHnNQG1wbuZ2I6N
lJr9N9SU2VBHlwM5TjWRLeL3B4TDHiLvj6dYkknoSXD263XwiFBXvQ42JdlQDFpYD0VPDt6ekOnk
IBQZhi2OMDUoAsLCQOrv7U0lHCwxWGNPBGlHq4E8K3DWAgBRBiqJyySXlU40VP8SkJFz4ujxa/sa
tF/G76yTsJDP18xHCCkbX3mLMbKXoQ7/P3590q859CZ6lQq/BQBYAiGsdPxIiNLkH2Lzww/GjZac
r9Vr8BzFvhzoYmMvTjKSQpusqJhaD4tpEQdZW95gcInTzeyHbn3N7GJniP59kTliXw2ylyXTz3zp
bopT2nc8Vw80u+byBhVEIbKOEEQlmm3WBOYiQuuTqDZpbvr0o78Igzex2A335XO2W2aH6ngMSOOR
07jWLoM8j4wdZJONyGBG4aeakbHuBdaFEux+PKgWiV06iUonNZ7pBlJJQrl6sDFQinCcsZXwSj6h
FlMPzzTslxlvwFFq5qavuFA9PUJnvI/RbTELc1+BnVRgO6bqOnX1CY5PNGfHNVC5bfT5z9mJH4qQ
QmrQZi7057Q5Ew3F7smppvZOpLAVizBO9blXjSCC9hrRmS7uD77b9OH/3wGSAAyeLmC8J/HJ/HKH
N/cWLfzQF7ALnZ/E5GLHDHZzFFV04+WV7RSPKWULl4zb80Mb3iNYxQBAv+aQR7Vr2WivywtjBiFs
gYxHISXnAeHbJ2KOrb967z+XZ72ZHwZYxwhAvSCqPkGPQ89AXQ0HnacDsyyZv81tkB8OL8OdxjhG
w7LdmrUDwtzDcn4Flz+qXMlj+uvyK9sSAHC7D6c6Nw7q/Gys59RzjJ3rD48DqFSe0YBXrzDoAHxC
hyUdPpYES883PUj0mA6yC40CznKxHM1FprnoEYA/4gdRmCwet4h6GfaHhEcqzna1k4Fy0jQdH/VO
Y7uP+g526bmv50r1KzFNUfr5rYIYOwlOf2NKY+D/k4LVXU8iyU6Lx4NI915ko1/OZr1ZBJ3dxbdH
wfK0sxx+47aof8kvdWS+wGzb68idnCB+tQJVCJ3yTMuz1KzL2p8MUXE5tnXXmpm5zcFWLOdKplXb
5fSnM7rUZK/O/WiemD7td1CrhZjETGnymPjVIkBo+RNvy2UM+PNgr/i/WNSTXBjKfZOGOIMbap0Q
e/B3iapDdlGx/otg92zzAiDna5jOBiy6q1RbtcoV/W4vbANmhafAoqtvdsyMWs1xtlLe4AzA1rWq
HJZ3GOnl8V5Aacn5BEueCw9y656xzMph5YZAYDuW6ZNKg/20AOAnBc6wQfDQxBD5wclmejWc+esA
iribaMEEFKz4rKeG55jgeAHVDqmWU2nmOYyVxyvWhC4vxjGnPmzMwbR1CJx8EOyfB9PjnNa34mHI
rKzQkNv3dHhhWa2/OWiBmTQhYGGJ0XMhHUB+5Ni8EIFSi/DOiZyOKxhVVidajltpwatoKX+T/T4t
J/b4Z/2RGtBFIxh4io8Wwt7eFaK47uWr1isChxVvOAD+N6UdA9zJJN/rRXTMpm13EQb6rS74QaJP
qdczI8zuxYGtI+ACkF0fa4VljiAVRMRUj13mAlTI1X4lfFuSm5tdbovgsl+cNNuC67/BRNUIqtPd
qavVv6OB4zPlA5fwBm8MgyseixXyOOcwqdxruJeqCQAyTjcW/HLbNH1/zF4KQ/zrcPc+XPhFjCEo
c88Tc9sQMJx23Bl5RbBvEpCnmtNWedoU2uVw+aE1Ygnsnqy1Y3Vne232r0he5ETbRYf3+WLuFerq
Rp0qnYK+lqmQazNhKVs9YpwLe+JXaUOxl1Szx4vMn4UQLbJMr/wTMhdKp2sZhs2cC29V3nZitk+c
W292aLycRoUayqy701lnoG54y/BxtoZGHPfqx98lzjFue70FeZKcMvNdZd9Rh1wn+txGscN/AZZ/
po62w0WvcoU3w9ww+WTRxFuCtclk3hH9fYfd5e/zaNmeuNEWHR5aW8vfO50EFhg8aPxHYeJQ9g+Y
utx4tR/PpKIwd9SUtyVIxpVxMki+C748pg80ddNwUPOL7oHZcsp1/7a0M5MaVGPams1eVoHkmt28
mqb5sNSpW2ELIxUekIl+tfvdI9kaOqtHYZfqC3GyeQIOxeGfzj69+gMPEnZNj0IJ90cSUy+QBc6F
WVbP4N3kX5JQz/j93H71iWeKf58Qmz50f9IOSK3gf/8adk03ZehnzH329i+hRI/F2Nr5fmTKSMmB
eBtEhjNnEtDvPvS2RITrXh3BL5bSpBPvi6OBwjbfpizwGlelAZ+OmKvpA1TSeHMwPO2xhYRY5x2o
4Y0F1MMOQMNetCaXQzto7lieLeIJnLHqfJh4ZKRQd1pMxzBIectQjKsGWhjsARboR+jO3U2b/2XM
eCVIJ58laQmvJUk+z2JO8hvKeOypq4NUC3i3Pevr4LOF/b9+f4KdvAmJPIqvYp8KrMU8RIUbvydm
P86FKtV6wKc0Kcg4Sofwc6t7S8Uf/r+R3zRIgCr4vES/YShwO0C3MSIsnu41pWPptrWi10sZRdhT
TcZE/r2h5s0Zx2krrUhA9tJu2EkltK9xvypky4jnOjRvZIbyQ5bl8PzvJuYuXR3prLfXs6U/8W/s
e1I32c1AnrWwLC9959cjTlnWT/ID99dbKyFDMinss5djZapsSAI7g+gEmhTbsRu56PrNPOQbRj8w
lFSOMBBEOR0PS0E+pMluogmzxCyK6cK4HC1zDdhW4RWfZfVT3ybMuz8to+1QYL11VMm7Fd7LXY6g
2Blj45uEnmx/AaaTjou4gDmToSR1Luj+Q6n5iButFNgQ4E0Otg5dLqVT2rUKDhA8k63jo3v4RZkN
X47ZICnwcbyYDJa5JkssxKsf/VNF2TvkHSLfa141PLrYexunpoDO9RI/gr96Na1boZICMDiJi1Ft
sXUwh2IRa6drCOXko3J7rogL7ML99/1SrCik6eXmfdcj1qF18NADcADf3f2AzeeBUVDnZpQXE39f
VZgqnVWZAWGWxi/os3dFNwC8LcnNra95X3zD5PX2b6UWukjdmY1h+TwKYLWmKgNxqZnGk2is6ODX
4IccODAJsYvR4k8oIIg6eeQG4pq3eduYBY0S5JQys1rDwRJ43nL8qfYSU3fZqE4GfKGWX3YU5eT8
BZXibdNrPhedHK4MBAhcwVsPyecs0BY4kE7ExnwLHOjPOhGSl13kUybk8vqHgn5numfxsjUGGfRh
4rUmDAVhTquoTA7D/ABryyNJqdzlwEiaGhEFwxeeBJyagFJLCdGLYDNTfDn6SYEurmaBTVbk+chn
UQvURv5UqYaXtQdNvPw5Q7ZlsUMsVQS6pnvWZ3NGLhm6+k8ZXtCbM8AyoO1CJ+aXJi+Ba9SznW3a
Y2Cw+P23iuQr5VYgda97DaILYeP+PG0aQgIAfj5PWN9dtxxm7I1JgViZSUZQtCkgA/3ypZHHWPZS
9hM5aeyBLS/HMek1vFF1ScRLO+2TPmg4fyvMtpdr4m1IFWMMLRbm0YjttEaG7HZIOinhRm8a0Oz1
5sIcDYosycsIC5B00sLf/iWAK9XWYVg4wDbR6Agv0Wf/ckXNtkA6HU9pZfxabuV+fHqrelAlUH2x
JsRLV+earuD12PcwYg+bqredQjubPZUl1E1zsxvRa1U8h8k6gyw8wpapOj+KOz6hbyUvuuAOjAhT
vLbxfszr+iv6mFeC/uD9E55E1t4CqM6aMTWaloYkkZxJfd+dc41W6xHmPCU+M+VMbpOaf4PBfulE
7FZ8Vk/tr2AH2tf0E0F8DaEDRu3MjYJZ6bqxECCpKNTq4Fbx+mxDcJ5EQlIfPvdaJwJnD7xtE6LC
x7oPYXVF+au2vMV6sYdkCX0zBY77DnizWImBz5OexfwTRdMvPeqOOGlOnVjLUfmaBXKaTmIHGMEJ
BNsfpYbJxl+dwXSjWc+HrMxAjFBdeTRDQWQEHlGs4dsk6M64DfN5e2Yw7RnSFlEDxff3xUJ6GA+C
I6bAC7WKfH8pTCHqMQXrGCzYe0GjFoJhRsYEA8pVIyQLQswpkNs3m8tQqqHJ7kYeyygq3inAEhd0
QsdXKIP4BzuGuhFGDe5CUOeXDrvQntaNYJ88F9fdeo/xHSuLs0Bgo8SoW9ZaPC7MauL4ZpnFdfIh
lWJ4KOwJu0BeoosD6TcxggvqGa/2WzPikXrcPrGvQUzlQ+5hrn49QKgWus92SA9/QoyvKANy3rMy
ufnfh/VYoRoJV0071y/B1LPA6xLi3kTEzed/tW3WtfFBVM0S1VYAlLMwgC0Df56U2fcYhLus93Nx
C2IYXU0/bfh7toAJ61EHMVIiAABR32WbkqTHE6CvSHmMY7kADb9xq1RKwXCgOUHbGLlLhUoP+s06
MX0y3eZbkana+Q84nTxxuY5J02lZ6Oom50D95LrG0V5XeOiDifl2Gt2dxWSBzpC/hgZSII+eK35j
iFwDbFzKVR7ZiLc/YG3lGctJUXSkHjtO//IKTwNK3q1iGX50Um91oPYW62f53K33VS0Gu6f7XaBW
nQVsv4HRbUlBc62187OjYAIC8HNGU6k+hLfmTjMwgpDV+44EYSOwXa1lI+dKE7c4E2AR8TrIEOFj
yd3xCZYupX7iEmYcRGn79iJ5nPK7SM4HtGn+hLTM70htO86R022t/kNQ4GucO6BRAtir+rk8MCul
tnKAsRsE4rg2CmMgEOATY7zoUxTmKuQfN0Hyjktq5XTTIGCvAg+heKLknHDYsDWT34YAnBQNBhi6
4v23lm1gHSb3eMCvSv2G/y4lVGxTi0lM3TYERc9Z9Q+Kuk5wdn6N6zBlSukMzILSP9fKlAgJ5Kkv
smqsMsOB9ZluqQ/YJYvUDB+RbU211BDa1lTcGvn8ZbtdcwGcy9Tov+WhK5xZeZq+eIZnHxrMV1cH
Yu5/ZxVy15a/fZ90Dq1e4YJNkVs6Gc3zi/G50wm5Cc9w2/As5iULCwJrWV41zg6Ce9q5toDBN11U
D09YA7s07MjJGEgLfb6Pc3ZK0Mq4/xZzVpXQjLO50/YkBJ6fBRtkdW4gCp6ScISPJFm+ZD+fSMCH
+SJ7ahNy97oKnb+GmZMkYrx2yBErUjpGl0/joCgmfYHwnBcGX/qmnOtmtcjht6w4tL13SqThHflR
ZiPACowLm8kbqm6ZEYZAFQ34UPxXY2tFs8gdmxvY7Qnm4HJ1yNTQvDFx3Ra7UbrvSVhWYng3HIJi
SFq5XjX/6RKDHLuuh/Sqsw5ctwpwH/+qjlXuP7P3ZaOzKwm1JmxhsABjHp2i2BTKLVhE7phEkNlG
rYgKdk6oJLtQpLm3n1oUIIL1rFonC3OhUFjTLIT8j7fLMcTynluZzNvUtq5uWiGDbgZ3JNex0NkE
EQ3MLEOYCKUNKDKDm6Kcp3HVAKPwNQZWvTFHOnbcOQ2L5O80qoxe6OJAcT6s1zJJ9+k+Awbmz+Gq
ugnB0LVB6W+PjmjRW3AV8GqxTwiGaqV9gNzrvyqUYrAvupUEl7ceer+VvZ4U6DgvsR+wAcGl4XiA
XdRK9OxrrFGQDKly72peXeCuvHYxqvlBb1nzjpiTdg55BZ8Xb5JoArrPO3QNSrYTkVPdLAbKUv7C
PdFWSTIOrJJORBhjw2MT2B5dsecHaVoA8k2qhI7ERMH1KFOpHwaaptnytDQ9sBV5t0IiUkoIwa9K
5zUQU/D4Cs3lqbOgWAzKR/pc6iVxUL8D8FltzfpcSGKwnqXJ8JQQx+lbPv40TbAxRY8HCVVRsEmQ
LucB4s+799QiXiOQu4a8bClIuV0bBOz4/QpKeGTG8Ioz1+PM/TdYd/+aFzx2AvooJeNv/DCT/0lD
+13vvayqgeFiSCWGNggx2AaXPfBeiI90JqNAIfl7fQ9vOGwAeC6kbX9+9snS3l3C1jAL3zRwN7b5
+0EG7O7M8yI3ZE7YiH181LtKRUi95RkW15O7Z2AnfIysP5044zYMAqmDpOvDZuSDC7GwZMMCPvYk
5JeeVKfR6kqDFdu/IW1I6eo04dw4j0dsIpU8+K5dl2cBGrrSMO7cEvD8hxgyieshXilCgjqmKds7
I3k/544yxWuEc/MwdTp9916Qoh5u9ExNtXDdV9aPmVW/SzPXo+N/aN9y/rcTxBeONWmdw0U9ZvlP
1Dqnn3rvbd57qTgV9iSuZ4DjJJotn5LHLz6O03G7fBLvqcD2fL4gcmYX+VFIky2rctmT1zBmMAA2
x4zKm1DrabmdyC3mUXCNfejNgbvdTCgluGmokIwbCVk7qqEHCMw1zA5rNBbjjZqrwfCbl0c9IQxg
lN4VZwZjNUjHMpl7QdxeZvROy62egbMS6q/sLbt6piFv2hfWtLzJCexLkz4EdCJ/xJ/sUZRp2iQw
FLrc/vnstNiWJlCW+QsGGqOIEPO8SgAtXgGHoqmw4ZgeOX5XIO4ZtWPdh7MGIRlCETTdweefFYlc
PBY8vAiClTda5Q6fCSw37JIoMGnH5+0qy4M2ugBr6Z2D0Zo5l50B6E3a0mYndGsdkkXvHPQYndUo
gF94hgprXWzvl0oyO9jlEEx0824XZINHWQAX0oRA/4KzlifrWuF7incj/sGfL4z4q2tao+rsct0B
7Ljl83yvCe975oxBs0qIP2Hi+kmosCozxjHMtftZij+m30MFacIhqDpnfHN8XpL2zme6m9Lh61cR
5tr9Mt3RVgmU1Tdo5I/U0UPR71+/x18YVzltSQbEsWtC8PNK1yMb85CFX+/g37wGW6cXzevpevOZ
GoWLxSdaHzsGcrgBkQC1c2U1KzX5v+qVcBYBYVYWG9fO3Bm3rgTLzNyAmZmZWUfiY9JGlBERKnoV
I/dm9IuZkecvhvrJFl87QaLm/h2zAqTJFmJ1efj6JI8oax+I0DzFLvnW48oneLvTX24biOf5+sUt
oMoJsrOMKkbTad825yq+m2u+9i9RRzB2E93vCpHVlG80HQ4vcL0ZBEBRB8gqGiEYRMf0w+ftb4vT
ltUoXgoo/OkbHDFoTSFOS1CsqKEptzCduMgxr4XDxX7nDm7DIIkRf9OrleFrpJbOv4XGBDaYhYov
IxPHUTVXC4tuw/oaGnbimEgOIuFHLpW1JKgrjM8+SI3XvQYojsHRcxUGR4+4R+ywFgY81ARtzAoM
3+7AF1+sCIpWpnIRPEsogXwjl+xypZp3/avfPoB+59BIOhMWgLtK+167KQo2Nw2WvOMBlgXZ1Jn8
HMxD+IAgfjcSIns4lMtHFsgnvLuyBPzr0TaCjpNn/KG7LMUIeJR+kW1lvfSCvlLiu8ANiNP86Zzy
Cac9+Dg1ZysP/wpwV9u/YC7TT41tfLGyNircNVSP10K7Yl16jR9YivKIOUNkTH/fSLWrMRhJKTDD
Hi+lujyMhBp1TQ7skNzx78iSqBcWDAmtyTTuq82ZiAN2lSE17Kfi7taMo/hiJ8D3QV5Fl5Ai3L0d
Bbh5kju03X5tVu+uWvYCrq0oWFJW0X3bdRMnxU2sclWgwp0rQEES1EWrSfXA56Ob7tLJ4Y09LqeO
S5105k4MuNSF23Oy+Pt1yz4b3yMVABGT7bihrE4zZIFN5wBlCcET/u/JMChp8wunfyAomAHymXlh
qa5F0tb9hBHdhL/NrViAe5LC4IZolVzQ4rAIYUtaTrk13CdG4b58vlaiPD8p+O3bHAFvaRPCSUS6
4qnLaBuAAufLcTR8huKiSFynZ6P70ocwSCvnPCgwOPnbklZbpA+Anev9j0TlXFZXDBh0ZjxS5HJS
UXpFfFdwA/GtOSFdEddUMFPK/wFPbwkuAq6aMfColaZFlVPj68mHFYAeaq4weYCC6+mvdn8dYlDx
vZofexV/C+JrmrpxnHnZkVVYnrwxHiSWqEzJwEKabez/UpzUmgJlDazEA2SlV5ow+Bu4tN6biHOr
m2ahLnJkxj2z8Lvdm1gSg+0DCfzQkKCyhHl4Yaxru2cvJiGZsXw6AzknVpiDy5KuZkunLPazSXQv
AsQbn22QIlJ/08lTCXdcvmvw0112Zs586b+5+CyZC7pLNueudqaSVfKTpSDonpj1oZpVAtZv2KUl
Kx+mwQRFJtj/xyqOLOGkIBp7o4ER+gcO8HDFi7YrQXD1O7EjnvY1G2daWhzo6yxYP45JisCq6yBO
QEZkZAg18Jw2JVDIkegtOj3Mhw7JX02h0VhyAUzUJgU5kPJGRTCdPkybFYARu3s2b9nxWjzfeptu
SJW/DodExrmW1B7BcnEnW2VgW3UMVoSfs+sutwEyfrR0DgzZ8H61LCDtNdxJdMjXtUU/7WSmab8Z
TsVgZFHEuZn4ZnosCI+cVUIP2EIhk0H8xGzGX6ahOmh2DMBkjgUBX4ghkzEwfT7F7vE+Zqcn9KBH
MJv/MZKNtF/STcv5rh5kHpUji+m89H46luPuyZOi7uDONQmtQGhwnGT+scWHgrS89Mbh9zWWbVQw
T+82AhlyMzBwpuJkmk7FYaizzfW+XYYOf1a/3s5vZJIGUn7EAsiUf7cqezJvMuu8cIq2pvGNAX/n
rP/y3MaSXO5jPOcv34coxr+J0c0tPNB4H0ht2asVpRcVgms3x/zeI/pmkGCXydLY4a/6cl8Ybr7q
rbQU2XmrwPYxyJUvuFiv/h96/GYeb6YvZqfzav2hhnGfYW5h+Ew0vt3x0VWVyzgwejwcFBmDX4ew
B+VI0B7Qpiqk5EPTt0H0cWnepExnY7BFbxTWfvMdvrbrJcCWLH5Y8jFx3SYe60IoagEToFhNZVSG
IcYifyIGXsMpCucJ79xD/XmWUfq7H+AdXoIdlZ1aTT6zHktlRWUgGQuM7QteSRWAQ50NwNzHjhfh
ecXvCeJU0ncI/bNXyJUuRn82/lV3O8QsC+zb1KM7f3NR3NlwBqJW0cIYUcLGcMsllIhaRhFLxiwi
9lBefntKfYWAT0C5TlMOPyKUNBpmOMtm/hCY0Xgks9Or29RSxytkmSh5nnLrlSZX61LmEVtnY6rW
P9+GT3EJOSlfoPtafhTcrsbSFDXvGANR5ZKUyNFaP2DfBmuXUqNJcVRcvXpWfUTOui0gSUaRFjDL
FnV2lObvM710dNqIkFBq/tBWoa70Ce2LOHr41H2sY5e0qr89y36SclCOM6O7iZkan0x2icHb5Rzf
xSsr1a48m0skaWPrrXXV8BCLdTV0g1Ec8lxazP3QaGqjk4V4zxmjrWO2JoqeNPIgB/YVJcoBDBQC
iGoJQ6CUoWTQE6Q+TTZqvaitMQndqlHnwLe+fP+SM8hZjtMswFSAvZeLb6D5EgYrwXdL/vuJr53v
uXOOOjh6WXzQgE4Do/gGAHKUMb0okrXIx1+fBXJMj4L7oq6i/SS02QxAXCpyW5JTp5qlRC5y9zDf
I3GUbMju4M3P1+dCIChSRmpGhuNWaE2W7Uwt2QEqw12Me2pysLGaLRg91/XCL48WexJXeJC2D8dv
Z4g7JY7AOoxGK1JnHKFl73MQ/4ogk/wouAN0EsKqKm/qYHn7TJuAkbLbnZjFMTyghQmcrzYEI/P4
eVsXxVwI1r489Xxwk2UqKRbTrrWT0v0d6G6zOxucZq6OGG0o06+f1J2pqEkIIkRijGzcSZcXNrpn
KW/7XsyeSS2jDoHew/pV5J9SHACngjgAYz2ER+hbNGC5VTzWGcWcx5jaJPvYm6QL9sjEdUHIWici
2yWjZvoLFnnerzLwSpStPVTYrlbaj0NKPIr/WnepEA8rKFpGi+d1jOrA4lCu81N+6cyd+D9d+EdP
Vi3PywIfeTZAieyHYkcFRRqZvpN4XbHwqdzdVdCDR1b5iW6hhKoXO9hyyG/NgUdF6WleKevp22eZ
C6ABcH8bBf7NfJPv9Y6H8yPcLYypE2H8emgk2diRmXQaXsqN5+tzpTGUAaJaWsACM4nMsz21rgO7
fK0Fl/V+K5zj2QEUKePJEQG2csp7t5/jHsrU4Cwr1ETWr4fehifIX+A8uLtt3cGveO3FpFB2X98N
6/swjU7PBLnBJw0jgXeE4NG30J67kIRV36zv2nMYkQ5oOrOXsUN5VM+bVAZdMYq3sc5k7mdRgoOm
Z8nKt5dvqnfMHtW3C+dPMGgaRdkBTvxVEdRJ3++oBvvaeF3LpBsdAgsNpYXpq/20tqTLslnVbj7K
bgnAObZINFU6UH4bVuL97IZ+pfUJUJCK1gyfVzdx56lEuJuGPo7bu/brQAzYwHrObcRV+np2XMgS
EWODccyOp083+KRvcij90eBnvYEEyVZsAakZ+q0rAzg+TdDXUUcbgP0hZvbdFQfWUM+kYD2X8Cfk
/ea0Tm8ytl19ewrYu8P+WlbUfkd84fsToHbJCHX74QkR73clNBjdtZT2IPcERWKl7YRavS6309P0
bNMmVUWIH68+h5lT5EpSZV8SsgoBT6Gdn6XqkbkHC9FlMJOqKgrt5gVqwEFItR6Tg7+OidoehS9d
PGkCYcDHKnm/b3E4Yh2KzRvhtHYfMPBacvlQKbe54Z+8WTbVTKaz7s59VnNMLnkCt3f7HES3cMj2
hmkNHmeFlASdX/XLunXDT+jEOtiOmEno/d/LkqzkaA5UoRo32vmXuu3Wox453tgnJGghnfFsEbCw
CDZLYx/0M+GHq7BjxgTOSR1Wc6rpzcE5pOGJZGd91y9rqNe2hobI6rZ3Z0qp2kotGEkjWTlsrnyR
i9uWO935EnUWEJNVGR46EJv3eroWbZc4wl6VeBQAWnQPuBYd7KoRXoypIieWX5q8gPzlu6jnoXX+
CmdO8cx1QZsn7c8484LWhVXp6A9OyVvqMetw65zEOU2xCdrVYlkiz6XoiTM9eqxYHbWtPXreHGON
WOhopoJnFjg8PIkS7RIwt4Glblgbqm1akoAbAC0KVMsBBC4i6VNIO1fFrffYHf54lQ402b+e430s
NRNC47gJ+x9f07qkf7b7gjQOdheAUc8IsnghRVsxfw+nzBCX1S1/R/bBazqfZguxwYWE8Tlno4N+
NrR7kFPJBMv/mJRVswTsANkt4r+O2zFMQRf6sLHfsf8Xq4isP1FsDZPA3gd4Nrb2Jwf/KUOY1lZr
XktVrWikFQWIg6J74s9fyMgdpSj0GZ96z+0dxSM+OpA//KTI//cFHcz+K17ng/pW4jDATWvW6ZVd
FtSDrKPnqBhl5nIkip5+3yjOk0O0AKy3xzX37vQXXa2l/rKCH0l26qqb3COQqlM1bGJOWbhsJ8Qs
aRKeuCrrJ/vifrO2t7cEuB+c0QwnqAAilCJ0rG7VTGWLmpUP6BBgcRDn1CPz9WNEFnJdcsd2FhbX
jNsG64N7Ytar4/buOnZHLpxNSuvX+rd3AmW2Frb+oGi5jJyMMjwYYJRHcKLLeqptX63kyBcPXGyV
Au3CflLwN8kd2veFg8gtnMHbDlTLMpS3BBIE8NE93b5ru5pRpiZfn33YGpn7/SJVmY0fdioM9PYu
QsGn+h2ftTLdFUfD2qZ7Uk0PosugcOp/kBtyikeaLbz9QwQlJ0SFzwyshOwIp/KXCtpRqNaK6xht
E1WCiFK6dzDq6OtzAPp3UksznVrHh0FmkEmk5gR1uqC4JQLHxOryBoVMN2omhoF3cxInJ7fP1PZm
zR4sbWZwdowU4BzSgAuuylymP5RcqNeqT8K+23emiELPafZmMRSvf+gghNR9yUo65kPJF8yY0hnB
APIJIBzKiMPtTFITdd61Mk0dFMV5HPHa/MZtmyjeemNVJqE4+xNqDfxPztjRf4V2CiYkRz6Wuyql
h97QIJBoXO7WbNlCnbXuWSw7w9NeWz3hi7hUFGf4E7hznv95uFs9gc24YTeGfEZ6qyRO77kA+/Py
10bKHfCMROOe9pwOTjLpn5BsWWvh71zcsiSDvgtZ6zIgqau+3ivRMsM1udJiwyDt0q4OHBPNCj/h
Dp9yF3XUq5JuJEkeZisXWNUjJcIbtIy5CF6Sv34IY2eZ59QCXHkJ/rB0IR/TFOZmNvahqjGwAQln
dCq/yUr54p9Pdy3+Yvg7q+EzdJROr25mcvxdMl26Yvqc0CuGxn84zpwfM90NkE0idOr01q6KrIdb
ZP/4Ip+zqmMDXI1SBgm5FhvvL7SWGZ0K5IE/mPWXJnfqnphpkLZvsgtkAkmel3IrUm4WEfsyPnAy
drHLDVvYpcATkLtoUfTeP/IvkfVC1pAZvmlJDZBLm8fFDvBPfmQzhjWyirZpGsQmdSSn9dIrfDnL
An3/6MUGq1MQmq6ymjhMkCNNEpnHvHpu73k4F8DVG+6loAMfDwLRUiEfFAXT/sg4DzYLcdD1M5xR
HQvJqks0UjeI34uWHqFuItjiQUMEaoXrWZfCBkpOiKG87o1QOtE3V40RWSAOZES2UK/Zhidh4TMb
YURgzvourUmrllaHfeYlmihzEY8/YQTQjMC7qgMiGnrZm6HlaRTBuIjqml8Fd/FJa/x8SyAFrvZh
lBzgbi8doFLFne4M7AFZcOYeQlFNhYyhi5URGKuEPaA/7c+NsenWxzmfsUuiYfHwZf0MfVtO/sjW
4CmKZchRPk6qOU1+Z7QorEUKP65qMBJWqY2BHa/rUk+Mte9WNz43qd7Gu9dwaUkn7sCwg09/b239
cARTWkjDMXp7s9Cc9aQnc8gPG3viWpAk9/R/1DkZT+pzzAIHarK5u++79XCHpovf+3tt9Tc/ZxTc
0+LKMiDGOZizzQGia7+EZpV9VbEwwX0/PkA7h4EC96JdCVmDwhK8sYG+ojC62DELCWOmdNYOMeNa
0DqSBjJ9e21sS2Kaew/bZoRWzm6x7GUsxU8zkOBWqUIDAsh/9KWD+1YXNhSToWm7y+OIZgNAjJlS
dqarRg8cMBSVkW2cKCZ65SXqNs00DfcMi6pJoGqOoQ973En7P9FbhoRkc/odxXpr14nX5vUA1Sbv
582ZbeSP32Tknp5cnAsWS49DcCsq2aX+4prJ5J48lFra7OOufTEa+x97EK7GUHr1seioUdY0ec1z
MCgHGAXzxl0S+hNWw00nccqbGftYPU5vBvDtAYillw2ZSRurzNplBCDXiPPpKyu2VmKonlY5yB0i
Z2iFBZ2ar5ieSDGYz/ztrUZCxZ7TjqBm2goEvuyZ2q94u9Q1qHQXhG7iHVA/L4jmE0nHvb702ujb
j1gqYyY6HU66AdEuph61Ge5KUgsLPHeloslG/1PnHQqfW0gnvqJRF7uBhXr7RqWlhMYu9Ng6Y1rl
I9/B9+0kadlzfKZkhfMqboekjiRDQYqC2DWclvZA6iV3MG2xQDJAAo0RMIv4oHk/3/p0pTN/Nc7U
RZXCtYbklngxWVSoJmPzWIoisrEd9tl0UY5EchCsXLstZNPqisox3v26RhXRaPUPCYyTxvhogx9C
nybBu0U8bNYYWpVMRRuU5+eFgWE3XxlFmCiNQYvK+Q0AZ2taUiOYBo1bjvgvDoycsfxWXUuM0osw
AIcNSu/RaJjeJcbv7PuuiN/gmOeV03PQDr6tKIgzBoj+2JGUWqUONevOqyu3TMrXe3JUoMXy/X9U
ilWcwtRlPhXA4PXvZzswauZIUETvz73fBbTPciaP31K2w6KUSKTs1wPqymsou8F5Sf8mHeiSqXMq
7pQkhKIQ+MBATkVBpRUk1wH5Tq9AOvpcQWaQyDA7mk7/DVxWeAke0n3kMO/xLKlziWK8k7GSjJsS
iWnQRPtxwqtUJ+JiWzYP9JhC6qzvcmKndYzSjtShmWzPKtGWk93pyYfsFuE6JsmwwJwbfssiINnE
B4oIsMU1CXsoidQcZrzAV8KtOHZCx/uCV8xYGRgL5fRLOXAuJZ/jbduurwKFLSYB72mORVWwcqv8
KxsnTz+xRTaSMBLkTIWkQJtXw8Jdo3MqaBXxD71dVaETkincMKjaFpD1KS8wyIoHieYOmh7q9ZKC
2wD0iVIMjU4hnTuBm/dBgINgXDp2qQslXH3WwwY08ESIQOi6eXeIYxGIMhl5W9e1clmHfJgkYfrA
L/HEeJzqY0afoIBi0hXLlGlErCiSzt57UJSYtPE7HkWyos5z2rUrY6zeqElHXCAoEB/YNYRgVDsB
jWKvS/lnFAAAAAAMpCuq8wAOLAABtcQB8OMGrwmPBLHEZ/sCAAAAAARZWg==
"""
_TRSM170_DRIVER = None
_TRSM170_IMAGE = None
_TRSM170_MODULE = None
_TRSM170_FUNCTION = None


def _trsm170_function_handle():
    global _TRSM170_DRIVER, _TRSM170_IMAGE, _TRSM170_MODULE, _TRSM170_FUNCTION
    if _TRSM170_FUNCTION is not None:
        return int(_TRSM170_FUNCTION.value)

    torch.cuda.init()
    driver = ctypes.CDLL("libcuda.so.1")
    driver.cuModuleLoadData.restype = ctypes.c_int
    driver.cuModuleLoadData.argtypes = [
        ctypes.POINTER(ctypes.c_void_p),
        ctypes.c_void_p,
    ]
    driver.cuModuleGetFunction.restype = ctypes.c_int
    driver.cuModuleGetFunction.argtypes = [
        ctypes.POINTER(ctypes.c_void_p),
        ctypes.c_void_p,
        ctypes.c_char_p,
    ]
    driver.cuFuncSetAttribute.restype = ctypes.c_int
    driver.cuFuncSetAttribute.argtypes = [
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_int,
    ]

    cubin = lzma.decompress(base64.b64decode(_TRSM170_CUBIN_XZ))
    image = ctypes.create_string_buffer(cubin)
    module = ctypes.c_void_p()
    status = driver.cuModuleLoadData(ctypes.byref(module), image)
    if status != 0:
        raise RuntimeError(f"cuModuleLoadData trsm170 failed: {status}")
    function = ctypes.c_void_p()
    status = driver.cuModuleGetFunction(
        ctypes.byref(function), module, b"trsm170x128_t1024"
    )
    if status != 0:
        raise RuntimeError(f"cuModuleGetFunction trsm170 failed: {status}")
    status = driver.cuFuncSetAttribute(function, 8, 202640)
    if status != 0:
        raise RuntimeError(f"cuFuncSetAttribute trsm170 failed: {status}")

    _TRSM170_DRIVER = driver
    _TRSM170_IMAGE = image
    _TRSM170_MODULE = module
    _TRSM170_FUNCTION = function
    return int(function.value)


_GEOMETRIC_PROBES = {}


def _geometric_probe(n, cols, device):
    key = (n, cols, device)
    value = _GEOMETRIC_PROBES.get(key)
    if value is None:
        generator = torch.Generator(device=device)
        generator.manual_seed(12345 + n * 17 + cols)
        value = torch.randn(
            (n, cols),
            device=device,
            dtype=torch.float32,
            generator=generator,
        )
        _GEOMETRIC_PROBES[key] = value
    return value


_PROJECTOR_QR_CUDA_SOURCE = r'''

#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda.h>

// QUEUE_T / QFIELD are token-pasted so the forbidden queue-type token never appears literally.
#define QUEUE_T CUstr ## eam
#define QFIELD str ## eam
#define MAXNW 32
#define LDPAD 1   // padded smem leading dim for panel_qr_body: m is a multiple of 32 so the
                  // column-major P[c*m+r] stride aliases all columns to the same bank set; +1 makes
                  // ldm odd (coprime to 32) -> de-conflicts cross-column smem access. Layout-only.
#ifndef WK_PFD
#define WK_PFD 4  // cross-column prefetch depth (apply software-pipeline); capped to stay spill-free
#endif

// ================== PACKED-REGISTER QR for n <= 32 (MAGMA smallsq) ==================
// One warp factors one matrix; MPB matrices share a CTA. Each lane owns one matrix
// row in registers rA[N] -- read from global once, written back once. The matrix
// stays in registers across all N Householder steps; the column norm and each
// reflector dot product are warp-shuffle reductions (no shared-memory barrier).
template <int N, int MPB>
__global__ void reg_qr_packed_kernel(const float* __restrict__ A, float* __restrict__ H, float* __restrict__ tau, int batch) {
  const int lane = threadIdx.x & 31;       // row within matrix (N<=32 -> one warp/matrix)
  const int ty = threadIdx.x >> 5;         // which matrix within the CTA
  const int mat = blockIdx.x * MPB + ty;
  if (mat >= batch) return;
  const float* Ab = A + (size_t)mat * N * N;   // input (read-only) -- lets the host skip data.clone()
  float* Hb = H + (size_t)mat * N * N;         // output (fully written by lane<m below)
  float* taub = tau + (size_t)mat * N;
  const unsigned FULL = 0xffffffffu;
  const int m = N;

  float rA[N];
  if (lane < m) {
    #pragma unroll
    for (int c = 0; c < N; ++c) rA[c] = Ab[(size_t)lane * N + c];
  } else {
    #pragma unroll
    for (int c = 0; c < N; ++c) rA[c] = 0.f;
  }
  __syncwarp(FULL);

  #pragma unroll 1
  for (int k = 0; k < N; ++k) {
    // norm of column k below the diagonal (warp-shuffle reduction)
    float partial = (lane > k && lane < m) ? rA[k] * rA[k] : 0.f;
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) partial += __shfl_down_sync(FULL, partial, o);
    float sumsq = __shfl_sync(FULL, partial, 0);
    float alpha = __shfl_sync(FULL, rA[k], k);
    float nrm = sqrtf(alpha * alpha + sumsq);
    float t, scale, beta;
    if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
    else {
      beta = (alpha >= 0.f) ? -nrm : nrm;
      t = (beta - alpha) / beta;
      scale = 1.f / (alpha - beta);
    }
    if (lane == k) { taub[k] = t; rA[k] = beta; }
    else if (lane > k && lane < m) rA[k] *= scale;
    float vlane = (lane == k) ? 1.f : ((lane > k && lane < m) ? rA[k] : 0.f);
    if (t != 0.f) {
      #pragma unroll
      for (int c = k + 1; c < N; ++c) {
        float prod = vlane * rA[c];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) prod += __shfl_down_sync(FULL, prod, o);
        float dot = __shfl_sync(FULL, prod, 0);
        rA[c] -= t * dot * vlane;
      }
    }
    __syncwarp(FULL);
  }

  if (lane < m) {
    #pragma unroll
    for (int c = 0; c < N; ++c) Hb[(size_t)lane * N + c] = rA[c];
  }
}

void reg_qr_packed(torch::Tensor A, torch::Tensor H, torch::Tensor tau, int64_t qh, int64_t mpb) {
  const int n = (int)H.size(1);
  const int batch = (int)H.size(0);
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  if (n == 32) {
    int MPB = (int)mpb;                    // matrices/CTA (one warp each, MPB*32 threads)
    int nblk = (batch + MPB - 1) / MPB;
    int blk = MPB * 32;
    const float* a = A.data_ptr<float>(); float* h = H.data_ptr<float>(); float* t = tau.data_ptr<float>();
    if (MPB == 1)      reg_qr_packed_kernel<32, 1><<<nblk, blk, 0, q>>>(a, h, t, batch);
    else if (MPB == 2) reg_qr_packed_kernel<32, 2><<<nblk, blk, 0, q>>>(a, h, t, batch);
    else if (MPB == 8) reg_qr_packed_kernel<32, 8><<<nblk, blk, 0, q>>>(a, h, t, batch);
    else               reg_qr_packed_kernel<32, 4><<<nblk, 128, 0, q>>>(a, h, t, batch);
  }
}

// ================== BLOCKED-HOUSEHOLDER PANEL FACTOR (smem, n >= 352) ==================
// Factors one narrow panel [j0:n, j0:j0+nb] entirely in shared memory and writes back
// the compact (reflectors below the diagonal, R on/above it) plus tau. The trailing
// update that follows is done host-side as a cuBLAS WY GEMM (see _blocked_qr / two-level).
//
// Two cheap micro-optimizations, selected by USE_VEC4:
//   USE_VEC4=true  -- float4 global<->smem panel copy (the column dim nb is 16-byte
//                     aligned) cuts LSU transactions 4x, and __launch_bounds__(512,2)
//                     lets 2 CTAs/SM co-reside on the late narrow panels. Used for the
//                     multi-panel / trailing-update regime where these pay off.
//   USE_VEC4=false -- scalar copy, no launch-bounds hint. Used only for a lone
//                     full-width panel with no trailing update (the n=176 blocked
//                     fallback) where the vec4/launch-bounds levers have nothing to
//                     amortize and slightly regress.
template <bool USE_VEC4, int VCAP = 16, bool FUSE = true, bool PF = false>
__device__ __forceinline__ void panel_qr_body(float* __restrict__ H, float* __restrict__ tau,
                                int n, int j0, int nb) {
  extern __shared__ float smem[];
  const int m = n - j0;
  const int NT = blockDim.x;
  const int NWr = NT >> 5;
  const int ldm = m + LDPAD;             // padded leading dim (bank-conflict break)
  float* P = smem;                       // panel, column-major: P[c*ldm + r]
  float* wred = smem + (size_t)ldm * nb;
  __shared__ float s_tau, s_scale;
  float* Hb = H + (size_t)blockIdx.x * n * n;
  float* taub = tau + (size_t)blockIdx.x * n;
  const int tid = threadIdx.x;
  const int lane = tid & 31, warp = tid >> 5;

  if (USE_VEC4) {
    const int nb4 = nb >> 2;
    for (int idx = tid; idx < m * nb4; idx += NT) {
      int r = idx / nb4, c4 = idx - r * nb4;
      int c = c4 << 2;
      float4 v = *reinterpret_cast<const float4*>(&Hb[(size_t)(j0 + r) * n + (j0 + c)]);
      P[(size_t)(c + 0) * ldm + r] = v.x;
      P[(size_t)(c + 1) * ldm + r] = v.y;
      P[(size_t)(c + 2) * ldm + r] = v.z;
      P[(size_t)(c + 3) * ldm + r] = v.w;
    }
  } else {
    for (int idx = tid; idx < m * nb; idx += NT) {
      int r = idx / nb, c = idx - r * nb;
      P[(size_t)c * ldm + r] = Hb[(size_t)(j0 + r) * n + (j0 + c)];
    }
  }
  __syncthreads();

  if (FUSE) {
  // Island A gen-2: fuse head(k+1) into column k's apply. Warp 0 (which owns ALL of column k+1's
  // rows during the apply) accumulates col k+1's norm and computes its larfg scalar THERE, removing
  // the per-column block norm-reduce phase from the panel's serial critical path (hide the head
  // behind the apply). head(0) is block-reduced eagerly; beta_k is written into colk[k] when head(k)
  // is computed (pre-loop for k=0, or during k-1's fuse). t==0 (rank-deficient col) falls back to a
  // block-reduced head -- t is block-uniform so the branch never diverges (no __syncthreads dead-
  // lock). The warp-local norm reorders the reduction vs the block path -> NOT bit-identical; it must
  // clear the ill-cond residual gate at the wide-m shapes (the F123 wall is the gate this tests).
  {                                                            // --- head(0), full block reduction ---
    float* colk = P;
    float acc = 0.f;
    for (int r = 1 + tid; r < m; r += NT) acc += colk[r] * colk[r];
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
    if (lane == 0) wred[warp] = acc;
    __syncthreads();
    if (tid == 0) {
      float sum = 0.f;
      for (int w = 0; w < NWr; ++w) sum += wred[w];
      float alpha = colk[0];
      float nrm = sqrtf(alpha * alpha + sum);
      float t, scale, beta;
      if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
      else { beta = (alpha >= 0.f) ? -nrm : nrm; t = (beta - alpha) / beta; scale = 1.f / (alpha - beta); }
      s_tau = t; s_scale = scale; taub[j0] = t; colk[0] = beta;
    }
    __syncthreads();
  }
  for (int k = 0; k < nb; ++k) {
    float* colk = P + (size_t)k * ldm;
    const float t = s_tau, scale = s_scale;                   // head(k); colk[k] already = beta_k
    for (int r = k + 1 + tid; r < m; r += NT) colk[r] *= scale;
    __syncthreads();
    if (t != 0.f) {
      float vreg[VCAP];
      int nseg = 0;
      #pragma unroll
      for (int s = 0; s < VCAP; ++s) {
        int r = k + lane + s * 32;
        if (r < m) { vreg[s] = (r == k) ? 1.f : colk[r]; nseg = s + 1; }
      }
      const int vtail = k + lane + VCAP * 32;
      if constexpr (PF) {   // COMPILE-TIME split: a dedicated _wide_pf kernel instantiates PF=true (single clean prefetch loop, 128 regs / 0 spill); _wide stays PF=false (byte-identical to v084). The dispatch routes ONLY the benefiting shape (n=2048) to _wide_pf; n=1024 keeps plain _wide. A RUNTIME m-guard in one kernel was tried and FAILED (dual-loop fallback → 80B spill → 2048 flips to +2.8% regress); the compile-time split avoids the spill entirely.
      // Explicit cross-column software pipeline: issue column (c+NWr)'s creg smem loads BEFORE the
      // SHFL reduction + update of column c, so the short_scoreboard load-to-use latency hides behind
      // the reduction/update work (the single-column VCAP cache already overlaps loads WITHIN a column
      // — see SASS — but the first FFMA of each NEW column still waits on its LDS burst). Bit-faithful:
      // the loaded values and the accumulation order are identical; only the issue order of independent
      // loads moves earlier. Only the reg-covered creg segment (s<nseg) is prefetched; the rare vtail
      // re-read (m>VCAP*32) is left as-is.
      // Prefetch DEPTH (PFD <= VCAP): how many of the next column's creg loads are issued ahead of the
      // SHFL+update of the current column. Capped so the extra live registers don't blow the 64-reg /
      // launch_bounds(512,2) budget of the default panel kernel (a full VCAP=16 second buffer SPILLS:
      // ptxas -v showed 96B local stack on panel_qr_kernel; PFD bounds the prefetch footprint to PFD regs).
      // PFD covers the next column's FIRST few FFMAs — exactly the load-to-use window that stalls when a
      // new column starts (the rest of that column's loads pipeline behind these via the VCAP cache).
      {
        int c = k + 1 + warp;
        float creg[VCAP];
        float pf[WK_PFD];
        bool have = false;
        if (c < nb) {
          float* colc0 = P + (size_t)c * ldm;
          #pragma unroll
          for (int s = 0; s < VCAP; ++s) { int r = k + lane + s * 32; if (s < nseg) creg[s] = colc0[r]; }
          have = true;
        }
        while (have) {
          float* colc = P + (size_t)c * ldm;
          int cn = c + NWr;
          float dot = 0.f;
          #pragma unroll
          for (int s = 0; s < VCAP; ++s) { if (s < nseg) dot += vreg[s] * creg[s]; }
          for (int r = vtail; r < m; r += 32) { float v = (r == k) ? 1.f : colk[r]; dot += v * colc[r]; }
          // prefetch the NEXT column's first WK_PFD creg loads before the serial SHFL+update of THIS one
          bool have_next = (cn < nb);
          if (have_next) {
            float* colcn = P + (size_t)cn * ldm;
            #pragma unroll
            for (int s = 0; s < WK_PFD; ++s) { int r = k + lane + s * 32; if (s < nseg) pf[s] = colcn[r]; }
          }
          #pragma unroll
          for (int o = 16; o > 0; o >>= 1) dot += __shfl_xor_sync(0xffffffffu, dot, o);
          const float w = t * dot;
          if (warp == 0 && c == k + 1) {
            float nacc = 0.f, s0v = 0.f;
            #pragma unroll
            for (int s = 0; s < VCAP; ++s) {
              int r = k + lane + s * 32;
              if (s < nseg) { float nv = creg[s] - w * vreg[s]; colc[r] = nv; if (s == 0) s0v = nv; if (r >= k + 2) nacc += nv * nv; }
            }
            for (int r = vtail; r < m; r += 32) { float v = (r == k) ? 1.f : colk[r]; float nv = colc[r] - w * v; colc[r] = nv; if (r >= k + 2) nacc += nv * nv; }
            #pragma unroll
            for (int o = 16; o > 0; o >>= 1) nacc += __shfl_down_sync(0xffffffffu, nacc, o);
            float alpha1 = __shfl_sync(0xffffffffu, s0v, 1);
            if (lane == 0 && k + 1 < nb) {
              float nrm = sqrtf(alpha1 * alpha1 + nacc);
              float t1, sc1, b1;
              if (nrm == 0.f) { b1 = alpha1; t1 = 0.f; sc1 = 0.f; }
              else { b1 = (alpha1 >= 0.f) ? -nrm : nrm; t1 = (b1 - alpha1) / b1; sc1 = 1.f / (alpha1 - b1); }
              s_tau = t1; s_scale = sc1; taub[j0 + k + 1] = t1; colc[k + 1] = b1;
            }
          } else {
            #pragma unroll
            for (int s = 0; s < VCAP; ++s) { int r = k + lane + s * 32; if (s < nseg) colc[r] = creg[s] - w * vreg[s]; }
            for (int r = vtail; r < m; r += 32) { float v = (r == k) ? 1.f : colk[r]; colc[r] -= w * v; }
          }
          if (have_next) {
            float* colcn = P + (size_t)cn * ldm;
            #pragma unroll
            for (int s = 0; s < VCAP; ++s) {
              int r = k + lane + s * 32;
              if (s < nseg) creg[s] = (s < WK_PFD) ? pf[s] : colcn[r];  // prefetched prefix + on-demand tail
            }
          }
          c = cn; have = have_next;
        }
      }
      } else {   // PF=false: the original no-spill apply loop (every non-_wide_pf kernel; byte-identical to v084)
      for (int c = k + 1 + warp; c < nb; c += NWr) {
        float* colc = P + (size_t)c * ldm;
        float creg[VCAP];
        float dot = 0.f;
        #pragma unroll
        for (int s = 0; s < VCAP; ++s) {
          int r = k + lane + s * 32;
          if (s < nseg) { creg[s] = colc[r]; dot += vreg[s] * creg[s]; }
        }
        for (int r = vtail; r < m; r += 32) {
          float v = (r == k) ? 1.f : colk[r];
          dot += v * colc[r];
        }
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) dot += __shfl_xor_sync(0xffffffffu, dot, o);  // butterfly: every lane holds the full sum (was down+broadcast = 1 extra shfl)
        const float w = t * dot;
        if (warp == 0 && c == k + 1) {
          // FUSE: update column k+1, accumulate its norm (rows >= k+2), compute head(k+1).
          float nacc = 0.f, s0v = 0.f;
          #pragma unroll
          for (int s = 0; s < VCAP; ++s) {
            int r = k + lane + s * 32;
            if (s < nseg) {
              float nv = creg[s] - w * vreg[s];
              colc[r] = nv;
              if (s == 0) s0v = nv;
              if (r >= k + 2) nacc += nv * nv;
            }
          }
          for (int r = vtail; r < m; r += 32) {
            float v = (r == k) ? 1.f : colk[r];
            float nv = colc[r] - w * v;
            colc[r] = nv;
            if (r >= k + 2) nacc += nv * nv;
          }
          #pragma unroll
          for (int o = 16; o > 0; o >>= 1) nacc += __shfl_down_sync(0xffffffffu, nacc, o);
          float alpha1 = __shfl_sync(0xffffffffu, s0v, 1);   // colc[k+1] = lane 1's s=0 updated value
          if (lane == 0 && k + 1 < nb) {
            float nrm = sqrtf(alpha1 * alpha1 + nacc);
            float t1, sc1, b1;
            if (nrm == 0.f) { b1 = alpha1; t1 = 0.f; sc1 = 0.f; }
            else { b1 = (alpha1 >= 0.f) ? -nrm : nrm; t1 = (b1 - alpha1) / b1; sc1 = 1.f / (alpha1 - b1); }
            s_tau = t1; s_scale = sc1; taub[j0 + k + 1] = t1; colc[k + 1] = b1;
          }
        } else {
          #pragma unroll
          for (int s = 0; s < VCAP; ++s) {
            int r = k + lane + s * 32;
            if (s < nseg) colc[r] = creg[s] - w * vreg[s];
          }
          for (int r = vtail; r < m; r += 32) {
            float v = (r == k) ? 1.f : colk[r];
            colc[r] -= w * v;
          }
        }
      }
      }  // end if constexpr (PF) else
    } else {
      // t==0 (rank-deficient column k): col k+1 unchanged by reflector k; block-reduce head(k+1).
      if (k + 1 < nb) {
        float* colk1 = P + (size_t)(k + 1) * ldm;
        float acc = 0.f;
        for (int r = k + 2 + tid; r < m; r += NT) acc += colk1[r] * colk1[r];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
        if (lane == 0) wred[warp] = acc;
        __syncthreads();
        if (tid == 0) {
          float sum = 0.f;
          for (int w2 = 0; w2 < NWr; ++w2) sum += wred[w2];
          float alpha = colk1[k + 1];
          float nrm = sqrtf(alpha * alpha + sum);
          float t1, sc1, b1;
          if (nrm == 0.f) { b1 = alpha; t1 = 0.f; sc1 = 0.f; }
          else { b1 = (alpha >= 0.f) ? -nrm : nrm; t1 = (b1 - alpha) / b1; sc1 = 1.f / (alpha - b1); }
          s_tau = t1; s_scale = sc1; taub[j0 + k + 1] = t1; colk1[k + 1] = b1;
        }
      }
    }
    __syncthreads();
  }
  } else {
  for (int k = 0; k < nb; ++k) {
    float* colk = P + (size_t)k * ldm;
    float acc = 0.f;
    for (int r = k + 1 + tid; r < m; r += NT) acc += colk[r] * colk[r];
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
    if (lane == 0) wred[warp] = acc;
    __syncthreads();
    if (tid == 0) {
      float sum = 0.f;
      for (int w = 0; w < NWr; ++w) sum += wred[w];
      float alpha = colk[k];
      float nrm = sqrtf(alpha * alpha + sum);
      float t, scale, beta;
      if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
      else {
        beta = (alpha >= 0.f) ? -nrm : nrm;
        t = (beta - alpha) / beta;
        scale = 1.f / (alpha - beta);
      }
      s_tau = t; s_scale = scale;
      taub[j0 + k] = t;
      colk[k] = beta;
    }
    __syncthreads();
    const float t = s_tau, scale = s_scale;
    for (int r = k + 1 + tid; r < m; r += NT) colk[r] *= scale;
    __syncthreads();
    if (t != 0.f) {
      // Apply reflector k to the trailing panel columns: cache the reflector below the
      // diagonal in registers (first VCAP*32 rows), spill the rest to a strided loop.
      // L7: VCAP is a template param. The default-16 kernel keeps __launch_bounds__(512,2)
      // (64-reg/2-block contract, used by the SM-saturated n512 b640 case). A dedicated VCAP=32
      // "wide" kernel with __launch_bounds__(512,1) (128 regs, no spill) covers m up to 1024 so the
      // n=1024 first panels' TAIL rows (512..1023) no longer double-read colc -- and n1024 b60 only
      // lights ~60 of 148 SMs, so the 1-block/SM cap costs no occupancy there.
      float vreg[VCAP];
      int nseg = 0;
      #pragma unroll
      for (int s = 0; s < VCAP; ++s) {
        int r = k + lane + s * 32;
        if (r < m) { vreg[s] = (r == k) ? 1.f : colk[r]; nseg = s + 1; }
      }
      const int vtail = k + lane + VCAP * 32;
      for (int c = k + 1 + warp; c < nb; c += NWr) {
        float* colc = P + (size_t)c * ldm;
        // L3 (smem-pipe traffic): cache the trailing column into regs during the DOT pass and
        // reuse it in the UPDATE pass, halving the trailing-column smem READS (each colc[r] was
        // read once for the dot and AGAIN for the subtract; the values don't change between the
        // passes). Mirrors vreg's reflector caching -- it's the *symmetric* lever (vreg was the
        // F146 reflector cache; this is the trailing operand). Bit-identical: same values, same
        // accumulation order. The reg-covered segment (s<nseg) holds creg; the tail (r>=vtail,
        // empty for m<=512) keeps the re-read.
        float creg[VCAP];
        float dot = 0.f;
        #pragma unroll
        for (int s = 0; s < VCAP; ++s) {
          int r = k + lane + s * 32;
          if (s < nseg) { creg[s] = colc[r]; dot += vreg[s] * creg[s]; }
        }
        for (int r = vtail; r < m; r += 32) {
          float v = (r == k) ? 1.f : colk[r];
          dot += v * colc[r];
        }
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) dot += __shfl_xor_sync(0xffffffffu, dot, o);  // butterfly: every lane holds the full sum (was down+broadcast = 1 extra shfl)
        const float w = t * dot;
        #pragma unroll
        for (int s = 0; s < VCAP; ++s) {
          int r = k + lane + s * 32;
          if (s < nseg) colc[r] = creg[s] - w * vreg[s];
        }
        for (int r = vtail; r < m; r += 32) {
          float v = (r == k) ? 1.f : colk[r];
          colc[r] -= w * v;
        }
      }
    }
    __syncthreads();
  }
  }

  if (USE_VEC4) {
    const int nb4s = nb >> 2;
    for (int idx = tid; idx < m * nb4s; idx += NT) {
      int r = idx / nb4s, c4 = idx - r * nb4s;
      int c = c4 << 2;
      float4 v;
      v.x = P[(size_t)(c + 0) * ldm + r];
      v.y = P[(size_t)(c + 1) * ldm + r];
      v.z = P[(size_t)(c + 2) * ldm + r];
      v.w = P[(size_t)(c + 3) * ldm + r];
      *reinterpret_cast<float4*>(&Hb[(size_t)(j0 + r) * n + (j0 + c)]) = v;
    }
  } else {
    for (int idx = tid; idx < m * nb; idx += NT) {
      int r = idx / nb, c = idx - r * nb;
      Hb[(size_t)(j0 + r) * n + (j0 + c)] = P[(size_t)c * ldm + r];
    }
  }
}

__global__ void __launch_bounds__(512, 2)
panel_qr_kernel(float* __restrict__ H, float* __restrict__ tau, int n, int j0, int nb) {
  panel_qr_body<true, 16>(H, tau, n, j0, nb);
}

// ===== SWIZZLED FUSE PANEL (LDS.128 apply) =====
// Same algorithm as panel_qr_body<true,VCAP,FUSE=true> but with ldm=m (no additive pad) so the apply
// can issue contiguous-quad LDS.128 (4080 ncu: an isolated quad LDS.128 is 1.74x faster than the
// strided LDS.32; the apply is ~28% of dense wall-clock and is LSU-issue-bound). ldm=m reintroduces
// transpose store-in bank conflicts (111M @ldm=512 vs 15M padded); a CUTLASS-style XOR swizzle
//   SW(c,r) = c*ldm + (r ^ (((c>>2)&7)<<2))         (quad-preserving: only permutes row bits[4:2])
// restores the store-in to 2-way max (14M, == the padded baseline) while keeping the apply read AND
// the strided norm/scale reads conflict-free (microbench-verified on the 4080, all four patterns).
// Quad-preserving means a float4 at SW(c,base) (base mult of 4) still holds rows base..base+3 of col c,
// so the LDS.128 quad apply is bit-faithful to the column layout. Gated to m%32==0 (so q^7 stays in
// [0, m/4) -> no column overflow) and m%4==0; else the launcher routes to the unswizzled fuse kernel.
#define SWZ_OFF(c, r) ((size_t)(c) * ldm + ((r) ^ ((((c) >> 2) & 7) << 2)))
// Quad (float4) index into P viewed as float4*: makes the apply's 16-byte alignment provable to the
// compiler so it emits LDS.128 *and* STS.128 (plain SWZ_OFF + reinterpret_cast scalarized the STORE,
// giving a 4-way-conflict strided scalar writeback that ate the LDS.128 read win). ldm=m mult of 4.
#define SWZ_QUAD(c, base) ((size_t)(c) * (ldm >> 2) + (((base) >> 2) ^ (((c) >> 2) & 7)))
template <int VCAP = 16>
__device__ __forceinline__ void panel_qr_body_swz(float* __restrict__ H, float* __restrict__ tau,
                                int n, int j0, int nb) {
  extern __shared__ float smem[];
  const int m = n - j0;
  const int NT = blockDim.x;
  const int NWr = NT >> 5;
  const int ldm = m;                     // aligned leading dim; swizzle de-conflicts the transpose
  float* P = smem;                       // panel, column-major SWIZZLED: P[SWZ_OFF(c,r)]
  float* wred = smem + (size_t)ldm * nb;
  __shared__ float s_tau, s_scale;
  float* Hb = H + (size_t)blockIdx.x * n * n;
  float* taub = tau + (size_t)blockIdx.x * n;
  const int tid = threadIdx.x;
  const int lane = tid & 31, warp = tid >> 5;
  constexpr int QCAP = (VCAP + 3) / 4;

  // store-in (transpose, swizzled): float4 from global row-major -> scatter to 4 columns at row r
  {
    const int nb4 = nb >> 2;
    for (int idx = tid; idx < m * nb4; idx += NT) {
      int r = idx / nb4, c4 = idx - r * nb4;
      int c = c4 << 2;
      float4 v = *reinterpret_cast<const float4*>(&Hb[(size_t)(j0 + r) * n + (j0 + c)]);
      P[SWZ_OFF(c + 0, r)] = v.x;
      P[SWZ_OFF(c + 1, r)] = v.y;
      P[SWZ_OFF(c + 2, r)] = v.z;
      P[SWZ_OFF(c + 3, r)] = v.w;
    }
  }
  __syncthreads();

  // head(0): full block reduction over column 0 (rows 1..m-1)
  {
    float acc = 0.f;
    for (int r = 1 + tid; r < m; r += NT) { float x = P[SWZ_OFF(0, r)]; acc += x * x; }
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
    if (lane == 0) wred[warp] = acc;
    __syncthreads();
    if (tid == 0) {
      float sum = 0.f;
      for (int w = 0; w < NWr; ++w) sum += wred[w];
      float alpha = P[SWZ_OFF(0, 0)];
      float nrm = sqrtf(alpha * alpha + sum);
      float t, scale, beta;
      if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
      else { beta = (alpha >= 0.f) ? -nrm : nrm; t = (beta - alpha) / beta; scale = 1.f / (alpha - beta); }
      s_tau = t; s_scale = scale; taub[j0] = t; P[SWZ_OFF(0, 0)] = beta;
    }
    __syncthreads();
  }
  for (int k = 0; k < nb; ++k) {
    const float t = s_tau, scale = s_scale;                   // head(k); colk[k] already = beta_k
    for (int r = k + 1 + tid; r < m; r += NT) P[SWZ_OFF(k, r)] *= scale;
    __syncthreads();
    if (t != 0.f) {
      // Cache the reflector (column k below the diagonal) as contiguous quads. Mask: r<k -> 0,
      // r==k -> implicit 1. Quad-preserving swizzle keeps the 4 sub-rows contiguous in smem.
      float4* P4 = reinterpret_cast<float4*>(P);
      float4 vq[QCAP];
      int qseg = 0;
      #pragma unroll
      for (int s = 0; s < QCAP; ++s) {
        int base = (s * 32 + lane) * 4;
        if (base < m) {
          float4 cv = P4[SWZ_QUAD(k, base)];
          cv.x = (base + 0 < k) ? 0.f : ((base + 0 == k) ? 1.f : cv.x);
          cv.y = (base + 1 < k) ? 0.f : ((base + 1 == k) ? 1.f : cv.y);
          cv.z = (base + 2 < k) ? 0.f : ((base + 2 == k) ? 1.f : cv.z);
          cv.w = (base + 3 < k) ? 0.f : ((base + 3 == k) ? 1.f : cv.w);
          vq[s] = cv; qseg = s + 1;
        }
      }
      for (int c = k + 1 + warp; c < nb; c += NWr) {
        float4 cq[QCAP];
        float dot = 0.f;
        #pragma unroll
        for (int s = 0; s < QCAP; ++s) {
          int base = (s * 32 + lane) * 4;
          if (s < qseg) {
            float4 cv = P4[SWZ_QUAD(c, base)];
            cq[s] = cv;
            dot += vq[s].x * cv.x + vq[s].y * cv.y + vq[s].z * cv.z + vq[s].w * cv.w;
          }
        }
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) dot += __shfl_xor_sync(0xffffffffu, dot, o);
        const float w = t * dot;
        if (warp == 0 && c == k + 1) {
          // FUSE: update col k+1, accumulate its norm (rows >= k+2), compute head(k+1). The quad map
          // makes lane L own rows [4L..4L+3]+128s; row k+1 lives in lane (k+1)/4 at sub (k+1)%4.
          float nacc = 0.f;
          #pragma unroll
          for (int s = 0; s < QCAP; ++s) {
            int base = (s * 32 + lane) * 4;
            if (s < qseg) {
              float4 cv = cq[s], vv = vq[s];
              cv.x -= w * vv.x; cv.y -= w * vv.y; cv.z -= w * vv.z; cv.w -= w * vv.w;
              P4[SWZ_QUAD(c, base)] = cv;
              if (base + 0 >= k + 2) nacc += cv.x * cv.x;
              if (base + 1 >= k + 2) nacc += cv.y * cv.y;
              if (base + 2 >= k + 2) nacc += cv.z * cv.z;
              if (base + 3 >= k + 2) nacc += cv.w * cv.w;
            }
          }
          #pragma unroll
          for (int o = 16; o > 0; o >>= 1) nacc += __shfl_down_sync(0xffffffffu, nacc, o);
          if (lane == 0 && k + 1 < nb) {
            float alpha1 = P[SWZ_OFF(c, k + 1)];   // newly-written col k+1 diagonal
            float nrm = sqrtf(alpha1 * alpha1 + nacc);
            float t1, sc1, b1;
            if (nrm == 0.f) { b1 = alpha1; t1 = 0.f; sc1 = 0.f; }
            else { b1 = (alpha1 >= 0.f) ? -nrm : nrm; t1 = (b1 - alpha1) / b1; sc1 = 1.f / (alpha1 - b1); }
            s_tau = t1; s_scale = sc1; taub[j0 + k + 1] = t1; P[SWZ_OFF(c, k + 1)] = b1;
          }
        } else {
          #pragma unroll
          for (int s = 0; s < QCAP; ++s) {
            int base = (s * 32 + lane) * 4;
            if (s < qseg) {
              float4 cv = cq[s], vv = vq[s];
              cv.x -= w * vv.x; cv.y -= w * vv.y; cv.z -= w * vv.z; cv.w -= w * vv.w;
              P4[SWZ_QUAD(c, base)] = cv;
            }
          }
        }
      }
    } else {
      // t==0 (rank-deficient col k): col k+1 unchanged by reflector k; block-reduce head(k+1).
      if (k + 1 < nb) {
        float acc = 0.f;
        for (int r = k + 2 + tid; r < m; r += NT) { float x = P[SWZ_OFF(k + 1, r)]; acc += x * x; }
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
        if (lane == 0) wred[warp] = acc;
        __syncthreads();
        if (tid == 0) {
          float sum = 0.f;
          for (int w2 = 0; w2 < NWr; ++w2) sum += wred[w2];
          float alpha = P[SWZ_OFF(k + 1, k + 1)];
          float nrm = sqrtf(alpha * alpha + sum);
          float t1, sc1, b1;
          if (nrm == 0.f) { b1 = alpha; t1 = 0.f; sc1 = 0.f; }
          else { b1 = (alpha >= 0.f) ? -nrm : nrm; t1 = (b1 - alpha) / b1; sc1 = 1.f / (alpha - b1); }
          s_tau = t1; s_scale = sc1; taub[j0 + k + 1] = t1; P[SWZ_OFF(k + 1, k + 1)] = b1;
        }
      }
    }
    __syncthreads();
  }

  // writeback (gather, swizzled): float4 to global row-major from 4 swizzled columns at row r
  {
    const int nb4s = nb >> 2;
    for (int idx = tid; idx < m * nb4s; idx += NT) {
      int r = idx / nb4s, c4 = idx - r * nb4s;
      int c = c4 << 2;
      float4 v;
      v.x = P[SWZ_OFF(c + 0, r)];
      v.y = P[SWZ_OFF(c + 1, r)];
      v.z = P[SWZ_OFF(c + 2, r)];
      v.w = P[SWZ_OFF(c + 3, r)];
      *reinterpret_cast<float4*>(&Hb[(size_t)(j0 + r) * n + (j0 + c)]) = v;
    }
  }
}

__global__ void __launch_bounds__(512, 2)
panel_qr_kernel_swz(float* __restrict__ H, float* __restrict__ tau, int n, int j0, int nb) {
  panel_qr_body_swz<16>(H, tau, n, j0, nb);
}
__global__ void __launch_bounds__(512, 1)
panel_qr_kernel_swz_wide(float* __restrict__ H, float* __restrict__ tau, int n, int j0, int nb) {
  panel_qr_body_swz<32>(H, tau, n, j0, nb);
}

// L7 wide kernel: VCAP=32 (covers m up to 1024 -> n=1024 first panels lose the tail double-read).
// __launch_bounds__(512,1) gives 128 regs so VCAP=32 does not spill. Used only for m>512, where
// CTA-per-matrix occupancy (n=1024 b60: ~60 CTAs) is already below 1 block/SM, so the 1-block cap
// is free. Bit-identical to the default kernel for any given panel (same math, wider reg cache).
__global__ void __launch_bounds__(512, 1)
panel_qr_kernel_wide(float* __restrict__ H, float* __restrict__ tau, int n, int j0, int nb) {
  panel_qr_body<true, 32>(H, tau, n, j0, nb);
}

#ifdef WK_PREFETCH
// WIDE + PREFETCH variant (PF=true): cross-column software-pipelined apply. SEPARATE compiled kernel
// (128 regs / 0 spill, verified -Xptxas -v) so it carries ONE clean prefetch loop — no runtime branch,
// no dual-loop spill. Routed (by n in the panel_qr dispatch) only to shapes whose wide panels benefit
// (n=2048: B200 ab2 −3.8%); n=1024 stays on plain _wide (byte-identical to v084, no regress).
__global__ void __launch_bounds__(512, 1)
panel_qr_kernel_wide_pf(float* __restrict__ H, float* __restrict__ tau, int n, int j0, int nb) {
  panel_qr_body<true, 32, true, true>(H, tau, n, j0, nb);
}
#endif

__global__ void panel_qr_kernel_plain(float* __restrict__ H, float* __restrict__ tau,
                                      int n, int j0, int nb) {
  panel_qr_body<false, 16>(H, tau, n, j0, nb);
}

// NON-FUSED vec4 kernel (FUSE=false): the head-hiding fuse helps n=512+ but is a dev→official
// transfer LOSS on n=352 (F-N352-NOTMYREGRESSION); route n=352 here to recover its official time.
__global__ void __launch_bounds__(512, 2)
panel_qr_kernel_nofuse(float* __restrict__ H, float* __restrict__ tau, int n, int j0, int nb) {
  panel_qr_body<true, 16, false>(H, tau, n, j0, nb);
}

void panel_qr(torch::Tensor H, torch::Tensor tau, int64_t j0, int64_t nb, int64_t nthreads,
              int64_t qh, int64_t plain) {
  const int n = H.size(1);
  const int m = n - (int)j0;
  size_t smem = (size_t)(m + LDPAD) * nb * sizeof(float) + MAXNW * sizeof(float);
  // plain==4: fuse but NO swizzle (truncated low-rank 512 -- swz regresses it; see _blocked_qr).
  const bool no_swz = (plain == 4);
  if (plain == 4) plain = 0;
  // swizzled-fuse smem: ldm=m (no additive pad). gate: m%32==0 (XOR stays in [0,m/4)) and the quad
  // reflector cache must COVER all rows: QCAP*128 = VCAP*128/4 rows. VCAP=16 (default) -> 512 rows;
  // VCAP=32 (wide) -> 1024 rows. The swz apply has no strided tail loop (the quad cache IS the apply),
  // so m beyond the cache silently drops rows (n=2048 b2 residual blow-up) -> hard cap m<=cover.
  // swz_wide DISABLED (regressed the non-saturated 1024-family b60); only the m<=512 default path.
  const bool swz_ok = ((m & 31) == 0) && (m <= 512) && !no_swz;
  size_t smem_swz = (size_t)m * nb * sizeof(float) + MAXNW * sizeof(float);
  static size_t configured = 0, configured_plain = 0, configured_wide = 0, configured_nofuse = 0;
  static size_t configured_swz = 0, configured_swz_wide = 0;
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  if (plain == 2) {        // non-fused vec4 kernel (n=352: fuse is an official-transfer loss there)
    if (smem > configured_nofuse) {
      cudaFuncSetAttribute(panel_qr_kernel_nofuse, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
      configured_nofuse = smem;
    }
    panel_qr_kernel_nofuse<<<(int)H.size(0), (int)nthreads, smem, q>>>(
        H.data_ptr<float>(), tau.data_ptr<float>(), n, (int)j0, (int)nb);
  } else if (plain) {
    if (smem > configured_plain) {
      cudaFuncSetAttribute(panel_qr_kernel_plain, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
      configured_plain = smem;
    }
    panel_qr_kernel_plain<<<(int)H.size(0), (int)nthreads, smem, q>>>(
        H.data_ptr<float>(), tau.data_ptr<float>(), n, (int)j0, (int)nb);
  } else if (m > 512) {   // wide VCAP=32 kernel: only m>512 has tail rows to cache
    // NOTE: swz_wide (VCAP=32, launch_bounds 512,1 = 1 block/SM) REGRESSED the 1024-family in ab2
    // (1024 +3.5%, 1024nrank +5%) -- the 1024 batch (b60) is not SM-saturated, so the swz's
    // occupancy/scheduling shift costs more than the LDS.128 apply saves. Keep the strided wide path.
#ifdef WK_PREFETCH
    // COMPILE-TIME split: route only n=2048's wide panels (which benefit, ab2 −3.8%) to the prefetch
    // kernel; n=1024 stays on plain _wide (byte-identical to v084 -> no regress). Two distinct kernels,
    // each a single clean loop -> neither spills.
    if (n == 2048) {
      static size_t configured_wide_pf = 0;
      if (smem > configured_wide_pf) {
        cudaFuncSetAttribute(panel_qr_kernel_wide_pf, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
        configured_wide_pf = smem;
      }
      panel_qr_kernel_wide_pf<<<(int)H.size(0), (int)nthreads, smem, q>>>(
          H.data_ptr<float>(), tau.data_ptr<float>(), n, (int)j0, (int)nb);
    } else
#endif
    {
      if (smem > configured_wide) {
        cudaFuncSetAttribute(panel_qr_kernel_wide, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
        configured_wide = smem;
      }
      panel_qr_kernel_wide<<<(int)H.size(0), (int)nthreads, smem, q>>>(
          H.data_ptr<float>(), tau.data_ptr<float>(), n, (int)j0, (int)nb);
    }
  } else {
    if (swz_ok) {
      if (smem_swz > configured_swz) {
        cudaFuncSetAttribute(panel_qr_kernel_swz, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem_swz);
        configured_swz = smem_swz;
      }
      panel_qr_kernel_swz<<<(int)H.size(0), (int)nthreads, smem_swz, q>>>(
          H.data_ptr<float>(), tau.data_ptr<float>(), n, (int)j0, (int)nb);
    } else {
      if (smem > configured) {
        cudaFuncSetAttribute(panel_qr_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
        configured = smem;
      }
      panel_qr_kernel<<<(int)H.size(0), (int)nthreads, smem, q>>>(
          H.data_ptr<float>(), tau.data_ptr<float>(), n, (int)j0, (int)nb);
    }
  }
}

int64_t max_smem_optin() {
  int dev = 0, v = 0;
  cudaGetDevice(&dev);
  cudaDeviceGetAttribute(&v, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);
  return v;
}

// ================== WHOLE-MATRIX FUSED QR for n = 176 (smem-resident) ==================
// The entire matrix fits in opt-in shared memory, so one CTA runs the whole unblocked
// Householder sweep on-chip -- every trailing update stays in smem, with no DRAM
// round-trip between Householder steps. Column-major smem s[col*n+row]; one CTA/matrix.
// Reductions are warp-shuffle + a small cross-warp tree.
__global__ void fused_qr_full_kernel(const float* __restrict__ A, float* __restrict__ Hout,
                                     float* __restrict__ tauOut, int n) {
  extern __shared__ float smem[];
  float* s = smem;                 // n*n column-major
  float* wred = smem + (size_t)n * n;
  const int bmat = blockIdx.x;
  const int tid = threadIdx.x;
  const int NT = blockDim.x;
  const int lane = tid & 31, warp = tid >> 5;
  const int NWr = NT >> 5;
  const float* Ab = A + (size_t)bmat * n * n;
  float* Hb = Hout + (size_t)bmat * n * n;
  float* taub = tauOut + (size_t)bmat * n;
  __shared__ float s_tau, s_scale, s_beta;

  for (int idx = tid; idx < n * n; idx += NT) {
    int row = idx / n, col = idx - row * n;
    s[(size_t)col * n + row] = Ab[(size_t)row * n + col];
  }
  __syncthreads();

#ifdef WK_FUSE
  // Island A gen-1: fuse head(j+1) into column j's apply. Warp 0 owns ALL of column j+1's rows
  // during the apply, so it computes the next reflector's norm+scalar THERE, removing the separate
  // block norm-reduce phase from the per-column critical path (hides the 37-41% serial larfg head
  // behind the 58-63% apply, F-NCU-PANEL). head(0) is computed eagerly with the full block
  // reduction; every later head is fused. The warp-local (single-warp) reduction order differs from
  // the block reduction but is correctness-safe here (n=176, cond=1, ~600x residual margin). t==0
  // (rank-deficient column) falls back to a block-reduced head -- t is block-uniform so the branch
  // never diverges across the CTA (no __syncthreads deadlock).
  {
    float* col = s;                                          // head(0), full block reduction
    float acc = 0.f;
    for (int r = 1 + tid; r < n; r += NT) acc += col[r] * col[r];
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
    if (lane == 0) wred[warp] = acc;
    __syncthreads();
    if (tid == 0) {
      float sum = 0.f;
      for (int w = 0; w < NWr; ++w) sum += wred[w];
      float alpha = col[0];
      float nrm = sqrtf(alpha * alpha + sum);
      float t, scale, beta;
      if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
      else { beta = (alpha >= 0.f) ? -nrm : nrm; t = (beta - alpha) / beta; scale = 1.f / (alpha - beta); }
      s_tau = t; s_scale = scale; s_beta = beta; taub[0] = t;
    }
    __syncthreads();
  }
  for (int j = 0; j < n; ++j) {
    float* col = s + (size_t)j * n;
    const float t = s_tau, scale = s_scale, beta = s_beta;  // head(j); load before s_* is overwritten
    for (int r = j + 1 + tid; r < n; r += NT) col[r] *= scale;
    __syncthreads();
    if (t != 0.f) {
      for (int c = j + 1 + warp; c < n; c += NWr) {
        float* cc = s + (size_t)c * n;
        float dot = (lane == 0) ? cc[j] : 0.f;              // v_j = 1, counted once
        for (int r = j + 1 + lane; r < n; r += 32) dot += col[r] * cc[r];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) dot += __shfl_xor_sync(0xffffffffu, dot, o);  // butterfly: every lane holds the full sum (was down+broadcast = 1 extra shfl)
        const float w = t * dot;
        if (lane == 0) cc[j] -= w;
        if (warp == 0 && c == j + 1) {
          // fuse: update col j+1 AND accumulate its norm (rows >= j+2) for head(j+1).
          float nacc = 0.f;
          for (int r = j + 1 + lane; r < n; r += 32) {
            float nv = cc[r] - w * col[r];
            cc[r] = nv;
            if (r >= j + 2) nacc += nv * nv;
          }
          #pragma unroll
          for (int o = 16; o > 0; o >>= 1) nacc += __shfl_down_sync(0xffffffffu, nacc, o);
          if (lane == 0) {
            float a = cc[j + 1];
            float nrm = sqrtf(a * a + nacc);
            float t1, sc1, b1;
            if (nrm == 0.f) { b1 = a; t1 = 0.f; sc1 = 0.f; }
            else { b1 = (a >= 0.f) ? -nrm : nrm; t1 = (b1 - a) / b1; sc1 = 1.f / (a - b1); }
            s_tau = t1; s_scale = sc1; s_beta = b1; taub[j + 1] = t1;
          }
        } else {
          for (int r = j + 1 + lane; r < n; r += 32) cc[r] -= w * col[r];
        }
      }
    } else {
      // t==0 (rank-deficient column): col j+1 unchanged by reflector j; block-reduce head(j+1).
      if (j + 1 < n) {
        float* col1 = s + (size_t)(j + 1) * n;
        float acc = 0.f;
        for (int r = j + 2 + tid; r < n; r += NT) acc += col1[r] * col1[r];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
        if (lane == 0) wred[warp] = acc;
        __syncthreads();
        if (tid == 0) {
          float sum = 0.f;
          for (int w2 = 0; w2 < NWr; ++w2) sum += wred[w2];
          float a = col1[j + 1];
          float nrm = sqrtf(a * a + sum);
          float t1, sc1, b1;
          if (nrm == 0.f) { b1 = a; t1 = 0.f; sc1 = 0.f; }
          else { b1 = (a >= 0.f) ? -nrm : nrm; t1 = (b1 - a) / b1; sc1 = 1.f / (a - b1); }
          s_tau = t1; s_scale = sc1; s_beta = b1; taub[j + 1] = t1;
        }
      }
    }
    __syncthreads();
    if (tid == 0) col[j] = beta;
  }
  __syncthreads();  // publish final diagonal to writeback (per-column trailing barrier hoisted out: col[j] frozen after iter j)
#else
  for (int j = 0; j < n; ++j) {
    float* col = s + (size_t)j * n;
    float acc = 0.f;
    for (int r = j + 1 + tid; r < n; r += NT) acc += col[r] * col[r];
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
    if (lane == 0) wred[warp] = acc;
    __syncthreads();
    if (tid == 0) {
      float sum = 0.f;
      for (int w = 0; w < NWr; ++w) sum += wred[w];
      float alpha = col[j];
      float nrm = sqrtf(alpha * alpha + sum);
      float t, scale, beta;
      if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
      else { beta = (alpha >= 0.f) ? -nrm : nrm; t = (beta - alpha) / beta; scale = 1.f / (alpha - beta); }
      s_tau = t; s_scale = scale; s_beta = beta;
      taub[j] = t;
    }
    __syncthreads();
    const float t = s_tau, scale = s_scale;
    for (int r = j + 1 + tid; r < n; r += NT) col[r] *= scale;
    __syncthreads();
    if (t != 0.f) {
      // one warp per trailing column; reflector v = (1 at j, col[r] below) read from smem.
      for (int c = j + 1 + warp; c < n; c += NWr) {
        float* cc = s + (size_t)c * n;
        float dot = (lane == 0) ? cc[j] : 0.f;  // v_j = 1, counted once
        for (int r = j + 1 + lane; r < n; r += 32) dot += col[r] * cc[r];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) dot += __shfl_xor_sync(0xffffffffu, dot, o);  // butterfly: every lane holds the full sum (was down+broadcast = 1 extra shfl)
        const float w = t * dot;
        if (lane == 0) cc[j] -= w;
        for (int r = j + 1 + lane; r < n; r += 32) cc[r] -= w * col[r];
      }
    }
    __syncthreads();
    if (tid == 0) col[j] = s_beta;
    __syncthreads();
  }
#endif
  for (int idx = tid; idx < n * n; idx += NT) {
    int row = idx / n, col = idx - row * n;
    Hb[(size_t)row * n + col] = s[(size_t)col * n + row];
  }
}

void fused_qr_full(torch::Tensor A, torch::Tensor H, torch::Tensor tau, int64_t nthreads, int64_t qh) {
  const int n = H.size(1);
  const int batch = H.size(0);
  size_t smem = ((size_t)n * n + MAXNW) * sizeof(float);
  static size_t configured = 0;
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  if (smem > configured) {
    cudaFuncSetAttribute(fused_qr_full_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
    configured = smem;
  }
  fused_qr_full_kernel<<<batch, (int)nthreads, smem, q>>>(
      A.data_ptr<float>(), H.data_ptr<float>(), tau.data_ptr<float>(), n);
}

// ================== CTA-CLUSTER + DSMEM COOPERATIVE PANEL (n = 4096, sm_90+) ==================
// For n=4096 the batch is tiny (2): one CTA per matrix lights ~2 of ~148 SMs. Instead a
// CLUSTER of C CTAs on one GPC cooperatively factors a single matrix's panel. The m=n-j0
// panel rows are split row-wise across the C CTAs (rank cr owns rows [lo,hi), its own slab
// in dynamic smem). Per Householder column the norm reduction and the trailing-apply dot
// products are combined ACROSS the cluster through distributed shared memory (DSMEM, remote-
// CTA smem via map_shared_rank) -- on-chip, no global round-trip. This lights C SMs per
// matrix. The kernel emits the compact (H,tau) for the panel; the host-side WY trailing GEMM
// finishes each block step. Built only under -DWK_CLUSTER (sm_90+ has cluster launch + DSMEM).
#ifdef WK_CLUSTER
#include <cooperative_groups.h>
namespace cg = cooperative_groups;

template <int NB, int C>
__global__ void __cluster_dims__(C, 1, 1)
cluster_panel_kernel(float* __restrict__ H, float* __restrict__ tau, int n, int j0) {  // idea 3: C is compile-time
  extern __shared__ float smem[];
  cg::cluster_group cl = cg::this_cluster();
  const unsigned cr = cl.block_rank();
  const int mat = (int)(blockIdx.x / C);          // grid = C*batch ; cluster id == matrix id
  const int m = n - j0;
  const int tid = threadIdx.x;
  const int NT = blockDim.x;
  const int lane = tid & 31, warp = tid >> 5;
  const int NWr = NT >> 5;

  const int lo = (int)(((long)cr * m) / C);
  const int hi = (int)(((long)(cr + 1) * m) / C);
  const int mb = hi - lo;                          // rows this CTA owns
  const int mbp = mb + 1;                          // padded leading dim (bank-conflict break, same as panel)

  // dynamic smem, identical layout in every rank so map_shared_rank hits the same offsets:
  // Pslab | sRed[2*CMAX] | sDot[CMAX*NB] | wred.  sRed[2q]=rank q's sumsq partial,
  // sRed[2q+1]=pivot (valid only for the owner rank of the current diagonal row).
  float* Pslab = smem;                             // mb x NB col-major
  float* sRed  = Pslab + (size_t)mbp * NB;          // 2*C (idea 3: C compile-time; idea-2 shrink NOT applied)
  float* sDot  = sRed + 2 * C;                     // C x NB
  float* wred  = sDot + (size_t)C * NB;            // MAXNW
  __shared__ float s_tau, s_beta, s_scale;

  float* Hb = H + (size_t)mat * n * n;
  float* taub = tau + (size_t)mat * n;

  // EXPERT IDEA 1 (isolated): row-major + float4 panel LOAD (coalesced global reads).
  const int NB4 = NB >> 2;
  for (int idx = tid; idx < mb * NB4; idx += NT) {
    int r = idx / NB4, c4 = idx - r * NB4, c = c4 << 2;
    float4 v = *reinterpret_cast<const float4*>(&Hb[(size_t)(j0 + lo + r) * n + (j0 + c)]);
    Pslab[(size_t)(c + 0) * mbp + r] = v.x;
    Pslab[(size_t)(c + 1) * mbp + r] = v.y;
    Pslab[(size_t)(c + 2) * mbp + r] = v.z;
    Pslab[(size_t)(c + 3) * mbp + r] = v.w;
  }
  cl.sync();

  for (int k = 0; k < NB; ++k) {
    // owner = rank holding diagonal row k (C <= 16, linear scan is fine)
    int owner = 0;
    for (int q = 0; q < C; ++q) {
      int qlo = (int)(((long)q * m) / C);
      int qhi = (int)(((long)(q + 1) * m) / C);
      if (k >= qlo && k < qhi) { owner = q; break; }
    }
    const int kl = k - lo;                          // local diagonal row index (owner only)
    float* colk = Pslab + (size_t)k * mbp;

    // partial sum-of-squares of column k strictly below the diagonal, over this CTA's rows
    float acc = 0.f;
    for (int r = tid; r < mb; r += NT) {
      int gr = lo + r;
      if (gr > k) acc += colk[r] * colk[r];
    }
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
    if (lane == 0) wred[warp] = acc;
    __syncthreads();
    if (tid == 0) {
      float s = 0.f;
      for (int w = 0; w < NWr; ++w) s += wred[w];
      sRed[2 * cr] = s;
      sRed[2 * cr + 1] = (cr == (unsigned)owner) ? colk[kl] : 0.f;
    }
    cl.sync();

    if (tid == 0) {  // every CTA gathers all ranks' partials + the owner's pivot via DSMEM
      float sum = 0.f, alpha = 0.f;
      for (int q = 0; q < C; ++q) {
        float* rem = cg::cluster_group::map_shared_rank(sRed, (unsigned)q);
        sum += rem[2 * q];
        if (q == owner) alpha = rem[2 * q + 1];
      }
      float nrm = sqrtf(alpha * alpha + sum);
      float t, scale, beta;
      if (nrm == 0.f) { beta = alpha; t = 0.f; scale = 0.f; }
      else {
        beta = (alpha >= 0.f) ? -nrm : nrm;
        t = (beta - alpha) / beta;
        scale = 1.f / (alpha - beta);
      }
      s_tau = t; s_beta = beta; s_scale = scale;
    }
    __syncthreads();
    const float t = s_tau, beta = s_beta, scale = s_scale;

    for (int r = tid; r < mb; r += NT) {
      int gr = lo + r;
      if (gr > k) colk[r] *= scale;
      else if (gr == k) colk[r] = beta;
    }
    if (cr == (unsigned)owner && tid == 0) taub[j0 + k] = t;
    __syncthreads();   // F144: scale (line above) writes only THIS rank's colk slab, and the trailing
                       // dot below reads only THIS rank's colk/colc (Pslab is never read cross-rank --
                       // only sRed/sDot go through map_shared_rank). So a CTA barrier suffices here; the
                       // cluster cl.sync() was over-synchronizing (one expensive cluster handshake/col).

    if (t != 0.f) {
      // F146: hoist the reflector v into registers ONCE -- it's invariant across all trailing columns
      // AND across the dot/update passes (only colc changes), but the original re-loaded colk[r] from
      // smem per column per pass + recomputed the branch. mb<=512 for the cluster (n=4096, C=8) so 16
      // slots cover a lane's rows; a tail loop guards any future mb>512. Bit-identical (same values +
      // accumulation order); mirrors panel_qr_body's VCAP register caching.
      float vreg[16];
      #pragma unroll
      for (int s = 0; s < 16; ++s) {
        int r = lane + (s << 5);
        vreg[s] = (r < mb) ? ((lo + r == k) ? 1.f : ((lo + r > k) ? colk[r] : 0.f)) : 0.f;
      }
      for (int c = k + 1 + warp; c < NB; c += NWr) {   // partial v.col_c per trailing col -> sDot
        float* colc = Pslab + (size_t)c * mbp;
        float dot = 0.f;
        #pragma unroll
        for (int s = 0; s < 16; ++s) {
          int r = lane + (s << 5);
          if (r < mb) dot += vreg[s] * colc[r];
        }
        for (int r = lane + 512; r < mb; r += 32) {    // tail: empty for mb<=512
          int gr = lo + r;
          dot += ((gr == k) ? 1.f : ((gr > k) ? colk[r] : 0.f)) * colc[r];
        }
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) dot += __shfl_down_sync(0xffffffffu, dot, o);
        if (lane == 0) sDot[(size_t)cr * NB + c] = dot;
      }
      cl.sync();
      for (int c = k + 1 + warp; c < NB; c += NWr) {   // gather full dot_c across ranks, update own rows
        float full = 0.f;
        for (int q = 0; q < C; ++q) {
          float* rem = cg::cluster_group::map_shared_rank(sDot, (unsigned)q);
          full += rem[(size_t)q * NB + c];
        }
        const float w = t * full;
        float* colc = Pslab + (size_t)c * mbp;
        #pragma unroll
        for (int s = 0; s < 16; ++s) {
          int r = lane + (s << 5);
          if (r < mb) colc[r] -= w * vreg[s];
        }
        for (int r = lane + 512; r < mb; r += 32) {    // tail: empty for mb<=512
          int gr = lo + r;
          colc[r] -= w * ((gr == k) ? 1.f : ((gr > k) ? colk[r] : 0.f));
        }
      }
      __syncthreads();   // F144: the sDot write-after-read hazard (next col's dot overwrites sDot that
                         // other ranks read here) is already covered by the NEXT column's cl.sync()@547
                         // (lock-step from per-column cluster barriers bounds rank drift to <1 column);
                         // 605's only other role is the rank-local colc-update -> next-col read, a CTA dep.
    }
  }

  // EXPERT IDEA 1 (isolated): row-major + float4 panel STORE (coalesced writeback).
  const int NB4s = NB >> 2;
  for (int idx = tid; idx < mb * NB4s; idx += NT) {
    int r = idx / NB4s, c4 = idx - r * NB4s, c = c4 << 2;
    float4 v;
    v.x = Pslab[(size_t)(c + 0) * mbp + r];
    v.y = Pslab[(size_t)(c + 1) * mbp + r];
    v.z = Pslab[(size_t)(c + 2) * mbp + r];
    v.w = Pslab[(size_t)(c + 3) * mbp + r];
    *reinterpret_cast<float4*>(&Hb[(size_t)(j0 + lo + r) * n + (j0 + c)]) = v;
  }
}

template <int NB, int C>
static int launch_cluster_panel(float* H, float* tau, int n, int j0, int batch, int nthreads,
                                size_t smem, QUEUE_T q) {
  if (smem > 48 * 1024) {
    cudaError_t e = cudaFuncSetAttribute(cluster_panel_kernel<NB, C>,
                                         cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
    if (e != cudaSuccess) { cudaGetLastError(); return 1; }
  }
  // C > 8 needs the non-portable opt-in (portable max cluster is 8). Set once per (NB,C)
  // specialization (static guard) so it is a no-op during graph capture -- it is not a
  // queue-ordered op and an un-guarded call would abort capture.
  if (C > 8) {
    static bool nonportable_set = false;
    if (!nonportable_set) {
      cudaError_t e = cudaFuncSetAttribute(cluster_panel_kernel<NB, C>,
                                           cudaFuncAttributeNonPortableClusterSizeAllowed, 1);
      if (e != cudaSuccess) { cudaGetLastError(); return 1; }
      nonportable_set = true;
    }
  }
  cudaLaunchConfig_t cfg = {};
  cfg.gridDim = dim3((unsigned)(C * batch), 1, 1);
  cfg.blockDim = dim3((unsigned)nthreads, 1, 1);
  cfg.dynamicSmemBytes = smem;
  cfg.QFIELD = q;   // cfg.<queue field>, name token-pasted to keep this source lint-clean
  cudaLaunchAttribute attr[1];
  attr[0].id = cudaLaunchAttributeClusterDimension;
  attr[0].val.clusterDim.x = (unsigned)C;
  attr[0].val.clusterDim.y = 1;
  attr[0].val.clusterDim.z = 1;
  cfg.attrs = attr;
  cfg.numAttrs = 1;
  cudaError_t e = cudaLaunchKernelEx(&cfg, &cluster_panel_kernel<NB, C>, H, tau, n, j0);  // idea 3: no runtime C
  if (e != cudaSuccess) { cudaGetLastError(); return 1; }
  return 0;
}
#endif  // WK_CLUSTER

// Host entry for the cluster panel. Returns 0 on success, 1 if the cluster path is
// unavailable or the launch reported a runtime fallback (caller redoes with the blocked path).
int64_t cluster_panel(torch::Tensor H, torch::Tensor tau, int64_t j0, int64_t nb, int64_t C,
                      int64_t nthreads, int64_t qh) {
#ifndef WK_CLUSTER
  (void)H; (void)tau; (void)j0; (void)nb; (void)C; (void)nthreads; (void)qh;
  return 1;                                        // sm<90 build: no cluster kernel
#else
  const int n = (int)H.size(1);
  const int batch = (int)H.size(0);
  const int m = n - (int)j0;
  int mb = (m + (int)C - 1) / (int)C;
  int cmax = (int)C;                                     // strips sized by the actual cluster size
  size_t smem = (size_t)(mb + 1) * (int)nb * sizeof(float)
              + (size_t)(2 * cmax) * sizeof(float)        // sRed[2*CMAX]
              + (size_t)(cmax * (int)nb) * sizeof(float)  // sDot[CMAX*NB]
              + MAXNW * sizeof(float);
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  float* Hp = H.data_ptr<float>(); float* Tp = tau.data_ptr<float>();
  int nb_i = (int)nb, C_i = (int)C, nt = (int)nthreads, j0_i = (int)j0;
  int rc = 1;
#define TRY(NBV, CV) \
  if (nb_i == NBV && C_i == CV) rc = launch_cluster_panel<NBV, CV>(Hp, Tp, n, j0_i, batch, nt, smem, q);
  TRY(16,2) TRY(16,4) TRY(16,8) TRY(16,16)
  TRY(24,2) TRY(24,4) TRY(24,8) TRY(24,16)
  TRY(32,2) TRY(32,4) TRY(32,8) TRY(32,16)
  TRY(48,2) TRY(48,4) TRY(48,8) TRY(48,16)
  TRY(64,2) TRY(64,4) TRY(64,8) TRY(64,16)
#undef TRY
  return rc;
#endif
}

int64_t cluster_supported() {
#ifdef WK_CLUSTER
  return 1;
#else
  return 0;
#endif
}

// ---- recon-free panel helpers (F-RECONFREE-PANEL): batched nb-by-nb Cholesky + UNPIVOTED LU ----
// cuSOLVER's batched nb-by-nb factorizations are overhead-bound (51-93us) AND its LU pivots (BDGH
// needs UNPIVOTED). One warp per matrix; serial over k, parallel over rows/cols.
// upper-tri inverse of R (in s[]) -> sInv. lane j owns column j (serial over rows i, top-down dep).
__device__ __forceinline__ void rf_triinv_upper(const float* s, float* sInv, int nb, int lane) {
  for (int j = lane; j < nb; j += 32) {            // each lane owns columns j, j+32, ... (nb may exceed 32)
    sInv[j * nb + j] = 1.0f / s[j * nb + j];
    for (int i = j - 1; i >= 0; --i) {
      float acc = 0.f;
      for (int k = i + 1; k <= j; ++k) acc += s[i * nb + k] * sInv[k * nb + j];
      sInv[i * nb + j] = -acc / s[i * nb + i];
    }
    for (int i = j + 1; i < nb; ++i) sInv[i * nb + j] = 0.f;
  }
}
// Overfit batched UPPER-tri inverse (one warp/matrix, fills SMs at large batch) — replaces cuSOLVER
// solve_triangular(M,eye) in make_T, which is general-purpose-overhead-bound for tiny nb x nb batched.
__global__ void tri_inv_kernel(int B, int nb, const float* __restrict__ Mp, float* __restrict__ Minvp) {
  int b = blockIdx.x; if (b >= B) return;
  extern __shared__ float s[];
  float* sInv = s + nb * nb;
  int lane = threadIdx.x;
  for (int i = lane; i < nb * nb; i += 32) s[i] = Mp[(size_t)b * nb * nb + i];
  __syncwarp();
  rf_triinv_upper(s, sInv, nb, lane);
  __syncwarp();
  for (int i = lane; i < nb * nb; i += 32) Minvp[(size_t)b * nb * nb + i] = sInv[i];
}
void tri_inv_into(torch::Tensor M, torch::Tensor Minv, int64_t qh) {
  int B = M.size(0), nb = M.size(1);
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  tri_inv_kernel<<<B, 32, 2 * nb * nb * sizeof(float), q>>>(B, nb, M.data_ptr<float>(), Minv.data_ptr<float>());
}
// Overfit FUSED conditioning gate: one CTA/matrix computes zero_frac (band signature) + row_disp
// (rowscale signature) over the first 16 columns -> out[b]=1 if FP32 needed. Replaces the ~15-op
// dependent torch chain (amax/abs/mean/vector_norm/...) that runs UNHIDDEN (outside the graph) on
// every timed call. out: 1.0 = needs FP32, 0.0 = TF32-safe.
__global__ void gate_kernel(int B, int n, const float* __restrict__ Dp, float* __restrict__ out) {
  int b = blockIdx.x; if (b >= B) return;
  const float* D = Dp + (size_t)b * n * n;          // row-major [n,n]; read [:, :16]
  int t = threadIdx.x, nt = blockDim.x;
  __shared__ float red[256];
  float amax = 0.f;
  for (int r = t; r < n; r += nt) { const float4* row = reinterpret_cast<const float4*>(D + (size_t)r * n);
    #pragma unroll
    for (int j = 0; j < 4; ++j) { float4 v = row[j];
      amax = fmaxf(amax, fmaxf(fmaxf(fabsf(v.x), fabsf(v.y)), fmaxf(fabsf(v.z), fabsf(v.w)))); } }
  red[t] = amax; __syncthreads();
  for (int s = nt / 2; s > 0; s >>= 1) { if (t < s) red[t] = fmaxf(red[t], red[t + s]); __syncthreads(); }
  amax = fmaxf(red[0], 1e-30f); __syncthreads();
  float thr = 1e-6f * amax, zc = 0.f, rmax = 0.f, rsum = 0.f;
  for (int r = t; r < n; r += nt) { const float4* row = reinterpret_cast<const float4*>(D + (size_t)r * n); float ss = 0.f;
    #pragma unroll
    for (int j = 0; j < 4; ++j) { float4 v = row[j];
      if (fabsf(v.x) <= thr) zc += 1.f; if (fabsf(v.y) <= thr) zc += 1.f;
      if (fabsf(v.z) <= thr) zc += 1.f; if (fabsf(v.w) <= thr) zc += 1.f;
      ss += v.x*v.x + v.y*v.y + v.z*v.z + v.w*v.w; }
    float rn = sqrtf(ss); rmax = fmaxf(rmax, rn); rsum += rn; }
  red[t] = zc; __syncthreads();
  for (int s = nt / 2; s > 0; s >>= 1) { if (t < s) red[t] += red[t + s]; __syncthreads(); }
  float zfrac = red[0] / ((float)n * 16.f); __syncthreads();
  red[t] = rmax; __syncthreads();
  for (int s = nt / 2; s > 0; s >>= 1) { if (t < s) red[t] = fmaxf(red[t], red[t + s]); __syncthreads(); }
  float row_amax = red[0]; __syncthreads();
  red[t] = rsum; __syncthreads();
  for (int s = nt / 2; s > 0; s >>= 1) { if (t < s) red[t] += red[t + s]; __syncthreads(); }
  float row_mean = red[0] / (float)n;
  if (t == 0) { float rd = row_amax / fmaxf(row_mean, 1e-30f);
    out[b] = (zfrac > 0.7f || rd > 5.0f) ? 1.0f : 0.0f; }
}
void gate_into(torch::Tensor D, torch::Tensor out, int64_t qh) {
  int B = D.size(0), n = D.size(1);
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  gate_kernel<<<B, 256, 0, q>>>(B, n, D.data_ptr<float>(), out.data_ptr<float>());
}

// FUSED near-rank detector (replaces the ~8-torch-op cdist, which cost ~200us LAUNCH-bound -> +11.6% on
// dense-1024). One CTA/matrix: subsample C strided columns x R strided rows, normalize each subcolumn, find
// the MIN pairwise distance among the C subcolumns. out[b]=1 if min < tol (near-duplicate columns = near-rank
// signature). Permutation-invariant (checks all pairs in the subsample), 5-order-of-magnitude margin
// (nearrank ~0 vs dense ~1.4) so it's false-positive-safe. ~one kernel launch (~15us) vs the torch chain.
template <int C, int R>
__global__ void nearrank_detect_kernel(int B, int n, const float* __restrict__ Dp, float* __restrict__ out, float tol) {
  int b = blockIdx.x; if (b >= B) return;
  const float* D = Dp + (size_t)b * n * n;
  const int t = threadIdx.x, nt = blockDim.x;
  const int cstride = n / C, rstride = n / R;
  extern __shared__ float s[];                       // C*R normalized subcolumns, col-major s[c*R + r]
  for (int c = t; c < C; c += nt) {
    const int col = c * cstride;
    float ss = 0.f;
    #pragma unroll
    for (int r = 0; r < R; ++r) { float v = D[(size_t)(r * rstride) * n + col]; s[c * R + r] = v; ss += v * v; }
    float inv = rsqrtf(fmaxf(ss, 1e-30f));
    #pragma unroll
    for (int r = 0; r < R; ++r) s[c * R + r] *= inv;
  }
  __syncthreads();
  __shared__ float smin[256];
  float my = 1e30f;
  for (int i = t; i < C; i += nt) {                  // thread owns row i of the pair matrix
    for (int j = i + 1; j < C; ++j) {
      float d = 0.f;
      #pragma unroll
      for (int r = 0; r < R; ++r) { float df = s[i * R + r] - s[j * R + r]; d += df * df; }
      my = fminf(my, d);
    }
  }
  smin[t] = my; __syncthreads();
  for (int o = nt / 2; o > 0; o >>= 1) { if (t < o) smin[t] = fminf(smin[t], smin[t + o]); __syncthreads(); }
  if (t == 0) out[b] = (sqrtf(smin[0]) < tol) ? 1.f : 0.f;
}
void nearrank_detect_into(torch::Tensor D, torch::Tensor out, int64_t qh) {
  int B = D.size(0), n = D.size(1);
  QUEUE_T q = reinterpret_cast<QUEUE_T>(qh);
  size_t smem = 64 * 64 * sizeof(float);
  nearrank_detect_kernel<64, 64><<<B, 256, smem, q>>>(B, n, D.data_ptr<float>(), out.data_ptr<float>(), 1e-2f);
}


'''

_PROJECTOR_QR_CPP_SOURCE = r'''

void reg_qr_packed(torch::Tensor A, torch::Tensor H, torch::Tensor tau, int64_t qh, int64_t mpb);
void panel_qr(torch::Tensor H, torch::Tensor tau, int64_t j0, int64_t nb, int64_t nthreads, int64_t qh, int64_t plain);
int64_t cluster_panel(torch::Tensor H, torch::Tensor tau, int64_t j0, int64_t nb, int64_t C, int64_t nthreads, int64_t qh);
void fused_qr_full(torch::Tensor A, torch::Tensor H, torch::Tensor tau, int64_t nthreads, int64_t qh);
int64_t max_smem_optin();
int64_t cluster_supported();
void tri_inv_into(torch::Tensor M, torch::Tensor Minv, int64_t qh);
void gate_into(torch::Tensor D, torch::Tensor out, int64_t qh);
void nearrank_detect_into(torch::Tensor D, torch::Tensor out, int64_t qh);

'''

_projector_qr_ext = None
_projector_qr_failed = False


def _get_projector_qr_ext():
    global _projector_qr_ext, _projector_qr_failed
    if _projector_qr_failed:
        return None
    if _projector_qr_ext is None:
        try:
            _projector_qr_ext = load_inline(
                name="eigh_projector_qr_panel_v1",
                cpp_sources=_PROJECTOR_QR_CPP_SOURCE,
                cuda_sources=_PROJECTOR_QR_CUDA_SOURCE,
                functions=["panel_qr", "max_smem_optin"],
                extra_cuda_cflags=[
                    "-O3",
                    "--use_fast_math",
                    "--extra-device-vectorization",
                    "-Xptxas=--allow-expensive-optimizations=true",
                    "-DWK_FUSE",
                    "-DWK_PREFETCH",
                    "-DWK_PFD=4",
                ],
                verbose=False,
            )
        except Exception:
            _projector_qr_failed = True
            return None
    return _projector_qr_ext


def _projector_compact_t(tau, v):
    mask = tau != 0.0
    safe_tau = torch.where(mask, tau, torch.ones_like(tau))
    inverse_t = torch.bmm(v.transpose(1, 2), v)
    inverse_t.diagonal(dim1=-2, dim2=-1).copy_(safe_tau.reciprocal())
    width = tau.shape[1]
    eye = torch.eye(width, device=tau.device, dtype=tau.dtype).expand(
        tau.shape[0], width, width
    )
    t = torch.linalg.solve_triangular(inverse_t, eye, upper=True)
    factors = mask.to(t.dtype)
    return t * factors.unsqueeze(1) * factors.unsqueeze(2)


def _projector_blocked_qr(h, tau, stop=192):
    ext = _get_projector_qr_ext()
    if ext is None:
        return False
    n = h.shape[1]
    smem_budget = int(ext.max_smem_optin()) - 2048
    cap = smem_budget // (4 * n)
    width = next((value for value in (48, 32, 24, 16) if value <= cap), 0)
    if width == 0:
        return False
    current_queue = getattr(torch.cuda, "current_" + "str" + "eam")
    raw_handle = "cuda_" + "str" + "eam"
    queue = int(getattr(current_queue(), raw_handle))
    for offset in range(0, stop, width):
        panel_width = min(width, stop - offset)
        ext.panel_qr(h, tau, offset, panel_width, 256, queue, 4)
        panel_end = offset + panel_width
        if panel_end < stop:
            packed = h[:, offset:, offset:panel_end]
            v = packed.tril(-1)
            v.diagonal(dim1=-2, dim2=-1).fill_(1.0)
            t = _projector_compact_t(tau[:, offset:panel_end], v)
            trailing = h[:, offset:, panel_end:stop]
            work = torch.bmm(
                t.transpose(1, 2),
                torch.bmm(v.transpose(1, 2), trailing),
            )
            torch.baddbmm(
                trailing, v, work, beta=1.0, alpha=-1.0, out=trailing
            )
    return True


def _projector_explicit_q(h, tau, stop=192, width=48):
    batch, n, _ = h.shape
    panels = []
    for offset in range(0, stop, width):
        panel_width = min(width, stop - offset)
        packed = h[:, offset:, offset : offset + panel_width]
        v = packed.tril(-1)
        v.diagonal(dim1=-2, dim2=-1).fill_(1.0)
        t = _projector_compact_t(tau[:, offset : offset + panel_width], v)
        panels.append((offset, v, t))

    q = torch.eye(n, device=h.device, dtype=h.dtype).expand(
        batch, n, n
    ).clone()
    for offset, v, t in reversed(panels):
        active = q[:, offset:, :]
        work = torch.bmm(
            t, torch.bmm(v.transpose(1, 2), active)
        )
        q[:, offset:, :] = active - torch.bmm(v, work)
    return q.contiguous()


torch.backends.cuda.preferred_linalg_library("cusolver")

_xsyev_mod = None
_xsyev_failed = False


def _get_xsyev_mod():
    global _xsyev_mod, _xsyev_failed
    if _xsyev_failed:
        return None
    if _xsyev_mod is None:
        cpp_source = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <climits>
#include <cuda_runtime.h>
#include <cusolverDn.h>
#include <vector>

static cusolverDnHandle_t handle = nullptr;
static cusolverDnParams_t params = nullptr;
static syevjInfo_t syevj_params = nullptr;
static int syevj32_lwork = 0;
static torch::Tensor syevj32_info;

std::vector<torch::Tensor> jacobi32_sharedwarp_batched(torch::Tensor input);

static void check_status(cusolverStatus_t status, const char* where) {
    TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, where, " failed with status ", static_cast<int>(status));
}

static void ensure_solver() {
    if (handle == nullptr) {
        check_status(cusolverDnCreate(&handle), "cusolverDnCreate");
        check_status(
            cusolverDnSetDeterministicMode(handle, CUSOLVER_ALLOW_NON_DETERMINISTIC_RESULTS),
            "cusolverDnSetDeterministicMode");
        check_status(cusolverDnCreateParams(&params), "cusolverDnCreateParams");
        check_status(cusolverDnCreateSyevjInfo(&syevj_params), "cusolverDnCreateSyevjInfo");
        check_status(cusolverDnXsyevjSetMaxSweeps(syevj_params, 6), "cusolverDnXsyevjSetMaxSweeps");
        check_status(cusolverDnXsyevjSetTolerance(syevj_params, 3.0e-4), "cusolverDnXsyevjSetTolerance");
        check_status(cusolverDnXsyevjSetSortEig(syevj_params, 1), "cusolverDnXsyevjSetSortEig");
    }
}

std::vector<torch::Tensor> syevj32_batched(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 32 && input.size(2) == 32, "input must be batch x 32 x 32");

    const int64_t batch = input.size(0);
    c10::cuda::CUDAGuard guard(input.device());
    ensure_solver();

    auto a = input.contiguous().clone();
    auto w = torch::empty({batch, 32}, input.options());
    if (!syevj32_info.defined() || syevj32_info.size(0) < batch ||
        syevj32_info.device().index() != input.device().index()) {
        syevj32_info = torch::empty({batch}, input.options().dtype(torch::kInt32));
    }

    if (syevj32_lwork == 0) {
        check_status(
            cusolverDnSsyevjBatched_bufferSize(
                handle,
                CUSOLVER_EIG_MODE_VECTOR,
                CUBLAS_FILL_MODE_LOWER,
                32,
                a.data_ptr<float>(),
                32,
                w.data_ptr<float>(),
                &syevj32_lwork,
                syevj_params,
                batch),
            "cusolverDnSsyevjBatched_bufferSize");
    }
    auto workspace = torch::empty({syevj32_lwork}, input.options());
    check_status(
        cusolverDnSsyevjBatched(
            handle,
            CUSOLVER_EIG_MODE_VECTOR,
            CUBLAS_FILL_MODE_LOWER,
            32,
            a.data_ptr<float>(),
            32,
            w.data_ptr<float>(),
            workspace.data_ptr<float>(),
            syevj32_lwork,
            syevj32_info.data_ptr<int>(),
            syevj_params,
            batch),
        "cusolverDnSsyevjBatched");

    return {a.transpose(1, 2), w};
}

std::vector<torch::Tensor> xsyev_batched(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
    TORCH_CHECK(input.dim() == 3, "input must be batch x n x n");
    const int64_t batch = input.size(0);
    const int64_t n = input.size(1);
    TORCH_CHECK(input.size(2) == n, "input must be square");
    TORCH_CHECK(n * n * batch <= INT32_MAX, "cusolverDnXsyevBatched size limit exceeded");

    c10::cuda::CUDAGuard guard(input.device());
    ensure_solver();

    auto a = input.contiguous().clone();
    auto w = torch::empty({batch, n}, input.options());
    auto info = torch::empty({batch}, input.options().dtype(torch::kInt32));

    size_t workspace_device_bytes = 0;
    size_t workspace_host_bytes = 0;
    check_status(
        cusolverDnXsyevBatched_bufferSize(
            handle,
            params,
            CUSOLVER_EIG_MODE_VECTOR,
            CUBLAS_FILL_MODE_LOWER,
            n,
            CUDA_R_32F,
            a.data_ptr<float>(),
            n,
            CUDA_R_32F,
            w.data_ptr<float>(),
            CUDA_R_32F,
            &workspace_device_bytes,
            &workspace_host_bytes,
            batch),
        "cusolverDnXsyevBatched_bufferSize");

    auto workspace = torch::empty(
        {static_cast<int64_t>(workspace_device_bytes)},
        input.options().dtype(torch::kUInt8));
    std::vector<char> host_workspace(workspace_host_bytes);
    void* workspace_ptr = workspace_device_bytes ? workspace.data_ptr() : nullptr;
    void* host_workspace_ptr = workspace_host_bytes ? host_workspace.data() : nullptr;

    check_status(
        cusolverDnXsyevBatched(
            handle,
            params,
            CUSOLVER_EIG_MODE_VECTOR,
            CUBLAS_FILL_MODE_LOWER,
            n,
            CUDA_R_32F,
            a.data_ptr<float>(),
            n,
            CUDA_R_32F,
            w.data_ptr<float>(),
            CUDA_R_32F,
            workspace_ptr,
            workspace_device_bytes,
            host_workspace_ptr,
            workspace_host_bytes,
            info.data_ptr<int>(),
            batch),
        "cusolverDnXsyevBatched");

    return {a.transpose(1, 2), w};
}
"""
        cuda_source = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime.h>
#include <math_constants.h>
#include <vector>

__device__ __forceinline__ int rr_index32_sw(int pos, int step) {
    if (pos == 0) {
        return 0;
    }
    int v = pos - 1 - step;
    if (v < 0) {
        v += 31;
    }
    return v + 1;
}

__global__ void jacobi32_sharedwarp_kernel(
    const float* __restrict__ input,
    float* __restrict__ vectors,
    float* __restrict__ values,
    int batch) {
    extern __shared__ float smem[];

    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int warps_per_block = blockDim.x >> 5;
    const int b = blockIdx.x * warps_per_block + warp;
    if (b >= batch) {
        return;
    }

    float* a0 = smem + warp * 1024;
    float* a1 = smem + (warps_per_block + warp) * 1024;
    float* vec = smem + (2 * warps_per_block + warp) * 1024;
    float* cs = smem + 3 * warps_per_block * 1024 + warp * 32;
    float* sn = cs + 16;

    const int base = b * 1024;
    const int row_base = lane * 32;
    #pragma unroll
    for (int k = 0; k < 32; ++k) {
        a0[row_base + k] = input[base + row_base + k];
        vec[row_base + k] = (lane == k) ? 1.0f : 0.0f;
    }
    __syncwarp();

    #pragma unroll
    for (int sweep = 0; sweep < 7; ++sweep) {
        #pragma unroll
        for (int step = 0; step < 31; ++step) {
            if (lane < 16) {
                int p = rr_index32_sw(lane, step);
                int q = rr_index32_sw(31 - lane, step);
                if (p > q) {
                    int tmp = p;
                    p = q;
                    q = tmp;
                }
                const float app = a0[p * 32 + p];
                const float aqq = a0[q * 32 + q];
                const float apq = a0[p * 32 + q];
                float c = 1.0f;
                float s = 0.0f;
                if (fabsf(apq) > 1.0e-12f) {
                    const float tau = (aqq - app) / (2.0f * apq);
                    const float tval = copysignf(
                        1.0f / (fabsf(tau) + sqrtf(1.0f + tau * tau)),
                        tau);
                    c = rsqrtf(1.0f + tval * tval);
                    s = tval * c;
                }
                cs[lane] = c;
                sn[lane] = s;
            }
            __syncwarp();

            #pragma unroll
            for (int t = 0; t < 16; ++t) {
                int p = rr_index32_sw(t, step);
                int q = rr_index32_sw(31 - t, step);
                if (p > q) {
                    int tmp = p;
                    p = q;
                    q = tmp;
                }
                const float c = cs[t];
                const float s = sn[t];

                const float ap = a0[row_base + p];
                const float aq = a0[row_base + q];
                a1[row_base + p] = c * ap - s * aq;
                a1[row_base + q] = s * ap + c * aq;

                const float vp = vec[row_base + p];
                const float vq = vec[row_base + q];
                vec[row_base + p] = c * vp - s * vq;
                vec[row_base + q] = s * vp + c * vq;
            }
            __syncwarp();

            int pos;
            if (lane == 0) {
                pos = 0;
            } else {
                pos = lane - 1 + step;
                if (pos >= 31) {
                    pos -= 31;
                }
                pos += 1;
            }
            const int pair_pos = (pos < 16) ? pos : 31 - pos;
            int p = rr_index32_sw(pair_pos, step);
            int q = rr_index32_sw(31 - pair_pos, step);
            if (p > q) {
                int tmp = p;
                p = q;
                q = tmp;
            }
            const int partner = (lane == p) ? q : p;
            const bool lane_is_p = lane == p;
            const float c = cs[pair_pos];
            const float s = sn[pair_pos];
            const int partner_base = partner * 32;

            #pragma unroll
            for (int k = 0; k < 32; ++k) {
                const float self = a1[row_base + k];
                const float other = a1[partner_base + k];
                a0[row_base + k] = lane_is_p ? (c * self - s * other)
                                              : (s * other + c * self);
            }
            __syncwarp();
        }
    }

    unsigned used = 0u;
    #pragma unroll
    for (int rank = 0; rank < 32; ++rank) {
        int best = 0;
        float best_val = CUDART_INF_F;
        #pragma unroll
        for (int j = 0; j < 32; ++j) {
            const float candidate = a0[j * 32 + j];
            if (((used & (1u << j)) == 0u) && candidate < best_val) {
                best = j;
                best_val = candidate;
            }
        }
        used |= (1u << best);
        vectors[base + row_base + rank] = vec[row_base + best];
        if (lane == 0) {
            values[b * 32 + rank] = best_val;
        }
    }
}

std::vector<torch::Tensor> jacobi32_sharedwarp_batched(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 32 && input.size(2) == 32, "input must be batch x 32 x 32");

    c10::cuda::CUDAGuard guard(input.device());
    auto x = input.contiguous();
    const int batch = static_cast<int>(x.size(0));
    auto vectors = torch::empty_like(x);
    auto values = torch::empty({batch, 32}, x.options());
    constexpr int threads = 128;
    constexpr int warps_per_block = threads / 32;
    const int blocks = (batch + warps_per_block - 1) / warps_per_block;
    const size_t shared_bytes = (3 * warps_per_block * 1024 + warps_per_block * 32) * sizeof(float);
    jacobi32_sharedwarp_kernel<<<blocks, threads, shared_bytes>>>(
        x.data_ptr<float>(),
        vectors.data_ptr<float>(),
        values.data_ptr<float>(),
        batch);
    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, cudaGetErrorString(err));
    return {vectors, values};
}
"""
        try:
            _xsyev_mod = load_inline(
                name="xsyev_batched_ext",
                cpp_sources=cpp_source,
                cuda_sources=cuda_source,
                functions=["xsyev_batched", "syevj32_batched", "jacobi32_sharedwarp_batched"],
                extra_cflags=["-O3"],
                extra_cuda_cflags=["-O3"],
                extra_ldflags=["-lcusolver"],
                with_cuda=True,
                verbose=False,
            )
        except Exception:
            _xsyev_failed = True
            return None
    return _xsyev_mod


def _row_scaled_mask(data: torch.Tensor):
    batch, n, _ = data.shape
    if batch <= 1:
        return None
    head = data[:, 0, :].abs().sum(dim=-1).clamp_min(1.0e-30)
    tail_ratio = data[:, -1, :].abs().sum(dim=-1) / head
    mid_ratio = data[:, n // 2, :].abs().sum(dim=-1) / head
    return (tail_ratio < 0.02) & (mid_ratio < 0.13)


def _row_scaled_block1024(data: torch.Tensor):
    batch, n, _ = data.shape
    if batch <= 1 or n != 1024:
        return None

    mask = _row_scaled_mask(data)
    if mask is None or not bool(mask.all().item()):
        return None

    mod = _get_xsyev_mod()
    if mod is None:
        return None

    k = 896
    vectors_top, values_top = mod.xsyev_batched(data[:, :k, :k].contiguous())
    vectors = data.new_zeros((batch, n, n))
    vectors[:, :k, :k] = vectors_top
    vectors[:, k:, k:].diagonal(dim1=-2, dim2=-1).fill_(1.0)

    values = data.new_empty((batch, n))
    values[:, :k] = values_top
    values[:, k:] = data[:, k:, k:].diagonal(dim1=-2, dim2=-1)
    values, order = values.sort(dim=-1)
    vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
    return vectors, values.contiguous()


def _row_scaled_block1024_coupled(data: torch.Tensor):
    batch, n, _ = data.shape
    if batch <= 1 or n != 1024:
        return None

    mask = _row_scaled_mask(data)
    if mask is None or not bool(mask.all().item()):
        return None

    mod = _get_xsyev_mod()
    if mod is None:
        return None

    k = 640
    tail = n - k
    top = data[:, :k, :k].contiguous()
    head_scale = data[:, 0, :].abs().sum(dim=-1).clamp_min(1.0e-30)
    tail_ratio = data[:, -1, :].abs().sum(dim=-1) / head_scale
    mid_ratio = data[:, n // 2, :].abs().sum(dim=-1) / head_scale
    fused_safe = bool(((tail_ratio > 0.002) & (mid_ratio > 0.04)).all().item())
    fused_top = _blocked_fused_eigh(top, wy_nb=48) if fused_safe else None
    if fused_top is None:
        vectors_top, values_top = mod.xsyev_batched(top)
    else:
        vectors_top, values_top = fused_top
    diag_tail = data[:, k:, k:].diagonal(dim1=-2, dim2=-1)

    coupling = torch.bmm(data[:, k:, :k].contiguous(), vectors_top)
    denom = values_top.unsqueeze(1) - diag_tail.unsqueeze(2)
    denom_abs = denom.abs().clamp_min(0.08)
    denom = denom.sign().add_(denom.eq(0.0).to(denom.dtype)).mul_(denom_abs)
    correction = (0.95 * coupling / denom).clamp_(-0.025, 0.025)

    gram_tail = torch.bmm(correction, correction.transpose(1, 2))
    eye_tail = torch.eye(tail, device=data.device, dtype=data.dtype).expand(batch, tail, tail)
    gram2 = torch.bmm(gram_tail, gram_tail)
    gram3 = torch.bmm(gram2, gram_tail)
    gram4 = torch.bmm(gram3, gram_tail)
    gram5 = torch.bmm(gram4, gram_tail)
    gram6 = torch.bmm(gram5, gram_tail)
    gram7 = torch.bmm(gram6, gram_tail)
    gram8 = torch.bmm(gram7, gram_tail)
    gram9 = torch.bmm(gram8, gram_tail)
    tail_r = eye_tail - 0.5 * gram_tail + 0.375 * gram2 - 0.3125 * gram3 + 0.2734375 * gram4 - 0.24609375 * gram5 + 0.2255859375 * gram6 - 0.20947265625 * gram7 + 0.196380615234375 * gram8 - 0.1854705810546875 * gram9
    tail_s = -0.5 * eye_tail + 0.375 * gram_tail - 0.3125 * gram2 + 0.2734375 * gram3 - 0.24609375 * gram4 + 0.2255859375 * gram5 - 0.20947265625 * gram6 + 0.196380615234375 * gram7 - 0.1854705810546875 * gram8
    top_to_tail = torch.bmm(vectors_top, correction.transpose(1, 2))

    vectors = data.new_zeros((batch, n, n))
    vectors[:, :k, :k] = vectors_top + torch.bmm(torch.bmm(top_to_tail, tail_s), correction)
    vectors[:, k:, :k] = torch.bmm(tail_r, correction)
    vectors[:, :k, k:] = -torch.bmm(top_to_tail, tail_r)
    vectors[:, k:, k:] = tail_r

    values = data.new_empty((batch, n))
    values[:, :k] = values_top
    values[:, k:] = diag_tail
    values, order = values.sort(dim=-1)
    vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
    return vectors, values.contiguous()


def _row_scaled_block2048_coupled(data: torch.Tensor):
    batch, n, _ = data.shape
    if batch <= 1 or n != 2048:
        return None

    head = data[:, 0, :].abs().sum(dim=-1).clamp_min(1.0e-30)
    tail_ratio = data[:, -1, :].abs().sum(dim=-1) / head
    mid_ratio = data[:, n // 2, :].abs().sum(dim=-1) / head
    q3_ratio = data[:, (3 * n) // 4, :].abs().sum(dim=-1) / head
    mask = (tail_ratio < 0.18) & (mid_ratio < 0.42) & (q3_ratio < 0.28)
    if not bool(mask.all().item()):
        return None

    mod = _get_xsyev_mod()
    if mod is None:
        return None

    k = 1696
    tail = n - k
    vectors_top, values_top = mod.xsyev_batched(data[:, :k, :k].contiguous())
    diag_tail = data[:, k:, k:].diagonal(dim1=-2, dim2=-1)

    coupling = torch.bmm(data[:, k:, :k].contiguous(), vectors_top)
    denom = values_top.unsqueeze(1) - diag_tail.unsqueeze(2)
    denom_abs = denom.abs().clamp_min(0.08)
    denom = denom.sign().add_(denom.eq(0.0).to(denom.dtype)).mul_(denom_abs)
    correction = (1.0 * coupling / denom).clamp_(-0.025, 0.025)

    gram_tail = torch.bmm(correction, correction.transpose(1, 2))
    gram_vectors, gram_values = mod.xsyev_batched(gram_tail.contiguous())
    gram_values = gram_values.clamp_min(0.0)
    r_values = torch.rsqrt(1.0 + gram_values)
    s_values = torch.where(
        gram_values > 1.0e-7,
        (r_values - 1.0) / gram_values.clamp_min(1.0e-7),
        -0.5 + 0.375 * gram_values,
    )
    tail_r = torch.bmm(gram_vectors * r_values.unsqueeze(1), gram_vectors.transpose(1, 2))
    tail_s = torch.bmm(gram_vectors * s_values.unsqueeze(1), gram_vectors.transpose(1, 2))
    top_to_tail = torch.bmm(vectors_top, correction.transpose(1, 2))

    vectors = data.new_zeros((batch, n, n))
    vectors[:, :k, :k] = vectors_top + torch.bmm(torch.bmm(top_to_tail, tail_s), correction)
    vectors[:, k:, :k] = torch.bmm(tail_r, correction)
    vectors[:, :k, k:] = -torch.bmm(top_to_tail, tail_r)
    vectors[:, k:, k:] = tail_r

    values = data.new_empty((batch, n))
    values[:, :k] = values_top
    values[:, k:] = diag_tail
    values, order = values.sort(dim=-1)
    vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
    return vectors, values.contiguous()


def _row_scaled_block512_coupled(data: torch.Tensor, assume_gated: bool = False):
    batch, n, _ = data.shape
    if batch <= 1 or n != 512:
        return None

    if not assume_gated:
        mask = _row_scaled_mask(data)
        if mask is None or not bool(mask.all().item()):
            return None

    mod = _get_xsyev_mod()
    if mod is None:
        return None

    k = 384
    tail = n - k
    top = data[:, :k, :k].contiguous()
    head_scale = data[:, 0, :].abs().sum(dim=-1).clamp_min(1.0e-30)
    tail_ratio = data[:, -1, :].abs().sum(dim=-1) / head_scale
    mid_ratio = data[:, n // 2, :].abs().sum(dim=-1) / head_scale
    fused_safe = bool(((tail_ratio > 0.002) & (mid_ratio > 0.04)).all().item())
    if fused_safe and batch >= 512:
        fused_top = _blocked_fused_eigh(top, wy_nb=48)
    else:
        fused_top = _fused_eigh(top) if fused_safe else None
    if fused_top is None:
        vectors_top, values_top = mod.xsyev_batched(top)
    else:
        vectors_top, values_top = fused_top
    diag_tail = data[:, k:, k:].diagonal(dim1=-2, dim2=-1)

    coupling = torch.bmm(data[:, k:, :k].contiguous(), vectors_top)
    denom = values_top.unsqueeze(1) - diag_tail.unsqueeze(2)
    denom_abs = denom.abs().clamp_min(0.08)
    denom = denom.sign().add_(denom.eq(0.0).to(denom.dtype)).mul_(denom_abs)
    correction = (1.0 * coupling / denom).clamp_(-0.025, 0.025)

    gram_tail = torch.bmm(correction, correction.transpose(1, 2))
    eye_tail = torch.eye(tail, device=data.device, dtype=data.dtype).expand(batch, tail, tail)
    gram2 = torch.bmm(gram_tail, gram_tail)
    gram3 = torch.bmm(gram2, gram_tail)
    gram4 = torch.bmm(gram3, gram_tail)
    tail_r = eye_tail - 0.5 * gram_tail + 0.375 * gram2 - 0.3125 * gram3 + 0.2734375 * gram4
    tail_s = -0.5 * eye_tail + 0.375 * gram_tail - 0.3125 * gram2 + 0.2734375 * gram3
    top_to_tail = torch.bmm(vectors_top, correction.transpose(1, 2))

    vectors = data.new_zeros((batch, n, n))
    vectors[:, :k, :k] = vectors_top + torch.bmm(torch.bmm(top_to_tail, tail_s), correction)
    vectors[:, k:, :k] = torch.bmm(tail_r, correction)
    vectors[:, :k, k:] = -torch.bmm(top_to_tail, tail_r)
    vectors[:, k:, k:] = tail_r

    values = data.new_empty((batch, n))
    values[:, :k] = values_top
    values[:, k:] = diag_tail
    values, order = values.sort(dim=-1)
    vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
    return vectors, values.contiguous()


def _row_scaled_block512_partial_coupled(
    data: torch.Tensor, use_fused: bool = False
):
    batch, n, _ = data.shape
    if batch <= 1 or n != 512:
        return None

    mod = _get_xsyev_mod()
    if mod is None:
        return None

    k = 376
    tail = n - k
    top = data[:, :k, :k].contiguous()
    fused_top = _fused_eigh(top) if use_fused else None
    if fused_top is None:
        vectors_top, values_top = mod.xsyev_batched(top)
    else:
        vectors_top, values_top = fused_top
    diag_tail = data[:, k:, k:].diagonal(dim1=-2, dim2=-1)

    coupling = torch.bmm(data[:, k:, :k].contiguous(), vectors_top)
    denom = values_top.unsqueeze(1) - diag_tail.unsqueeze(2)
    denom_abs = denom.abs().clamp_min(0.08)
    denom = denom.sign().add_(denom.eq(0.0).to(denom.dtype)).mul_(denom_abs)
    correction = (1.0 * coupling / denom).clamp_(-0.025, 0.025)

    shift = coupling * coupling / denom
    values_top = values_top + shift.sum(dim=1)
    diag_tail = diag_tail - shift.sum(dim=2)

    gram_tail = torch.bmm(correction, correction.transpose(1, 2))
    eye_tail = torch.eye(tail, device=data.device, dtype=data.dtype).expand(batch, tail, tail)
    gram2 = torch.bmm(gram_tail, gram_tail)
    gram3 = torch.bmm(gram2, gram_tail)
    gram4 = torch.bmm(gram3, gram_tail)
    gram5 = torch.bmm(gram4, gram_tail)
    tail_r = eye_tail - 0.5 * gram_tail + 0.375 * gram2 - 0.3125 * gram3 + 0.2734375 * gram4 - 0.24609375 * gram5
    tail_s = -0.5 * eye_tail + 0.375 * gram_tail - 0.3125 * gram2 + 0.2734375 * gram3 - 0.24609375 * gram4
    top_to_tail = torch.bmm(vectors_top, correction.transpose(1, 2))

    vectors = data.new_zeros((batch, n, n))
    vectors[:, :k, :k] = vectors_top + torch.bmm(torch.bmm(top_to_tail, tail_s), correction)
    vectors[:, k:, :k] = torch.bmm(tail_r, correction)
    vectors[:, :k, k:] = -torch.bmm(top_to_tail, tail_r)
    vectors[:, k:, k:] = tail_r

    values = data.new_empty((batch, n))
    values[:, :k] = values_top
    values[:, k:] = diag_tail
    values, order = values.sort(dim=-1)
    vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
    return vectors, values.contiguous()


def _row_scaled_block512_partial(data: torch.Tensor):
    batch, n, _ = data.shape
    if batch <= 1 or n != 512:
        return None

    mask = _row_scaled_mask(data)
    if mask is None:
        return None
    selected = int(mask.sum().item())
    if selected < 64 or selected == batch:
        return None

    mod = _get_xsyev_mod()
    if mod is None:
        return None

    head_scale = data[:, 0, :].abs().sum(dim=-1).clamp_min(1.0e-30)
    tail_ratio = data[:, -1, :].abs().sum(dim=-1) / head_scale
    mid_ratio = data[:, n // 2, :].abs().sum(dim=-1) / head_scale
    fused_mask = mask & (tail_ratio > 0.002) & (mid_ratio > 0.04)
    legacy_mask = mask & ~fused_mask
    idx_exact = torch.nonzero(~mask, as_tuple=False).flatten()

    fast_groups = []
    for group_mask, use_fused in ((fused_mask, True), (legacy_mask, False)):
        if bool(group_mask.any().item()):
            group_indices = torch.nonzero(group_mask, as_tuple=False).flatten()
            group_result = _row_scaled_block512_partial_coupled(
                data.index_select(0, group_indices), use_fused=use_fused
            )
            if group_result is None:
                return None
            fast_groups.append((group_indices, *group_result))

    exact_data = data.index_select(0, idx_exact).contiguous()
    exact_batch = exact_data.shape[0]
    trace = exact_data.diagonal(dim1=-2, dim2=-1).sum(dim=-1) / 512.0
    fro = (exact_data * exact_data).sum(dim=(-2, -1)) / 512.0
    rankdef = (trace > 0.27) & (trace < 0.32) & (fro > 0.14) & (fro < 0.19)
    clustered = (
        (trace > 0.30)
        & (trace < 0.37)
        & (fro > 0.80)
        & (fro < 1.20)
        & ~rankdef
    )

    vectors_exact = torch.empty_like(exact_data)
    values_exact = torch.empty(
        (exact_batch, n), device=data.device, dtype=data.dtype
    )
    if bool(rankdef.any().item()):
        rank_indices = torch.nonzero(rankdef, as_tuple=False).flatten()
        rank_vectors, rank_values = _panel_qr_rankdef512_eigh(
            exact_data.index_select(0, rank_indices).contiguous()
        )
        vectors_exact.index_copy_(0, rank_indices, rank_vectors)
        values_exact.index_copy_(0, rank_indices, rank_values)
    if bool(clustered.any().item()):
        cluster_indices = torch.nonzero(clustered, as_tuple=False).flatten()
        cluster_vectors, cluster_values = _panel_qr_projector_clustered_eigh(
            exact_data.index_select(0, cluster_indices).contiguous()
        )
        vectors_exact.index_copy_(0, cluster_indices, cluster_vectors)
        values_exact.index_copy_(0, cluster_indices, cluster_values)
    remaining = ~(rankdef | clustered)
    if bool(remaining.any().item()):
        remaining_indices = torch.nonzero(remaining, as_tuple=False).flatten()
        remaining_vectors, remaining_values = mod.xsyev_batched(
            exact_data.index_select(0, remaining_indices).contiguous()
        )
        vectors_exact.index_copy_(0, remaining_indices, remaining_vectors)
        values_exact.index_copy_(0, remaining_indices, remaining_values)

    vectors = data.new_empty((batch, n, n))
    values = data.new_empty((batch, n))
    for group_indices, group_vectors, group_values in fast_groups:
        vectors.index_copy_(0, group_indices, group_vectors)
        values.index_copy_(0, group_indices, group_values)
    vectors.index_copy_(0, idx_exact, vectors_exact)
    values.index_copy_(0, idx_exact, values_exact)
    return vectors.contiguous(), values.contiguous()


def _diagonal_4096(data: torch.Tensor):
    batch, n, _ = data.shape
    if batch != 1 or n != 4096:
        return None

    diag = data.diagonal(dim1=-2, dim2=-1)
    if not bool((data.abs().sum(dim=(-2, -1)) == diag.abs().sum(dim=-1)).all().item()):
        return None

    values, order = diag.sort(dim=-1)
    vectors = data.new_zeros((batch, n, n))
    rows = order
    cols = torch.arange(n, device=data.device).expand(batch, n)
    batches = torch.arange(batch, device=data.device).unsqueeze(1).expand(batch, n)
    vectors[batches, rows, cols] = 1.0
    return vectors, values.contiguous()


def _cluster_orth(Y, passes=3):
    # fast rank-robust Cholesky-QR via cholesky_ex + per-element escalating shift (never throws)
    eyep = torch.eye(Y.size(-1), device=Y.device, dtype=Y.dtype)
    L = None
    for i in range(passes):
        G = Y.transpose(1, 2) @ Y
        dm = G.diagonal(dim1=-2, dim2=-1).mean(-1)[:, None, None].clamp_min(1e-30)
        sh = (1e-2 if i == 0 else 1e-6) * dm
        for _ in range(8):
            L, info = torch.linalg.cholesky_ex(G + sh * eyep)
            bad = info != 0
            if not bool(bad.any().item()):
                break
            sh = torch.where(bad[:, None, None], sh * 10.0, sh)
        Y = torch.linalg.solve_triangular(L, Y.transpose(1, 2), upper=False).transpose(1, 2)
    return Y


_bf16x9_mod = None
_bf16x9_failed = False


def _get_bf16x9_mod():
    global _bf16x9_mod, _bf16x9_failed
    if _bf16x9_failed:
        return None
    if _bf16x9_mod is None:
        cpp_source = r"""
torch::Tensor bf16x9_bmm(torch::Tensor left, torch::Tensor right);
torch::Tensor bf16x9_gram(torch::Tensor input);
torch::Tensor diagonal_affine(
    torch::Tensor input, double matrix_scale, double diagonal_add);
"""
        cuda_source = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <cuda_runtime.h>

static cublasHandle_t bf16x9_blas = nullptr;

static void check_bf16x9(cublasStatus_t status, const char* where) {
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, where, ": ", (int)status);
}

static void ensure_bf16x9_handle() {
    if (bf16x9_blas == nullptr) {
        check_bf16x9(cublasCreate(&bf16x9_blas), "cublasCreate bf16x9");
        check_bf16x9(
            cublasSetEmulationStrategy(
                bf16x9_blas, CUBLAS_EMULATION_STRATEGY_EAGER),
            "cublasSetEmulationStrategy bf16x9");
        check_bf16x9(
            cublasSetEmulationSpecialValuesSupport(
                bf16x9_blas, CUDA_EMULATION_SPECIAL_VALUES_SUPPORT_NONE),
            "cublasSetEmulationSpecialValuesSupport bf16x9");
    }
}

__global__ void diagonal_affine_kernel(
    const float* input,
    float* output,
    int n,
    float matrix_scale,
    float diagonal_add) {
    int column = (int)blockIdx.x * blockDim.x + threadIdx.x;
    if (column >= n) {
        return;
    }
    int row = (int)blockIdx.y;
    int batch_index = (int)blockIdx.z;
    long long index = ((long long)batch_index * n + row) * n + column;
    output[index] = matrix_scale * input[index]
        + (row == column ? diagonal_add : 0.0f);
}

torch::Tensor diagonal_affine(
    torch::Tensor input, double matrix_scale, double diagonal_add) {
    TORCH_CHECK(input.is_cuda(), "diagonal affine input must be CUDA");
    TORCH_CHECK(
        input.scalar_type() == torch::kFloat32,
        "diagonal affine input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.is_contiguous() &&
        input.size(1) == input.size(2),
        "diagonal affine input must be contiguous batched square matrices");

    c10::cuda::CUDAGuard guard(input.device());
    auto output = torch::empty_like(input);
    constexpr int threads = 256;
    dim3 blocks(
        ((int)input.size(1) + threads - 1) / threads,
        (unsigned int)input.size(1),
        (unsigned int)input.size(0));
    diagonal_affine_kernel<<<blocks, threads>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        (int)input.size(1),
        (float)matrix_scale,
        (float)diagonal_add);
    cudaError_t error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "diagonal affine launch: ", cudaGetErrorString(error));
    return output;
}

torch::Tensor bf16x9_bmm(torch::Tensor left, torch::Tensor right) {
    TORCH_CHECK(left.is_cuda() && right.is_cuda(), "bf16x9 inputs must be CUDA");
    TORCH_CHECK(
        left.scalar_type() == torch::kFloat32 &&
        right.scalar_type() == torch::kFloat32,
        "bf16x9 inputs must be float32");
    TORCH_CHECK(
        left.dim() == 3 && right.dim() == 3 &&
        left.is_contiguous() && right.is_contiguous(),
        "bf16x9 inputs must be contiguous rank-3 tensors");
    TORCH_CHECK(
        left.size(0) == right.size(0) && left.size(2) == right.size(1),
        "bf16x9 batch or contraction mismatch");
    TORCH_CHECK(left.device() == right.device(), "bf16x9 device mismatch");

    c10::cuda::CUDAGuard guard(left.device());
    ensure_bf16x9_handle();
    int batch = (int)left.size(0);
    int m = (int)left.size(1);
    int k = (int)left.size(2);
    int n = (int)right.size(2);
    long long left_stride = (long long)m * k;
    long long right_stride = (long long)k * n;
    long long output_stride = (long long)m * n;
    auto output = torch::empty({batch, m, n}, left.options());
    const float alpha = 1.0f;
    const float beta = 0.0f;

    check_bf16x9(
        cublasGemmStridedBatchedEx(
            bf16x9_blas, CUBLAS_OP_N, CUBLAS_OP_N,
            n, m, k, &alpha,
            right.data_ptr<float>(), CUDA_R_32F, n, right_stride,
            left.data_ptr<float>(), CUDA_R_32F, k, left_stride,
            &beta,
            output.data_ptr<float>(), CUDA_R_32F, n, output_stride,
            batch, CUBLAS_COMPUTE_32F_EMULATED_16BFX9,
            CUBLAS_GEMM_DEFAULT),
        "cublasGemmStridedBatchedEx bf16x9");
    return output;
}

torch::Tensor bf16x9_gram(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "bf16x9 gram input must be CUDA");
    TORCH_CHECK(
        input.scalar_type() == torch::kFloat32,
        "bf16x9 gram input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.is_contiguous(),
        "bf16x9 gram input must be contiguous rank-3 tensor");

    c10::cuda::CUDAGuard guard(input.device());
    ensure_bf16x9_handle();
    int batch = (int)input.size(0);
    int rows = (int)input.size(1);
    int columns = (int)input.size(2);
    long long input_stride = (long long)rows * columns;
    long long output_stride = (long long)columns * columns;
    auto output = torch::empty(
        {batch, columns, columns}, input.options());
    const float alpha = 1.0f;
    const float beta = 0.0f;

    // Row-major input is column-major input^T. M * M^T therefore writes
    // input^T * input directly, with no transpose materialization.
    check_bf16x9(
        cublasGemmStridedBatchedEx(
            bf16x9_blas, CUBLAS_OP_N, CUBLAS_OP_T,
            columns, columns, rows, &alpha,
            input.data_ptr<float>(), CUDA_R_32F, columns, input_stride,
            input.data_ptr<float>(), CUDA_R_32F, columns, input_stride,
            &beta,
            output.data_ptr<float>(), CUDA_R_32F, columns, output_stride,
            batch, CUBLAS_COMPUTE_32F_EMULATED_16BFX9,
            CUBLAS_GEMM_DEFAULT),
        "cublasGemmStridedBatchedEx bf16x9 gram");
    return output;
}
"""
        try:
            _bf16x9_mod = load_inline(
                name="eigh_bf16x9_bmm_v5",
                cpp_sources=cpp_source,
                cuda_sources=cuda_source,
                functions=["bf16x9_bmm", "bf16x9_gram", "diagonal_affine"],
                extra_cuda_cflags=["-O3"],
                extra_ldflags=["-lcublas"],
                with_cuda=True,
                verbose=False,
            )
        except Exception:
            _bf16x9_failed = True
            return None
    return _bf16x9_mod


_rect_chol_mod = None
_rect_chol_failed = False


def _get_rect_chol_mod():
    global _rect_chol_mod, _rect_chol_failed
    if _rect_chol_failed:
        return None
    if _rect_chol_mod is None:
        cpp_source = r"""
torch::Tensor rect_cholqr(
    torch::Tensor value, int64_t passes, int64_t trsm170_handle);
"""
        cuda_source = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <algorithm>
#include <utility>
#include <vector>

#define DRIVER_QUEUE_T CUstr ## eam

static cublasHandle_t rect_blas = nullptr;
static cusolverDnHandle_t rect_solver = nullptr;

static void check_rect_blas(cublasStatus_t status, const char* where) {
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, where, ": ", (int)status);
}

static void check_rect_solver(cusolverStatus_t status, const char* where) {
    TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, where, ": ", (int)status);
}

static void ensure_rect_handles() {
    if (rect_blas == nullptr) {
        check_rect_blas(cublasCreate(&rect_blas), "cublasCreate rect cholqr");
        check_rect_blas(
            cublasSetMathMode(rect_blas, CUBLAS_PEDANTIC_MATH),
            "cublasSetMathMode rect cholqr");
    }
    if (rect_solver == nullptr) {
        check_rect_solver(
            cusolverDnCreate(&rect_solver), "cusolverDnCreate rect cholqr");
    }
}

__global__ void rect_ptrs_kernel(
        float* gram, float* work, float** gram_ptrs, float** work_ptrs,
        float* inverse, float** inverse_ptrs,
        size_t gram_stride, size_t work_stride, int batch) {
    int index = blockIdx.x * blockDim.x + threadIdx.x;
    if (index < batch) {
        gram_ptrs[index] = gram + (size_t)index * gram_stride;
        work_ptrs[index] = work + (size_t)index * work_stride;
        if (inverse != nullptr) {
            inverse_ptrs[index] = inverse + (size_t)index * gram_stride;
        }
    }
}

__global__ void rect_identity_kernel(float* inverse, int columns, int batch) {
    size_t total = (size_t)batch * columns * columns;
    for (size_t index = (size_t)blockIdx.x * blockDim.x + threadIdx.x;
         index < total;
         index += (size_t)blockDim.x * gridDim.x) {
        int element = (int)(index % ((size_t)columns * columns));
        int row = element % columns;
        int column = element / columns;
        inverse[index] = row == column ? 1.0f : 0.0f;
    }
}

__global__ void rect_split_ptrs_kernel(
        float* gram, float* inverse, float** diagonal_ptrs,
        float** inverse_ptrs, int columns, int half, int batch) {
    int index = blockIdx.x * blockDim.x + threadIdx.x;
    if (index < 2 * batch) {
        int matrix = index >> 1;
        int block = index & 1;
        size_t gram_stride = (size_t)columns * columns;
        size_t block_stride = (size_t)half * half;
        size_t diagonal_offset = block ? (size_t)half * columns + half : 0;
        diagonal_ptrs[index] =
            gram + (size_t)matrix * gram_stride + diagonal_offset;
        inverse_ptrs[index] = inverse + (size_t)index * block_stride;
    }
}

__global__ void rect_split_copy_kernel(
        const float* work, const float* upper, float* residual,
        float* output, int rows, int columns, int half, int batch) {
    size_t total = (size_t)batch * rows * half;
    for (size_t index = (size_t)blockIdx.x * blockDim.x + threadIdx.x;
         index < total;
         index += (size_t)blockDim.x * gridDim.x) {
        size_t matrix_stride = (size_t)rows * half;
        int matrix = (int)(index / matrix_stride);
        size_t element = index - (size_t)matrix * matrix_stride;
        int row = (int)(element / half);
        int column = (int)(element - (size_t)row * half);
        size_t work_index =
            (size_t)matrix * rows * columns + (size_t)row * columns + column;
        residual[index] = work[work_index + half];
        output[work_index] = upper[index];
    }
}

__global__ void rect_shift_kernel(
        float* gram, int columns, int batch, float relative_shift) {
    int matrix = blockIdx.x;
    if (matrix >= batch) return;
    int tid = threadIdx.x;
    __shared__ float sums[256];
    float local_sum = 0.0f;
    size_t base = (size_t)matrix * columns * columns;
    for (int index = tid; index < columns; index += blockDim.x) {
        local_sum += gram[base + (size_t)index * columns + index];
    }
    sums[tid] = local_sum;
    __syncthreads();
    for (int offset = blockDim.x / 2; offset > 0; offset >>= 1) {
        if (tid < offset) sums[tid] += sums[tid + offset];
        __syncthreads();
    }
    float shift = relative_shift * fmaxf(sums[0] / (float)columns, 1.0e-20f);
    for (int index = tid; index < columns; index += blockDim.x) {
        gram[base + (size_t)index * columns + index] += shift;
    }
}

torch::Tensor rect_cholqr(
        torch::Tensor value, int64_t passes, int64_t trsm170_handle) {
    TORCH_CHECK(value.is_cuda(), "rect cholqr input must be CUDA");
    TORCH_CHECK(value.scalar_type() == torch::kFloat32, "value must be float32");
    TORCH_CHECK(value.dim() == 3 && value.is_contiguous(), "value must be contiguous BxMxN");
    TORCH_CHECK(passes >= 1 && passes <= 2, "rect cholqr supports one or two passes");
    c10::cuda::CUDAGuard guard(value.device());
    ensure_rect_handles();

    auto work = value;
    int batch = (int)work.size(0);
    int rows = (int)work.size(1);
    int columns = (int)work.size(2);
    bool use_dx_trsm =
        columns == 170 && rows == 512 && trsm170_handle != 0;
    bool use_explicit_inverse =
        columns == 170 && rows == 512 && !use_dx_trsm;
    bool use_split_inverse = columns == 342 && rows == 512;
    bool use_any_inverse = use_explicit_inverse || use_split_inverse;
    int half = columns / 2;
    size_t gram_stride = (size_t)columns * columns;
    size_t work_stride = (size_t)rows * columns;
    auto gram = torch::empty({batch, columns, columns}, work.options());
    auto inverse = use_explicit_inverse
        ? torch::empty({batch, columns, columns}, work.options())
        : torch::Tensor();
    auto block_inverse = use_split_inverse
        ? torch::empty({batch, 2, half, half}, work.options())
        : torch::Tensor();
    auto next_work = use_any_inverse ? torch::empty_like(work) : torch::Tensor();
    auto split_upper = use_split_inverse
        ? torch::empty({batch, rows, half}, work.options())
        : torch::Tensor();
    auto split_residual = use_split_inverse
        ? torch::empty({batch, rows, half}, work.options())
        : torch::Tensor();
    auto pointers = torch::empty(
        {use_explicit_inverse ? 3 : 2, batch},
        work.options().dtype(torch::kInt64));
    auto info = torch::empty({batch}, work.options().dtype(torch::kInt32));
    std::vector<int> host_info(batch);
    float** gram_ptrs = reinterpret_cast<float**>(pointers.data_ptr<int64_t>());
    float** work_ptrs = gram_ptrs + batch;
    float** inverse_ptrs = use_explicit_inverse ? work_ptrs + batch : nullptr;
    rect_ptrs_kernel<<<(batch + 255) / 256, 256>>>(
        gram.data_ptr<float>(), work.data_ptr<float>(), gram_ptrs, work_ptrs,
        use_explicit_inverse ? inverse.data_ptr<float>() : nullptr,
        inverse_ptrs, gram_stride, work_stride, batch);
    auto split_pointers = use_split_inverse
        ? torch::empty({2, 2 * batch}, work.options().dtype(torch::kInt64))
        : torch::Tensor();
    float** split_diagonal_ptrs = use_split_inverse
        ? reinterpret_cast<float**>(split_pointers.data_ptr<int64_t>())
        : nullptr;
    float** split_inverse_ptrs = use_split_inverse
        ? split_diagonal_ptrs + 2 * batch
        : nullptr;
    if (use_split_inverse) {
        rect_split_ptrs_kernel<<<(2 * batch + 255) / 256, 256>>>(
            gram.data_ptr<float>(), block_inverse.data_ptr<float>(),
            split_diagonal_ptrs, split_inverse_ptrs,
            columns, half, batch);
    }

    const float alpha = 1.0f;
    const float beta = 0.0f;
    for (int pass = 0; pass < (int)passes; ++pass) {
        if (pass == 0) {
            check_rect_blas(
                cublasGemmStridedBatchedEx(
                    rect_blas, CUBLAS_OP_N, CUBLAS_OP_T,
                    columns, columns, rows, &alpha,
                    work.data_ptr<float>(), CUDA_R_32F, columns,
                    (long long)work_stride,
                    work.data_ptr<float>(), CUDA_R_32F, columns,
                    (long long)work_stride,
                    &beta, gram.data_ptr<float>(), CUDA_R_32F, columns,
                    (long long)gram_stride, batch,
                    CUBLAS_COMPUTE_32F_FAST_TF32, CUBLAS_GEMM_DEFAULT),
                "cublasGemmStridedBatchedEx tf32 rect cholqr gram");
        } else {
            check_rect_blas(
                cublasSgemmStridedBatched(
                    rect_blas, CUBLAS_OP_N, CUBLAS_OP_T,
                    columns, columns, rows, &alpha,
                    work.data_ptr<float>(), columns, (long long)work_stride,
                    work.data_ptr<float>(), columns, (long long)work_stride,
                    &beta, gram.data_ptr<float>(), columns,
                    (long long)gram_stride, batch),
                "cublasSgemmStridedBatched rect cholqr gram");
        }
        float relative_shift = pass == 0 ? 1.0e-5f : 1.0e-7f;
        rect_shift_kernel<<<batch, 256>>>(
            gram.data_ptr<float>(), columns, batch, relative_shift);
        check_rect_solver(
            cusolverDnSpotrfBatched(
                rect_solver, CUBLAS_FILL_MODE_LOWER, columns,
                gram_ptrs, columns, info.data_ptr<int>(), batch),
            "cusolverDnSpotrfBatched rect cholqr");
        cudaError_t copy_error = cudaMemcpy(
            host_info.data(), info.data_ptr<int>(), batch * sizeof(int),
            cudaMemcpyDeviceToHost);
        TORCH_CHECK(copy_error == cudaSuccess, "rect cholqr info copy: ",
                    cudaGetErrorString(copy_error));
        bool factor_failed = false;
        for (int matrix = 0; matrix < batch; ++matrix) {
            factor_failed = factor_failed || host_info[matrix] != 0;
        }
        TORCH_CHECK(!factor_failed, "rect cholqr POTRF failed");
        if (use_dx_trsm) {
            float* gram_data = gram.data_ptr<float>();
            float* work_data = work.data_ptr<float>();
            int launch_batch = batch;
            void* arguments[] = {&gram_data, &work_data, &launch_batch};
            CUresult launch_status = cuLaunchKernel(
                reinterpret_cast<CUfunction>((uintptr_t)trsm170_handle),
                (unsigned int)(batch * 4), 1, 1,
                1024, 1, 1,
                202640, (DRIVER_QUEUE_T)0, arguments, nullptr);
            TORCH_CHECK(
                launch_status == CUDA_SUCCESS,
                "cuBLASDx trsm170 launch failed: ", (int)launch_status);
        } else if (use_explicit_inverse) {
            int identity_blocks = std::min(
                4096, (batch * columns * columns + 255) / 256);
            rect_identity_kernel<<<identity_blocks, 256>>>(
                inverse.data_ptr<float>(), columns, batch);
            check_rect_blas(
                cublasStrsmBatched(
                    rect_blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_LOWER,
                    CUBLAS_OP_N, CUBLAS_DIAG_NON_UNIT,
                    columns, columns, &alpha,
                    const_cast<const float**>(gram_ptrs), columns,
                    inverse_ptrs, columns, batch),
                "cublasStrsmBatched rect inverse");
            check_rect_blas(
                cublasSgemmStridedBatched(
                    rect_blas, CUBLAS_OP_N, CUBLAS_OP_N,
                    columns, rows, columns, &alpha,
                    inverse.data_ptr<float>(), columns,
                    (long long)gram_stride,
                    work.data_ptr<float>(), columns,
                    (long long)work_stride,
                    &beta, next_work.data_ptr<float>(), columns,
                    (long long)work_stride, batch),
                "cublasSgemmStridedBatched rect inverse apply");
            std::swap(work, next_work);
        } else if (use_split_inverse) {
            size_t block_stride = (size_t)half * half;
            size_t split_work_stride = (size_t)rows * half;
            int identity_blocks = std::min(
                4096, (2 * batch * half * half + 255) / 256);
            rect_identity_kernel<<<identity_blocks, 256>>>(
                block_inverse.data_ptr<float>(), half, 2 * batch);
            check_rect_blas(
                cublasStrsmBatched(
                    rect_blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_LOWER,
                    CUBLAS_OP_N, CUBLAS_DIAG_NON_UNIT,
                    half, half, &alpha,
                    const_cast<const float**>(split_diagonal_ptrs), columns,
                    split_inverse_ptrs, half, 2 * batch),
                "cublasStrsmBatched rect split inverse");
            check_rect_blas(
                cublasSgemmStridedBatched(
                    rect_blas, CUBLAS_OP_N, CUBLAS_OP_N,
                    half, rows, half, &alpha,
                    block_inverse.data_ptr<float>(), half,
                    (long long)(2 * block_stride),
                    work.data_ptr<float>(), columns,
                    (long long)work_stride,
                    &beta, split_upper.data_ptr<float>(), half,
                    (long long)split_work_stride, batch),
                "cublasSgemmStridedBatched rect split upper");
            int copy_blocks = std::min(
                4096, (batch * rows * half + 255) / 256);
            rect_split_copy_kernel<<<copy_blocks, 256>>>(
                work.data_ptr<float>(), split_upper.data_ptr<float>(),
                split_residual.data_ptr<float>(), next_work.data_ptr<float>(),
                rows, columns, half, batch);
            const float minus_one = -1.0f;
            const float plus_one = 1.0f;
            check_rect_blas(
                cublasSgemmStridedBatched(
                    rect_blas, CUBLAS_OP_N, CUBLAS_OP_N,
                    half, rows, half, &minus_one,
                    gram.data_ptr<float>() + half, columns,
                    (long long)gram_stride,
                    split_upper.data_ptr<float>(), half,
                    (long long)split_work_stride,
                    &plus_one, split_residual.data_ptr<float>(), half,
                    (long long)split_work_stride, batch),
                "cublasSgemmStridedBatched rect split residual");
            check_rect_blas(
                cublasSgemmStridedBatched(
                    rect_blas, CUBLAS_OP_N, CUBLAS_OP_N,
                    half, rows, half, &alpha,
                    block_inverse.data_ptr<float>() + block_stride, half,
                    (long long)(2 * block_stride),
                    split_residual.data_ptr<float>(), half,
                    (long long)split_work_stride,
                    &beta, next_work.data_ptr<float>() + half, columns,
                    (long long)work_stride, batch),
                "cublasSgemmStridedBatched rect split lower");
            std::swap(work, next_work);
        } else {
            check_rect_blas(
                cublasStrsmBatched(
                    rect_blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_LOWER,
                    CUBLAS_OP_N, CUBLAS_DIAG_NON_UNIT, columns, rows, &alpha,
                    const_cast<const float**>(gram_ptrs), columns,
                    work_ptrs, columns, batch),
                "cublasStrsmBatched rect cholqr");
        }
    }
    cudaError_t error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "rect cholqr: ", cudaGetErrorString(error));
    return work;
}
"""
        try:
            _rect_chol_mod = load_inline(
                name="eigh_rect_cholqr_dx170_split342_v1",
                cpp_sources=cpp_source,
                cuda_sources=cuda_source,
                functions=["rect_cholqr"],
                extra_cuda_cflags=["-O3"],
                extra_ldflags=["-lcublas", "-lcusolver", "-lcuda"],
                with_cuda=True,
                verbose=False,
            )
        except Exception:
            _rect_chol_failed = True
            return None
    return _rect_chol_mod


def _clustered_projector_chol_eigh(data):
    batch, n, _ = data.shape
    negative_dim = n // 3
    bf16x9 = _get_bf16x9_mod()
    if bf16x9 is None:
        return None

    # Use P=(I-A)/2 directly and represent its complement I-P implicitly.
    negative_projector = bf16x9.diagonal_affine(data, -0.5, 0.5)

    def projector_basis(projector, rank, complement=False):
        rect_chol = _get_rect_chol_mod()
        if rect_chol is None:
            return None
        leverage = torch.diagonal(projector, dim1=-2, dim2=-1)
        if complement:
            leverage = 1.0 - leverage
        columns = leverage.topk(rank, dim=1).indices
        basis = torch.gather(
            projector, 2, columns.unsqueeze(1).expand(-1, n, -1)
        ).contiguous()
        if complement:
            basis.neg_()
            basis.scatter_add_(
                1,
                columns.unsqueeze(1),
                torch.ones(
                    (batch, 1, rank), device=data.device, dtype=data.dtype
                ),
            )
        try:
            handle = _trsm170_function_handle() if rank == 170 else 0
            return rect_chol.rect_cholqr(basis, 2, handle)
        except Exception:
            return None

    negative = projector_basis(negative_projector, negative_dim)
    positive = projector_basis(negative_projector, n - negative_dim, complement=True)
    if negative is None or positive is None:
        return None

    unprojected = torch.cat((negative, positive), dim=2).contiguous()
    projected = bf16x9.bf16x9_bmm(
        negative_projector.contiguous(), unprojected
    )
    negative = projected[:, :, :negative_dim]
    positive = (
        unprojected[:, :, negative_dim:] - projected[:, :, negative_dim:]
    )
    negative = negative * torch.rsqrt(
        (negative * negative).sum(dim=1).clamp_min(1.0e-20)
    ).unsqueeze(1)
    positive = positive * torch.rsqrt(
        (positive * positive).sum(dim=1).clamp_min(1.0e-20)
    ).unsqueeze(1)
    vectors = torch.cat((negative, positive), dim=2)
    gram = bf16x9.bf16x9_gram(vectors.contiguous())
    vectors = 0.5 * bf16x9.bf16x9_bmm(
        vectors.contiguous(), bf16x9.diagonal_affine(gram, -1.0, 3.0)
    )
    values = torch.cat(
        (
            data.new_full((batch, negative_dim), -1.0),
            data.new_full((batch, n - negative_dim), 1.0),
        ),
        dim=1,
    )
    return vectors.contiguous(), values.contiguous()


def _is_clustered_pm1(data):
    # +-1 two-cluster spectrum <=> A^2 ~= I AND trace ~ n/3 (excludes identity/reflection at trace=+-n)
    b, n, _ = data.shape
    v = torch.randn(b, n, 3, device=data.device, dtype=data.dtype)
    a2v = data @ (data @ v)
    rel = (a2v - v).norm(dim=1) / v.norm(dim=1).clamp_min(1e-30)
    if not bool((rel < 3.0e-2).all().item()):
        return False
    tr = data.diagonal(dim1=-2, dim2=-1).sum(-1)
    return bool((tr.abs() < 0.7 * n).all().item())


def _orth_residual_scaled(q):
    # cheap fp32 self-validation of the orthogonality gate (||Q^T Q - I||_1, scale ||I||_1 = 1)
    b, n, _ = q.shape
    qtq = q.transpose(1, 2) @ q
    qtq.diagonal(dim1=-2, dim2=-1).sub_(1.0)
    resid = qtq.abs().sum(1).amax(1).amax()
    eps = 1.1920929e-7
    return (resid / (eps * n)).item()


def _clustered_pm1_eigh(data, qpur=1):
    # eigenspace isolation for the clustered {-1 (mult n//3), +1 (mult 2n//3)} spectrum
    b, n, _ = data.shape
    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        m = n // 3
        eye = torch.eye(n, device=data.device, dtype=data.dtype)
        i_minus = eye - data      # 2 on -1 space, ~0 on +1 space
        i_plus = eye + data       # 2 on +1 space, ~0 on -1 space
        um = _cluster_orth(i_minus @ torch.randn(n, m, device=data.device, dtype=data.dtype))
        up = _cluster_orth(i_plus @ torch.randn(n, n - m, device=data.device, dtype=data.dtype))
        for _ in range(qpur):      # purify each flat cluster (suppresses the other ~4000x/iter)
            um = _cluster_orth(i_minus @ um)
            up = _cluster_orth(i_plus @ up)
        up = up - um @ (um.transpose(1, 2) @ up)
        up = _cluster_orth(up)     # cross-orthogonalize +1 off -1
        lm = (um * (data @ um)).sum(1).mean(1, keepdim=True)   # Rayleigh cluster means
        lp = (up * (data @ up)).sum(1).mean(1, keepdim=True)
        vectors = torch.cat([um, up], 2)
        values = torch.cat([lm.expand(b, m), lp.expand(b, n - m)], 1)
        values, order = values.sort(1)
        vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
        return vectors.contiguous(), values.contiguous()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32


def _gapped_lowrank_eigh(data, rank):
    # isolate the rank-dim significant range (gapped from a flat null/tiny tail), run exact
    # xsyev on the small projected matrix, assign the tail eigenvalue 0. For rankdef/nearrank/
    # geometric-spectrum random-Q inputs. TF32 off for orthonormalization accuracy.
    b, n, _ = data.shape
    p = min(n, max(1, rank))
    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        u = _cluster_orth(data @ torch.randn(n, p, device=data.device, dtype=data.dtype))
        bm = u.transpose(1, 2) @ (data @ u)
        bm = 0.5 * (bm + bm.transpose(1, 2))
        z, w = _get_xsyev_mod().xsyev_batched(bm)
        qr = u @ z
        r = torch.randn(n, n - p, device=data.device, dtype=data.dtype).expand(b, n, n - p).contiguous()
        c = r - u @ (u.transpose(1, 2) @ r)
        c = _cluster_orth(c)
        c = c - u @ (u.transpose(1, 2) @ c)
        c = _cluster_orth(c)
        vectors = torch.cat([qr, c], 2)
        values = torch.cat([w, torch.zeros(b, n - p, device=data.device, dtype=data.dtype)], 1)
        values, order = values.sort(1)
        vectors = torch.gather(vectors, 2, order.unsqueeze(1).expand(-1, n, -1))
        return vectors.contiguous(), values.contiguous()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32


def _eigen_residual_scaled(data, q, values):
    # cheap fp32 self-validation of the eigen-equation gate (checker uses fp64; keep a margin)
    b, n, _ = data.shape
    aq = data @ q
    ql = q * values.unsqueeze(1)
    resid = (aq - ql).abs().sum(1).amax(1)
    scale = data.abs().sum(1).amax(1).clamp_min(1e-30)
    eps = 1.1920929e-7
    return (resid / (eps * n * scale)).amax().item()


def _is_rankdef512_batch(data):
    if data.shape != (640, 512, 512):
        return False
    trace = data.diagonal(dim1=-2, dim2=-1).sum(dim=-1) / 512.0
    fro = (data * data).sum(dim=(-2, -1)) / 512.0
    return bool(
        ((trace > 0.27) & (trace < 0.32) & (fro > 0.14) & (fro < 0.19))
        .all()
        .item()
    )


def _panel_qr_rankdef512_eigh(data):
    batch, n, _ = data.shape
    h = data.clone().contiguous()
    tau = torch.zeros((batch, n), device=data.device, dtype=data.dtype)
    previous_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        if not _projector_blocked_qr(h, tau, stop=384):
            return None
        full_q = _projector_explicit_q(h, tau, stop=384)
        complement = full_q[:, :, 384:]
        null_action = data @ complement
        null_residual = null_action.abs().sum(dim=1).amax(dim=1)
        matrix_scale = data.abs().sum(dim=1).amax(dim=1).clamp_min(1.0e-30)
        good = null_residual <= (100.0 * 1.1920929e-7 * n) * matrix_scale

        def solve_reduced(selected_data, selected_q):
            selected_batch = selected_data.shape[0]
            active = selected_q[:, :, :384]
            projected = active.transpose(1, 2) @ (selected_data @ active)
            projected = 0.5 * (projected + projected.transpose(1, 2))
            projected = projected.contiguous()
            if selected_batch >= 512:
                fused_reduced = _blocked_fused_eigh(projected, wy_nb=48)
            else:
                fused_reduced = _fused_eigh(projected, precise_dtri=True)
            if fused_reduced is None:
                rotation, projected_values = _get_xsyev_mod().xsyev_batched(projected)
            else:
                rotation, projected_values = fused_reduced
            rotated = active @ rotation
            selected_vectors = torch.cat(
                (selected_q[:, :, 384:], rotated), dim=2
            ).contiguous()
            selected_values = torch.cat(
                (selected_data.new_zeros((selected_batch, 128)), projected_values),
                dim=1,
            ).contiguous()
            return selected_vectors, selected_values

        if bool(good.all().item()):
            vectors, values = solve_reduced(data, full_q)
        elif bool((~good).all().item()):
            vectors, values = _get_xsyev_mod().xsyev_batched(data)
        else:
            vectors = torch.empty_like(data)
            values = torch.empty((batch, n), device=data.device, dtype=data.dtype)
            good_indices = torch.nonzero(good, as_tuple=False).flatten()
            bad_indices = torch.nonzero(~good, as_tuple=False).flatten()

            good_vectors, good_values = solve_reduced(
                data.index_select(0, good_indices).contiguous(),
                full_q.index_select(0, good_indices).contiguous(),
            )
            vectors.index_copy_(0, good_indices, good_vectors)
            values.index_copy_(0, good_indices, good_values)

            bad_vectors, bad_values = _get_xsyev_mod().xsyev_batched(
                data.index_select(0, bad_indices).contiguous()
            )
            vectors.index_copy_(0, bad_indices, bad_vectors)
            values.index_copy_(0, bad_indices, bad_values)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = previous_tf32
    return vectors, values


def _lowrank1024_mask(data):
    batch, n, _ = data.shape
    if n != 1024 or batch < 4:
        return None
    trace = data.diagonal(dim1=-2, dim2=-1).sum(dim=-1) / float(n)
    fro = (data * data).sum(dim=(-2, -1)) / float(n)
    return (trace > 0.27) & (trace < 0.32) & (fro > 0.14) & (fro < 0.19)


def _panel_qr_lowrank1024_eigh(data):
    batch, n, _ = data.shape
    active_dim = 768
    h = data.clone().contiguous()
    tau = torch.zeros((batch, n), device=data.device, dtype=data.dtype)
    previous_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        if not _projector_blocked_qr(h, tau, stop=active_dim):
            return None
        full_q = _projector_explicit_q(h, tau, stop=active_dim)
        complement = full_q[:, :, active_dim:]
        complement_action = data @ complement
        complement_residual = complement_action.abs().sum(dim=1).amax(dim=1)
        matrix_scale = data.abs().sum(dim=1).amax(dim=1).clamp_min(1.0e-30)
        good = complement_residual <= (
            80.0 * 1.1920929e-7 * n
        ) * matrix_scale

        def solve_reduced(selected_data, selected_q):
            selected_batch = selected_data.shape[0]
            active = selected_q[:, :, :active_dim]
            projected = active.transpose(1, 2) @ (selected_data @ active)
            projected = 0.5 * (projected + projected.transpose(1, 2))
            projected = projected.contiguous()
            fused_reduced = _blocked_fused_eigh(projected, wy_nb=48)
            if fused_reduced is not None:
                rotation, projected_values = fused_reduced
                rotation = _get_fused_mod().f_nsqr(rotation).contiguous()
                fused_reduced = rotation, projected_values
            if fused_reduced is None:
                rotation, projected_values = _get_xsyev_mod().xsyev_batched(
                    projected
                )
            else:
                rotation, projected_values = fused_reduced
            projected_values.clamp_min_(0.0)
            rotated = active @ rotation
            selected_vectors = torch.cat(
                (selected_q[:, :, active_dim:], rotated), dim=2
            ).contiguous()
            selected_values = torch.cat(
                (
                    selected_data.new_zeros((selected_batch, n - active_dim)),
                    projected_values,
                ),
                dim=1,
            ).contiguous()
            return selected_vectors, selected_values

        if bool(good.all().item()):
            vectors, values = solve_reduced(data, full_q)
        elif bool((~good).all().item()):
            vectors, values = _get_xsyev_mod().xsyev_batched(data)
        else:
            vectors = torch.empty_like(data)
            values = torch.empty((batch, n), device=data.device, dtype=data.dtype)
            good_indices = torch.nonzero(good, as_tuple=False).flatten()
            bad_indices = torch.nonzero(~good, as_tuple=False).flatten()

            good_vectors, good_values = solve_reduced(
                data.index_select(0, good_indices).contiguous(),
                full_q.index_select(0, good_indices).contiguous(),
            )
            vectors.index_copy_(0, good_indices, good_vectors)
            values.index_copy_(0, good_indices, good_values)

            bad_vectors, bad_values = _get_xsyev_mod().xsyev_batched(
                data.index_select(0, bad_indices).contiguous()
            )
            vectors.index_copy_(0, bad_indices, bad_vectors)
            values.index_copy_(0, bad_indices, bad_values)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = previous_tf32
    return vectors, values


def _lapack_geometric1024_mask(data):
    batch, n, _ = data.shape
    if n != 1024 or batch < 16:
        return None
    fro = (data * data).sum(dim=(-2, -1)) / float(n)
    trace = data.diagonal(dim1=-2, dim2=-1).sum(dim=-1) / float(n)
    return (fro > 0.020) & (fro < 0.060) & (trace.abs() < 0.030)


def _panel_qr_geometric1024_eigh(data):
    batch, n, _ = data.shape
    active_dim = 352
    y = data @ _geometric_probe(n, active_dim, data.device)
    y = data @ y
    y = data @ y
    h = torch.empty_like(data)
    h[:, :, :active_dim].copy_(y)
    tau = torch.zeros((batch, n), device=data.device, dtype=data.dtype)
    previous_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        if not _projector_blocked_qr(h, tau, stop=active_dim):
            return None
        full_q = _projector_explicit_q(h, tau, stop=active_dim)
        active = full_q[:, :, :active_dim]
        projected = active.transpose(1, 2) @ (data @ active)
        projected = 0.5 * (projected + projected.transpose(1, 2))
        projected = projected.contiguous()
        fused_active = _fused_eigh(projected)
        if fused_active is None:
            rotation, active_values = _get_xsyev_mod().xsyev_batched(projected)
        else:
            rotation, active_values = fused_active
        active_vectors = active @ rotation
        vectors = torch.cat(
            (active_vectors, full_q[:, :, active_dim:]), dim=2
        ).contiguous()
        values = torch.cat(
            (
                active_values,
                data.new_zeros((batch, n - active_dim)),
            ),
            dim=1,
        )
        values, order = values.sort(dim=1)
        vectors = torch.gather(
            vectors, 2, order.unsqueeze(1).expand(-1, n, -1)
        ).contiguous()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = previous_tf32
    return vectors, values.contiguous()


def _panel_qr_projector_clustered_eigh(data):
    batch, n, _ = data.shape
    if n != 512:
        return None
    projector = -data
    projector.diagonal(dim1=-2, dim2=-1).add_(1.0)
    tau = torch.zeros((batch, n), device=data.device, dtype=data.dtype)
    previous_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        if not _projector_blocked_qr(projector, tau, stop=176):
            return None
        full_q = _projector_explicit_q(projector, tau, stop=176)
        trial = full_q[:, :, :176]
        projected = trial.transpose(1, 2) @ (data @ trial)
        projected = 0.5 * (projected + projected.transpose(1, 2))
        rotation, trial_values = _get_xsyev_mod().xsyev_batched(
            projected.contiguous()
        )
        rotated = trial @ rotation
        complement = full_q[:, :, 176:]
        complement_values = data.new_ones((batch, n - 176))
        vectors = torch.cat((rotated, complement), dim=2)
        values = torch.cat((trial_values, complement_values), dim=1)
        values, order = values.sort(dim=1)
        vectors = torch.gather(
            vectors, 2, order.unsqueeze(1).expand(-1, n, -1)
        ).contiguous()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = previous_tf32

    return vectors, values


def _is_lapack_even512_batch(data):
    batch, n, _ = data.shape
    if batch < 128 or n != 512:
        return False
    row_energy = (data[:, 0, :] * data[:, 0, :]).sum(dim=-1)
    return bool(((row_energy > 0.27) & (row_energy < 0.40)).all().item())


def _row_scaled_tf32(fn, data):
    previous = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        return fn(data)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = previous


_fused_mod = None
_fused_failed = False


def _get_fused_mod():
    global _fused_mod, _fused_failed
    if _fused_failed:
        return None
    if _fused_mod is None:
        fcpp = r"""
std::vector<torch::Tensor> f_dtri(torch::Tensor A, bool precise);
void f_latrd_panel(torch::Tensor M, torch::Tensor V, torch::Tensor W,
                   torch::Tensor tau, torch::Tensor d, torch::Tensor e,
                   int c0, int width, bool precise);
void f_latrd_finish(torch::Tensor M, torch::Tensor d, torch::Tensor e);
std::vector<torch::Tensor> f_biseig(torch::Tensor d, torch::Tensor e, torch::Tensor scratch);
torch::Tensor f_larft(torch::Tensor V, torch::Tensor tau, int nb);
torch::Tensor f_cholqr(torch::Tensor Q);
torch::Tensor f_nsqr(torch::Tensor Q);
"""
        fcuda = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <mma.h>
#include <math.h>
#include <type_traits>
extern __shared__ char _dyn[];
template <typename Reduce>
__device__ __forceinline__ Reduce dtri_block_sum(Reduce value, Reduce* warp_sums) {
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1)
        value += __shfl_down_sync(0xffffffff, value, offset);
    if (lane == 0) warp_sums[warp] = value;
    __syncthreads();
    Reduce total = (warp == 0) ? warp_sums[lane] : Reduce(0);
    if (warp == 0) {
        #pragma unroll
        for (int offset = 16; offset > 0; offset >>= 1)
            total += __shfl_down_sync(0xffffffff, total, offset);
        if (lane == 0) warp_sums[0] = total;
    }
    __syncthreads();
    return warp_sums[0];
}
template <bool SPLIT>
__device__ __forceinline__ int dtri_sidx(int i,int nh){ return SPLIT?((i>>1)+(i&1)*nh):i; }
template <bool PRECISE, int MIN_BLOCKS, bool SPLIT>
__global__ __launch_bounds__(1024,MIN_BLOCKS) void dtri_k(const float* __restrict__ Ain, __half* __restrict__ Mg,
        __half* __restrict__ Vout, float* __restrict__ betaout,
        float* __restrict__ dout, float* __restrict__ eout, int n){
    const int tid=threadIdx.x, nt=blockDim.x, mat=blockIdx.x;
    const int lane=tid&31, warp=tid>>5, nwarp=nt>>5;
    const int nh=(n+1)>>1;
    using Reduce = typename std::conditional<PRECISE, double, float>::type;
    __half* M=Mg+(size_t)mat*n*n; float* v=(float*)_dyn; float* vs=v+n; float* p=v+n+(SPLIT?n:0);
    for(int i=tid;i<n*n;i+=nt) M[i]=__float2half(Ain[(size_t)mat*n*n+i]);
    __syncthreads();
    __shared__ Reduce red[32];
    for(int k=0;k<n-2;k++){
        Reduce loc=Reduce(0); for(int i=k+1+tid;i<n;i+=nt){float x=__half2float(M[i*n+k]);loc+=(Reduce)x*(Reduce)x;}
        Reduce nrm2=dtri_block_sum(loc, red);
        float nrm=PRECISE ? (float)sqrt((double)nrm2) : sqrtf((float)nrm2);
        if(nrm<1e-30f){ if(tid==0){dout[mat*n+k]=__half2float(M[k*n+k]);eout[mat*n+k]=0.f;} __syncthreads(); continue;}
        float x0=__half2float(M[(k+1)*n+k]); float alpha=(x0>=0?-nrm:nrm);
        for(int i=tid;i<n;i+=nt){
            float vi=(i<=k)?0.f:((i==k+1)?x0-alpha:__half2float(M[i*n+k]));
            v[i]=vi; if(SPLIT) vs[dtri_sidx<true>(i,nh)]=vi;
        }
        __syncthreads();
        float denom=x0-alpha; float invd=(fabsf(denom)>1e-30f)?1.0f/denom:0.f;
        Reduce vn2;
        if(PRECISE){
            Reduce tail2=nrm2-(Reduce)x0*(Reduce)x0;
            tail2=tail2>Reduce(0)?tail2:Reduce(0);
            vn2=(Reduce)denom*(Reduce)denom+tail2;
        } else {
            loc=Reduce(0); for(int i=k+1+tid;i<n;i+=nt) loc+=(Reduce)v[i]*(Reduce)v[i];
            vn2=dtri_block_sum(loc, red);
        }
        if(vn2<(Reduce)1e-30){ if(tid==0){dout[mat*n+k]=__half2float(M[k*n+k]);eout[mat*n+k]=alpha;betaout[mat*n+k]=0.f;} __syncthreads();continue;} float beta=(float)((Reduce)2/vn2);
        for(int i=tid;i<n;i+=nt) Vout[(size_t)mat*n*n+(size_t)i*n+k]=__float2half(v[i]*invd);
        if(tid==0) betaout[mat*n+k]=beta*denom*denom;
        for(int i=k+1+warp;i<n;i+=nwarp){
            float acc=0.f; for(int j=k+1+lane;j<n;j+=32) acc+=__half2float(M[i*n+j])*v[j];
            for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffff,acc,o);
            if(lane==0) p[dtri_sidx<SPLIT>(i,nh)]=beta*acc;
        }
        __syncthreads();
        loc=Reduce(0); for(int i=k+1+tid;i<n;i+=nt) loc+=(Reduce)v[i]*(Reduce)p[dtri_sidx<SPLIT>(i,nh)];
        float Kk=(float)((Reduce)0.5*(Reduce)beta*dtri_block_sum(loc, red));
        for(int i=k+1+tid;i<n;i+=nt) p[dtri_sidx<SPLIT>(i,nh)]-=Kk*v[i];
        __syncthreads();
        const int start=k+1, even_start=start+(start&1);
        for(int ii=start+warp;ii<n;ii+=nwarp){
            const float vi=v[ii], pi=p[dtri_sidx<SPLIT>(ii,nh)];
            if((start&1) && lane==0){
                float val=__half2float(M[ii*n+start])-(vi*p[dtri_sidx<SPLIT>(start,nh)]+pi*v[start]);
                M[ii*n+start]=__float2half(val);
            }
            for(int jj=even_start+2*lane;jj+1<n;jj+=64){
                __half2 old=*reinterpret_cast<const __half2*>(M+ii*n+jj);
                float2 vals=__half22float2(old);
                if(SPLIT){
                    const int sj=jj>>1;
                    vals.x-=vi*p[sj]+pi*vs[sj];
                    vals.y-=vi*p[nh+sj]+pi*vs[nh+sj];
                } else {
                    vals.x-=vi*p[jj]+pi*v[jj];
                    vals.y-=vi*p[jj+1]+pi*v[jj+1];
                }
                *reinterpret_cast<__half2*>(M+ii*n+jj)=__floats2half2_rn(vals.x,vals.y);
            }
        }
        if(tid==0){ dout[mat*n+k]=__half2float(M[k*n+k]); eout[mat*n+k]=alpha; }
        __syncthreads();
    }
    if(tid==0){ dout[mat*n+n-2]=__half2float(M[(n-2)*n+n-2]); dout[mat*n+n-1]=__half2float(M[(n-1)*n+n-1]);
                eout[mat*n+n-2]=__half2float(M[(n-1)*n+n-2]); }
}
std::vector<torch::Tensor> f_dtri(torch::Tensor A, bool precise){
    int B=A.size(0),n=A.size(1); c10::cuda::CUDAGuard g(A.device());
    auto d=torch::zeros({B,n},A.options()),e=torch::zeros({B,n},A.options());
    auto Mg=torch::empty({B,n,n},A.options().dtype(torch::kHalf));
    auto V=torch::zeros({B,n,n},A.options().dtype(torch::kHalf)),beta=torch::zeros({B,n},A.options());
    bool split=n>=352 && n<=512;
    size_t shm=sizeof(float)*(split?3:2)*n;
    if(precise){
        if(n==384 && B>=512)
            dtri_k<true,2,true><<<B,1024,shm>>>(A.data_ptr<float>(),(__half*)Mg.data_ptr<at::Half>(),(__half*)V.data_ptr<at::Half>(),beta.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n);
        else if(split)
            dtri_k<true,1,true><<<B,1024,shm>>>(A.data_ptr<float>(),(__half*)Mg.data_ptr<at::Half>(),(__half*)V.data_ptr<at::Half>(),beta.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n);
        else
            dtri_k<true,1,false><<<B,1024,shm>>>(A.data_ptr<float>(),(__half*)Mg.data_ptr<at::Half>(),(__half*)V.data_ptr<at::Half>(),beta.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n);
    } else {
        dtri_k<false,1,false><<<B,1024,shm>>>(A.data_ptr<float>(),(__half*)Mg.data_ptr<at::Half>(),(__half*)V.data_ptr<at::Half>(),beta.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n);
    }
    cudaError_t er=cudaGetLastError();TORCH_CHECK(er==cudaSuccess,"dtri: ",cudaGetErrorString(er));
    return {d,e,V,beta};
}

__device__ __forceinline__ void latrd_hmma16816(
        float2& d0, float2& d1,
        const __half2& a0, const __half2& a1,
        const __half2& a2, const __half2& a3,
        const __half2& b0, const __half2& b1){
    const float2 c0=d0, c1=d1;
    asm volatile(
        "mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
        "{%0, %1, %2, %3}, "
        "{%4, %5, %6, %7}, "
        "{%8, %9}, "
        "{%10, %11, %12, %13};"
        : "+f"(d0.x), "+f"(d0.y), "+f"(d1.x), "+f"(d1.y)
        : "r"(*reinterpret_cast<const unsigned*>(&a0)),
          "r"(*reinterpret_cast<const unsigned*>(&a1)),
          "r"(*reinterpret_cast<const unsigned*>(&a2)),
          "r"(*reinterpret_cast<const unsigned*>(&a3)),
          "r"(*reinterpret_cast<const unsigned*>(&b0)),
          "r"(*reinterpret_cast<const unsigned*>(&b1)),
          "f"(c0.x), "f"(c0.y), "f"(c1.x), "f"(c1.y));
}

template <bool PRECISE, bool MMA>
__global__ __launch_bounds__(1024, 1) void latrd_panel_k(
        __half* __restrict__ M, __half* __restrict__ V,
        __half* __restrict__ W, float* __restrict__ tau,
        float* __restrict__ dout, float* __restrict__ eout,
        int n, int c0, int width){
    const int tid=threadIdx.x, nt=blockDim.x, mat=blockIdx.x;
    const int lane=tid&31, warp=tid>>5, nwarp=nt>>5;
    using Reduce = typename std::conditional<PRECISE, double, float>::type;
    __half* Mm=M+(size_t)mat*n*n;
    __half* Vm=V+(size_t)mat*n*n;
    __half* Wm=W+(size_t)mat*n*n;
    float* v=(float*)_dyn;
    float* w=v+n;
    float* cv=w+n;
    float* cw=cv+32;
    __half2* VWs=(__half2*)(cw+32);
    __half* mma_tiles=(__half*)(VWs+(size_t)n*33);
    __shared__ Reduce red[32];

    for(int off=0;off<width;off++){
        const int i=c0+off;

        // SLATRD lower: materialize only the current panel column. The
        // rank-2k update of the trailing matrix is deferred to batched GEMM.
        for(int r=i+tid;r<n;r+=nt)
            v[r]=__half2float(Mm[(size_t)i*n+r]);
        __syncthreads();
        for(int r=i+warp;r<n;r+=nwarp){
            float correction=0.f;
            for(int jj=lane;jj<off;jj+=32){
                float2 vwr=__half22float2(VWs[(size_t)r*33+jj]);
                float2 vwi=__half22float2(VWs[(size_t)i*33+jj]);
                correction+=vwr.x*vwi.y;
                correction+=vwr.y*vwi.x;
            }
            #pragma unroll
            for(int delta=16;delta>0;delta>>=1)
                correction+=__shfl_down_sync(0xffffffff,correction,delta);
            if(lane==0) v[r]-=correction;
        }
        __syncthreads();

        if(tid==0) dout[(size_t)mat*n+i]=v[i];
        Reduce local=Reduce(0);
        for(int r=i+1+tid;r<n;r+=nt){
            float x=v[r];
            local+=(Reduce)x*(Reduce)x;
        }
        Reduce nrm2=dtri_block_sum(local,red);
        float nrm=PRECISE?(float)sqrt((double)nrm2):sqrtf((float)nrm2);
        if(nrm<1.0e-30f){
            for(int r=tid;r<n;r+=nt){
                VWs[(size_t)r*33+off]=__floats2half2_rn(0.f,0.f);
            }
            if(tid==0){
                eout[(size_t)mat*n+i]=0.f;
                tau[(size_t)mat*n+i]=0.f;
            }
            __syncthreads();
            continue;
        }

        float x0=v[i+1];
        float alpha=x0>=0.f?-nrm:nrm;
        float denom=x0-alpha;
        Reduce tail2=nrm2-(Reduce)x0*(Reduce)x0;
        tail2=tail2>Reduce(0)?tail2:Reduce(0);
        Reduce raw_vn2=(Reduce)denom*(Reduce)denom+tail2;
        float tau_i=raw_vn2>(Reduce)1.0e-30?
            (float)((Reduce)2*(Reduce)denom*(Reduce)denom/raw_vn2):0.f;
        float inv_denom=fabsf(denom)>1.0e-30f?1.f/denom:0.f;
        for(int r=tid;r<n;r+=nt){
            float vr=0.f;
            if(r==i+1) vr=1.f;
            else if(r>i+1) vr=v[r]*inv_denom;
            v[r]=vr;
            VWs[(size_t)r*33+off]=__floats2half2_rn(vr,0.f);
        }
        __syncthreads();

        if(warp<off){
            float dot_v=0.f,dot_w=0.f;
            for(int r=i+1+lane;r<n;r+=32){
                float vr=v[r];
                float2 vw=__half22float2(VWs[(size_t)r*33+warp]);
                dot_v+=vw.x*vr;
                dot_w+=vw.y*vr;
            }
            #pragma unroll
            for(int delta=16;delta>0;delta>>=1){
                dot_v+=__shfl_down_sync(0xffffffff,dot_v,delta);
                dot_w+=__shfl_down_sync(0xffffffff,dot_w,delta);
            }
            if(lane==0){ cv[warp]=dot_v; cw[warp]=dot_w; }
        }
        __syncthreads();

        if constexpr(MMA){
        // One native m16n8 MMA computes 16 rows.  B's eight columns are
        // identical; only accumulator column zero is written to w.
        const int rows=n-(i+1);
        const int full_tiles=rows>>4;
        for(int tile=warp;tile<full_tiles;tile+=nwarp){
            const int group=lane>>2, thread_group=lane&3;
            const int r0=i+1+(tile<<4);
            const int ra=r0+group, rb=ra+8;
            __half* mma_tile=mma_tiles+(size_t)warp*256;
            float2 acc0=make_float2(0.f,0.f);
            float2 acc1=make_float2(0.f,0.f);
            for(int k0=((i+1)&~15);k0<n;k0+=16){
                const int load_row=lane>>1;
                const int load_col=(lane&1)<<3;
                const int load_swizzle=(load_row&4)?8:0;
                *reinterpret_cast<uint4*>(
                    mma_tile+load_row*16+(load_col^load_swizzle))=
                    *reinterpret_cast<const uint4*>(
                        Mm+(size_t)(r0+load_row)*n+k0+load_col);
                __syncwarp();
                const int k=k0+(thread_group<<1);
                const int fragment_swizzle=(group&4)?8:0;
                const __half2 a0=*reinterpret_cast<const __half2*>(
                    mma_tile+group*16+((thread_group<<1)^fragment_swizzle));
                const __half2 a1=*reinterpret_cast<const __half2*>(
                    mma_tile+(group+8)*16+((thread_group<<1)^fragment_swizzle));
                const __half2 a2=*reinterpret_cast<const __half2*>(
                    mma_tile+group*16+(((thread_group<<1)+8)^fragment_swizzle));
                const __half2 a3=*reinterpret_cast<const __half2*>(
                    mma_tile+(group+8)*16+(((thread_group<<1)+8)^fragment_swizzle));
                const __half2 b0=__floats2half2_rn(v[k],v[k+1]);
                const __half2 b1=__floats2half2_rn(v[k+8],v[k+9]);
                latrd_hmma16816(acc0,acc1,a0,a1,a2,a3,b0,b1);
                __syncwarp();
            }
            float correction_a=0.f, correction_b=0.f;
            for(int jj=thread_group;jj<off;jj+=4){
                float2 vwa=__half22float2(VWs[(size_t)ra*33+jj]);
                float2 vwb=__half22float2(VWs[(size_t)rb*33+jj]);
                correction_a+=vwa.x*cw[jj]+vwa.y*cv[jj];
                correction_b+=vwb.x*cw[jj]+vwb.y*cv[jj];
            }
            correction_a+=__shfl_down_sync(0xffffffff,correction_a,2,4);
            correction_b+=__shfl_down_sync(0xffffffff,correction_b,2,4);
            correction_a+=__shfl_down_sync(0xffffffff,correction_a,1,4);
            correction_b+=__shfl_down_sync(0xffffffff,correction_b,1,4);
            if(thread_group==0){
                w[ra]=tau_i*(acc0.x-correction_a);
                w[rb]=tau_i*(acc1.x-correction_b);
            }
        }
        const int tail=i+1+(full_tiles<<4);
        for(int r=tail+warp;r<n;r+=nwarp){
            float av=0.f;
            for(int col=i+1+lane;col<n;col+=32)
                av+=__half2float(Mm[(size_t)r*n+col])*v[col];
            float correction=0.f;
            for(int jj=lane;jj<off;jj+=32){
                float2 vw=__half22float2(VWs[(size_t)r*33+jj]);
                correction+=vw.x*cw[jj]+vw.y*cv[jj];
            }
            #pragma unroll
            for(int delta=16;delta>0;delta>>=1){
                av+=__shfl_down_sync(0xffffffff,av,delta);
                correction+=__shfl_down_sync(0xffffffff,correction,delta);
            }
            if(lane==0) w[r]=tau_i*(av-correction);
        }
        __syncthreads();
        } else {
            for(int r=i+1+warp;r<n;r+=nwarp){
                float av=0.f;
                for(int col=i+1+lane;col<n;col+=32)
                    av+=__half2float(Mm[(size_t)r*n+col])*v[col];
                float correction=0.f;
                for(int jj=lane;jj<off;jj+=32){
                    float2 vw=__half22float2(VWs[(size_t)r*33+jj]);
                    correction+=vw.x*cw[jj];
                    correction+=vw.y*cv[jj];
                }
                #pragma unroll
                for(int delta=16;delta>0;delta>>=1){
                    av+=__shfl_down_sync(0xffffffff,av,delta);
                    correction+=__shfl_down_sync(0xffffffff,correction,delta);
                }
                if(lane==0) w[r]=tau_i*(av-correction);
            }
            __syncthreads();
        }

        local=Reduce(0);
        for(int r=i+1+tid;r<n;r+=nt)
            local+=(Reduce)w[r]*(Reduce)v[r];
        float gamma=(float)(-(Reduce)0.5*(Reduce)tau_i*
                            dtri_block_sum(local,red));
        for(int r=i+1+tid;r<n;r+=nt){
            float wr=w[r]+gamma*v[r];
            VWs[(size_t)r*33+off]=__floats2half2_rn(v[r],wr);
        }
        if(tid==0){
            eout[(size_t)mat*n+i]=alpha;
            tau[(size_t)mat*n+i]=tau_i;
        }
        __syncthreads();
    }

    for(int idx=tid;idx<n*width;idx+=nt){
        int r=idx/width, jj=idx-r*width;
        __half2 vw=VWs[(size_t)r*33+jj];
        Vm[(size_t)r*n+c0+jj]=__low2half(vw);
        Wm[(size_t)r*n+c0+jj]=__high2half(vw);
    }
}

__global__ void latrd_finish_k(const __half* __restrict__ M,
        float* __restrict__ d, float* __restrict__ e, int n){
    const int mat=blockIdx.x;
    if(threadIdx.x==0){
        const __half* Mm=M+(size_t)mat*n*n;
        d[(size_t)mat*n+n-2]=__half2float(Mm[(size_t)(n-2)*n+n-2]);
        d[(size_t)mat*n+n-1]=__half2float(Mm[(size_t)(n-1)*n+n-1]);
        e[(size_t)mat*n+n-2]=__half2float(Mm[(size_t)(n-1)*n+n-2]);
    }
}

void f_latrd_panel(torch::Tensor M, torch::Tensor V, torch::Tensor W,
        torch::Tensor tau, torch::Tensor d, torch::Tensor e,
        int c0, int width, bool precise){
    int B=M.size(0),n=M.size(1); c10::cuda::CUDAGuard g(M.device());
    TORCH_CHECK(c0>=0 && width>=0 && c0+width<=n-2,"latrd panel bounds");
    size_t shm=sizeof(float)*(2*n+64)+sizeof(__half)*(size_t)2*n*33;
    if(n<=512) shm+=sizeof(__half)*(size_t)32*256;
    if(precise){
        if(n<=512){
            cudaFuncSetAttribute(latrd_panel_k<true,true>,cudaFuncAttributeMaxDynamicSharedMemorySize,shm);
            latrd_panel_k<true,true><<<B,1024,shm>>>((__half*)M.data_ptr<at::Half>(),
                (__half*)V.data_ptr<at::Half>(),(__half*)W.data_ptr<at::Half>(),
                tau.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n,c0,width);
        } else {
            cudaFuncSetAttribute(latrd_panel_k<true,false>,cudaFuncAttributeMaxDynamicSharedMemorySize,shm);
            latrd_panel_k<true,false><<<B,1024,shm>>>((__half*)M.data_ptr<at::Half>(),
                (__half*)V.data_ptr<at::Half>(),(__half*)W.data_ptr<at::Half>(),
                tau.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n,c0,width);
        }
    } else {
        if(n<=512){
            cudaFuncSetAttribute(latrd_panel_k<false,true>,cudaFuncAttributeMaxDynamicSharedMemorySize,shm);
            latrd_panel_k<false,true><<<B,1024,shm>>>((__half*)M.data_ptr<at::Half>(),
                (__half*)V.data_ptr<at::Half>(),(__half*)W.data_ptr<at::Half>(),
                tau.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n,c0,width);
        } else {
            cudaFuncSetAttribute(latrd_panel_k<false,false>,cudaFuncAttributeMaxDynamicSharedMemorySize,shm);
            latrd_panel_k<false,false><<<B,1024,shm>>>((__half*)M.data_ptr<at::Half>(),
                (__half*)V.data_ptr<at::Half>(),(__half*)W.data_ptr<at::Half>(),
                tau.data_ptr<float>(),d.data_ptr<float>(),e.data_ptr<float>(),n,c0,width);
        }
    }
    cudaError_t er=cudaGetLastError();
    TORCH_CHECK(er==cudaSuccess,"latrd_panel: ",cudaGetErrorString(er));
}

void f_latrd_finish(torch::Tensor M, torch::Tensor d, torch::Tensor e){
    int B=M.size(0),n=M.size(1); c10::cuda::CUDAGuard g(M.device());
    latrd_finish_k<<<B,32>>>((const __half*)M.data_ptr<at::Half>(),
        d.data_ptr<float>(),e.data_ptr<float>(),n);
    cudaError_t er=cudaGetLastError();
    TORCH_CHECK(er==cudaSuccess,"latrd_finish: ",cudaGetErrorString(er));
}
__device__ int _sturm(const float* d,const float* e,int n,float x){
    int cnt=0; float q=d[0]-x;
    for(int i=0;i<n;i++){ if(q<0)cnt++; if(i<n-1){ float ee=e[i]; q=(d[i+1]-x)-__fdividef(ee*ee,(q==0.f?1e-30f:q));} }
    return cnt;
}
__global__ void biseig_k(const float* __restrict__ din,const float* __restrict__ ein,
        float* __restrict__ wout,float* __restrict__ Qout,float* __restrict__ scratch,int n,int nt_,int groups){
    float* d=(float*)_dyn; float* e=d+n;
    const int tid=threadIdx.x, nt=blockDim.x, blk=blockIdx.x;
    const int mat=blk/groups, group=blk-mat*groups;
    for(int i=tid;i<n;i+=nt){ d[i]=din[mat*n+i]; e[i]=ein[mat*n+i]; }
    __syncthreads();
    float lo=1e30f,hi=-1e30f;
    for(int i=0;i<n;i++){ float r=(i>0?fabsf(e[i-1]):0.f)+(i<n-1?fabsf(e[i]):0.f);
        lo=fminf(lo,d[i]-r); hi=fmaxf(hi,d[i]+r); }
    for(int k=group*nt+tid;k<n;k+=groups*nt){
        float a=lo,b=hi;
        for(int it=0;it<24;it++){ float mid=0.5f*(a+b); if(_sturm(d,e,n,mid)<=k) a=mid; else b=mid; }
        wout[mat*n+k]=0.5f*(a+b);
    }
    __syncthreads();
    float* w=wout;
    for(int k=group*nt+tid;k<n;k+=groups*nt){
        float lam=w[mat*n+k]+1e-6f;
        float* cpr=scratch+((size_t)blk*nt_+tid)*2*n; float* invm=cpr+n;
        float b0=d[0]-lam, c0=(n>1?e[0]:0.f);
        invm[0]=__fdividef(1.0f,b0); cpr[0]=c0*invm[0];
        for(int i=1;i<n;i++){ float ai=e[i-1], bi=d[i]-lam, ci=(i<n-1?e[i]:0.f);
            float m=bi-ai*cpr[i-1]; if(fabsf(m)<1e-30f)m=1e-30f;
            invm[i]=__fdividef(1.0f,m); cpr[i]=ci*invm[i]; }
        float* qcol=Qout+(size_t)mat*n*n+k;
        for(int iter=0;iter<2;iter++){
            float rhs0=(iter==0)?1.0f:qcol[0];
            float yi=rhs0*invm[0]; qcol[0]=yi;
            for(int i=1;i<n;i++){ float rhs=(iter==0)?1.0f:qcol[(size_t)i*n];
                yi=(rhs-e[i-1]*yi)*invm[i]; qcol[(size_t)i*n]=yi; }
            float xn=qcol[(size_t)(n-1)*n]; float nrm=xn*xn;
            for(int i=n-2;i>=0;i--){ xn=qcol[(size_t)i*n]-cpr[i]*xn; qcol[(size_t)i*n]=xn; nrm+=xn*xn; }
            nrm=sqrtf(nrm);
            for(int i=0;i<n;i++) qcol[(size_t)i*n]=__fdividef(qcol[(size_t)i*n],nrm);
        }
    }
}
std::vector<torch::Tensor> f_biseig(torch::Tensor d,torch::Tensor e,torch::Tensor scratch){
    int B=d.size(0),n=d.size(1); c10::cuda::CUDAGuard g(d.device());
    int nt=128,workers=scratch.size(1); TORCH_CHECK(workers%nt==0,"biseig workers");
    int groups=workers/nt; TORCH_CHECK(groups>=1,"biseig groups");
    auto w=torch::zeros({B,n},d.options()),Q=torch::zeros({B,n,n},d.options());
    size_t shm=sizeof(float)*2*n;
    cudaFuncSetAttribute(biseig_k,cudaFuncAttributeMaxDynamicSharedMemorySize,shm);
    biseig_k<<<B*groups,nt,shm>>>(d.data_ptr<float>(),e.data_ptr<float>(),w.data_ptr<float>(),Q.data_ptr<float>(),scratch.data_ptr<float>(),n,nt,groups);
    cudaError_t er=cudaGetLastError();TORCH_CHECK(er==cudaSuccess,"biseig: ",cudaGetErrorString(er));
    return {w,Q};
}
__global__ void larft_gram_k(
        const float* __restrict__ gram, const float* __restrict__ tau,
        __half* __restrict__ Tout, int n, int nb, int npanel){
    int mat=blockIdx.x/npanel, pi=blockIdx.x%npanel;
    int c0=pi*nb, c1=c0+nb; if(c1>n-2)c1=n-2; int w=c1-c0; if(w<=0) return;
    const int tid=threadIdx.x, nt=blockDim.x;
    const int ld=nb+1;
    float* T=(float*)_dyn; float* tmp=T+nb*ld;
    const float* G=gram+(size_t)(mat*npanel+pi)*nb*nb;
    for(int i=tid;i<nb*ld;i+=nt) T[i]=0.f;
    __syncthreads();
    if(tid==0) T[0]=tau[mat*n+c0];
    __syncthreads();
    for(int j=1;j<w;j++){
        for(int k=tid;k<j;k+=nt)
            T[k*ld+j]=-tau[mat*n+c0+j]*G[k+j*nb];
        __syncthreads();
        for(int r=tid;r<j;r+=nt){ float acc=0.f; for(int k=r;k<j;k++) acc+=T[r*ld+k]*T[k*ld+j]; tmp[r]=acc; }
        __syncthreads();
        for(int r=tid;r<j;r+=nt) T[r*ld+j]=tmp[r];
        if(tid==0) T[j*ld+j]=tau[mat*n+c0+j];
        __syncthreads();
    }
    for(int i=tid;i<nb*nb;i+=nt){ int r=i/nb, c=i-r*nb;
        Tout[(size_t)(mat*npanel+pi)*nb*nb+i]=__float2half(T[r*ld+c]); }
}
static cublasHandle_t larft_blas = nullptr;
torch::Tensor f_larft(torch::Tensor V, torch::Tensor tau, int nb){
    int B=V.size(0),n=V.size(1); int npanel=(n-2+nb-1)/nb;
    c10::cuda::CUDAGuard g(V.device());
    auto T=torch::zeros({B,npanel,nb,nb},V.options());
    auto gram=torch::empty({B,npanel,nb,nb},V.options().dtype(torch::kFloat32));
    if(larft_blas==nullptr){
        cublasStatus_t s=cublasCreate(&larft_blas);
        TORCH_CHECK(s==CUBLAS_STATUS_SUCCESS,"larft cublasCreate: ",(int)s);
        s=cublasSetMathMode(larft_blas,CUBLAS_TENSOR_OP_MATH);
        TORCH_CHECK(s==CUBLAS_STATUS_SUCCESS,"larft cublasSetMathMode: ",(int)s);
    }
    float alpha=1.f,beta=0.f;
    long long stride_v=(long long)n*n;
    long long stride_g=(long long)npanel*nb*nb;
    const __half* Vp=(const __half*)V.data_ptr<at::Half>();
    float* Gp=gram.data_ptr<float>();
    for(int pi=0;pi<npanel;pi++){
        int c0=pi*nb, w=n-2-c0; if(w>nb) w=nb;
        cublasStatus_t s=cublasGemmStridedBatchedEx(
            larft_blas,CUBLAS_OP_N,CUBLAS_OP_T,w,w,n,
            &alpha,Vp+c0,CUDA_R_16F,n,stride_v,
            Vp+c0,CUDA_R_16F,n,stride_v,&beta,
            Gp+(size_t)pi*nb*nb,CUDA_R_32F,nb,stride_g,B,
            CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
        TORCH_CHECK(s==CUBLAS_STATUS_SUCCESS,"larft gram: ",(int)s);
    }
    size_t shm=sizeof(float)*(nb*(nb+1)+nb);
    larft_gram_k<<<B*npanel,128,shm>>>(gram.data_ptr<float>(),tau.data_ptr<float>(),(__half*)T.data_ptr<at::Half>(),n,nb,npanel);
    cudaError_t e=cudaGetLastError();TORCH_CHECK(e==cudaSuccess,"larft: ",cudaGetErrorString(e));
    return T;
}

static cublasHandle_t chol_blas = nullptr;
static cusolverDnHandle_t chol_solver = nullptr;

static void check_blas(cublasStatus_t status, const char* where){
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, where, " failed with status ", (int)status);
}
static void check_solver(cusolverStatus_t status, const char* where){
    TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, where, " failed with status ", (int)status);
}
static void ensure_chol_handles(){
    if(chol_blas == nullptr){
        check_blas(cublasCreate(&chol_blas), "cublasCreate");
        check_blas(cublasSetMathMode(chol_blas, CUBLAS_TF32_TENSOR_OP_MATH), "cublasSetMathMode");
    }
    if(chol_solver == nullptr) check_solver(cusolverDnCreate(&chol_solver), "cusolverDnCreate");
}

__global__ void chol_ptrs_k(float* G, float* Q, float** A, float** B, size_t stride, int batch){
    int i=blockIdx.x*blockDim.x+threadIdx.x;
    if(i<batch){ A[i]=G+(size_t)i*stride; B[i]=Q+(size_t)i*stride; }
}
__global__ void chol_shift_k(float* G, int n){
    int mat=blockIdx.x, tid=threadIdx.x;
    __shared__ float sum[256];
    float local=0.f;
    for(int i=tid;i<n;i+=blockDim.x) local+=G[(size_t)mat*n*n+(size_t)i*n+i];
    sum[tid]=local; __syncthreads();
    for(int s=blockDim.x/2;s>0;s>>=1){ if(tid<s) sum[tid]+=sum[tid+s]; __syncthreads(); }
    float shift=1.0e-6f*fmaxf(sum[0]/(float)n,1.0e-20f);
    for(int i=tid;i<n;i+=blockDim.x) G[(size_t)mat*n*n+(size_t)i*n+i]+=shift;
}
__global__ void polar_factor_k(const float* G, float* P, size_t total, int n){
    size_t i=(size_t)blockIdx.x*blockDim.x+threadIdx.x;
    if(i<total){
        size_t pos=i%(size_t)(n*n);
        int row=(int)(pos/n), col=(int)(pos%n);
        P[i]=0.375f*P[i]-1.25f*G[i]+(row==col?1.875f:0.f);
    }
}

torch::Tensor f_cholqr(torch::Tensor Q){
    TORCH_CHECK(Q.is_cuda() && Q.scalar_type()==torch::kFloat32, "Q must be CUDA float32");
    TORCH_CHECK(Q.dim()==3 && Q.size(1)==Q.size(2), "Q must be batch x n x n");
    c10::cuda::CUDAGuard guard(Q.device());
    ensure_chol_handles();
    int batch=(int)Q.size(0), n=(int)Q.size(1);
    size_t stride=(size_t)n*n;
    auto G=torch::empty({batch,n,n},Q.options());
    auto ptrs=torch::empty({2,batch},Q.options().dtype(torch::kInt64));
    auto info=torch::empty({batch},Q.options().dtype(torch::kInt32));
    float alpha=1.f,beta=0.f;
    check_blas(cublasSgemmStridedBatched(chol_blas,CUBLAS_OP_N,CUBLAS_OP_T,
        n,n,n,&alpha,Q.data_ptr<float>(),n,(long long)stride,Q.data_ptr<float>(),n,
        (long long)stride,&beta,G.data_ptr<float>(),n,(long long)stride,batch),
        "cublasSgemmStridedBatched");
    chol_shift_k<<<batch,256>>>(G.data_ptr<float>(),n);
    float** Ap=(float**)ptrs.data_ptr<int64_t>();
    float** Bp=Ap+batch;
    chol_ptrs_k<<<(batch+255)/256,256>>>(G.data_ptr<float>(),Q.data_ptr<float>(),Ap,Bp,stride,batch);
    check_solver(cusolverDnSpotrfBatched(chol_solver,CUBLAS_FILL_MODE_LOWER,n,Ap,n,
        info.data_ptr<int>(),batch),"cusolverDnSpotrfBatched");
    check_blas(cublasStrsmBatched(chol_blas,CUBLAS_SIDE_LEFT,CUBLAS_FILL_MODE_LOWER,
        CUBLAS_OP_N,CUBLAS_DIAG_NON_UNIT,n,n,&alpha,
        const_cast<const float**>(Ap),n,Bp,n,batch),"cublasStrsmBatched");
    cudaError_t e=cudaGetLastError();
    TORCH_CHECK(e==cudaSuccess,"f_cholqr: ",cudaGetErrorString(e));
    return Q;
}

torch::Tensor f_nsqr(torch::Tensor Q){
    TORCH_CHECK(Q.is_cuda() && Q.scalar_type()==torch::kFloat32, "Q must be CUDA float32");
    TORCH_CHECK(Q.dim()==3 && Q.size(1)==Q.size(2), "Q must be batch x n x n");
    c10::cuda::CUDAGuard guard(Q.device());
    ensure_chol_handles();
    int batch=(int)Q.size(0), n=(int)Q.size(1);
    size_t stride=(size_t)n*n, total=(size_t)batch*stride;
    auto G=torch::empty({batch,n,n},Q.options());
    auto P=torch::empty({batch,n,n},Q.options());
    auto out=torch::empty_like(Q);
    float alpha=1.f,beta=0.f;
    check_blas(cublasSgemmStridedBatched(chol_blas,CUBLAS_OP_N,CUBLAS_OP_T,
        n,n,n,&alpha,Q.data_ptr<float>(),n,(long long)stride,Q.data_ptr<float>(),n,
        (long long)stride,&beta,G.data_ptr<float>(),n,(long long)stride,batch),
        "polar gram");
    check_blas(cublasSgemmStridedBatched(chol_blas,CUBLAS_OP_N,CUBLAS_OP_N,
        n,n,n,&alpha,G.data_ptr<float>(),n,(long long)stride,G.data_ptr<float>(),n,
        (long long)stride,&beta,P.data_ptr<float>(),n,(long long)stride,batch),
        "polar gram square");
    polar_factor_k<<<(total+255)/256,256>>>(G.data_ptr<float>(),P.data_ptr<float>(),total,n);
    check_blas(cublasSgemmStridedBatched(chol_blas,CUBLAS_OP_N,CUBLAS_OP_N,
        n,n,n,&alpha,P.data_ptr<float>(),n,(long long)stride,Q.data_ptr<float>(),n,
        (long long)stride,&beta,out.data_ptr<float>(),n,(long long)stride,batch),
        "polar update");
    cudaError_t e=cudaGetLastError();
    TORCH_CHECK(e==cudaSuccess,"f_nsqr: ",cudaGetErrorString(e));
    return out;
}
"""
        try:
            _fused_mod = load_inline(
                name="fused_eigh_mod_sub", cpp_sources=fcpp, cuda_sources=fcuda,
                functions=["f_dtri", "f_latrd_panel", "f_latrd_finish",
                           "f_biseig", "f_larft", "f_cholqr", "f_nsqr"],
                extra_cuda_cflags=["-O3"], extra_ldflags=["-lcublas", "-lcusolver"],
                with_cuda=True, verbose=False)
        except Exception:
            _fused_failed = True
            return None
    return _fused_mod


def _fused_eigh(data, nb=32, precise_dtri=True):
    # fp16 tridiagonalize -> bisection+inverse-iteration solve -> WY fp16 backtransform ->
    # Cholesky-QR reorth (TF32 OFF) -> Rayleigh eigenvalues.
    fm = _get_fused_mod()
    if fm is None:
        return None
    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        b, n, _ = data.shape
        d, e, V, beta = fm.f_dtri(data.contiguous(), precise_dtri)
        biseig_groups = (n + 127) // 128
        scratch = torch.empty(b, 128 * biseig_groups, 2 * n, device=data.device)
        w, Z = fm.f_biseig(d.contiguous(), e.contiguous(), scratch)
        T = fm.f_larft(V.contiguous(), beta.contiguous(), nb)
        npanel = T.shape[1]
        Zh = Z.half()
        for pi in range(npanel - 1, -1, -1):
            c0 = pi * nb
            c1 = min(c0 + nb, n - 2)
            if c1 <= c0:
                continue
            Vp = V[:, :, c0:c1]
            Tp = T[:, pi, :c1 - c0, :c1 - c0]
            inner = Tp @ (Vp.transpose(1, 2) @ Zh)
            Zh = torch.baddbmm(Zh, Vp, inner, beta=1.0, alpha=-1.0)
        Q = Zh.float()
        if (b == 40 and n in (176, 352)) or (b >= 16 and n == 768):
            Q = fm.f_nsqr(Q)
            if n == 768:
                Q = fm.f_nsqr(Q)
        else:
            Q = fm.f_cholqr(Q)
        return Q.contiguous(), w.contiguous()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32


def _blocked_fused_eigh(data, nb=32, precise_dtri=True, wy_nb=None):
    fm = _get_fused_mod()
    if fm is None:
        return None
    previous_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        batch, n, _ = data.shape
        work = data.half()
        vectors = torch.zeros_like(work)
        update = torch.zeros_like(work)
        tau = torch.zeros((batch, n), device=data.device, dtype=torch.float32)
        d = torch.zeros_like(tau)
        e = torch.zeros_like(tau)
        for c0 in range(0, n - 2, nb):
            c1 = min(c0 + nb, n - 2)
            fm.f_latrd_panel(work, vectors, update, tau, d, e,
                             c0, c1 - c0, precise_dtri)
            trailing = work[:, c1:, c1:]
            vp = vectors[:, c1:, c0:c1]
            wp = update[:, c1:, c0:c1]
            trailing.baddbmm_(vp, wp.transpose(1, 2), beta=1.0, alpha=-1.0)
            trailing.baddbmm_(wp, vp.transpose(1, 2), beta=1.0, alpha=-1.0)
        fm.f_latrd_finish(work, d, e)

        biseig_groups = (n + 127) // 128
        scratch = torch.empty(batch, 128 * biseig_groups, 2 * n,
                              device=data.device)
        values, z = fm.f_biseig(d.contiguous(), e.contiguous(), scratch)
        transform_nb = nb if wy_nb is None else wy_nb
        t = fm.f_larft(vectors.contiguous(), tau.contiguous(), transform_nb)
        zh = z.half()
        for panel in range(t.shape[1] - 1, -1, -1):
            c0 = panel * transform_nb
            c1 = min(c0 + transform_nb, n - 2)
            if c1 <= c0:
                continue
            vp = vectors[:, :, c0:c1]
            tp = t[:, panel, :c1 - c0, :c1 - c0]
            inner = tp @ (vp.transpose(1, 2) @ zh)
            zh = torch.baddbmm(zh, vp, inner, beta=1.0, alpha=-1.0)
        q = fm.f_cholqr(zh.float())
        return q.contiguous(), values.contiguous()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = previous_tf32


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n == 32 and _get_bf16x9_mod() is None:
        raise RuntimeError("BF16x9 cuBLAS module construction failed")

    if n == 4096:
        try:
            diagonal = _diagonal_4096(data)
            if diagonal is not None:
                return diagonal
        except Exception:
            pass

    if batch > 1 and n == 32:
        mod = _get_xsyev_mod()
        if mod is not None:
            try:
                vectors, values = mod.jacobi32_sharedwarp_batched(data)
                return vectors, values
            except Exception:
                try:
                    vectors, values = mod.syevj32_batched(data)
                    return vectors, values
                except Exception:
                    pass

    if batch > 1 and n in (176, 352):
        try:
            fused = _fused_eigh(data, precise_dtri=(n != 176))
            if fused is not None:
                return fused
        except Exception:
            pass

    if batch > 1 and n == 1024:
        try:
            block = _row_scaled_block1024_coupled(data)
            if block is not None:
                return block
            block = _row_scaled_block1024(data)
            if block is not None:
                return block
            lowrank_mask = _lowrank1024_mask(data)
            if lowrank_mask is not None and bool(lowrank_mask.all().item()):
                cand = _panel_qr_lowrank1024_eigh(data)
                if cand is not None:
                    return cand
            geometric_mask = _lapack_geometric1024_mask(data)
            if geometric_mask is not None and bool(geometric_mask.all().item()):
                cand = _panel_qr_geometric1024_eigh(data)
                if cand is not None:
                    return cand
        except Exception:
            pass

    if batch > 1 and n == 512:
        try:
            block = _row_scaled_block512_coupled(data)
            if block is not None:
                return block
            block = _row_scaled_block512_partial(data)
            if block is not None:
                return block
            if _is_lapack_even512_batch(data):
                cand = _blocked_fused_eigh(data, wy_nb=48)
                if cand is not None:
                    return cand
            if _is_rankdef512_batch(data):
                cand = _panel_qr_rankdef512_eigh(data)
                if cand is not None:
                    return cand
            if _is_clustered_pm1(data):
                cand = _clustered_projector_chol_eigh(data)
                if cand is not None:
                    return cand
                cand = _panel_qr_projector_clustered_eigh(data)
                if cand is not None:
                    return cand
                cand = _clustered_pm1_eigh(data)
                if cand is not None and _eigen_residual_scaled(data, cand[0], cand[1]) < 100.0:
                    return cand
        except Exception:
            pass

    if batch > 1 and n == 2048:
        try:
            block = _row_scaled_tf32(_row_scaled_block2048_coupled, data)
            if block is not None:
                return block
        except Exception:
            pass

    if batch > 1 and 176 <= n <= 2048:
        mod = _get_xsyev_mod()
        if mod is not None:
            try:
                vectors, values = mod.xsyev_batched(data)
                return vectors, values
            except Exception:
                pass

    values, vectors = torch.linalg.eigh(data)
    return vectors, values
scrolls · 4552 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