Details
CUDA: XOR swizzle flash attn K,V smem fp16 tiles (#25635)
- CUDA: XOR swizzle flash attn K,V smem fp16 tiles
Signed-off-by: ynankani ynankani@nvidia.com
- Fix use 64bit generic pointer instead of 32bit shared pointer
Signed-off-by: ynankani ynankani@nvidia.com
-
fix shared memory race in FA on DGX Spark
-
Handle corener case
Signed-off-by: ynankani ynankani@nvidia.com
- Add swizzle test cases and gate sync for swizzled path only
Signed-off-by: ynankani ynankani@nvidia.com
- gate CUDA PTX
Signed-off-by: ynankani ynankani@nvidia.com
- offset calculation specific for swizzle branch
Signed-off-by: ynankani ynankani@nvidia.com
- Reafctor code
Signed-off-by: ynankani ynankani@nvidia.com
- Refactor FA swizzle ldmatrix if/else into helpers (K row/col, V offset)
Signed-off-by: ynankani ynankani@nvidia.com
- rebase and update test case args
Signed-off-by: ynankani ynankani@nvidia.com
- Allow swizzle for non-pow2 shapes, for which nbatch_2%32==0
Signed-off-by: ynankani ynankani@nvidia.com
Signed-off-by: ynankani ynankani@nvidia.com
Website:
Attestations:
macOS/iOS:
- macOS Apple Silicon (arm64)
- macOS Apple Silicon (arm64, KleidiAI enabled) DISABLED
- macOS Intel (x64)
- iOS XCFramework
Linux:
- Ubuntu x64 (CPU)
- Ubuntu arm64 (CPU)
- Ubuntu s390x (CPU)
- Ubuntu x64 (Vulkan)
- Ubuntu arm64 (Vulkan)
- Ubuntu x64 (ROCm 7.14)
- Ubuntu x64 (OpenVINO)
- Ubuntu x64 (SYCL FP32)
- Ubuntu x64 (SYCL FP16)
Android:
Windows:
- Windows x64 (CPU)
- Windows arm64 (CPU)
- Windows arm64 (OpenCL Adreno)
- Windows x64 (CUDA 12) - CUDA 12.4 DLLs
- Windows x64 (CUDA 13) - CUDA 13.3 DLLs
- Windows arm64 (CUDA 13) (preview) - CUDA 13.4 DLLs
- Windows x64 (Vulkan)
- Windows x64 (OpenVINO)
- Windows x64 (SYCL)
- Windows x64 (ROCm 7.14)
openEuler:
- DISABLED
- openEuler x86 (310p)
- openEuler x86 (910b, ACL Graph)
- openEuler aarch64 (310p)
- openEuler aarch64 (910b, ACL Graph)
UI: