submission 930373
Pranshu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 14951 lines, June 9 Researcher Reciprocity License v1.0.
x38_sprint_lt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-930373?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:93eb76b7f37cf2a943d54713395f79c8bc4d211b2e51f0a9243847a99495070c
license declaredunknown
license concludedunknown
authorsPranshu
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"one-rail NVFP4 outer board exceeds 88 KiB");mbarrier
"fence.mbarrier_init.release.cluster;"mma
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "persistent-kernel
bool use_persistent256 =shared-memory
extern __shared__ __align__(16) float smem[];tcgen05
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "tma
const CUtensorMap& map,vector-width = float4
__device__ __forceinline__ float4 loco_ld_global_v4(Kernel source
x38_sprint_lt.py14951 lines
import hashlib
from pathlib import Path
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
_major, _minor = torch.cuda.get_device_capability()
if (_major, _minor) not in ((10, 0), (12, 0)):
raise RuntimeError("xRCholo1024 requires SM100a or SM120a")
_SM_ARCH = "100a" if _major == 10 else "120a"
_MAX_G = 2
_PAYLOAD_PIPELINE = 1
_NVFP4_OUTER = 0
_LEFT_MACRO = 1024
_K256_FRONT = 0
_HEXLIFT_OUTER = 0
_HEXLIFT_MIN_CTAS = 96
_HEXLIFT_MIN_HISTORY = 256
_HEX6_ALL = 0
_HEX4_GE8192 = 0
_B60_ISSUES = 8
_B60_N128 = 0
_TMA_F16_B60 = 0
CPP = r"""
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContextLight.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <cublasLt.h>
#include <torch/library.h>
#include <cuda_runtime_api.h>
#ifndef XCAL_XRCHOLV18_NVFP4_OUTER
#define XCAL_XRCHOLV18_NVFP4_OUTER 0
#endif
constexpr int64_t XCAL_LARGE3_LT_WORKSPACE_BYTES = 128ll << 20;
extern "C" int xrcholv18_cholesky_launch(
const float* input,
float* factor,
float* output,
uint8_t* plane_values,
uint8_t* plane_scales,
void* lt_handle,
void* blas_handle,
void* lt_workspace,
uint64_t lt_workspace_bytes,
int batch,
int input_n,
int work_n,
int resident,
int use_tcgen,
int use_e2m1);
at::Tensor xrchol_final(const at::Tensor& input) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch,n,n]");
TORCH_CHECK(input.size(1) == input.size(2), "input matrices must be square");
c10::cuda::CUDAGuard guard(input.device());
auto contiguous = input.contiguous();
int batch = static_cast<int>(contiguous.size(0));
int n = static_cast<int>(contiguous.size(1));
if (!batch || !n) return at::empty_like(contiguous);
int device = contiguous.get_device();
thread_local int cached_device = -1;
thread_local cudaDeviceProp cached_prop{};
if (cached_device != device) {
cudaError_t status = cudaGetDeviceProperties(&cached_prop, device);
TORCH_CHECK(
status == cudaSuccess,
"cudaGetDeviceProperties failed: ",
cudaGetErrorString(status));
cached_device = device;
}
int resident = cached_prop.multiProcessorCount;
if (resident > 148) resident = 148;
TORCH_CHECK(
(cached_prop.major == 10 && cached_prop.minor == 0)
|| (cached_prop.major == 12 && cached_prop.minor == 0),
"xRCholo1024 requires SM100a or SM120a");
bool use_tcgen = n > 128;
int work_n = use_tcgen ? ((n + 127) & ~127) : n;
bool use_persistent256 =
use_tcgen && n == 256 && work_n == 256 && batch != 64;
// Exact B640 is a batch/frontier board. Keep its FP32 factor plus the
// one-rail FP16 sidecar alive for the K128 MMA front.
bool use_persistent512 = false;
bool use_fp16x3 =
use_tcgen
&& n < 2048
&& batch < 16
&& !(n == 1024 && batch == 4);
bool use_e2m1 = use_tcgen && work_n >= 8192;
bool use_lt_large3 =
use_tcgen
&& batch == 1
&& work_n == n
&& (n == 8192 || n == 16384 || n == 32768);
bool use_blas32 =
use_tcgen
&& batch == 1
&& work_n == n
&& n == 32768;
bool use_blas_mid =
use_tcgen
&& work_n == n
&& ((n == 512 && (batch == 16 || batch == 640))
|| (n == 1024 && (batch == 4 || batch == 60))
|| (n == 2048 && (batch == 2 || batch == 8))
|| (n == 4096 && (batch == 1 || batch == 2)));
bool use_lt_tuned = use_lt_large3 || use_blas_mid;
void* lt_handle = nullptr;
if (use_lt_tuned) {
static thread_local void* cached_lt_handle = nullptr;
if (!cached_lt_handle) {
cached_lt_handle = reinterpret_cast<void*>(
at::cuda::getCurrentCUDABlasLtHandle());
}
lt_handle = cached_lt_handle;
}
void* blas_handle = nullptr;
if (use_blas32 || use_blas_mid || use_lt_large3) {
static thread_local void* cached_blas_handle = nullptr;
if (!cached_blas_handle) {
cached_blas_handle = reinterpret_cast<void*>(
at::cuda::getCurrentCUDABlasHandle());
}
blas_handle = cached_blas_handle;
}
at::Tensor output;
at::Tensor factor;
at::Tensor factor_workspace;
at::Tensor plane_scale_storage;
uint8_t* plane_values = nullptr;
uint8_t* plane_scales = nullptr;
void* lt_workspace = nullptr;
uint64_t lt_workspace_bytes = 0;
if (use_persistent256 || use_persistent512) {
output = at::empty_like(contiguous);
factor = output;
} else if (use_tcgen) {
int64_t factor_elements =
(int64_t)batch * (int64_t)work_n * (int64_t)work_n;
int64_t fp16_words = (factor_elements + 1) >> 1;
int64_t fp16_shadow_words =
use_fp16x3 ? factor_elements : fp16_words;
int64_t cells = use_e2m1
? (int64_t)batch * 2 * (work_n - 128)
: 0;
int64_t plane_value_bytes = cells * 4 * 2 * 16;
int64_t plane_value_words = (plane_value_bytes + 3) >> 2;
int64_t lt_workspace_words = use_lt_tuned
? (XCAL_LARGE3_LT_WORKSPACE_BYTES + 3) >> 2
: 0;
if (use_e2m1) {
int64_t scale_bytes = cells * 4;
plane_scale_storage = at::empty(
{scale_bytes},
contiguous.options().dtype(at::kByte));
plane_scales = plane_scale_storage.data_ptr<uint8_t>();
}
factor_workspace = at::empty(
{factor_elements
+ fp16_shadow_words
+ plane_value_words
+ lt_workspace_words},
contiguous.options());
factor = factor_workspace
.narrow(0, 0, factor_elements)
.view({batch, work_n, work_n});
output = work_n == n
? factor.view_as(contiguous)
: at::empty_like(contiguous);
plane_values = reinterpret_cast<uint8_t*>(
factor_workspace.data_ptr<float>() + factor_elements);
if (lt_workspace_words) {
lt_workspace = factor_workspace.data_ptr<float>()
+ factor_elements
+ fp16_shadow_words
+ plane_value_words;
lt_workspace_bytes =
(uint64_t)lt_workspace_words * sizeof(float);
}
} else {
output = at::empty_like(contiguous);
factor = output;
}
int status = xrcholv18_cholesky_launch(
contiguous.data_ptr<float>(),
factor.data_ptr<float>(),
output.data_ptr<float>(),
plane_values,
plane_scales,
lt_handle,
blas_handle,
lt_workspace,
lt_workspace_bytes,
batch,
n,
work_n,
resident,
use_tcgen ? 1 : 0,
use_e2m1 ? 1 : 0);
TORCH_CHECK(
status == static_cast<int>(cudaSuccess),
"CUDA launch failed: ",
cudaGetErrorString(static_cast<cudaError_t>(status)));
return output;
}
TORCH_LIBRARY(xcaliber_x38_sprint_lt, m) {
m.def("chol(Tensor input) -> Tensor");
m.impl("chol", &xrchol_final);
}
"""
CUDA = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <cublasLt.h>
#include <stdint.h>
#ifndef XCAL_XRCHOLV18_MAX_G
#define XCAL_XRCHOLV18_MAX_G 16
#endif
#ifndef XCAL_XRCHOLV18_PAYLOAD_PIPELINE
#define XCAL_XRCHOLV18_PAYLOAD_PIPELINE 1
#endif
#ifndef XCAL_XRCHOLV18_B60_ISSUES
#define XCAL_XRCHOLV18_B60_ISSUES 8
#endif
#ifndef XCAL_XRCHOLV18_B60_N128
#define XCAL_XRCHOLV18_B60_N128 0
#endif
#ifndef XCAL_XRTS_TMA_F16_B60
#define XCAL_XRTS_TMA_F16_B60 0
#endif
#if XCAL_XRCHOLV18_B60_ISSUES == 4
#define XCAL_XRCHOLV18_B60_RESIDENT 3
#else
#define XCAL_XRCHOLV18_B60_RESIDENT 2
#endif
#ifndef XCAL_XRCHOLV18_NVFP4_OUTER
#define XCAL_XRCHOLV18_NVFP4_OUTER 0
#endif
#ifndef XCAL_XRCHOLV18_LEFT_MACRO
#define XCAL_XRCHOLV18_LEFT_MACRO 2048
#endif
#ifndef XCAL_XRCHOLV18_K256_FRONT
#define XCAL_XRCHOLV18_K256_FRONT 1
#endif
#ifndef XCAL_XRCHOLV18_HEXLIFT_OUTER
#define XCAL_XRCHOLV18_HEXLIFT_OUTER 1
#endif
#ifndef XCAL_XRCHOLV18_HEXLIFT_MIN_CTAS
#define XCAL_XRCHOLV18_HEXLIFT_MIN_CTAS 96
#endif
#ifndef XCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY
#define XCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY 32
#endif
#ifndef XCAL_XRCHOLV18_HEX6_ALL
#define XCAL_XRCHOLV18_HEX6_ALL 0
#endif
#ifndef XCAL_XRCHOLV18_HEX4_GE8192
#define XCAL_XRCHOLV18_HEX4_GE8192 1
#endif
static_assert(
XCAL_XRCHOLV18_MAX_G == 1
|| XCAL_XRCHOLV18_MAX_G == 2
|| XCAL_XRCHOLV18_MAX_G == 4
|| XCAL_XRCHOLV18_MAX_G == 8
|| XCAL_XRCHOLV18_MAX_G == 16);
static_assert(
XCAL_XRCHOLV18_PAYLOAD_PIPELINE == 0
|| XCAL_XRCHOLV18_PAYLOAD_PIPELINE == 1);
static_assert(
XCAL_XRCHOLV18_B60_ISSUES == 4
|| XCAL_XRCHOLV18_B60_ISSUES == 8);
static_assert(
XCAL_XRCHOLV18_B60_N128 == 0
|| XCAL_XRCHOLV18_B60_N128 == 1);
static_assert(
XCAL_XRTS_TMA_F16_B60 == 0
|| XCAL_XRTS_TMA_F16_B60 == 1);
static_assert(
XCAL_XRCHOLV18_NVFP4_OUTER == 0
|| XCAL_XRCHOLV18_NVFP4_OUTER == 1);
static_assert(
XCAL_XRCHOLV18_LEFT_MACRO == 256
|| XCAL_XRCHOLV18_LEFT_MACRO == 512
|| XCAL_XRCHOLV18_LEFT_MACRO == 1024
|| XCAL_XRCHOLV18_LEFT_MACRO == 2048);
static_assert(
XCAL_XRCHOLV18_K256_FRONT == 0
|| XCAL_XRCHOLV18_K256_FRONT == 1);
static_assert(
XCAL_XRCHOLV18_HEXLIFT_OUTER == 0
|| XCAL_XRCHOLV18_HEXLIFT_OUTER == 1);
static_assert(
XCAL_XRCHOLV18_HEXLIFT_MIN_CTAS >= 1
&& XCAL_XRCHOLV18_HEXLIFT_MIN_CTAS <= 148);
static_assert(
XCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY >= 0
&& XCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY <= 512);
static_assert(
XCAL_XRCHOLV18_HEX6_ALL == 0
|| XCAL_XRCHOLV18_HEX6_ALL == 1);
static_assert(
XCAL_XRCHOLV18_HEX4_GE8192 == 0
|| XCAL_XRCHOLV18_HEX4_GE8192 == 1);
// A[b,n,n] -> L[b,n,n], lower-only FP32.
//
// V15 exact mid-shape board: N256 left-looking strided-batched GEMM.
// Native K128 POTRF/MMA-TRSM and sibling edge remain unchanged.
//
// V14 clean pipeline board:
// - 32-byte power-of-two lower preamble, no per-packet div/mod
// - pair-local 64-thread MMA-TRSM handoffs
// - next-task TMEM seed overlap in the generic trailing owner
// - B60 split packet formation overlapped with TCGEN issue
//
// V18 board:
// n <= 128: exact scalar POTRF; there is no matrix-product instruction
// n > 128: N128 dependency fronts, N256 update panels
// n512, B >= resident: persistent K64 factor fronts, K64 L10, TCGEN Gram
// n1024, B >= 16: K128 fronts, incremental lower N256 outer state
// n < 2048 && B < 16 sibling: FP16 hi/lo x3, true M128xN128
// non-E2M1 fast sibling: rounded FP16, true M128xN128
// large K128 sibling: measured s2f6 x4 dependency edge
// n>=4096 outer: sealed P2048 history -> scaled lower-N residual owners
// n < 2048 && B < 16 history: two K128 FP16 hi/lo x3 groups
// other fat history: FP16 K256 -> blocked-lower right-looking update
// n > 8192 outer: one-rail NVFP4 K64 -> one M128xN256 issue
// D payload: FP32 TMEM, imported/exported without a warp-MMA route
// arbitrary n: internal ceil(n/128) block-diagonal padding, exact crop
// large macro: P2048; shape-scaled N board, K128 packet cadence
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200
// Native no-swizzle K-major TF32 packet for M/N64 x K8.
// One packet is 512 b32 words (2 KiB):
// index(mn,k) = 4*mn + (k&3) + 256*(k>>2).
__device__ __forceinline__ uint32_t loco_bf16x2(float x0, float x1) {
uint32_t packed;
asm volatile(
"cvt.rn.bf16x2.f32 %0, %2, %1;"
: "=r"(packed)
: "f"(x0), "f"(x1));
return packed;
}
__device__ __forceinline__ void loco_bf16x2_hi_lo(
float x0,
float x1,
uint32_t& high,
uint32_t& low
) {
high = loco_bf16x2(x0, x1);
float h0 = __uint_as_float(high << 16);
float h1 = __uint_as_float(high & 0xffff0000u);
low = loco_bf16x2(x0 - h0, x1 - h1);
}
__device__ __forceinline__ uint32_t loco_f16x2(float x0, float x1) {
uint32_t packed;
asm volatile(
"cvt.rn.f16x2.f32 %0, %2, %1;"
: "=r"(packed)
: "f"(x0), "f"(x1));
return packed;
}
__device__ __forceinline__ float loco_f16_value(uint16_t x) {
float value;
asm volatile(
"cvt.f32.f16 %0, %1;"
: "=f"(value)
: "h"(x));
return value;
}
__device__ __forceinline__ void loco_f16x2_hi_lo(
float x0,
float x1,
uint32_t& high,
uint32_t& low
) {
high = loco_f16x2(x0, x1);
float h0 = loco_f16_value((uint16_t)high);
float h1 = loco_f16_value((uint16_t)(high >> 16));
low = loco_f16x2(x0 - h0, x1 - h1);
}
__device__ __forceinline__ uint64_t loco_fp16_packed_word(
int b,
int n,
int row,
int col,
int rail
) {
// M64 x K16 tiles, two K8 banks per tile. Each K8 bank follows
// the no-swizzle K-major atom consumed by TCGEN:
//
// [8 high rows][8 residual-low rows]
//
// Adjacent rows stay 16 B apart. Rail bases are 128 B apart, and
// consecutive 8-row atoms are 256 B apart.
uint64_t row_tiles = (uint64_t)n >> 6;
uint64_t k_tiles = (uint64_t)n >> 4;
return ((((uint64_t)b * row_tiles + (uint64_t)(row >> 6))
* k_tiles
+ (uint64_t)(col >> 4))
<< 10)
+ (uint64_t)(((col >> 3) & 1) << 9)
+ (uint64_t)(((row & 63) >> 3) << 6)
+ (uint64_t)(rail << 5)
+ (uint64_t)((row & 7) << 2)
+ (uint64_t)((col & 7) >> 1);
}
__device__ __forceinline__ float4 loco_ld_global_v4(
const float* src
) {
uint32_t x0;
uint32_t x1;
uint32_t x2;
uint32_t x3;
asm volatile(
"ld.global.ca.nc.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(x0), "=r"(x1), "=r"(x2), "=r"(x3)
: "l"(src)
: "memory");
return make_float4(
__uint_as_float(x0),
__uint_as_float(x1),
__uint_as_float(x2),
__uint_as_float(x3));
}
__device__ __forceinline__ float4 loco_ld_global_v4_cs(
const float* src
) {
uint32_t x0;
uint32_t x1;
uint32_t x2;
uint32_t x3;
asm volatile(
"ld.global.cs.nc.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(x0), "=r"(x1), "=r"(x2), "=r"(x3)
: "l"(src)
: "memory");
return make_float4(
__uint_as_float(x0),
__uint_as_float(x1),
__uint_as_float(x2),
__uint_as_float(x3));
}
__device__ __forceinline__ uint4 loco_ld_global_u4(
const uint8_t* src
) {
uint4 value;
asm volatile(
"ld.global.cs.nc.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(value.x), "=r"(value.y),
"=r"(value.z), "=r"(value.w)
: "l"(src)
: "memory");
return value;
}
__device__ __forceinline__ void loco_ld_global_u8(
const uint32_t* src,
uint4& first,
uint4& second
) {
uint32_t r0;
uint32_t r1;
uint32_t r2;
uint32_t r3;
uint32_t r4;
uint32_t r5;
uint32_t r6;
uint32_t r7;
asm volatile(
"ld.global.nc.v8.b32.L2::256B.L1::no_allocate "
"{%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3),
"=r"(r4), "=r"(r5), "=r"(r6), "=r"(r7)
: "l"(src)
: "memory");
first = make_uint4(r0, r1, r2, r3);
second = make_uint4(r4, r5, r6, r7);
}
__device__ __forceinline__ uint32_t loco_ld_global_u32_ca(
const uint32_t* src
) {
uint32_t value;
asm volatile(
"ld.global.ca.u32 %0, [%1];"
: "=r"(value)
: "l"(src)
: "memory");
return value;
}
__device__ __forceinline__ void loco_ld_global_u8_ca(
const uint32_t* src,
uint4& first,
uint4& second
) {
uint32_t r0;
uint32_t r1;
uint32_t r2;
uint32_t r3;
uint32_t r4;
uint32_t r5;
uint32_t r6;
uint32_t r7;
asm volatile(
"ld.global.nc.v8.b32.L2::256B.L1::no_allocate "
"{%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3),
"=r"(r4), "=r"(r5), "=r"(r6), "=r"(r7)
: "l"(src)
: "memory");
first = make_uint4(r0, r1, r2, r3);
second = make_uint4(r4, r5, r6, r7);
}
__device__ __forceinline__ void loco_st_shared_u8(
uint32_t dst,
const uint4& first,
const uint4& second
) {
// PTX v8 stores are global-only. Keep the source transaction intact and
// land its two contiguous halves without the old descriptor-bank scatter.
asm volatile(
"st.shared.v4.b32 "
"[%0], {%1,%2,%3,%4};\n\t"
"st.shared.v4.b32 "
"[%0+16], {%5,%6,%7,%8};"
:
: "r"(dst),
"r"(first.x), "r"(first.y),
"r"(first.z), "r"(first.w),
"r"(second.x), "r"(second.y),
"r"(second.z), "r"(second.w)
: "memory");
}
__device__ __forceinline__ void loco_st_shared_u4(
uint32_t dst,
const uint4& value
) {
asm volatile(
"st.shared.v4.b32 [%0], {%1,%2,%3,%4};"
:
: "r"(dst),
"r"(value.x), "r"(value.y),
"r"(value.z), "r"(value.w)
: "memory");
}
__device__ __forceinline__ void loco_mma_bf16(
uint32_t a0,
uint32_t a1,
uint32_t a2,
uint32_t a3,
uint32_t b0,
uint32_t b1,
float& d0,
float& d1,
float& d2,
float& d3
) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3},{%4,%5,%6,%7},{%8,%9},{%0,%1,%2,%3};"
: "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
"r"(b0), "r"(b1));
}
__device__ __forceinline__ void loco_mma_f16(
uint32_t a0,
uint32_t a1,
uint32_t a2,
uint32_t a3,
uint32_t b0,
uint32_t b1,
float& d0,
float& d1,
float& d2,
float& d3
) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
"{%0,%1,%2,%3},{%4,%5,%6,%7},{%8,%9},{%0,%1,%2,%3};"
: "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
"r"(b0), "r"(b1));
}
__device__ __forceinline__ void loco_mma_bf16x3(
float2 at0,
float2 ab0,
float2 at8,
float2 ab8,
float2 bv0,
float2 bv8,
float& d0,
float& d1,
float& d2,
float& d3
) {
uint32_t ah0, ah1, ah2, ah3;
uint32_t al0, al1, al2, al3;
uint32_t bh0, bh1, bl0, bl1;
loco_bf16x2_hi_lo(at0.x, at0.y, ah0, al0);
loco_bf16x2_hi_lo(ab0.x, ab0.y, ah1, al1);
loco_bf16x2_hi_lo(at8.x, at8.y, ah2, al2);
loco_bf16x2_hi_lo(ab8.x, ab8.y, ah3, al3);
loco_bf16x2_hi_lo(bv0.x, bv0.y, bh0, bl0);
loco_bf16x2_hi_lo(bv8.x, bv8.y, bh1, bl1);
bh0 ^= 0x80008000u;
bh1 ^= 0x80008000u;
bl0 ^= 0x80008000u;
bl1 ^= 0x80008000u;
loco_mma_bf16(
ah0, ah1, ah2, ah3, bh0, bh1,
d0, d1, d2, d3);
loco_mma_bf16(
ah0, ah1, ah2, ah3, bl0, bl1,
d0, d1, d2, d3);
loco_mma_bf16(
al0, al1, al2, al3, bh0, bh1,
d0, d1, d2, d3);
}
#endif
__device__ __forceinline__ float xcal_sqrt_approx(float x) {
float y;
asm volatile("sqrt.approx.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
__device__ __forceinline__ float xcal_rcp_approx(float x) {
float y;
asm volatile("rcp.approx.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
__device__ __forceinline__ float xcal_rsqrt_approx(float x) {
float y;
asm volatile("rsqrt.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
constexpr int LOCO_FRONTIER128_SMEM =
(128 * 129 + 128) * (int)sizeof(float);
constexpr int LOCO_FACTOR256_SMEM =
(256 * 257 / 2 + 256 + 16 * 256) * (int)sizeof(float);
constexpr int LOCO_FACTOR256_K16_SMEM =
(256 * 257 / 2 + 256) * (int)sizeof(float);
constexpr int LOCO_TRSM256_SMEM = LOCO_FACTOR256_SMEM;
constexpr int LOCO_INNER_HEXLIFT_PLANES = 4;
constexpr int LOCO_HEXLIFT_OUTER_FIRST_PLANE = 1;
constexpr int LOCO_HEXLIFT_OUTER_ACTIVE_PLANES =
LOCO_INNER_HEXLIFT_PLANES - LOCO_HEXLIFT_OUTER_FIRST_PLANE;
static_assert(
LOCO_HEXLIFT_OUTER_ACTIVE_PLANES == 3,
"outer HexLift plane count changed");
constexpr int LOCO_INNER_HEXLIFT_STAGE_ROWS = 384;
constexpr int LOCO_INNER_HEXLIFT_SCALE_BYTES =
LOCO_INNER_HEXLIFT_STAGE_ROWS * (int)sizeof(uint32_t);
constexpr int LOCO_INNER_HEXLIFT_MAX_SHADOW_BYTES =
LOCO_INNER_HEXLIFT_PLANES
* 2 * LOCO_INNER_HEXLIFT_STAGE_ROWS * 16
+ LOCO_INNER_HEXLIFT_SCALE_BYTES;
static_assert(
LOCO_INNER_HEXLIFT_MAX_SHADOW_BYTES == 50688,
"K256 diagonal-B shadow layout changed");
struct __align__(128) LocoHexLiftInnerShared {
// Throughput K256 board only. Persistent N256/N512 retain their measured
// baseline shared-memory layouts and BF16 update paths.
float tile[256 * 257 / 2];
float inverse[256];
uint8_t packet_a[LOCO_INNER_HEXLIFT_PLANES][4096];
uint8_t packet_b[LOCO_INNER_HEXLIFT_PLANES][8192];
uint8_t scale_a[LOCO_INNER_HEXLIFT_PLANES][512];
uint8_t scale_b[LOCO_INNER_HEXLIFT_PLANES][1024];
float scratch[8192];
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
LOCO_FACTOR256_SMEM == 148992,
"persistent N256 shared-memory request changed");
static_assert(
sizeof(LocoHexLiftInnerShared) == 220800,
"K256 HexLift shared-memory board changed");
constexpr int LOCO_HEXLIFT256_SMEM =
(int)sizeof(LocoHexLiftInnerShared);
__device__ __forceinline__ void loco_mbar_init(uint32_t addr);
__device__ __forceinline__ uint32_t loco_elect_one();
__device__ __forceinline__ void loco_hexlift_factor_update(
LocoHexLiftInnerShared& shared,
uint8_t* __restrict__ diagonal_shadow,
int panel,
int start,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase);
__device__ __forceinline__ void loco_hexlift_trsm_update(
LocoHexLiftInnerShared& shared,
const uint8_t* __restrict__ diagonal_shadow,
float* __restrict__ rows,
int row_stride,
int row_count,
int row0,
int panel,
int start,
bool load_b,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase);
__device__ __forceinline__ int loco_lower256(int row, int col) {
return (row * (row + 1) >> 1) + col;
}
__device__ __forceinline__ void loco_potrf64_packed(
float* tile,
int panel
) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int end = panel + 64;
// One exact FP32 K64 front. The tail is not burned by scalar FFMA here;
// loco_factor256_bf16 publishes it with tensor MMA after TRSM.
#pragma unroll 1
for (int p = panel; p < end; ++p) {
if (warp == 0) {
float reciprocal = 0.0f;
if (lane == 0) {
int diagonal = loco_lower256(p, p);
float value = xcal_sqrt_approx(
fmaxf(tile[diagonal], 1.0e-20f));
tile[diagonal] = value;
reciprocal = xcal_rcp_approx(value);
}
reciprocal = __shfl_sync(0xffffffffu, reciprocal, 0);
for (int row = p + 1 + lane; row < end; row += 32) {
int row_base = row * (row + 1) >> 1;
tile[row_base + p] *= reciprocal;
}
}
__syncthreads();
for (int row = p + 1 + warp; row < end; row += 32) {
int row_base = row * (row + 1) >> 1;
float lip = tile[row_base + p];
for (int col = p + 1 + lane; col <= row; col += 32) {
int dst = row_base + col;
tile[dst] = fmaf(
-lip,
tile[(col * (col + 1) >> 1) + p],
tile[dst]);
}
}
__syncthreads();
}
}
__device__ __forceinline__ void loco_trsm64_packed_warp(
float* tile,
const float* inverse,
int row,
int panel
) {
int lane = threadIdx.x & 31;
int row_base = row * (row + 1) >> 1;
float x0 = tile[row_base + panel + lane];
float x1 = tile[row_base + panel + lane + 32];
#pragma unroll 1
for (int c = 0; c < 64; ++c) {
int diagonal_row = panel + c;
int diagonal_base = diagonal_row * (diagonal_row + 1) >> 1;
float partial = lane < c
? x0 * tile[diagonal_base + panel + lane]
: 0.0f;
if (lane + 32 < c) {
partial = fmaf(
x1,
tile[diagonal_base + panel + lane + 32],
partial);
}
partial += __shfl_down_sync(0xffffffffu, partial, 16);
partial += __shfl_down_sync(0xffffffffu, partial, 8);
partial += __shfl_down_sync(0xffffffffu, partial, 4);
partial += __shfl_down_sync(0xffffffffu, partial, 2);
partial += __shfl_down_sync(0xffffffffu, partial, 1);
float sum = __shfl_sync(0xffffffffu, partial, 0);
int owner = c & 31;
float rhs = c < 32
? __shfl_sync(0xffffffffu, x0, owner)
: __shfl_sync(0xffffffffu, x1, owner);
float value = (rhs - sum) * inverse[diagonal_row];
if (lane == owner) {
if (c < 32) x0 = value;
else x1 = value;
}
}
tile[row_base + panel + lane] = x0;
tile[row_base + panel + lane + 32] = x1;
}
__device__ __forceinline__ void loco_factor256_bf16(
float* tile,
float* inverse,
float* a_smem
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int group = lane >> 2;
int rank = lane & 3;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 64) {
loco_potrf64_packed(tile, panel);
if (tid < 64) {
int diagonal = panel + tid;
inverse[diagonal] = xcal_rcp_approx(
tile[loco_lower256(diagonal, diagonal)]);
}
__syncthreads();
int start = panel + 64;
#pragma unroll
for (int wave = 0; wave < 6; ++wave) {
int row = start + (wave << 5) + warp;
if (row < 256) {
loco_trsm64_packed_warp(
tile,
inverse,
row,
panel);
}
}
__syncthreads();
// K64 update cone. A[16,64] is one 4KB coalesced load cycle;
// warp w owns C[16,8], four BF16 MMA issues consume the whole K64.
#pragma unroll 1
for (int row0 = start; row0 < 256; row0 += 16) {
int source_row = tid >> 6;
int source_col = tid & 63;
int source_base =
(row0 + source_row) * (row0 + source_row + 1) >> 1;
a_smem[source_row * 64 + source_col] =
tile[source_base + panel + source_col];
__syncthreads();
int col0 = start + (warp << 3);
if (col0 < 256 && row0 + 15 >= col0) {
int top = row0 + group;
int bottom = top + 8;
int col = col0 + (rank << 1);
int b_row = col0 + group;
float d0 = top >= col
? tile[loco_lower256(top, col)]
: 0.0f;
float d1 = top >= col + 1
? tile[loco_lower256(top, col + 1)]
: 0.0f;
float d2 = bottom < 256 && bottom >= col
? tile[loco_lower256(bottom, col)]
: 0.0f;
float d3 = bottom < 256 && bottom >= col + 1
? tile[loco_lower256(bottom, col + 1)]
: 0.0f;
#pragma unroll
for (int kk = 0; kk < 64; kk += 16) {
int k0 = kk + (rank << 1);
int k8 = k0 + 8;
float2 at0 = *reinterpret_cast<const float2*>(
a_smem + group * 64 + k0);
float2 ab0 = *reinterpret_cast<const float2*>(
a_smem + (group + 8) * 64 + k0);
float2 at8 = *reinterpret_cast<const float2*>(
a_smem + group * 64 + k8);
float2 ab8 = *reinterpret_cast<const float2*>(
a_smem + (group + 8) * 64 + k8);
float2 bv0 = {0.0f, 0.0f};
float2 bv8 = {0.0f, 0.0f};
if (b_row < 256) {
int b_base = b_row * (b_row + 1) >> 1;
bv0 = make_float2(
tile[b_base + panel + k0],
tile[b_base + panel + k0 + 1]);
bv8 = make_float2(
tile[b_base + panel + k8],
tile[b_base + panel + k8 + 1]);
}
loco_mma_bf16(
loco_bf16x2(at0.x, at0.y),
loco_bf16x2(ab0.x, ab0.y),
loco_bf16x2(at8.x, at8.y),
loco_bf16x2(ab8.x, ab8.y),
loco_bf16x2(bv0.x, bv0.y) ^ 0x80008000u,
loco_bf16x2(bv8.x, bv8.y) ^ 0x80008000u,
d0, d1, d2, d3);
}
if (top >= col) {
tile[loco_lower256(top, col)] = d0;
}
if (top >= col + 1) {
tile[loco_lower256(top, col + 1)] = d1;
}
if (bottom < 256 && bottom >= col) {
tile[loco_lower256(bottom, col)] = d2;
}
if (bottom < 256 && bottom >= col + 1) {
tile[loco_lower256(bottom, col + 1)] = d3;
}
}
__syncthreads();
}
}
}
__device__ __forceinline__ void loco_factor256_hexlift(
LocoHexLiftInnerShared& shared,
uint8_t* __restrict__ diagonal_shadow,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase
) {
float* tile = shared.tile;
float* inverse = shared.inverse;
int tid = threadIdx.x;
int warp = tid >> 5;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 64) {
loco_potrf64_packed(tile, panel);
if (tid < 64) {
int diagonal = panel + tid;
inverse[diagonal] = xcal_rcp_approx(
tile[loco_lower256(diagonal, diagonal)]);
}
__syncthreads();
int start = panel + 64;
#pragma unroll
for (int wave = 0; wave < 6; ++wave) {
int row = start + (wave << 5) + warp;
if (row < 256) {
loco_trsm64_packed_warp(
tile,
inverse,
row,
panel);
}
}
__syncthreads();
// Preserve the exact FP32 K64 front, then update only its live
// residual cone through the route-adaptive HexLift term group.
if (start < 256) {
loco_hexlift_factor_update(
shared,
diagonal_shadow,
panel,
start,
hex10,
issuer,
mma_phase);
}
}
}
__device__ __forceinline__ void loco_factor256_bf16x3(
float* tile,
float* inverse,
float* a_smem
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int group = lane >> 2;
int rank = lane & 3;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 64) {
loco_potrf64_packed(tile, panel);
if (tid < 64) {
int diagonal = panel + tid;
inverse[diagonal] = xcal_rcp_approx(
tile[loco_lower256(diagonal, diagonal)]);
}
__syncthreads();
int start = panel + 64;
#pragma unroll
for (int wave = 0; wave < 6; ++wave) {
int row = start + (wave << 5) + warp;
if (row < 256) {
loco_trsm64_packed_warp(
tile,
inverse,
row,
panel);
}
}
__syncthreads();
// K64 update cone with BF16 high/residual-low compensation. This is
// the direct N256 path's precision rail; the scalar fronts remain FP32.
#pragma unroll 1
for (int row0 = start; row0 < 256; row0 += 16) {
int source_row = tid >> 6;
int source_col = tid & 63;
int source_base =
(row0 + source_row) * (row0 + source_row + 1) >> 1;
a_smem[source_row * 64 + source_col] =
tile[source_base + panel + source_col];
__syncthreads();
int col0 = start + (warp << 3);
if (col0 < 256 && row0 + 15 >= col0) {
int top = row0 + group;
int bottom = top + 8;
int col = col0 + (rank << 1);
int b_row = col0 + group;
float d0 = top >= col
? tile[loco_lower256(top, col)]
: 0.0f;
float d1 = top >= col + 1
? tile[loco_lower256(top, col + 1)]
: 0.0f;
float d2 = bottom < 256 && bottom >= col
? tile[loco_lower256(bottom, col)]
: 0.0f;
float d3 = bottom < 256 && bottom >= col + 1
? tile[loco_lower256(bottom, col + 1)]
: 0.0f;
#pragma unroll
for (int kk = 0; kk < 64; kk += 16) {
int k0 = kk + (rank << 1);
int k8 = k0 + 8;
float2 at0 = *reinterpret_cast<const float2*>(
a_smem + group * 64 + k0);
float2 ab0 = *reinterpret_cast<const float2*>(
a_smem + (group + 8) * 64 + k0);
float2 at8 = *reinterpret_cast<const float2*>(
a_smem + group * 64 + k8);
float2 ab8 = *reinterpret_cast<const float2*>(
a_smem + (group + 8) * 64 + k8);
float2 bv0 = {0.0f, 0.0f};
float2 bv8 = {0.0f, 0.0f};
if (b_row < 256) {
int b_base = b_row * (b_row + 1) >> 1;
bv0 = make_float2(
tile[b_base + panel + k0],
tile[b_base + panel + k0 + 1]);
bv8 = make_float2(
tile[b_base + panel + k8],
tile[b_base + panel + k8 + 1]);
}
loco_mma_bf16x3(
at0, ab0, at8, ab8, bv0, bv8,
d0, d1, d2, d3);
}
if (top >= col) {
tile[loco_lower256(top, col)] = d0;
}
if (top >= col + 1) {
tile[loco_lower256(top, col + 1)] = d1;
}
if (bottom < 256 && bottom >= col) {
tile[loco_lower256(bottom, col)] = d2;
}
if (bottom < 256 && bottom >= col + 1) {
tile[loco_lower256(bottom, col + 1)] = d3;
}
}
__syncthreads();
}
}
}
__device__ __forceinline__ void loco_factor256_k16_bf16(
float* tile,
float* inverse
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int group = lane >> 2;
int rank = lane & 3;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 16) {
int end = panel + 16;
// Warp zero factors the exact FP32 K16 diagonal block. One lane owns
// each row below the pivot, so the two warp edges publish the scaled
// column and its rank-one update without a CTA-wide pivot edge.
if (warp == 0) {
#pragma unroll 1
for (int p = panel; p < end; ++p) {
float reciprocal = 0.0f;
if (lane == 0) {
int diagonal = loco_lower256(p, p);
float value = xcal_sqrt_approx(
fmaxf(tile[diagonal], 1.0e-20f));
tile[diagonal] = value;
reciprocal = xcal_rcp_approx(value);
inverse[p] = reciprocal;
}
reciprocal = __shfl_sync(
0xffffffffu, reciprocal, 0);
int row = p + 1 + lane;
if (row < end) {
int row_base = row * (row + 1) >> 1;
tile[row_base + p] *= reciprocal;
}
__syncwarp(0xffffffffu);
if (row < end) {
int row_base = row * (row + 1) >> 1;
float lip = tile[row_base + p];
for (int col = p + 1; col <= row; ++col) {
int dst = row_base + col;
tile[dst] = fmaf(
-lip,
tile[loco_lower256(col, p)],
tile[dst]);
}
}
__syncwarp(0xffffffffu);
}
}
__syncthreads();
int start = end;
if (start < 256) {
// One warp owns one complete row solve. The first sixteen lanes
// retain the mutable K16 RHS while every pivot reduction stays
// exact FP32.
for (int row = start + warp; row < 256; row += 32) {
int row_base = row * (row + 1) >> 1;
float x = lane < 16
? tile[row_base + panel + lane]
: 0.0f;
#pragma unroll 1
for (int c = 0; c < 16; ++c) {
int diagonal_row = panel + c;
int diagonal_base =
diagonal_row * (diagonal_row + 1) >> 1;
float partial = lane < c
? x * tile[diagonal_base + panel + lane]
: 0.0f;
partial += __shfl_down_sync(
0xffffffffu, partial, 16);
partial += __shfl_down_sync(
0xffffffffu, partial, 8);
partial += __shfl_down_sync(
0xffffffffu, partial, 4);
partial += __shfl_down_sync(
0xffffffffu, partial, 2);
partial += __shfl_down_sync(
0xffffffffu, partial, 1);
float sum = __shfl_sync(
0xffffffffu, partial, 0);
float rhs = __shfl_sync(
0xffffffffu, x, c);
float value =
(rhs - sum) * inverse[diagonal_row];
if (lane == c) x = value;
}
if (lane < 16) {
tile[row_base + panel + lane] = x;
}
}
__syncthreads();
// Each warp owns one N8 stripe for every M16 row block. Packed
// triangular row bases can be odd, so all operand pairs are built
// from scalar loads before the single BF16 K16 issue.
int col0 = start + (warp << 3);
#pragma unroll 1
for (int row0 = start; row0 < 256; row0 += 16) {
if (col0 < 256 && row0 + 15 >= col0) {
int top = row0 + group;
int bottom = top + 8;
int col = col0 + (rank << 1);
int b_row = col0 + group;
int top_base = top * (top + 1) >> 1;
int bottom_base = bottom * (bottom + 1) >> 1;
int b_base = b_row * (b_row + 1) >> 1;
float d0 = top >= col
? tile[top_base + col]
: 0.0f;
float d1 = top >= col + 1
? tile[top_base + col + 1]
: 0.0f;
float d2 = bottom >= col
? tile[bottom_base + col]
: 0.0f;
float d3 = bottom >= col + 1
? tile[bottom_base + col + 1]
: 0.0f;
int k0 = panel + (rank << 1);
int k8 = k0 + 8;
float2 at0 = make_float2(
tile[top_base + k0],
tile[top_base + k0 + 1]);
float2 ab0 = make_float2(
tile[bottom_base + k0],
tile[bottom_base + k0 + 1]);
float2 at8 = make_float2(
tile[top_base + k8],
tile[top_base + k8 + 1]);
float2 ab8 = make_float2(
tile[bottom_base + k8],
tile[bottom_base + k8 + 1]);
float2 bv0 = make_float2(
tile[b_base + k0],
tile[b_base + k0 + 1]);
float2 bv8 = make_float2(
tile[b_base + k8],
tile[b_base + k8 + 1]);
loco_mma_bf16(
loco_bf16x2(at0.x, at0.y),
loco_bf16x2(ab0.x, ab0.y),
loco_bf16x2(at8.x, at8.y),
loco_bf16x2(ab8.x, ab8.y),
loco_bf16x2(bv0.x, bv0.y) ^ 0x80008000u,
loco_bf16x2(bv8.x, bv8.y) ^ 0x80008000u,
d0, d1, d2, d3);
if (top >= col) {
tile[top_base + col] = d0;
}
if (top >= col + 1) {
tile[top_base + col + 1] = d1;
}
if (bottom >= col) {
tile[bottom_base + col] = d2;
}
if (bottom >= col + 1) {
tile[bottom_base + col + 1] = d3;
}
}
}
__syncthreads();
}
}
}
template <int N, int WARPS, bool F16x1 = false>
__device__ __forceinline__ void loco_factor_k16(
float* tile,
float* inverse,
uint32_t* f16_packet = nullptr
) {
static_assert((N & 15) == 0, "K16 factor requires N16 alignment");
static_assert(WARPS * 8 == N, "one warp must own each N8 stripe");
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int group = lane >> 2;
int rank = lane & 3;
#pragma unroll 1
for (int panel = 0; panel < N; panel += 16) {
int end = panel + 16;
// The scalar edge is warp-local. Only the K16 publication and the
// completed update cone need a CTA edge, replacing the old two
// CTA-wide barriers at every scalar pivot.
if (warp == 0) {
#pragma unroll 1
for (int p = panel; p < end; ++p) {
float reciprocal = 0.0f;
if (lane == 0) {
int diagonal = loco_lower256(p, p);
float pivot =
fmaxf(tile[diagonal], 1.0e-20f);
reciprocal = xcal_rsqrt_approx(pivot);
float value = pivot * reciprocal;
tile[diagonal] = value;
inverse[p] = reciprocal;
}
reciprocal = __shfl_sync(
0xffffffffu, reciprocal, 0);
int row = p + 1 + lane;
if (row < end) {
int row_base = row * (row + 1) >> 1;
tile[row_base + p] *= reciprocal;
}
__syncwarp(0xffffffffu);
if (row < end) {
int row_base = row * (row + 1) >> 1;
float lip = tile[row_base + p];
#pragma unroll
for (int col = p + 1; col <= row; ++col) {
int dst = row_base + col;
tile[dst] = fmaf(
-lip,
tile[loco_lower256(col, p)],
tile[dst]);
}
}
__syncwarp(0xffffffffu);
}
}
__syncthreads();
int first = end;
if (first < N) {
// One warp owns one complete K16 RHS row at a time. The mutable
// row stays in RMEM until all sixteen pivots have been consumed.
for (int row = first + warp; row < N; row += WARPS) {
int row_base = row * (row + 1) >> 1;
float x = lane < 16
? tile[row_base + panel + lane]
: 0.0f;
#pragma unroll
for (int c = 0; c < 16; ++c) {
int diagonal_row = panel + c;
int diagonal_base =
diagonal_row * (diagonal_row + 1) >> 1;
float partial = lane < c
? x * tile[diagonal_base + panel + lane]
: 0.0f;
partial += __shfl_down_sync(
0xffffffffu, partial, 16);
partial += __shfl_down_sync(
0xffffffffu, partial, 8);
partial += __shfl_down_sync(
0xffffffffu, partial, 4);
partial += __shfl_down_sync(
0xffffffffu, partial, 2);
partial += __shfl_down_sync(
0xffffffffu, partial, 1);
float sum = __shfl_sync(
0xffffffffu, partial, 0);
float rhs = __shfl_sync(
0xffffffffu, x, c);
float value =
(rhs - sum) * inverse[diagonal_row];
if (lane == c) x = value;
}
if constexpr (F16x1) {
float x1 = __shfl_down_sync(
0xffffffffu, x, 1);
if (lane < 16 && !(lane & 1)) {
// Two K8 planes make every m16n8 operand load
// bank-perfect: bank = 4*row_group + rank.
f16_packet[
((lane & 8) << 7)
+ (row << 2)
+ ((lane >> 1) & 3)
] = loco_f16x2(x, x1);
}
}
if (lane < 16) {
tile[row_base + panel + lane] = x;
}
}
__syncthreads();
// Exact B640 takes one rounded-FP16 issue. Other ranked owners
// retain BF16 high/residual-low compensation.
int col0 = first + (warp << 3);
#pragma unroll 1
for (int row0 = first; row0 < N; row0 += 16) {
if (col0 < N && row0 + 15 >= col0) {
int top = row0 + group;
int bottom = top + 8;
int col = col0 + (rank << 1);
int b_row = col0 + group;
int top_base = top * (top + 1) >> 1;
int bottom_base = bottom * (bottom + 1) >> 1;
int b_base = b_row * (b_row + 1) >> 1;
float d0 = top >= col
? tile[top_base + col]
: 0.0f;
float d1 = top >= col + 1
? tile[top_base + col + 1]
: 0.0f;
float d2 = bottom < N && bottom >= col
? tile[bottom_base + col]
: 0.0f;
float d3 = bottom < N && bottom >= col + 1
? tile[bottom_base + col + 1]
: 0.0f;
if constexpr (F16x1) {
loco_mma_f16(
f16_packet[(top << 2) + rank],
bottom < N
? f16_packet[(bottom << 2) + rank]
: 0u,
f16_packet[1024 + (top << 2) + rank],
bottom < N
? f16_packet[
1024 + (bottom << 2) + rank]
: 0u,
f16_packet[(b_row << 2) + rank]
^ 0x80008000u,
f16_packet[
1024 + (b_row << 2) + rank]
^ 0x80008000u,
d0, d1, d2, d3);
} else {
int k0 = panel + (rank << 1);
int k8 = k0 + 8;
float2 at0 = make_float2(
tile[top_base + k0],
tile[top_base + k0 + 1]);
float2 at8 = make_float2(
tile[top_base + k8],
tile[top_base + k8 + 1]);
float2 ab0 = make_float2(0.0f, 0.0f);
float2 ab8 = make_float2(0.0f, 0.0f);
if (bottom < N) {
ab0 = make_float2(
tile[bottom_base + k0],
tile[bottom_base + k0 + 1]);
ab8 = make_float2(
tile[bottom_base + k8],
tile[bottom_base + k8 + 1]);
}
float2 bv0 = make_float2(
tile[b_base + k0],
tile[b_base + k0 + 1]);
float2 bv8 = make_float2(
tile[b_base + k8],
tile[b_base + k8 + 1]);
loco_mma_bf16x3(
at0, ab0, at8, ab8, bv0, bv8,
d0, d1, d2, d3);
}
if (top >= col) {
tile[top_base + col] = d0;
}
if (top >= col + 1) {
tile[top_base + col + 1] = d1;
}
if (bottom < N && bottom >= col) {
tile[bottom_base + col] = d2;
}
if (bottom < N && bottom >= col + 1) {
tile[bottom_base + col + 1] = d3;
}
}
}
// Every row0 owns a disjoint N16 stripe. Publish the complete
// rank-K16 cone once before the next scalar dependency front.
__syncthreads();
}
}
}
__device__ __forceinline__ void loco_trsm64_dense_warp(
float* __restrict__ row,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
int panel
) {
int lane = threadIdx.x & 31;
float x0 = row[panel + lane];
float x1 = row[panel + lane + 32];
#pragma unroll 1
for (int c = 0; c < 64; ++c) {
int diagonal_row = panel + c;
int diagonal_base = diagonal_row * (diagonal_row + 1) >> 1;
float partial = lane < c
? x0 * diagonal[diagonal_base + panel + lane]
: 0.0f;
if (lane + 32 < c) {
partial = fmaf(
x1,
diagonal[diagonal_base + panel + lane + 32],
partial);
}
partial += __shfl_down_sync(0xffffffffu, partial, 16);
partial += __shfl_down_sync(0xffffffffu, partial, 8);
partial += __shfl_down_sync(0xffffffffu, partial, 4);
partial += __shfl_down_sync(0xffffffffu, partial, 2);
partial += __shfl_down_sync(0xffffffffu, partial, 1);
float sum = __shfl_sync(0xffffffffu, partial, 0);
int owner = c & 31;
float rhs = c < 32
? __shfl_sync(0xffffffffu, x0, owner)
: __shfl_sync(0xffffffffu, x1, owner);
float value = (rhs - sum) * inverse[diagonal_row];
if (lane == owner) {
if (c < 32) x0 = value;
else x1 = value;
}
}
row[panel + lane] = x0;
row[panel + lane + 32] = x1;
}
__device__ __forceinline__ void loco_trsm64_dense_shadow_warp(
float* __restrict__ row,
uint32_t* __restrict__ shadow_row,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
int panel
) {
int lane = threadIdx.x & 31;
float x0 = row[panel + lane];
float x1 = row[panel + lane + 32];
#pragma unroll 1
for (int c = 0; c < 64; ++c) {
int diagonal_row = panel + c;
int diagonal_base = diagonal_row * (diagonal_row + 1) >> 1;
float partial = lane < c
? x0 * diagonal[diagonal_base + panel + lane]
: 0.0f;
if (lane + 32 < c) {
partial = fmaf(
x1,
diagonal[diagonal_base + panel + lane + 32],
partial);
}
partial += __shfl_down_sync(0xffffffffu, partial, 16);
partial += __shfl_down_sync(0xffffffffu, partial, 8);
partial += __shfl_down_sync(0xffffffffu, partial, 4);
partial += __shfl_down_sync(0xffffffffu, partial, 2);
partial += __shfl_down_sync(0xffffffffu, partial, 1);
float sum = __shfl_sync(0xffffffffu, partial, 0);
int owner = c & 31;
float rhs = c < 32
? __shfl_sync(0xffffffffu, x0, owner)
: __shfl_sync(0xffffffffu, x1, owner);
float value = (rhs - sum) * inverse[diagonal_row];
if (lane == owner) {
if (c < 32) x0 = value;
else x1 = value;
}
}
float y0 = __shfl_down_sync(0xffffffffu, x0, 1);
float y1 = __shfl_down_sync(0xffffffffu, x1, 1);
row[panel + lane] = x0;
row[panel + lane + 32] = x1;
if (!(lane & 1)) {
int word = (panel >> 1) + (lane >> 1);
shadow_row[word] = loco_f16x2(x0, y0);
shadow_row[word + 16] = loco_f16x2(x1, y1);
}
}
__device__ __forceinline__ void loco_trsm256_bf16(
float* __restrict__ rows,
uint32_t* __restrict__ shadow_rows,
int row_stride,
int row_count,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
float* a_smem
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int group = lane >> 2;
int rank = lane & 3;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 64) {
for (int row = warp; row < row_count; row += 32) {
loco_trsm64_dense_shadow_warp(
rows + (uint64_t)row * row_stride,
shadow_rows + (uint64_t)row * (row_stride >> 1),
diagonal,
inverse,
panel);
}
__syncthreads();
int start = panel + 64;
#pragma unroll 1
for (int row0 = 0;
row0 < row_count && start < 256;
row0 += 16) {
if (tid < 256) {
int source_row = tid >> 4;
int source_col = (tid & 15) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (row0 + source_row < row_count) {
value = loco_ld_global_v4(
rows
+ (uint64_t)(row0 + source_row) * row_stride
+ panel
+ source_col);
}
*reinterpret_cast<float4*>(
a_smem + source_row * 64 + source_col) = value;
}
__syncthreads();
int col0 = start + (warp << 3);
if (col0 < 256) {
int top = row0 + group;
int bottom = top + 8;
int col = col0 + (rank << 1);
int b_row = col0 + group;
int b_base = b_row * (b_row + 1) >> 1;
float d0 = 0.0f;
float d1 = 0.0f;
float d2 = 0.0f;
float d3 = 0.0f;
if (top < row_count) {
float* top_row =
rows + (uint64_t)top * row_stride;
d0 = top_row[col];
d1 = top_row[col + 1];
}
if (bottom < row_count) {
float* bottom_row =
rows + (uint64_t)bottom * row_stride;
d2 = bottom_row[col];
d3 = bottom_row[col + 1];
}
#pragma unroll
for (int kk = 0; kk < 64; kk += 16) {
int k0 = kk + (rank << 1);
int k8 = k0 + 8;
float2 at0 = *reinterpret_cast<const float2*>(
a_smem + group * 64 + k0);
float2 ab0 = *reinterpret_cast<const float2*>(
a_smem + (group + 8) * 64 + k0);
float2 at8 = *reinterpret_cast<const float2*>(
a_smem + group * 64 + k8);
float2 ab8 = *reinterpret_cast<const float2*>(
a_smem + (group + 8) * 64 + k8);
float2 bv0 = make_float2(
diagonal[b_base + panel + k0],
diagonal[b_base + panel + k0 + 1]);
float2 bv8 = make_float2(
diagonal[b_base + panel + k8],
diagonal[b_base + panel + k8 + 1]);
loco_mma_bf16(
loco_bf16x2(at0.x, at0.y),
loco_bf16x2(ab0.x, ab0.y),
loco_bf16x2(at8.x, at8.y),
loco_bf16x2(ab8.x, ab8.y),
loco_bf16x2(bv0.x, bv0.y) ^ 0x80008000u,
loco_bf16x2(bv8.x, bv8.y) ^ 0x80008000u,
d0, d1, d2, d3);
}
if (top < row_count) {
float* top_row =
rows + (uint64_t)top * row_stride;
top_row[col] = d0;
top_row[col + 1] = d1;
}
if (bottom < row_count) {
float* bottom_row =
rows + (uint64_t)bottom * row_stride;
bottom_row[col] = d2;
bottom_row[col + 1] = d3;
}
}
__syncthreads();
}
}
}
__device__ __forceinline__ void loco_trsm256_hexlift(
LocoHexLiftInnerShared& shared,
const uint8_t* __restrict__ diagonal_shadow,
float* __restrict__ rows,
uint32_t* __restrict__ shadow_rows,
int row_stride,
int row_count,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase
) {
int warp = threadIdx.x >> 5;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 64) {
for (int row = warp; row < row_count; row += 32) {
loco_trsm64_dense_shadow_warp(
rows + (uint64_t)row * row_stride,
shadow_rows + (uint64_t)row * (row_stride >> 1),
diagonal,
inverse,
panel);
}
__syncthreads();
int start = panel + 64;
for (int row0 = 0;
row0 < row_count && start < 256;
row0 += 128) {
loco_hexlift_trsm_update(
shared,
diagonal_shadow,
rows,
row_stride,
row_count,
row0,
panel,
start,
row0 == 0,
hex10,
issuer,
mma_phase);
}
}
}
template <bool DirectPublish, bool F16x1 = false>
__device__ __forceinline__ void loco_persistent_trsm64_dense_warp(
float* __restrict__ row,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
int panel,
float* a_smem,
int packet_row
) {
int lane = threadIdx.x & 31;
float x0 = row[panel + lane];
float x1 = row[panel + lane + 32];
#pragma unroll 1
for (int c = 0; c < 64; ++c) {
int diagonal_row = panel + c;
int diagonal_base = diagonal_row * (diagonal_row + 1) >> 1;
float partial = lane < c
? x0 * diagonal[diagonal_base + panel + lane]
: 0.0f;
if (lane + 32 < c) {
partial = fmaf(
x1,
diagonal[diagonal_base + panel + lane + 32],
partial);
}
partial += __shfl_down_sync(0xffffffffu, partial, 16);
partial += __shfl_down_sync(0xffffffffu, partial, 8);
partial += __shfl_down_sync(0xffffffffu, partial, 4);
partial += __shfl_down_sync(0xffffffffu, partial, 2);
partial += __shfl_down_sync(0xffffffffu, partial, 1);
float sum = __shfl_sync(0xffffffffu, partial, 0);
int owner = c & 31;
float rhs = c < 32
? __shfl_sync(0xffffffffu, x0, owner)
: __shfl_sync(0xffffffffu, x1, owner);
float value = (rhs - sum) * inverse[diagonal_row];
if (lane == owner) {
if (c < 32) x0 = value;
else x1 = value;
}
}
row[panel + lane] = x0;
row[panel + lane + 32] = x1;
if constexpr (DirectPublish) {
if (panel < 192) {
if constexpr (F16x1) {
float y0 = __shfl_down_sync(
0xffffffffu, x0, 1);
float y1 = __shfl_down_sync(
0xffffffffu, x1, 1);
if (!(lane & 1)) {
// Four K16 chunks, each split into two bank-perfect K8
// planes. One solved K64 row becomes 32 packed words.
reinterpret_cast<uint32_t*>(a_smem)[
((lane & 24) << 7)
+ (packet_row << 2)
+ ((lane >> 1) & 3)
] = loco_f16x2(x0, y0);
reinterpret_cast<uint32_t*>(a_smem)[
4096
+ ((lane & 24) << 7)
+ (packet_row << 2)
+ ((lane >> 1) & 3)
] = loco_f16x2(x1, y1);
}
} else {
a_smem[packet_row * 72 + lane] = x0;
a_smem[packet_row * 72 + lane + 32] = x1;
}
}
}
}
template <bool DirectPublish, bool F16x1 = false>
__device__ __forceinline__ void loco_persistent_trsm256(
float* __restrict__ rows,
int row_stride,
int row_count,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
float* a_smem,
uint32_t* b_packet = nullptr
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int group = lane >> 2;
int rank = lane & 3;
#pragma unroll 1
for (int panel = 0; panel < 256; panel += 64) {
if constexpr (F16x1) {
if (panel < 192) {
// Immutable L00 is packed once per K64 front. Bits
// [12:11,10,9:2,1:0] = [K16,K8,row,rank], so each later
// operand instruction sees all 32 banks exactly once.
#pragma unroll
for (int word = tid; word < 8192; word += 1024) {
int row = (word >> 2) & 255;
int k =
panel
+ ((word >> 7) & 56)
+ ((word & 3) << 1);
b_packet[word] = row >= panel + 64
? loco_f16x2(
diagonal[loco_lower256(row, k)],
diagonal[loco_lower256(row, k + 1)])
: 0u;
}
}
}
// The dependency front is exact FP32. B640 rounds the solved K64
// product once to FP16; other ranked paths retain BF16x3.
for (int row = warp; row < row_count; row += 32) {
loco_persistent_trsm64_dense_warp<DirectPublish, F16x1>(
rows + (uint64_t)row * row_stride,
diagonal,
inverse,
panel,
a_smem,
row);
}
__syncthreads();
int start = panel + 64;
if (start < 256) {
if constexpr (!DirectPublish) {
for (int vector = tid;
vector < row_count * 16;
vector += 1024) {
int source_row = vector >> 4;
int source_col = (vector & 15) << 2;
*reinterpret_cast<float4*>(
a_smem + source_row * 72 + source_col) =
loco_ld_global_v4_cs(
rows
+ (uint64_t)source_row * row_stride
+ panel
+ source_col);
}
__syncthreads();
}
// K64 is resident before its future-column cone consumes it.
// +8 words gives each eight-row MMA group a new bank phase.
int col0 = start + (warp << 3);
#pragma unroll 1
for (int row0 = 0; row0 < row_count; row0 += 16) {
int top = row0 + group;
int bottom = top + 8;
int col = col0 + (rank << 1);
int b_row = col0 + group;
float d0 = 0.0f;
float d1 = 0.0f;
float d2 = 0.0f;
float d3 = 0.0f;
if (col0 < 256 && top < row_count) {
float* top_row =
rows + (uint64_t)top * row_stride;
d0 = top_row[col];
d1 = top_row[col + 1];
}
if (col0 < 256 && bottom < row_count) {
float* bottom_row =
rows + (uint64_t)bottom * row_stride;
d2 = bottom_row[col];
d3 = bottom_row[col + 1];
}
if (col0 < 256 && top < row_count) {
int b_base = b_row * (b_row + 1) >> 1;
#pragma unroll
for (int kk = 0; kk < 64; kk += 16) {
if constexpr (F16x1) {
loco_mma_f16(
reinterpret_cast<uint32_t*>(a_smem)[
(kk << 7)
+ (top << 2)
+ rank],
bottom < row_count
? reinterpret_cast<uint32_t*>(a_smem)[
(kk << 7)
+ (bottom << 2)
+ rank]
: 0u,
reinterpret_cast<uint32_t*>(a_smem)[
(kk << 7)
+ 1024
+ (top << 2)
+ rank],
bottom < row_count
? reinterpret_cast<uint32_t*>(a_smem)[
(kk << 7)
+ 1024
+ (bottom << 2)
+ rank]
: 0u,
b_packet[
(kk << 7)
+ (b_row << 2)
+ rank]
^ 0x80008000u,
b_packet[
(kk << 7)
+ 1024
+ (b_row << 2)
+ rank]
^ 0x80008000u,
d0, d1, d2, d3);
} else {
int k0 = kk + (rank << 1);
int k8 = k0 + 8;
float2 at0 =
*reinterpret_cast<const float2*>(
a_smem + top * 72 + k0);
float2 at8 =
*reinterpret_cast<const float2*>(
a_smem + top * 72 + k8);
float2 ab0 = make_float2(0.0f, 0.0f);
float2 ab8 = make_float2(0.0f, 0.0f);
if (bottom < row_count) {
ab0 =
*reinterpret_cast<const float2*>(
a_smem + bottom * 72 + k0);
ab8 =
*reinterpret_cast<const float2*>(
a_smem + bottom * 72 + k8);
}
float2 bv0 = make_float2(
diagonal[b_base + panel + k0],
diagonal[b_base + panel + k0 + 1]);
float2 bv8 = make_float2(
diagonal[b_base + panel + k8],
diagonal[b_base + panel + k8 + 1]);
loco_mma_bf16x3(
at0, ab0, at8, ab8, bv0, bv8,
d0, d1, d2, d3);
}
}
}
if (col0 < 256 && top < row_count) {
float* top_row =
rows + (uint64_t)top * row_stride;
top_row[col] = d0;
top_row[col + 1] = d1;
}
if (col0 < 256 && bottom < row_count) {
float* bottom_row =
rows + (uint64_t)bottom * row_stride;
bottom_row[col] = d2;
bottom_row[col + 1] = d3;
}
}
// All N16 row stripes are disjoint; one edge publishes the full
// K64 product before the next exact solve front.
__syncthreads();
}
}
}
// Ranked N256 is a 64-matrix latency/throughput hybrid. One CTA owns one
// matrix end-to-end: direct input -> packed lower board -> compensated factor
// -> lower-only output. This removes the global factor workspace, FP16 shadow,
// and every inter-kernel dependency edge from the V9 route.
__global__ __launch_bounds__(1024, 1)
void loco_persistent256_x3_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
extern __shared__ __align__(16) float smem[];
float* tile = smem;
float* inverse = tile + 256 * 257 / 2;
float* a_smem = inverse + 256;
int tid = threadIdx.x;
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
uint64_t base = (uint64_t)b * 256u * 256u;
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int column = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
input + base + (uint64_t)row * 256u + column);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
tile[loco_lower256(row, column + e)] = word[e];
}
}
}
__syncthreads();
loco_factor_k16<256, 32>(tile, inverse);
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int column = (vector & 63) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
word[e] = tile[loco_lower256(row, column + e)];
}
}
*reinterpret_cast<float4*>(
output + base + (uint64_t)row * 256u + column) = value;
}
__syncthreads();
}
}
__global__ __launch_bounds__(1024, 1)
void loco_factor256_kernel(
float* __restrict__ matrix,
uint32_t* __restrict__ packed_fp16,
int batch,
int n,
int panel
) {
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoHexLiftInnerShared& shared =
*reinterpret_cast<LocoHexLiftInnerShared*>(dynamic_smem);
float* tile = shared.tile;
int tid = threadIdx.x;
int warp = tid >> 5;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t mma_phase = 0u;
bool hex10 = n < 2048;
if (tid == 0) {
loco_mbar_init(done);
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
uint64_t base = (uint64_t)b * (uint64_t)n * (uint64_t)n;
// K256 never writes the first 256 rows of the canonical FP16
// history. Reuse that dead prefix as one rolling diagonal-B packet;
// the largest Hex10 board occupies only 50,688 bytes.
uint8_t* diagonal_shadow =
reinterpret_cast<uint8_t*>(packed_fp16)
+ (base << 1);
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = loco_ld_global_v4(
matrix
+ base
+ (uint64_t)(panel + row) * n
+ panel
+ col);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
tile[loco_lower256(row, col + e)] = word[e];
}
}
}
__syncthreads();
loco_factor256_hexlift(
shared,
diagonal_shadow,
hex10,
issuer,
mma_phase);
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
word[e] = tile[loco_lower256(row, col + e)];
}
}
*reinterpret_cast<float4*>(
matrix
+ base
+ (uint64_t)(panel + row) * n
+ panel
+ col) = value;
}
__syncthreads();
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
}
__global__ __launch_bounds__(1024, 1)
void loco_trsm256_kernel(
float* __restrict__ matrix,
uint32_t* __restrict__ packed_fp16,
int batch,
int n,
int panel,
int row_groups,
int rows_per_group
) {
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoHexLiftInnerShared& shared =
*reinterpret_cast<LocoHexLiftInnerShared*>(dynamic_smem);
float* tile = shared.tile;
float* inverse = shared.inverse;
int tid = threadIdx.x;
int warp = tid >> 5;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t mma_phase = 0u;
bool hex10 = n < 2048;
int tasks = batch * row_groups;
if (tid == 0) {
loco_mbar_init(done);
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
for (int task = blockIdx.x; task < tasks; task += gridDim.x) {
int b = task / row_groups;
int group = task - b * row_groups;
uint64_t base = (uint64_t)b * (uint64_t)n * (uint64_t)n;
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = loco_ld_global_v4(
matrix
+ base
+ (uint64_t)(panel + row) * n
+ panel
+ col);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
tile[loco_lower256(row, col + e)] = word[e];
}
}
}
__syncthreads();
if (tid < 256) {
inverse[tid] = xcal_rcp_approx(
tile[loco_lower256(tid, tid)]);
}
__syncthreads();
int first = panel + 256 + group * rows_per_group;
int last = first + rows_per_group;
if (last > n) last = n;
if (first < last) {
uint64_t first_element =
base + (uint64_t)first * n + panel;
const uint8_t* diagonal_shadow =
reinterpret_cast<const uint8_t*>(packed_fp16)
+ (base << 1);
loco_trsm256_hexlift(
shared,
diagonal_shadow,
matrix + first_element,
packed_fp16 + (first_element >> 1),
n,
last - first,
tile,
inverse,
hex10,
issuer,
mma_phase);
}
__syncthreads();
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
}
__global__ __launch_bounds__(1024, 1)
void lower_copy_kernel(
const float* __restrict__ input,
float* __restrict__ output,
uint64_t total,
int n
) {
// Every ranked N is a power of two and N128 aligned. Move one aligned
// 32-byte row packet per lane and derive (row,column) with shifts/masks;
// this removes the old 64-bit divide/modulo from every 16-byte packet.
if (!(n & (n - 1)) && !(n & 7)) {
int log2_n = 31 - __clz((unsigned int)n);
int row_shift = log2_n - 3;
uint64_t row_packets = (uint64_t)n >> 3;
uint64_t row_mask = row_packets - 1u;
uint64_t total_packets = total >> 3;
uint64_t stride = (uint64_t)gridDim.x * blockDim.x;
for (uint64_t packet =
(uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
packet < total_packets;
packet += stride) {
uint64_t linear_row = packet >> row_shift;
int row = (int)(linear_row & (uint64_t)(n - 1));
int column = (int)(packet & row_mask) << 3;
float4 first = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float4 second = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (column <= row) {
const float* source = input + (packet << 3);
first = loco_ld_global_v4_cs(source);
second = loco_ld_global_v4_cs(source + 4);
float* first_word = reinterpret_cast<float*>(&first);
float* second_word = reinterpret_cast<float*>(&second);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (column + e > row) first_word[e] = 0.0f;
if (column + 4 + e > row) second_word[e] = 0.0f;
}
}
float* destination = output + (packet << 3);
*reinterpret_cast<float4*>(destination) = first;
*reinterpret_cast<float4*>(destination + 4) = second;
}
return;
}
// Arbitrary padded shapes retain the original safe 16-byte path.
uint64_t row_vectors = (uint64_t)n >> 2;
uint64_t total_vectors = total >> 2;
for (uint64_t vector =
(uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
vector < total_vectors;
vector += (uint64_t)gridDim.x * blockDim.x) {
uint64_t linear_row = vector / row_vectors;
int row = (int)(linear_row % (uint64_t)n);
int column = (int)(vector - linear_row * row_vectors) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (column <= row) {
value = loco_ld_global_v4_cs(input + (vector << 2));
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (column + e > row) word[e] = 0.0f;
}
}
*reinterpret_cast<float4*>(output + (vector << 2)) = value;
}
}
__global__ __launch_bounds__(1024, 1)
void lower_pad_copy_kernel(
const float* __restrict__ input,
float* __restrict__ factor,
uint64_t total,
int input_n,
int work_n
) {
uint64_t input_stride =
(uint64_t)input_n * (uint64_t)input_n;
uint64_t work_stride =
(uint64_t)work_n * (uint64_t)work_n;
for (uint64_t x = (uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
x < total;
x += (uint64_t)gridDim.x * blockDim.x) {
uint64_t b = x / work_stride;
uint64_t local = x - b * work_stride;
int row = (int)(local / (uint64_t)work_n);
int col = (int)(local - (uint64_t)row * (uint64_t)work_n);
float value = 0.0f;
if (row < input_n && col < input_n) {
if (row >= col) {
value = input[
b * input_stride
+ (uint64_t)row * input_n
+ col];
}
} else if (row == col) {
// diag(A, I) is SPD and leaves chol(A) unchanged in the crop.
value = 1.0f;
}
factor[x] = value;
}
}
__global__ __launch_bounds__(1024, 1)
void lower_crop_kernel(
const float* __restrict__ factor,
float* __restrict__ output,
uint64_t total,
int output_n,
int work_n
) {
uint64_t output_stride =
(uint64_t)output_n * (uint64_t)output_n;
uint64_t work_stride =
(uint64_t)work_n * (uint64_t)work_n;
for (uint64_t x = (uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
x < total;
x += (uint64_t)gridDim.x * blockDim.x) {
uint64_t b = x / output_stride;
uint64_t local = x - b * output_stride;
int row = (int)(local / (uint64_t)output_n);
int col = (int)(local - (uint64_t)row * (uint64_t)output_n);
output[x] = row >= col
? factor[
b * work_stride
+ (uint64_t)row * work_n
+ col]
: 0.0f;
}
}
// Macro-left GEMM dirties upper entries only inside each macro's diagonal
// square. The rest of the upper triangle is still the zero written by the
// initial lower copy, so clear O(N*macro) cells instead of O(N^2).
__global__ __launch_bounds__(1024, 1)
void zero_macro_upper_kernel(
float* __restrict__ factor,
int batch,
int n,
int macro_width
) {
int macro_blocks = (n + macro_width - 1) / macro_width;
int dirty_blocks = macro_blocks - 1;
if (dirty_blocks <= 0) return;
uint64_t row_vectors = (uint64_t)macro_width >> 2;
uint64_t tile_vectors = (uint64_t)macro_width * row_vectors;
uint64_t total_vectors =
(uint64_t)batch * (uint64_t)dirty_blocks * tile_vectors;
uint64_t matrix_stride = (uint64_t)n * (uint64_t)n;
for (uint64_t vector =
(uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
vector < total_vectors;
vector += (uint64_t)gridDim.x * blockDim.x) {
uint64_t tile_linear = vector / tile_vectors;
uint64_t local = vector - tile_linear * tile_vectors;
int b = (int)(tile_linear / (uint64_t)dirty_blocks);
int block =
(int)(tile_linear - (uint64_t)b * dirty_blocks) + 1;
int row = (int)(local / row_vectors);
int column = (int)(local - (uint64_t)row * row_vectors) << 2;
int base = block * macro_width;
int extent = n - base;
if (extent > macro_width) extent = macro_width;
if (row >= extent || column >= extent || row >= column + 4) {
continue;
}
float* dst = factor
+ (uint64_t)b * matrix_stride
+ (uint64_t)(base + row) * n
+ base
+ column;
if (row < column && column + 3 < extent) {
*reinterpret_cast<float4*>(dst) =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
} else {
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (column + e < extent && row < column + e) {
dst[e] = 0.0f;
}
}
}
}
}
// A P256 left-looking GEMM writes the full 256-column slab. The two
// K128 factors clean their own diagonal squares, leaving only the 128x128
// cross-half rectangle above the diagonal dirty. Clear exactly those cells;
// do not sweep the full O(N^2) upper triangle.
__global__ __launch_bounds__(256, 4)
void zero_mid_blas_cross_upper_kernel(
float* __restrict__ factor,
int batch,
int n
) {
int dirty_panels = (n >> 8) - 1;
if (dirty_panels <= 0) return;
// One task is four adjacent FP32 values in the 128x128 top-right half.
uint64_t vectors_per_panel = 128u * 32u;
uint64_t total =
(uint64_t)batch * (uint64_t)dirty_panels * vectors_per_panel;
uint64_t matrix_stride = (uint64_t)n * (uint64_t)n;
for (uint64_t task =
(uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
task < total;
task += (uint64_t)gridDim.x * blockDim.x) {
uint64_t panel_linear = task / vectors_per_panel;
uint64_t local = task - panel_linear * vectors_per_panel;
int b = (int)(panel_linear / (uint64_t)dirty_panels);
int panel_id =
(int)(panel_linear - (uint64_t)b * dirty_panels) + 1;
int row = (int)(local >> 5);
int col4 = (int)(local & 31u) << 2;
int panel = panel_id << 8;
float* dst = factor
+ (uint64_t)b * matrix_stride
+ (uint64_t)(panel + row) * n
+ panel + 128 + col4;
*reinterpret_cast<float4*>(dst) =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
}
__global__ __launch_bounds__(1024, 1)
void zero_upper_kernel(
float* __restrict__ factor,
uint64_t total,
int n,
int start
) {
int extent = n - start;
uint64_t region_stride =
(uint64_t)extent * (uint64_t)extent;
uint64_t matrix_stride =
(uint64_t)n * (uint64_t)n;
for (uint64_t x = (uint64_t)blockIdx.x * blockDim.x + threadIdx.x;
x < total;
x += (uint64_t)gridDim.x * blockDim.x) {
uint64_t b = x / region_stride;
uint64_t local = x - b * region_stride;
int row = (int)(local / (uint64_t)extent);
int col =
(int)(local - (uint64_t)row * (uint64_t)extent);
if (row < col) {
factor[
b * matrix_stride
+ (uint64_t)(start + row) * n
+ start
+ col] = 0.0f;
}
}
}
/*
Ranked shape 3: B256 x N128.
One CTA owns one matrix in the xR1 row x N16 register map:
tid = (q << 7) | r
q = 0..7 one N16 register packet
r = 0..127 one matrix row
The old xR1 panel used CTA-wide barriers at every scalar pivot even though
only q's 128-thread group touched that pivot. Here barrier IDs 1..8 are
private panel edges; barrier 0 is reserved for the four inter-panel handoff
edges. Fresh input is consumed on every call and the complete lower-only
output is written on every call.
*/
template<int Q>
static __device__ __forceinline__ void xrshape128_panel_barrier() {
asm volatile("barrier.cta.sync %0, 128;"
:
: "n"(Q + 1)
: "memory");
}
template<int Q, int C, int K>
static __device__ __forceinline__ void xrshape128_panel_update(
float& fc,
float fk,
float* smem
) {
if constexpr (C > K) {
constexpr int column_row = (Q << 4) + C;
if ((int)(threadIdx.x & 127u) >= column_row) {
asm volatile(
"{\n\t"
".reg .f32 l;\n\t"
"ld.shared.f32 l, [%1];\n\t"
"fma.rn.f32 %0, %2, l, %0;\n\t"
"}"
: "+f"(fc)
: "r"((unsigned int)__cvta_generic_to_shared(
smem + (K << 7) + column_row)),
"f"(fk)
: "memory");
}
}
}
template<int Q, int K, bool Rsqrt>
static __device__ __forceinline__ void xrshape128_pivot(
float& fk,
float& f0,
float& f1,
float& f2,
float& f3,
float& f4,
float& f5,
float& f6,
float& f7,
float& f8,
float& f9,
float& f10,
float& f11,
float& f12,
float& f13,
float& f14,
float& f15,
float* smem
) {
if ((threadIdx.x >> 7) != Q) return;
int row = (int)(threadIdx.x & 127u);
constexpr int diagonal = (Q << 4) + K;
if (row == diagonal) {
if constexpr (Rsqrt) {
asm volatile(
"{\n\t"
".reg .f32 ri;\n\t"
"rsqrt.approx.ftz.f32 ri, %0;\n\t"
"mul.rn.f32 %0, %0, ri;\n\t"
"xor.b32 ri, ri, 0x80000000;\n\t"
"st.shared.f32 [%1], ri;\n\t"
"}"
: "+f"(fk)
: "r"((unsigned int)__cvta_generic_to_shared(
smem + 2048))
: "memory");
} else {
asm volatile(
"{\n\t"
".reg .f32 ri;\n\t"
"sqrt.approx.ftz.f32 %0, %0;\n\t"
"rcp.approx.ftz.f32 ri, %0;\n\t"
"xor.b32 ri, ri, 0x80000000;\n\t"
"st.shared.f32 [%1], ri;\n\t"
"}"
: "+f"(fk)
: "r"((unsigned int)__cvta_generic_to_shared(
smem + 2048))
: "memory");
}
}
xrshape128_panel_barrier<Q>();
if (row == diagonal) {
asm volatile("st.shared.f32 [%0], %1;"
:
: "r"((unsigned int)__cvta_generic_to_shared(
smem + (K << 7) + row)),
"f"(fk)
: "memory");
} else if (row > diagonal) {
asm volatile(
"{\n\t"
".reg .f32 l;\n\t"
"ld.shared.f32 l, [%2];\n\t"
"fma.rn.f32 %0, %0, l, 0f00000000;\n\t"
"xor.b32 l, %0, 0x80000000;\n\t"
"st.shared.f32 [%1], l;\n\t"
"}"
: "+f"(fk)
: "r"((unsigned int)__cvta_generic_to_shared(
smem + (K << 7) + row)),
"r"((unsigned int)__cvta_generic_to_shared(smem + 2048))
: "memory");
} else {
asm volatile("st.shared.b32 [%0], %1;"
:
: "r"((unsigned int)__cvta_generic_to_shared(
smem + (K << 7) + row)),
"r"(0u)
: "memory");
}
xrshape128_panel_barrier<Q>();
// The shared panel is local K0..K15, so its solved row is Q*16+C.
xrshape128_panel_update<Q, 0, K>(f0, fk, smem);
xrshape128_panel_update<Q, 1, K>(f1, fk, smem);
xrshape128_panel_update<Q, 2, K>(f2, fk, smem);
xrshape128_panel_update<Q, 3, K>(f3, fk, smem);
xrshape128_panel_update<Q, 4, K>(f4, fk, smem);
xrshape128_panel_update<Q, 5, K>(f5, fk, smem);
xrshape128_panel_update<Q, 6, K>(f6, fk, smem);
xrshape128_panel_update<Q, 7, K>(f7, fk, smem);
xrshape128_panel_update<Q, 8, K>(f8, fk, smem);
xrshape128_panel_update<Q, 9, K>(f9, fk, smem);
xrshape128_panel_update<Q,10, K>(f10, fk, smem);
xrshape128_panel_update<Q,11, K>(f11, fk, smem);
xrshape128_panel_update<Q,12, K>(f12, fk, smem);
xrshape128_panel_update<Q,13, K>(f13, fk, smem);
xrshape128_panel_update<Q,14, K>(f14, fk, smem);
xrshape128_panel_update<Q,15, K>(f15, fk, smem);
if (row > diagonal) {
asm volatile("xor.b32 %0, %0, 0x80000000;" : "+f"(fk));
}
}
template<int HM, int HN>
static __device__ __forceinline__ void xrshape128_mma(
float& r0,
float& r1,
float& r2,
float& r3,
float& r4,
float& r5,
float& r6,
float& r7,
float* smem
) {
if ((threadIdx.x & 96u) + (HM << 4) + 15u
< ((threadIdx.x >> 7) << 4) + (HN << 3))
return;
asm volatile(
"{\n\t"
".reg .pred row8, take;\n\t"
".reg .b32 a0, a1, a2, a3, b0, b1, lane, src;\n\t"
".reg .f32 d0, d1, d2, d3, x, y;\n\t"
"mov.b32 d0, 0;\n\t"
"mov.b32 d1, 0;\n\t"
"mov.b32 d2, 0;\n\t"
"mov.b32 d3, 0;\n\t"
"ld.shared.b32 a0, [%8];\n\t"
"ld.shared.b32 a1, [%8+128];\n\t"
"ld.shared.b32 a2, [%8+2048];\n\t"
"ld.shared.b32 a3, [%8+2176];\n\t"
"xor.b32 a0, a0, 0x80008000;\n\t"
"xor.b32 a1, a1, 0x80008000;\n\t"
"xor.b32 a2, a2, 0x80008000;\n\t"
"xor.b32 a3, a3, 0x80008000;\n\t"
"ld.shared.b32 b0, [%9+4096];\n\t"
"ld.shared.b32 b1, [%9+6144];\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, {d0,d1,d2,d3};\n\t"
"ld.shared.b32 b0, [%9];\n\t"
"ld.shared.b32 b1, [%9+2048];\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, {d0,d1,d2,d3};\n\t"
"ld.shared.b32 a0, [%8+4096];\n\t"
"ld.shared.b32 a1, [%8+4224];\n\t"
"ld.shared.b32 a2, [%8+6144];\n\t"
"ld.shared.b32 a3, [%8+6272];\n\t"
"xor.b32 a0, a0, 0x80008000;\n\t"
"xor.b32 a1, a1, 0x80008000;\n\t"
"xor.b32 a2, a2, 0x80008000;\n\t"
"xor.b32 a3, a3, 0x80008000;\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, {d0,d1,d2,d3};\n\t"
"mov.u32 lane, %%laneid;\n\t"
"and.b32 src, lane, 16;\n\t"
"setp.eq.u32 take, src, %10;\n\t"
"and.b32 src, lane, 8;\n\t"
"setp.ne.u32 row8, src, 0;\n\t"
"and.b32 src, lane, 7;\n\t"
"shl.b32 src, src, 2;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %0, %0, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %1, %1, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %2, %2, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %3, %3, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %4, %4, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %5, %5, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %6, %6, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %7, %7, x;\n\t"
"}"
: "+f"(r0), "+f"(r1), "+f"(r2), "+f"(r3),
"+f"(r4), "+f"(r5), "+f"(r6), "+f"(r7)
: "r"((unsigned int)__cvta_generic_to_shared(
smem + ((threadIdx.x & 96u) << 2) + (HM << 6)
+ (threadIdx.x & 31u))),
"r"((unsigned int)__cvta_generic_to_shared(
smem + ((threadIdx.x >> 7) << 6) + (HN << 5)
+ (threadIdx.x & 31u))),
"n"(HM << 4)
: "memory");
}
template<int HM, int HN>
static __device__ __forceinline__ void xrshape128_mma_f16(
float& r0,
float& r1,
float& r2,
float& r3,
float& r4,
float& r5,
float& r6,
float& r7,
float* smem
) {
if ((threadIdx.x & 96u) + (HM << 4) + 15u
< ((threadIdx.x >> 7) << 4) + (HN << 3))
return;
asm volatile(
"{\n\t"
".reg .pred row8, take;\n\t"
".reg .b32 a0, a1, a2, a3, b0, b1, lane, src;\n\t"
".reg .f32 d0, d1, d2, d3, x, y;\n\t"
"mov.b32 d0, 0;\n\t"
"mov.b32 d1, 0;\n\t"
"mov.b32 d2, 0;\n\t"
"mov.b32 d3, 0;\n\t"
"ld.shared.b32 a0, [%8];\n\t"
"ld.shared.b32 a1, [%8+128];\n\t"
"ld.shared.b32 a2, [%8+2048];\n\t"
"ld.shared.b32 a3, [%8+2176];\n\t"
"xor.b32 a0, a0, 0x80008000;\n\t"
"xor.b32 a1, a1, 0x80008000;\n\t"
"xor.b32 a2, a2, 0x80008000;\n\t"
"xor.b32 a3, a3, 0x80008000;\n\t"
"ld.shared.b32 b0, [%9];\n\t"
"ld.shared.b32 b1, [%9+2048];\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, {d0,d1,d2,d3};\n\t"
"mov.u32 lane, %%laneid;\n\t"
"and.b32 src, lane, 16;\n\t"
"setp.eq.u32 take, src, %10;\n\t"
"and.b32 src, lane, 8;\n\t"
"setp.ne.u32 row8, src, 0;\n\t"
"and.b32 src, lane, 7;\n\t"
"shl.b32 src, src, 2;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %0, %0, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %1, %1, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %2, %2, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %3, %3, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %4, %4, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %5, %5, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %6, %6, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %7, %7, x;\n\t"
"}"
: "+f"(r0), "+f"(r1), "+f"(r2), "+f"(r3),
"+f"(r4), "+f"(r5), "+f"(r6), "+f"(r7)
: "r"((unsigned int)__cvta_generic_to_shared(
smem + ((threadIdx.x & 96u) << 2) + (HM << 6)
+ (threadIdx.x & 31u))),
"r"((unsigned int)__cvta_generic_to_shared(
smem + ((threadIdx.x >> 7) << 6) + (HN << 5)
+ (threadIdx.x & 31u))),
"n"(HM << 4)
: "memory");
}
/*
Large-front sidecar, produced once while the K128 factor is still resident.
The one-rail FP16 shadow for the diagonal K128 is a block-triangular hybrid:
block row > block col : solved L[q,p]
block row = block col : S[q] = inv(L[q,q])
block row < block col : unused
TRSM consumes B rows directly:
Xq = (Bq - sum_p Xp * L[q,p]^T) * S[q]^T.
Every stored row is contiguous K16, matching mma.sync's B row fragment.
*/
template<int Q>
static __device__ __forceinline__ void xrshape128_publish_mma_sidecar(
float f0,
float f1,
float f2,
float f3,
float f4,
float f5,
float f6,
float f7,
float f8,
float f9,
float f10,
float f11,
float f12,
float f13,
float f14,
float f15,
float* __restrict__ smem,
uint32_t* __restrict__ packed_fp16,
uint64_t matrix_base,
int n,
int panel
) {
if (!packed_fp16 || (threadIdx.x >> 7) != Q) return;
int row = (int)(threadIdx.x & 127u);
int row_block = row >> 4;
uint64_t dst = (
matrix_base
+ (uint64_t)(panel + row) * (uint64_t)n
+ panel
+ (Q << 4)) >> 1;
if (row_block > Q) {
packed_fp16[dst + 0] = loco_f16x2(f0, f1);
packed_fp16[dst + 1] = loco_f16x2(f2, f3);
packed_fp16[dst + 2] = loco_f16x2(f4, f5);
packed_fp16[dst + 3] = loco_f16x2(f6, f7);
packed_fp16[dst + 4] = loco_f16x2(f8, f9);
packed_fp16[dst + 5] = loco_f16x2(f10, f11);
packed_fp16[dst + 6] = loco_f16x2(f12, f13);
packed_fp16[dst + 7] = loco_f16x2(f14, f15);
return;
}
if (row_block != Q) return;
// One thread owns one row of S=inv(Lqq). The recurrence is row-local:
// S[i,i] = 1/L[i,i]
// S[i,j] = -sum_{t=j+1..i} S[i,t]L[t,j] / L[j,j].
int i = row & 15;
float inv_row[16];
#pragma unroll
for (int j = 0; j < 16; ++j) inv_row[j] = 0.0f;
#pragma unroll
for (int j = 15; j >= 0; --j) {
if (j > i) continue;
float diagonal = smem[(j << 7) + (Q << 4) + j];
float reciprocal = xcal_rcp_approx(diagonal);
if (j == i) {
inv_row[j] = reciprocal;
} else {
float sum = 0.0f;
#pragma unroll
for (int t = 0; t < 16; ++t) {
if (t > j && t <= i) {
sum = fmaf(
inv_row[t],
smem[(j << 7) + (Q << 4) + t],
sum);
}
}
inv_row[j] = -sum * reciprocal;
}
}
packed_fp16[dst + 0] = loco_f16x2(inv_row[0], inv_row[1]);
packed_fp16[dst + 1] = loco_f16x2(inv_row[2], inv_row[3]);
packed_fp16[dst + 2] = loco_f16x2(inv_row[4], inv_row[5]);
packed_fp16[dst + 3] = loco_f16x2(inv_row[6], inv_row[7]);
packed_fp16[dst + 4] = loco_f16x2(inv_row[8], inv_row[9]);
packed_fp16[dst + 5] = loco_f16x2(inv_row[10], inv_row[11]);
packed_fp16[dst + 6] = loco_f16x2(inv_row[12], inv_row[13]);
packed_fp16[dst + 7] = loco_f16x2(inv_row[14], inv_row[15]);
}
#define XRSHAPE128_PUBLISH_MMA(Q) \
if constexpr (PublishMmaSidecar) { \
xrshape128_publish_mma_sidecar<Q>( \
f0,f1,f2,f3,f4,f5,f6,f7, \
f8,f9,f10,f11,f12,f13,f14,f15, \
smem, packed_fp16, base, n, panel); \
}
#define XRSHAPE128_PIVOTS(Q, R) \
xrshape128_pivot<Q, 0,R>(f0, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 1,R>(f1, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 2,R>(f2, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 3,R>(f3, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 4,R>(f4, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 5,R>(f5, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 6,R>(f6, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 7,R>(f7, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 8,R>(f8, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q, 9,R>(f9, f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q,10,R>(f10,f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q,11,R>(f11,f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q,12,R>(f12,f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q,13,R>(f13,f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q,14,R>(f14,f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_pivot<Q,15,R>(f15,f0,f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13,f14,f15,smem)
#define XRSHAPE128_HANDOFF(Q) \
__syncthreads(); \
asm volatile( \
"ld.shared.f32 %0, [%2];\n\t" \
"ld.shared.f32 %1, [%3];" \
: "=f"(x), "=f"(y) \
: "r"((unsigned int)__cvta_generic_to_shared( \
smem + ((threadIdx.x >> 7) << 8) \
+ (threadIdx.x & 127u))), \
"r"((unsigned int)__cvta_generic_to_shared( \
smem + ((threadIdx.x >> 7) << 8) \
+ (threadIdx.x & 127u) + 128)) \
: "memory"); \
__syncthreads(); \
asm volatile( \
"{\n\t" \
".reg .f32 he, ho, le, lo;\n\t" \
".reg .b32 bf;\n\t" \
"lop3.b32 he, %1, 0xffff0000, 0, 0xc0;\n\t" \
"lop3.b32 ho, %2, 0xffff0000, 0, 0xc0;\n\t" \
"sub.rn.f32 le, %1, he;\n\t" \
"sub.rn.f32 lo, %2, ho;\n\t" \
"cvt.rn.bf16x2.f32 bf, ho, he;\n\t" \
"st.shared.b32 [%0], bf;\n\t" \
"cvt.rn.bf16x2.f32 bf, lo, le;\n\t" \
"st.shared.b32 [%0+4096], bf;\n\t" \
"}" \
: \
: "r"((unsigned int)__cvta_generic_to_shared( \
smem + ((threadIdx.x & 127u) << 2) \
+ ((threadIdx.x >> 7) & 3u) \
+ ((threadIdx.x >> 9) << 9))), \
"f"(x), "f"(y) \
: "memory"); \
__syncthreads(); \
if ((threadIdx.x >> 7) > Q) { \
xrshape128_mma<0,0>(f0,f1,f2,f3,f4,f5,f6,f7,smem); \
xrshape128_mma<0,1>(f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_mma<1,0>(f0,f1,f2,f3,f4,f5,f6,f7,smem); \
xrshape128_mma<1,1>(f8,f9,f10,f11,f12,f13,f14,f15,smem); \
} \
__syncthreads()
#define XRSHAPE128_HANDOFF_F16(Q) \
__syncthreads(); \
asm volatile( \
"ld.shared.f32 %0, [%2];\n\t" \
"ld.shared.f32 %1, [%3];" \
: "=f"(x), "=f"(y) \
: "r"((unsigned int)__cvta_generic_to_shared( \
smem + ((threadIdx.x >> 7) << 8) \
+ (threadIdx.x & 127u))), \
"r"((unsigned int)__cvta_generic_to_shared( \
smem + ((threadIdx.x >> 7) << 8) \
+ (threadIdx.x & 127u) + 128)) \
: "memory"); \
__syncthreads(); \
asm volatile( \
"cvt.rn.f16x2.f32 %0, %2, %1;\n\t" \
"st.shared.b32 [%3], %0;" \
: "=r"(packet_word) \
: "f"(x), "f"(y), \
"r"((unsigned int)__cvta_generic_to_shared( \
smem + ((threadIdx.x & 127u) << 2) \
+ ((threadIdx.x >> 7) & 3u) \
+ ((threadIdx.x >> 9) << 9))) \
: "memory"); \
__syncthreads(); \
if ((threadIdx.x >> 7) > Q) { \
xrshape128_mma_f16<0,0>(f0,f1,f2,f3,f4,f5,f6,f7,smem); \
xrshape128_mma_f16<0,1>(f8,f9,f10,f11,f12,f13,f14,f15,smem); \
xrshape128_mma_f16<1,0>(f0,f1,f2,f3,f4,f5,f6,f7,smem); \
xrshape128_mma_f16<1,1>(f8,f9,f10,f11,f12,f13,f14,f15,smem); \
} \
__syncthreads()
__global__ __launch_bounds__(1024, 1)
void xrshape128_b256_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
float f0 = 0.0f, f1 = 0.0f, f2 = 0.0f, f3 = 0.0f;
float f4 = 0.0f, f5 = 0.0f, f6 = 0.0f, f7 = 0.0f;
float f8 = 0.0f, f9 = 0.0f, f10 = 0.0f, f11 = 0.0f;
float f12 = 0.0f, f13 = 0.0f, f14 = 0.0f, f15 = 0.0f;
float x, y;
uint32_t packet_word = 0u;
__shared__ __align__(128) float smem[2049];
int row = (int)(threadIdx.x & 127u);
int column = (int)(threadIdx.x >> 7) << 4;
uint64_t base = (uint64_t)blockIdx.x << 14;
const float* src = input + base + (uint64_t)row * 128u + column;
if (row >= column + 15) {
float4 a = *reinterpret_cast<const float4*>(src + 0);
float4 b = *reinterpret_cast<const float4*>(src + 4);
float4 c = *reinterpret_cast<const float4*>(src + 8);
float4 d = *reinterpret_cast<const float4*>(src + 12);
f0=a.x; f1=a.y; f2=a.z; f3=a.w;
f4=b.x; f5=b.y; f6=b.z; f7=b.w;
f8=c.x; f9=c.y; f10=c.z; f11=c.w;
f12=d.x; f13=d.y; f14=d.z; f15=d.w;
} else if (row >= column) {
if (column + 0 <= row) f0 = src[ 0];
if (column + 1 <= row) f1 = src[ 1];
if (column + 2 <= row) f2 = src[ 2];
if (column + 3 <= row) f3 = src[ 3];
if (column + 4 <= row) f4 = src[ 4];
if (column + 5 <= row) f5 = src[ 5];
if (column + 6 <= row) f6 = src[ 6];
if (column + 7 <= row) f7 = src[ 7];
if (column + 8 <= row) f8 = src[ 8];
if (column + 9 <= row) f9 = src[ 9];
if (column + 10 <= row) f10 = src[10];
if (column + 11 <= row) f11 = src[11];
if (column + 12 <= row) f12 = src[12];
if (column + 13 <= row) f13 = src[13];
if (column + 14 <= row) f14 = src[14];
if (column + 15 <= row) f15 = src[15];
}
XRSHAPE128_PIVOTS(0, true); XRSHAPE128_HANDOFF_F16(0);
XRSHAPE128_PIVOTS(1, true); XRSHAPE128_HANDOFF_F16(1);
XRSHAPE128_PIVOTS(2, true); XRSHAPE128_HANDOFF_F16(2);
XRSHAPE128_PIVOTS(3, true); XRSHAPE128_HANDOFF_F16(3);
XRSHAPE128_PIVOTS(4, true); XRSHAPE128_HANDOFF_F16(4);
XRSHAPE128_PIVOTS(5, true); XRSHAPE128_HANDOFF_F16(5);
XRSHAPE128_PIVOTS(6, true); XRSHAPE128_HANDOFF_F16(6);
XRSHAPE128_PIVOTS(7, true);
if (row < column + 0) f0 = 0.0f;
if (row < column + 1) f1 = 0.0f;
if (row < column + 2) f2 = 0.0f;
if (row < column + 3) f3 = 0.0f;
if (row < column + 4) f4 = 0.0f;
if (row < column + 5) f5 = 0.0f;
if (row < column + 6) f6 = 0.0f;
if (row < column + 7) f7 = 0.0f;
if (row < column + 8) f8 = 0.0f;
if (row < column + 9) f9 = 0.0f;
if (row < column + 10) f10 = 0.0f;
if (row < column + 11) f11 = 0.0f;
if (row < column + 12) f12 = 0.0f;
if (row < column + 13) f13 = 0.0f;
if (row < column + 14) f14 = 0.0f;
if (row < column + 15) f15 = 0.0f;
float* dst = output + base + (uint64_t)row * 128u + column;
*reinterpret_cast<float4*>(dst + 0) = make_float4(f0,f1,f2,f3);
*reinterpret_cast<float4*>(dst + 4) = make_float4(f4,f5,f6,f7);
*reinterpret_cast<float4*>(dst + 8) = make_float4(f8,f9,f10,f11);
*reinterpret_cast<float4*>(dst + 12) = make_float4(f12,f13,f14,f15);
}
/*
First-principles K128 diagonal front.
This is the passing B256/N128 register board, made strided and in-place.
The board is independent of matrix extent; only the storage stride changes:
matrix[b, panel:panel+128, panel:panel+128]
Eight named 128-thread barriers own the scalar dependency groups. The
seven inter-group handoffs remain CTA-wide because they publish a complete
compensated BF16 K16 product to the next group.
*/
template <bool Rsqrt, bool PublishMmaSidecar, bool FastF16Handoff>
__global__ __launch_bounds__(1024, 1)
void xrshape128_diag_mid_kernel(
float* __restrict__ matrix,
uint32_t* __restrict__ packed_fp16,
int batch,
int n,
int panel
) {
float f0 = 0.0f, f1 = 0.0f, f2 = 0.0f, f3 = 0.0f;
float f4 = 0.0f, f5 = 0.0f, f6 = 0.0f, f7 = 0.0f;
float f8 = 0.0f, f9 = 0.0f, f10 = 0.0f, f11 = 0.0f;
float f12 = 0.0f, f13 = 0.0f, f14 = 0.0f, f15 = 0.0f;
float x, y;
uint32_t packet_word = 0u;
__shared__ __align__(128) float smem[2049];
int b = (int)blockIdx.x;
if (b >= batch) return;
int row = (int)(threadIdx.x & 127u);
int column = (int)(threadIdx.x >> 7) << 4;
uint64_t base = (uint64_t)b * (uint64_t)n * (uint64_t)n;
float* src = matrix
+ base
+ (uint64_t)(panel + row) * (uint64_t)n
+ panel
+ column;
if (row >= column + 15) {
float4 a = *reinterpret_cast<const float4*>(src + 0);
float4 b4 = *reinterpret_cast<const float4*>(src + 4);
float4 c = *reinterpret_cast<const float4*>(src + 8);
float4 d = *reinterpret_cast<const float4*>(src + 12);
f0=a.x; f1=a.y; f2=a.z; f3=a.w;
f4=b4.x; f5=b4.y; f6=b4.z; f7=b4.w;
f8=c.x; f9=c.y; f10=c.z; f11=c.w;
f12=d.x; f13=d.y; f14=d.z; f15=d.w;
} else if (row >= column) {
if (column + 0 <= row) f0 = src[ 0];
if (column + 1 <= row) f1 = src[ 1];
if (column + 2 <= row) f2 = src[ 2];
if (column + 3 <= row) f3 = src[ 3];
if (column + 4 <= row) f4 = src[ 4];
if (column + 5 <= row) f5 = src[ 5];
if (column + 6 <= row) f6 = src[ 6];
if (column + 7 <= row) f7 = src[ 7];
if (column + 8 <= row) f8 = src[ 8];
if (column + 9 <= row) f9 = src[ 9];
if (column + 10 <= row) f10 = src[10];
if (column + 11 <= row) f11 = src[11];
if (column + 12 <= row) f12 = src[12];
if (column + 13 <= row) f13 = src[13];
if (column + 14 <= row) f14 = src[14];
if (column + 15 <= row) f15 = src[15];
}
XRSHAPE128_PIVOTS(0, Rsqrt); XRSHAPE128_PUBLISH_MMA(0);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(0); } else { XRSHAPE128_HANDOFF(0); }
XRSHAPE128_PIVOTS(1, Rsqrt); XRSHAPE128_PUBLISH_MMA(1);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(1); } else { XRSHAPE128_HANDOFF(1); }
XRSHAPE128_PIVOTS(2, Rsqrt); XRSHAPE128_PUBLISH_MMA(2);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(2); } else { XRSHAPE128_HANDOFF(2); }
XRSHAPE128_PIVOTS(3, Rsqrt); XRSHAPE128_PUBLISH_MMA(3);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(3); } else { XRSHAPE128_HANDOFF(3); }
XRSHAPE128_PIVOTS(4, Rsqrt); XRSHAPE128_PUBLISH_MMA(4);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(4); } else { XRSHAPE128_HANDOFF(4); }
XRSHAPE128_PIVOTS(5, Rsqrt); XRSHAPE128_PUBLISH_MMA(5);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(5); } else { XRSHAPE128_HANDOFF(5); }
XRSHAPE128_PIVOTS(6, Rsqrt); XRSHAPE128_PUBLISH_MMA(6);
if constexpr (FastF16Handoff) { XRSHAPE128_HANDOFF_F16(6); } else { XRSHAPE128_HANDOFF(6); }
XRSHAPE128_PIVOTS(7, Rsqrt); XRSHAPE128_PUBLISH_MMA(7);
if (row < column + 0) f0 = 0.0f;
if (row < column + 1) f1 = 0.0f;
if (row < column + 2) f2 = 0.0f;
if (row < column + 3) f3 = 0.0f;
if (row < column + 4) f4 = 0.0f;
if (row < column + 5) f5 = 0.0f;
if (row < column + 6) f6 = 0.0f;
if (row < column + 7) f7 = 0.0f;
if (row < column + 8) f8 = 0.0f;
if (row < column + 9) f9 = 0.0f;
if (row < column + 10) f10 = 0.0f;
if (row < column + 11) f11 = 0.0f;
if (row < column + 12) f12 = 0.0f;
if (row < column + 13) f13 = 0.0f;
if (row < column + 14) f14 = 0.0f;
if (row < column + 15) f15 = 0.0f;
float* dst = matrix
+ base
+ (uint64_t)(panel + row) * (uint64_t)n
+ panel
+ column;
*reinterpret_cast<float4*>(dst + 0) = make_float4(f0,f1,f2,f3);
*reinterpret_cast<float4*>(dst + 4) = make_float4(f4,f5,f6,f7);
*reinterpret_cast<float4*>(dst + 8) = make_float4(f8,f9,f10,f11);
*reinterpret_cast<float4*>(dst + 12) = make_float4(f12,f13,f14,f15);
}
/*
Ranked shape 4: B64 x N256.
A = [A00 * ] L = [L00 0 ]
[A10 A11] [L10 L11]
Four ordered launches keep every call honest while matching B200 residency:
64 CTAs factor A00
128 CTAs solve L10, two 64-row teams per matrix
128 CTAs form A11-L10*L10^T, two balanced tile teams per matrix
64 CTAs factor the resulting A11 block
Every launch reads the supplied matrix or an immediately preceding result.
There is no cross-call retention and no numerical-library path.
*/
constexpr int XRSHAPE256_STRIDE = 132;
constexpr int XRSHAPE256_TRSM_SMEM =
(128 * XRSHAPE256_STRIDE + 128) * (int)sizeof(float);
constexpr int XRSHAPE256_UPDATE_SMEM =
128 * XRSHAPE256_STRIDE * (int)sizeof(float);
__global__ __launch_bounds__(1024, 1)
void xrshape256_diag_b64_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int diagonal
) {
float f0 = 0.0f, f1 = 0.0f, f2 = 0.0f, f3 = 0.0f;
float f4 = 0.0f, f5 = 0.0f, f6 = 0.0f, f7 = 0.0f;
float f8 = 0.0f, f9 = 0.0f, f10 = 0.0f, f11 = 0.0f;
float f12 = 0.0f, f13 = 0.0f, f14 = 0.0f, f15 = 0.0f;
float x, y;
__shared__ __align__(128) float smem[2049];
int row = (int)(threadIdx.x & 127u);
int column = (int)(threadIdx.x >> 7) << 4;
uint64_t base = (uint64_t)blockIdx.x << 16;
const float* src = input
+ base
+ (uint64_t)(diagonal + row) * 256u
+ diagonal
+ column;
if (row >= column + 15) {
float4 a = *reinterpret_cast<const float4*>(src + 0);
float4 b = *reinterpret_cast<const float4*>(src + 4);
float4 c = *reinterpret_cast<const float4*>(src + 8);
float4 d = *reinterpret_cast<const float4*>(src + 12);
f0=a.x; f1=a.y; f2=a.z; f3=a.w;
f4=b.x; f5=b.y; f6=b.z; f7=b.w;
f8=c.x; f9=c.y; f10=c.z; f11=c.w;
f12=d.x; f13=d.y; f14=d.z; f15=d.w;
} else if (row >= column) {
if (column + 0 <= row) f0 = src[ 0];
if (column + 1 <= row) f1 = src[ 1];
if (column + 2 <= row) f2 = src[ 2];
if (column + 3 <= row) f3 = src[ 3];
if (column + 4 <= row) f4 = src[ 4];
if (column + 5 <= row) f5 = src[ 5];
if (column + 6 <= row) f6 = src[ 6];
if (column + 7 <= row) f7 = src[ 7];
if (column + 8 <= row) f8 = src[ 8];
if (column + 9 <= row) f9 = src[ 9];
if (column + 10 <= row) f10 = src[10];
if (column + 11 <= row) f11 = src[11];
if (column + 12 <= row) f12 = src[12];
if (column + 13 <= row) f13 = src[13];
if (column + 14 <= row) f14 = src[14];
if (column + 15 <= row) f15 = src[15];
}
XRSHAPE128_PIVOTS(0, false); XRSHAPE128_HANDOFF(0);
XRSHAPE128_PIVOTS(1, false); XRSHAPE128_HANDOFF(1);
XRSHAPE128_PIVOTS(2, false); XRSHAPE128_HANDOFF(2);
XRSHAPE128_PIVOTS(3, false); XRSHAPE128_HANDOFF(3);
XRSHAPE128_PIVOTS(4, false); XRSHAPE128_HANDOFF(4);
XRSHAPE128_PIVOTS(5, false); XRSHAPE128_HANDOFF(5);
XRSHAPE128_PIVOTS(6, false); XRSHAPE128_HANDOFF(6);
XRSHAPE128_PIVOTS(7, false);
if (row < column + 0) f0 = 0.0f;
if (row < column + 1) f1 = 0.0f;
if (row < column + 2) f2 = 0.0f;
if (row < column + 3) f3 = 0.0f;
if (row < column + 4) f4 = 0.0f;
if (row < column + 5) f5 = 0.0f;
if (row < column + 6) f6 = 0.0f;
if (row < column + 7) f7 = 0.0f;
if (row < column + 8) f8 = 0.0f;
if (row < column + 9) f9 = 0.0f;
if (row < column + 10) f10 = 0.0f;
if (row < column + 11) f11 = 0.0f;
if (row < column + 12) f12 = 0.0f;
if (row < column + 13) f13 = 0.0f;
if (row < column + 14) f14 = 0.0f;
if (row < column + 15) f15 = 0.0f;
float* dst = output
+ base
+ (uint64_t)(diagonal + row) * 256u
+ diagonal
+ column;
*reinterpret_cast<float4*>(dst + 0) = make_float4(f0,f1,f2,f3);
*reinterpret_cast<float4*>(dst + 4) = make_float4(f4,f5,f6,f7);
*reinterpret_cast<float4*>(dst + 8) = make_float4(f8,f9,f10,f11);
*reinterpret_cast<float4*>(dst + 12) = make_float4(f12,f13,f14,f15);
// A00's companion quadrant is the strict upper-right block of L.
if (!diagonal) {
float* upper = output
+ base
+ (uint64_t)row * 256u
+ 128u
+ column;
float4 zero = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
*reinterpret_cast<float4*>(upper + 0) = zero;
*reinterpret_cast<float4*>(upper + 4) = zero;
*reinterpret_cast<float4*>(upper + 8) = zero;
*reinterpret_cast<float4*>(upper + 12) = zero;
}
}
static __device__ __forceinline__ void xrshape256_trsm_pair(
const float* __restrict__ input,
float* __restrict__ output,
const float* __restrict__ diagonal,
const float* __restrict__ inverse,
uint64_t base,
int row0,
int row1
) {
int lane = threadIdx.x & 31;
const float* src0 = input
+ base + (uint64_t)(128 + row0) * 256u;
const float* src1 = input
+ base + (uint64_t)(128 + row1) * 256u;
float a0 = src0[lane];
float a1 = src0[lane + 32];
float a2 = src0[lane + 64];
float a3 = src0[lane + 96];
float b0 = src1[lane];
float b1 = src1[lane + 32];
float b2 = src1[lane + 64];
float b3 = src1[lane + 96];
#pragma unroll 1
for (int c = 0; c < 128; ++c) {
const float* drow = diagonal
+ (uint64_t)c * XRSHAPE256_STRIDE;
float l0 = drow[lane];
float l1 = drow[lane + 32];
float l2 = drow[lane + 64];
float l3 = drow[lane + 96];
float pa = lane < c ? a0 * l0 : 0.0f;
float pb = lane < c ? b0 * l0 : 0.0f;
if (lane + 32 < c) {
pa = fmaf(a1, l1, pa);
pb = fmaf(b1, l1, pb);
}
if (lane + 64 < c) {
pa = fmaf(a2, l2, pa);
pb = fmaf(b2, l2, pb);
}
if (lane + 96 < c) {
pa = fmaf(a3, l3, pa);
pb = fmaf(b3, l3, pb);
}
pa += __shfl_down_sync(0xffffffffu, pa, 16);
pb += __shfl_down_sync(0xffffffffu, pb, 16);
pa += __shfl_down_sync(0xffffffffu, pa, 8);
pb += __shfl_down_sync(0xffffffffu, pb, 8);
pa += __shfl_down_sync(0xffffffffu, pa, 4);
pb += __shfl_down_sync(0xffffffffu, pb, 4);
pa += __shfl_down_sync(0xffffffffu, pa, 2);
pb += __shfl_down_sync(0xffffffffu, pb, 2);
pa += __shfl_down_sync(0xffffffffu, pa, 1);
pb += __shfl_down_sync(0xffffffffu, pb, 1);
pa = __shfl_sync(0xffffffffu, pa, 0);
pb = __shfl_sync(0xffffffffu, pb, 0);
int owner = c & 31;
int quarter = c >> 5;
float ra = quarter == 0
? __shfl_sync(0xffffffffu, a0, owner)
: quarter == 1
? __shfl_sync(0xffffffffu, a1, owner)
: quarter == 2
? __shfl_sync(0xffffffffu, a2, owner)
: __shfl_sync(0xffffffffu, a3, owner);
float rb = quarter == 0
? __shfl_sync(0xffffffffu, b0, owner)
: quarter == 1
? __shfl_sync(0xffffffffu, b1, owner)
: quarter == 2
? __shfl_sync(0xffffffffu, b2, owner)
: __shfl_sync(0xffffffffu, b3, owner);
float reciprocal = inverse[c];
float va = (ra - pa) * reciprocal;
float vb = (rb - pb) * reciprocal;
if (lane == owner) {
if (quarter == 0) { a0 = va; b0 = vb; }
else if (quarter == 1) { a1 = va; b1 = vb; }
else if (quarter == 2) { a2 = va; b2 = vb; }
else { a3 = va; b3 = vb; }
}
}
float* dst0 = output
+ base + (uint64_t)(128 + row0) * 256u;
float* dst1 = output
+ base + (uint64_t)(128 + row1) * 256u;
dst0[lane] = a0;
dst0[lane + 32] = a1;
dst0[lane + 64] = a2;
dst0[lane + 96] = a3;
dst1[lane] = b0;
dst1[lane + 32] = b1;
dst1[lane + 64] = b2;
dst1[lane + 96] = b3;
}
__global__ __launch_bounds__(1024, 1)
void xrshape256_trsm_b64_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
extern __shared__ __align__(128) float smem[];
float* diagonal = smem;
float* inverse = diagonal + 128 * XRSHAPE256_STRIDE;
int tid = threadIdx.x;
int warp = tid >> 5;
int task = (int)blockIdx.x;
int matrix = task >> 1;
int half = task & 1;
uint64_t base = (uint64_t)matrix << 16;
#pragma unroll
for (int vector = tid; vector < 128 * 32; vector += 1024) {
int row = vector >> 5;
int column = (vector & 31) << 2;
float4 value = *reinterpret_cast<const float4*>(
output + base + (uint64_t)row * 256u + column);
*reinterpret_cast<float4*>(
diagonal + row * XRSHAPE256_STRIDE + column) = value;
}
__syncthreads();
if (tid < 128) {
inverse[tid] = xcal_rcp_approx(
diagonal[tid * XRSHAPE256_STRIDE + tid]);
}
__syncthreads();
int row0 = (half << 6) + warp;
int row1 = row0 + 32;
xrshape256_trsm_pair(
input, output, diagonal, inverse, base, row0, row1);
}
static __device__ __forceinline__ void xrshape256_decode_tile(
int tile,
int& row0,
int& column0
) {
int row_block;
int prefix;
if (tile < 2) { row_block = 0; prefix = 0; }
else if (tile < 6) { row_block = 1; prefix = 2; }
else if (tile < 12) { row_block = 2; prefix = 6; }
else if (tile < 20) { row_block = 3; prefix = 12; }
else if (tile < 30) { row_block = 4; prefix = 20; }
else if (tile < 42) { row_block = 5; prefix = 30; }
else if (tile < 56) { row_block = 6; prefix = 42; }
else { row_block = 7; prefix = 56; }
row0 = row_block << 4;
column0 = (tile - prefix) << 3;
}
static __device__ __forceinline__ void xrshape256_update_tile(
const float* __restrict__ input,
float* __restrict__ output,
const float* __restrict__ l10,
uint64_t base,
int tile
) {
int row0, column0;
xrshape256_decode_tile(tile, row0, column0);
int lane = threadIdx.x & 31;
int group = lane >> 2;
int rank = lane & 3;
int top = row0 + group;
int bottom = top + 8;
int column = column0 + (rank << 1);
int b_row = column0 + group;
const float* a11_top = input
+ base + (uint64_t)(128 + top) * 256u + 128u;
const float* a11_bottom = input
+ base + (uint64_t)(128 + bottom) * 256u + 128u;
float d0 = top >= column ? a11_top[column] : 0.0f;
float d1 = top >= column + 1 ? a11_top[column + 1] : 0.0f;
float d2 = bottom >= column ? a11_bottom[column] : 0.0f;
float d3 = bottom >= column + 1 ? a11_bottom[column + 1] : 0.0f;
#pragma unroll
for (int kk = 0; kk < 128; kk += 16) {
int k0 = kk + (rank << 1);
int k8 = k0 + 8;
const float* top_row = l10
+ top * XRSHAPE256_STRIDE;
const float* bottom_row = l10
+ bottom * XRSHAPE256_STRIDE;
const float* b_values = l10
+ b_row * XRSHAPE256_STRIDE;
float2 at0 = *reinterpret_cast<const float2*>(top_row + k0);
float2 ab0 = *reinterpret_cast<const float2*>(bottom_row + k0);
float2 at8 = *reinterpret_cast<const float2*>(top_row + k8);
float2 ab8 = *reinterpret_cast<const float2*>(bottom_row + k8);
float2 bv0 = *reinterpret_cast<const float2*>(b_values + k0);
float2 bv8 = *reinterpret_cast<const float2*>(b_values + k8);
loco_mma_f16(
loco_f16x2(at0.x, at0.y),
loco_f16x2(ab0.x, ab0.y),
loco_f16x2(at8.x, at8.y),
loco_f16x2(ab8.x, ab8.y),
loco_f16x2(bv0.x, bv0.y) ^ 0x80008000u,
loco_f16x2(bv8.x, bv8.y) ^ 0x80008000u,
d0, d1, d2, d3);
}
float* dst_top = output
+ base + (uint64_t)(128 + top) * 256u + 128u + column;
float* dst_bottom = output
+ base + (uint64_t)(128 + bottom) * 256u + 128u + column;
if (top >= column + 1) {
*reinterpret_cast<float2*>(dst_top) = make_float2(d0, d1);
} else if (top == column) {
dst_top[0] = d0;
}
if (bottom >= column + 1) {
*reinterpret_cast<float2*>(dst_bottom) = make_float2(d2, d3);
} else if (bottom == column) {
dst_bottom[0] = d2;
}
}
__global__ __launch_bounds__(1024, 1)
void xrshape256_update_b64_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
extern __shared__ __align__(128) float l10[];
int tid = threadIdx.x;
int warp = tid >> 5;
int task = (int)blockIdx.x;
int matrix = task >> 1;
int rank = task & 1;
uint64_t base = (uint64_t)matrix << 16;
#pragma unroll
for (int vector = tid; vector < 128 * 32; vector += 1024) {
int row = vector >> 5;
int column = (vector & 31) << 2;
float4 value = *reinterpret_cast<const float4*>(
output
+ base
+ (uint64_t)(128 + row) * 256u
+ column);
*reinterpret_cast<float4*>(
l10 + row * XRSHAPE256_STRIDE + column) = value;
}
__syncthreads();
int tile = rank + (warp << 1);
xrshape256_update_tile(input, output, l10, base, tile);
tile += 64;
if (tile < 72) {
xrshape256_update_tile(input, output, l10, base, tile);
}
}
#undef XRSHAPE128_PIVOTS
#undef XRSHAPE128_HANDOFF
constexpr int LOCO_N32_WARPS = 28;
constexpr int LOCO_N32_THREADS = LOCO_N32_WARPS * 32;
constexpr int LOCO_N64_WARPS = 7;
constexpr int LOCO_N64_THREADS = LOCO_N64_WARPS * 32;
constexpr int LOCO_N64_SMEM =
LOCO_N64_WARPS * 64 * 65 * (int)sizeof(float);
// One warp owns one complete N32 matrix. Every lane retains one source row
// in RMEM, so the entire factorization has no CTA-wide edge and no SMEM.
__global__ __launch_bounds__(LOCO_N32_THREADS, 1)
void loco_potrf32_warp_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int b = (int)blockIdx.x * LOCO_N32_WARPS + warp;
if (b >= batch) return;
uint64_t base = (uint64_t)b * 32u * 32u;
uint64_t row_base = base + (uint64_t)lane * 32u;
float row[32];
#pragma unroll
for (int vector = 0; vector < 8; ++vector) {
float4 packet = *reinterpret_cast<const float4*>(
input + row_base + (vector << 2));
float* word = reinterpret_cast<float*>(&packet);
#pragma unroll
for (int e = 0; e < 4; ++e) {
row[(vector << 2) + e] = word[e];
}
}
constexpr uint32_t mask = 0xffffffffu;
#pragma unroll
for (int pivot = 0; pivot < 32; ++pivot) {
float value = row[pivot];
#pragma unroll
for (int h = 0; h < pivot; ++h) {
float diagonal_row = __shfl_sync(mask, row[h], pivot);
value = fmaf(-row[h], diagonal_row, value);
}
float diagonal = lane == pivot
? xcal_sqrt_approx(fmaxf(value, 1.0e-20f))
: 0.0f;
float reciprocal = lane == pivot
? xcal_rcp_approx(diagonal)
: 0.0f;
diagonal = __shfl_sync(mask, diagonal, pivot);
reciprocal = __shfl_sync(mask, reciprocal, pivot);
if (lane >= pivot) {
row[pivot] = lane == pivot
? diagonal
: value * reciprocal;
}
}
#pragma unroll
for (int vector = 0; vector < 8; ++vector) {
int column = vector << 2;
float4 packet = make_float4(
column + 0 <= lane ? row[column + 0] : 0.0f,
column + 1 <= lane ? row[column + 1] : 0.0f,
column + 2 <= lane ? row[column + 2] : 0.0f,
column + 3 <= lane ? row[column + 3] : 0.0f);
*reinterpret_cast<float4*>(
output + row_base + column) = packet;
}
}
// Seven warp-private N64 tiles make exactly 147 CTAs for the ranked B1024
// board: one CTA per SM on a 148-SM B200 without any cross-warp barrier.
__global__ __launch_bounds__(LOCO_N64_THREADS, 1)
void loco_potrf64_warp_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
extern __shared__ __align__(128) float tiles[];
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int b = (int)blockIdx.x * LOCO_N64_WARPS + warp;
float* tile = tiles + warp * 64 * 65;
constexpr uint32_t mask = 0xffffffffu;
if (b >= batch) return;
uint64_t base = (uint64_t)b * 64u * 64u;
for (int x = lane; x < 64 * 64; x += 32) {
int row = x >> 6;
int column = x & 63;
tile[row * 65 + column] = column <= row
? input[base + x]
: 0.0f;
}
__syncwarp(mask);
int row0 = lane;
int row1 = lane + 32;
#pragma unroll 1
for (int pivot = 0; pivot < 64; ++pivot) {
float value0 = tile[row0 * 65 + pivot];
float value1 = tile[row1 * 65 + pivot];
#pragma unroll 1
for (int h = 0; h < pivot; ++h) {
float diagonal_row = tile[pivot * 65 + h];
if (row0 >= pivot) {
value0 = fmaf(
-tile[row0 * 65 + h], diagonal_row, value0);
}
if (row1 >= pivot) {
value1 = fmaf(
-tile[row1 * 65 + h], diagonal_row, value1);
}
}
float diagonal = 0.0f;
float reciprocal = 0.0f;
int owner = pivot & 31;
if (lane == owner) {
float value = pivot < 32 ? value0 : value1;
diagonal = xcal_sqrt_approx(fmaxf(value, 1.0e-20f));
reciprocal = xcal_rcp_approx(diagonal);
}
diagonal = __shfl_sync(mask, diagonal, owner);
reciprocal = __shfl_sync(mask, reciprocal, owner);
if (row0 >= pivot) {
tile[row0 * 65 + pivot] = row0 == pivot
? diagonal
: value0 * reciprocal;
}
if (row1 >= pivot) {
tile[row1 * 65 + pivot] = row1 == pivot
? diagonal
: value1 * reciprocal;
}
__syncwarp(mask);
}
for (int x = lane; x < 64 * 64; x += 32) {
int row = x >> 6;
int column = x & 63;
output[base + x] = column <= row
? tile[row * 65 + column]
: 0.0f;
}
}
constexpr int LOCO_N64_X3_MATRICES = 7;
constexpr int LOCO_N64_X3_WARPS_PER_MATRIX = 4;
constexpr int LOCO_N64_X3_THREADS =
LOCO_N64_X3_MATRICES * LOCO_N64_X3_WARPS_PER_MATRIX * 32;
constexpr int LOCO_N64_X3_PACKET_WORDS = 2048;
constexpr int LOCO_N64_X3_SMEM =
LOCO_N64_X3_MATRICES
* LOCO_N64_X3_PACKET_WORDS
* (int)sizeof(uint32_t);
struct LocoN64X3Rows {
float f0, f1, f2, f3;
float f4, f5, f6, f7;
float f8, f9, f10, f11;
float f12, f13, f14, f15;
};
template<int Q>
static __device__ __forceinline__ void loco_n64_x3_load_row(
const float* __restrict__ input,
uint64_t base,
int row,
LocoN64X3Rows& r
) {
constexpr int column = Q << 4;
r.f0 = 0.0f; r.f1 = 0.0f; r.f2 = 0.0f; r.f3 = 0.0f;
r.f4 = 0.0f; r.f5 = 0.0f; r.f6 = 0.0f; r.f7 = 0.0f;
r.f8 = 0.0f; r.f9 = 0.0f; r.f10 = 0.0f; r.f11 = 0.0f;
r.f12 = 0.0f; r.f13 = 0.0f; r.f14 = 0.0f; r.f15 = 0.0f;
if (row >= column + 15) {
const float* src =
input + base + (uint64_t)row * 64u + column;
float4 a = *reinterpret_cast<const float4*>(src + 0);
float4 b = *reinterpret_cast<const float4*>(src + 4);
float4 c = *reinterpret_cast<const float4*>(src + 8);
float4 d = *reinterpret_cast<const float4*>(src + 12);
r.f0=a.x; r.f1=a.y; r.f2=a.z; r.f3=a.w;
r.f4=b.x; r.f5=b.y; r.f6=b.z; r.f7=b.w;
r.f8=c.x; r.f9=c.y; r.f10=c.z; r.f11=c.w;
r.f12=d.x; r.f13=d.y; r.f14=d.z; r.f15=d.w;
} else if (row >= column) {
const float* src =
input + base + (uint64_t)row * 64u + column;
if (column + 0 <= row) r.f0 = src[ 0];
if (column + 1 <= row) r.f1 = src[ 1];
if (column + 2 <= row) r.f2 = src[ 2];
if (column + 3 <= row) r.f3 = src[ 3];
if (column + 4 <= row) r.f4 = src[ 4];
if (column + 5 <= row) r.f5 = src[ 5];
if (column + 6 <= row) r.f6 = src[ 6];
if (column + 7 <= row) r.f7 = src[ 7];
if (column + 8 <= row) r.f8 = src[ 8];
if (column + 9 <= row) r.f9 = src[ 9];
if (column + 10 <= row) r.f10 = src[10];
if (column + 11 <= row) r.f11 = src[11];
if (column + 12 <= row) r.f12 = src[12];
if (column + 13 <= row) r.f13 = src[13];
if (column + 14 <= row) r.f14 = src[14];
if (column + 15 <= row) r.f15 = src[15];
}
}
template<int Q>
static __device__ __forceinline__ void loco_n64_x3_store_row(
float* __restrict__ output,
uint64_t base,
int row,
LocoN64X3Rows& r
) {
constexpr int column = Q << 4;
if (row < column + 0) r.f0 = 0.0f;
if (row < column + 1) r.f1 = 0.0f;
if (row < column + 2) r.f2 = 0.0f;
if (row < column + 3) r.f3 = 0.0f;
if (row < column + 4) r.f4 = 0.0f;
if (row < column + 5) r.f5 = 0.0f;
if (row < column + 6) r.f6 = 0.0f;
if (row < column + 7) r.f7 = 0.0f;
if (row < column + 8) r.f8 = 0.0f;
if (row < column + 9) r.f9 = 0.0f;
if (row < column + 10) r.f10 = 0.0f;
if (row < column + 11) r.f11 = 0.0f;
if (row < column + 12) r.f12 = 0.0f;
if (row < column + 13) r.f13 = 0.0f;
if (row < column + 14) r.f14 = 0.0f;
if (row < column + 15) r.f15 = 0.0f;
float* dst = output + base + (uint64_t)row * 64u + column;
*reinterpret_cast<float4*>(dst + 0) =
make_float4(r.f0, r.f1, r.f2, r.f3);
*reinterpret_cast<float4*>(dst + 4) =
make_float4(r.f4, r.f5, r.f6, r.f7);
*reinterpret_cast<float4*>(dst + 8) =
make_float4(r.f8, r.f9, r.f10, r.f11);
*reinterpret_cast<float4*>(dst + 12) =
make_float4(r.f12, r.f13, r.f14, r.f15);
}
template<int Q, int C, int K>
static __device__ __forceinline__ void loco_n64_x3_panel_update(
float& c0,
float& c1,
float k0,
float k1
) {
if constexpr (C > K) {
constexpr int column = (Q << 4) + C;
constexpr int owner = column & 31;
float solved;
if constexpr (column < 32) {
solved = __shfl_sync(0xffffffffu, k0, owner);
} else {
solved = __shfl_sync(0xffffffffu, k1, owner);
}
int lane = threadIdx.x & 31;
if (lane >= column) {
c0 = fmaf(-k0, solved, c0);
}
if (lane + 32 >= column) {
c1 = fmaf(-k1, solved, c1);
}
}
}
template<int Q, int K>
static __device__ __forceinline__ void loco_n64_x3_column(
LocoN64X3Rows& r0,
LocoN64X3Rows& r1,
float& k0,
float& k1
) {
constexpr int diagonal = (Q << 4) + K;
constexpr int owner = diagonal & 31;
int lane = threadIdx.x & 31;
float value;
if constexpr (diagonal < 32) {
value = __shfl_sync(0xffffffffu, k0, owner);
} else {
value = __shfl_sync(0xffffffffu, k1, owner);
}
float root = 0.0f;
float inverse = 0.0f;
if (lane == owner) {
root = xcal_sqrt_approx(fmaxf(value, 1.0e-20f));
inverse = xcal_rcp_approx(root);
}
root = __shfl_sync(0xffffffffu, root, owner);
inverse = __shfl_sync(0xffffffffu, inverse, owner);
if (lane == diagonal) {
k0 = root;
} else if (lane > diagonal) {
k0 *= inverse;
}
if (lane + 32 == diagonal) {
k1 = root;
} else if (lane + 32 > diagonal) {
k1 *= inverse;
}
loco_n64_x3_panel_update<Q, 0,K>(r0.f0, r1.f0, k0, k1);
loco_n64_x3_panel_update<Q, 1,K>(r0.f1, r1.f1, k0, k1);
loco_n64_x3_panel_update<Q, 2,K>(r0.f2, r1.f2, k0, k1);
loco_n64_x3_panel_update<Q, 3,K>(r0.f3, r1.f3, k0, k1);
loco_n64_x3_panel_update<Q, 4,K>(r0.f4, r1.f4, k0, k1);
loco_n64_x3_panel_update<Q, 5,K>(r0.f5, r1.f5, k0, k1);
loco_n64_x3_panel_update<Q, 6,K>(r0.f6, r1.f6, k0, k1);
loco_n64_x3_panel_update<Q, 7,K>(r0.f7, r1.f7, k0, k1);
loco_n64_x3_panel_update<Q, 8,K>(r0.f8, r1.f8, k0, k1);
loco_n64_x3_panel_update<Q, 9,K>(r0.f9, r1.f9, k0, k1);
loco_n64_x3_panel_update<Q,10,K>(r0.f10,r1.f10,k0, k1);
loco_n64_x3_panel_update<Q,11,K>(r0.f11,r1.f11,k0, k1);
loco_n64_x3_panel_update<Q,12,K>(r0.f12,r1.f12,k0, k1);
loco_n64_x3_panel_update<Q,13,K>(r0.f13,r1.f13,k0, k1);
loco_n64_x3_panel_update<Q,14,K>(r0.f14,r1.f14,k0, k1);
loco_n64_x3_panel_update<Q,15,K>(r0.f15,r1.f15,k0, k1);
}
template<int Q>
static __device__ __forceinline__ void loco_n64_x3_factor(
LocoN64X3Rows& r0,
LocoN64X3Rows& r1
) {
loco_n64_x3_column<Q, 0>(r0,r1,r0.f0, r1.f0);
loco_n64_x3_column<Q, 1>(r0,r1,r0.f1, r1.f1);
loco_n64_x3_column<Q, 2>(r0,r1,r0.f2, r1.f2);
loco_n64_x3_column<Q, 3>(r0,r1,r0.f3, r1.f3);
loco_n64_x3_column<Q, 4>(r0,r1,r0.f4, r1.f4);
loco_n64_x3_column<Q, 5>(r0,r1,r0.f5, r1.f5);
loco_n64_x3_column<Q, 6>(r0,r1,r0.f6, r1.f6);
loco_n64_x3_column<Q, 7>(r0,r1,r0.f7, r1.f7);
loco_n64_x3_column<Q, 8>(r0,r1,r0.f8, r1.f8);
loco_n64_x3_column<Q, 9>(r0,r1,r0.f9, r1.f9);
loco_n64_x3_column<Q,10>(r0,r1,r0.f10,r1.f10);
loco_n64_x3_column<Q,11>(r0,r1,r0.f11,r1.f11);
loco_n64_x3_column<Q,12>(r0,r1,r0.f12,r1.f12);
loco_n64_x3_column<Q,13>(r0,r1,r0.f13,r1.f13);
loco_n64_x3_column<Q,14>(r0,r1,r0.f14,r1.f14);
loco_n64_x3_column<Q,15>(r0,r1,r0.f15,r1.f15);
}
static __device__ __forceinline__ void loco_n64_x3_pack_pair(
uint32_t destination,
float x,
float y
) {
asm volatile(
"{\n\t"
".reg .f32 hx, hy, lx, ly;\n\t"
".reg .b32 bf;\n\t"
"lop3.b32 hx, %1, 0xffff0000, 0, 0xc0;\n\t"
"lop3.b32 hy, %2, 0xffff0000, 0, 0xc0;\n\t"
"sub.rn.f32 lx, %1, hx;\n\t"
"sub.rn.f32 ly, %2, hy;\n\t"
"cvt.rn.bf16x2.f32 bf, hy, hx;\n\t"
"st.shared.b32 [%0], bf;\n\t"
"cvt.rn.bf16x2.f32 bf, ly, lx;\n\t"
"st.shared.b32 [%0+4096], bf;\n\t"
"}"
:
: "r"(destination), "f"(x), "f"(y)
: "memory");
}
static __device__ __forceinline__ void loco_n64_x3_pack(
uint32_t packet,
LocoN64X3Rows& r0,
LocoN64X3Rows& r1
) {
uint32_t top = packet + ((threadIdx.x & 31) << 4);
uint32_t bottom = top + 512;
loco_n64_x3_pack_pair(top + 0, r0.f0, r0.f1);
loco_n64_x3_pack_pair(top + 4, r0.f2, r0.f3);
loco_n64_x3_pack_pair(top + 8, r0.f4, r0.f5);
loco_n64_x3_pack_pair(top + 12, r0.f6, r0.f7);
loco_n64_x3_pack_pair(top + 2048, r0.f8, r0.f9);
loco_n64_x3_pack_pair(top + 2052, r0.f10, r0.f11);
loco_n64_x3_pack_pair(top + 2056, r0.f12, r0.f13);
loco_n64_x3_pack_pair(top + 2060, r0.f14, r0.f15);
loco_n64_x3_pack_pair(bottom + 0, r1.f0, r1.f1);
loco_n64_x3_pack_pair(bottom + 4, r1.f2, r1.f3);
loco_n64_x3_pack_pair(bottom + 8, r1.f4, r1.f5);
loco_n64_x3_pack_pair(bottom + 12, r1.f6, r1.f7);
loco_n64_x3_pack_pair(bottom + 2048, r1.f8, r1.f9);
loco_n64_x3_pack_pair(bottom + 2052, r1.f10, r1.f11);
loco_n64_x3_pack_pair(bottom + 2056, r1.f12, r1.f13);
loco_n64_x3_pack_pair(bottom + 2060, r1.f14, r1.f15);
}
template<int M, int N, int Q>
static __device__ __forceinline__ void loco_n64_x3_mma(
float& r0,
float& r1,
float& r2,
float& r3,
float& r4,
float& r5,
float& r6,
float& r7,
uint32_t packet
) {
if constexpr ((M << 4) + 15 < (Q << 4) + (N << 3)) {
return;
}
asm volatile(
"{\n\t"
".reg .pred row8, take;\n\t"
".reg .b32 a0, a1, a2, a3, b0, b1, lane, src;\n\t"
".reg .f32 d0, d1, d2, d3, x, y;\n\t"
"mov.b32 d0, 0;\n\t"
"mov.b32 d1, 0;\n\t"
"mov.b32 d2, 0;\n\t"
"mov.b32 d3, 0;\n\t"
"ld.shared.b32 a0, [%8];\n\t"
"ld.shared.b32 a1, [%8+128];\n\t"
"ld.shared.b32 a2, [%8+2048];\n\t"
"ld.shared.b32 a3, [%8+2176];\n\t"
"xor.b32 a0, a0, 0x80008000;\n\t"
"xor.b32 a1, a1, 0x80008000;\n\t"
"xor.b32 a2, a2, 0x80008000;\n\t"
"xor.b32 a3, a3, 0x80008000;\n\t"
"ld.shared.b32 b0, [%9+4096];\n\t"
"ld.shared.b32 b1, [%9+6144];\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, "
"{d0,d1,d2,d3};\n\t"
"ld.shared.b32 b0, [%9];\n\t"
"ld.shared.b32 b1, [%9+2048];\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, "
"{d0,d1,d2,d3};\n\t"
"ld.shared.b32 a0, [%8+4096];\n\t"
"ld.shared.b32 a1, [%8+4224];\n\t"
"ld.shared.b32 a2, [%8+6144];\n\t"
"ld.shared.b32 a3, [%8+6272];\n\t"
"xor.b32 a0, a0, 0x80008000;\n\t"
"xor.b32 a1, a1, 0x80008000;\n\t"
"xor.b32 a2, a2, 0x80008000;\n\t"
"xor.b32 a3, a3, 0x80008000;\n\t"
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{d0,d1,d2,d3}, {a0,a1,a2,a3}, {b0,b1}, "
"{d0,d1,d2,d3};\n\t"
"mov.u32 lane, %%laneid;\n\t"
"and.b32 src, lane, 16;\n\t"
"setp.eq.u32 take, src, %10;\n\t"
"and.b32 src, lane, 8;\n\t"
"setp.ne.u32 row8, src, 0;\n\t"
"and.b32 src, lane, 7;\n\t"
"shl.b32 src, src, 2;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %0, %0, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %1, %1, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %2, %2, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %3, %3, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %4, %4, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %5, %5, x;\n\t"
"add.u32 src, src, 1;\n\t"
"shfl.sync.idx.b32 x, d0, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d2, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %6, %6, x;\n\t"
"shfl.sync.idx.b32 x, d1, src, 0x1f, 0xffffffff;\n\t"
"shfl.sync.idx.b32 y, d3, src, 0x1f, 0xffffffff;\n\t"
"selp.f32 x, y, x, row8;\n\t"
"@take add.rn.f32 %7, %7, x;\n\t"
"}"
: "+f"(r0), "+f"(r1), "+f"(r2), "+f"(r3),
"+f"(r4), "+f"(r5), "+f"(r6), "+f"(r7)
: "r"(packet + (M << 8) + ((threadIdx.x & 31) << 2)),
"r"(packet + (Q << 8) + (N << 7)
+ ((threadIdx.x & 31) << 2)),
"n"((M & 1) << 4)
: "memory");
}
template<int Q>
static __device__ __forceinline__ void loco_n64_x3_update(
LocoN64X3Rows& r0,
LocoN64X3Rows& r1,
uint32_t packet
) {
loco_n64_x3_mma<0,0,Q>(
r0.f0,r0.f1,r0.f2,r0.f3,r0.f4,r0.f5,r0.f6,r0.f7,
packet);
loco_n64_x3_mma<0,1,Q>(
r0.f8,r0.f9,r0.f10,r0.f11,r0.f12,r0.f13,r0.f14,r0.f15,
packet);
loco_n64_x3_mma<1,0,Q>(
r0.f0,r0.f1,r0.f2,r0.f3,r0.f4,r0.f5,r0.f6,r0.f7,
packet);
loco_n64_x3_mma<1,1,Q>(
r0.f8,r0.f9,r0.f10,r0.f11,r0.f12,r0.f13,r0.f14,r0.f15,
packet);
loco_n64_x3_mma<2,0,Q>(
r1.f0,r1.f1,r1.f2,r1.f3,r1.f4,r1.f5,r1.f6,r1.f7,
packet);
loco_n64_x3_mma<2,1,Q>(
r1.f8,r1.f9,r1.f10,r1.f11,r1.f12,r1.f13,r1.f14,r1.f15,
packet);
loco_n64_x3_mma<3,0,Q>(
r1.f0,r1.f1,r1.f2,r1.f3,r1.f4,r1.f5,r1.f6,r1.f7,
packet);
loco_n64_x3_mma<3,1,Q>(
r1.f8,r1.f9,r1.f10,r1.f11,r1.f12,r1.f13,r1.f14,r1.f15,
packet);
}
static __device__ __forceinline__ void loco_n64_x3_edge(int matrix) {
asm volatile(
"barrier.cta.sync %0, 128;"
:
: "r"(matrix)
: "memory");
}
// B1024 -> CTA147 -> one CTA/SM
//
// CTA896 = 7 matrices x 4 panel warps
// warp = 4*matrix + q, q owns N[q*16:q*16+16]
// lane = rows lane and lane+32, two FP32 N16 register packets
//
// The active q factors only its local K16 history in FP32. Its solved M64xK16
// packet is published once as BF16 high/residual-low; q+1..3 consume
// high*high + high*low + low*high through the native M16N8K16 map. Six
// matrix-private 128-thread edges replace every CTA-wide pivot edge.
__global__ __launch_bounds__(LOCO_N64_X3_THREADS, 1)
void loco_potrf64_bf16x3_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
extern __shared__ __align__(128) uint32_t packets[];
int warp = threadIdx.x >> 5;
int matrix = warp >> 2;
int panel = warp & 3;
int lane = threadIdx.x & 31;
int b = (int)blockIdx.x * LOCO_N64_X3_MATRICES + matrix;
if (b >= batch) return;
uint64_t base = (uint64_t)b * 64u * 64u;
LocoN64X3Rows r0;
LocoN64X3Rows r1;
if (panel == 0) {
loco_n64_x3_load_row<0>(input, base, lane, r0);
loco_n64_x3_load_row<0>(input, base, lane + 32, r1);
} else if (panel == 1) {
loco_n64_x3_load_row<1>(input, base, lane, r0);
loco_n64_x3_load_row<1>(input, base, lane + 32, r1);
} else if (panel == 2) {
loco_n64_x3_load_row<2>(input, base, lane, r0);
loco_n64_x3_load_row<2>(input, base, lane + 32, r1);
} else {
loco_n64_x3_load_row<3>(input, base, lane, r0);
loco_n64_x3_load_row<3>(input, base, lane + 32, r1);
}
uint32_t packet = (uint32_t)__cvta_generic_to_shared(
packets + matrix * LOCO_N64_X3_PACKET_WORDS);
if (panel == 0) {
loco_n64_x3_factor<0>(r0, r1);
loco_n64_x3_pack(packet, r0, r1);
}
loco_n64_x3_edge(matrix);
if (panel == 1) {
loco_n64_x3_update<1>(r0, r1, packet);
} else if (panel == 2) {
loco_n64_x3_update<2>(r0, r1, packet);
} else if (panel == 3) {
loco_n64_x3_update<3>(r0, r1, packet);
}
loco_n64_x3_edge(matrix);
if (panel == 1) {
loco_n64_x3_factor<1>(r0, r1);
loco_n64_x3_pack(packet, r0, r1);
}
loco_n64_x3_edge(matrix);
if (panel == 2) {
loco_n64_x3_update<2>(r0, r1, packet);
} else if (panel == 3) {
loco_n64_x3_update<3>(r0, r1, packet);
}
loco_n64_x3_edge(matrix);
if (panel == 2) {
loco_n64_x3_factor<2>(r0, r1);
loco_n64_x3_pack(packet, r0, r1);
}
loco_n64_x3_edge(matrix);
if (panel == 3) {
loco_n64_x3_update<3>(r0, r1, packet);
}
loco_n64_x3_edge(matrix);
if (panel == 3) {
loco_n64_x3_factor<3>(r0, r1);
}
if (panel == 0) {
loco_n64_x3_store_row<0>(output, base, lane, r0);
loco_n64_x3_store_row<0>(output, base, lane + 32, r1);
} else if (panel == 1) {
loco_n64_x3_store_row<1>(output, base, lane, r0);
loco_n64_x3_store_row<1>(output, base, lane + 32, r1);
} else if (panel == 2) {
loco_n64_x3_store_row<2>(output, base, lane, r0);
loco_n64_x3_store_row<2>(output, base, lane + 32, r1);
} else {
loco_n64_x3_store_row<3>(output, base, lane, r0);
loco_n64_x3_store_row<3>(output, base, lane + 32, r1);
}
}
constexpr int LOCO_N128_K16_THREADS = 16 * 32;
constexpr int LOCO_N128_K16_SMEM =
(128 * 129 / 2 + 128) * (int)sizeof(float);
// One CTA owns one N128 matrix. K16 scalar fronts stay exact FP32; all
// trailing cones use the compensated BF16x3 rail.
__global__ __launch_bounds__(LOCO_N128_K16_THREADS, 2)
void loco_potrf128_k16x3_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
extern __shared__ __align__(128) float smem[];
float* tile = smem;
float* inverse = tile + 128 * 129 / 2;
int tid = threadIdx.x;
int b = (int)blockIdx.x;
if (b >= batch) return;
uint64_t base = (uint64_t)b * 128u * 128u;
for (int vector = tid; vector < 128 * 32; vector += blockDim.x) {
int row = vector >> 5;
int column = (vector & 31) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (column <= row) {
value = loco_ld_global_v4_cs(
input + base + (uint64_t)row * 128u + column);
}
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
tile[loco_lower256(row, column + e)] = word[e];
}
}
}
__syncthreads();
loco_factor_k16<128, 16>(tile, inverse);
for (int vector = tid; vector < 128 * 32; vector += blockDim.x) {
int row = vector >> 5;
int column = (vector & 31) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
word[e] = tile[loco_lower256(row, column + e)];
}
}
*reinterpret_cast<float4*>(
output + base + (uint64_t)row * 128u + column) = value;
}
}
__global__ __launch_bounds__(1024, 1)
void potrf_small_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch,
int n
) {
// n rows / (n/4 warps) = four rows per warp.
// 32/64/128 use CTA256/512/1024; expose 64 ordinary warps/SM.
extern __shared__ float smem[];
float* tile = smem;
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int cta_warps = blockDim.x >> 5;
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
uint64_t base = (uint64_t)b * (uint64_t)n * (uint64_t)n;
for (int x = threadIdx.x; x < n * n; x += blockDim.x) {
int row = x / n;
int col = x - row * n;
tile[row * (n + 1) + col] = row >= col
? input[base + x]
: 0.0f;
}
__syncthreads();
for (int p = 0; p < n; ++p) {
// W0: diagonal -> reciprocal -> whole L[:,p] column.
// One warp-local handoff replaces diag barrier + sparse lane0 wave.
if (warp == 0) {
float reciprocal = 0.0f;
if (lane == 0) {
float diagonal = xcal_sqrt_approx(
fmaxf(tile[p * (n + 1) + p], 1.0e-20f));
tile[p * (n + 1) + p] = diagonal;
reciprocal = xcal_rcp_approx(diagonal);
}
reciprocal = __shfl_sync(
0xffffffffu, reciprocal, 0);
for (int row = p + 1 + lane; row < n; row += 32) {
tile[row * (n + 1) + p] *= reciprocal;
}
}
__syncthreads();
for (int row = p + 1 + warp; row < n; row += cta_warps) {
float lip = tile[row * (n + 1) + p];
for (int col = p + 1 + lane; col <= row; col += 32) {
tile[row * (n + 1) + col] = fmaf(
-lip,
tile[col * (n + 1) + p],
tile[row * (n + 1) + col]);
}
}
__syncthreads();
}
// This is the small-n copy too: every output word is defined.
for (int x = threadIdx.x; x < n * n; x += blockDim.x) {
int row = x / n;
int col = x - row * n;
output[base + x] = row >= col
? tile[row * (n + 1) + col]
: 0.0f;
}
__syncthreads();
}
}
__global__ __launch_bounds__(1024, 1)
void potrf128_kernel(
float* __restrict__ matrix,
int batch,
int n,
int k,
int bs
) {
extern __shared__ float smem[];
float* tile = smem;
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int cta_warps = blockDim.x >> 5;
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
uint64_t base = (uint64_t)b * (uint64_t)n * (uint64_t)n;
for (int x = threadIdx.x; x < bs * bs; x += blockDim.x) {
int row = x / bs;
int col = x - row * bs;
tile[row * 129 + col] = row >= col
? matrix[base + (uint64_t)(k + row) * n + k + col]
: 0.0f;
}
__syncthreads();
for (int p = 0; p < bs; ++p) {
// W0 performs POTRF's scalar edge and coalesces the column wave.
// Barrier count per pivot: 3 -> 2, with no approximation change.
if (warp == 0) {
float reciprocal = 0.0f;
if (lane == 0) {
float diagonal = xcal_sqrt_approx(
fmaxf(tile[p * 129 + p], 1.0e-20f));
tile[p * 129 + p] = diagonal;
reciprocal = xcal_rcp_approx(diagonal);
}
reciprocal = __shfl_sync(
0xffffffffu, reciprocal, 0);
for (int row = p + 1 + lane; row < bs; row += 32) {
tile[row * 129 + p] *= reciprocal;
}
}
__syncthreads();
for (int row = p + 1 + warp; row < bs; row += cta_warps) {
float lip = tile[row * 129 + p];
for (int col = p + 1 + lane; col <= row; col += 32) {
tile[row * 129 + col] = fmaf(
-lip,
tile[col * 129 + p],
tile[row * 129 + col]);
}
}
__syncthreads();
}
for (int x = threadIdx.x; x < bs * bs; x += blockDim.x) {
int row = x / bs;
int col = x - row * bs;
if (row >= col) {
matrix[base + (uint64_t)(k + row) * n + k + col]
= tile[row * 129 + col];
}
}
__syncthreads();
}
}
__device__ __forceinline__ void
loco_publish_hexlift_outer_k128_warp(
uint32_t fp16_0,
uint32_t fp16_1,
uint32_t fp16_2,
uint32_t fp16_3,
uint8_t* __restrict__ packed_values,
uint8_t* __restrict__ packed_scales,
int b,
int n,
int k,
int row,
int first_plane);
__device__ __forceinline__ void loco_bar_sync_64(int barrier_id) {
// One M16 row tile is owned by exactly two adjacent warps. Its Y/X
// exchange is private to that pair, so a 64-thread named barrier replaces
// the old CTA-wide edge. M256 uses IDs 0..15 exactly once per phase.
asm volatile(
"bar.sync %0, 64;"
:: "r"(barrier_id)
: "memory");
}
/*
Ranked one-rail TRSM: eight K16 block solves instead of 128 scalar rounds.
CTA ownership:
MRows/16 row tiles x two N8 warps = MRows/8 compute warps.
The diagonal sidecar is already FP16 and block-native. Each q step is:
Yq = Bq - sum_{p<q} Xp * L[q,p]^T
Xq = Yq * inv(L[q,q])^T
`mma.sync.m16n8k16` performs both products with FP32 accumulators. The
solved X block is published once to FP32 factor state and the row-major
one-rail FP16 history used by the existing outer board.
*/
template<int MRows, int RowBase = 0>
__global__ __launch_bounds__(1024, 1)
void trsm128_mma_f16_kernel(
float* __restrict__ matrix,
uint32_t* __restrict__ packed_fp16,
int batch,
int n,
int k,
int rows
) {
static_assert(
MRows >= 16 && MRows <= 256 && !(MRows & 15),
"unsupported block-TRSM M tile");
static_assert(!(RowBase & 15), "TRSM row base must be M16 aligned");
extern __shared__ __align__(128) uint32_t shared_words[];
uint32_t* diagonal = shared_words; // 128 x 128 FP16
uint32_t* x_stage = diagonal + 128 * 64; // MRows x 128 FP16
uint32_t* y_stage = x_stage + MRows * 64; // MRows x 16 FP16
int tid = (int)threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int group = lane >> 2;
int rank = lane & 3;
constexpr int compute_warps = MRows >> 3;
int row_tiles = (rows - RowBase + MRows - 1) / MRows;
int b = (int)blockIdx.x / row_tiles;
int row_tile = (int)blockIdx.x - b * row_tiles;
if (b >= batch) return;
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
// The POTRF producer stored the hybrid L / inv(Lqq) sidecar row-major.
for (int word = tid; word < 128 * 64; word += blockDim.x) {
int row = word >> 6;
int pair = word & 63;
uint64_t source = (
matrix_base
+ (uint64_t)(k + row) * (uint64_t)n
+ k) >> 1;
diagonal[word] = packed_fp16[source + pair];
}
__syncthreads();
int local_row_base = RowBase + row_tile * MRows;
int global_row_base = k + 128 + local_row_base;
// The diagonal stage is CTA-wide; everything below is M16 pair-local.
// Retire launch padding and a fully OOB pair before its eight MMA rounds.
if (warp >= compute_warps) return;
if (local_row_base + ((warp >> 1) << 4) >= rows) return;
#pragma unroll
for (int q = 0; q < 8; ++q) {
float d0 = 0.0f;
float d1 = 0.0f;
float d2 = 0.0f;
float d3 = 0.0f;
int top = 0;
int bottom = 0;
int n_half = 0;
if (warp < compute_warps) {
int m16 = warp >> 1;
n_half = warp & 1;
top = (m16 << 4) + group;
bottom = top + 8;
int column = (q << 4) + (n_half << 3) + (rank << 1);
int global_top = global_row_base + top;
int global_bottom = global_row_base + bottom;
if (local_row_base + top < rows) {
const float* source = matrix
+ matrix_base
+ (uint64_t)global_top * (uint64_t)n
+ k
+ column;
d0 = source[0];
d1 = source[1];
}
if (local_row_base + bottom < rows) {
const float* source = matrix
+ matrix_base
+ (uint64_t)global_bottom * (uint64_t)n
+ k
+ column;
d2 = source[0];
d3 = source[1];
}
#pragma unroll
for (int p = 0; p < 8; ++p) {
if (p >= q) continue;
int a_pair = (p << 3) + rank;
int b_row = (q << 4) + (n_half << 3) + group;
int b_pair = (p << 3) + rank;
uint32_t a0 = x_stage[top * 64 + a_pair];
uint32_t a1 = x_stage[bottom * 64 + a_pair];
uint32_t a2 = x_stage[top * 64 + a_pair + 4];
uint32_t a3 = x_stage[bottom * 64 + a_pair + 4];
uint32_t b0 = diagonal[b_row * 64 + b_pair];
uint32_t b1 = diagonal[b_row * 64 + b_pair + 4];
a0 ^= 0x80008000u;
a1 ^= 0x80008000u;
a2 ^= 0x80008000u;
a3 ^= 0x80008000u;
loco_mma_f16(
a0, a1, a2, a3, b0, b1,
d0, d1, d2, d3);
}
int y_pair = (n_half << 2) + rank;
y_stage[top * 8 + y_pair] = loco_f16x2(d0, d1);
y_stage[bottom * 8 + y_pair] = loco_f16x2(d2, d3);
}
if (warp < compute_warps) {
loco_bar_sync_64(warp >> 1);
}
if (warp < compute_warps) {
int m16 = warp >> 1;
n_half = warp & 1;
top = (m16 << 4) + group;
bottom = top + 8;
int a_pair = rank;
int b_row = (q << 4) + (n_half << 3) + group;
int b_pair = (q << 3) + rank;
uint32_t a0 = y_stage[top * 8 + a_pair];
uint32_t a1 = y_stage[bottom * 8 + a_pair];
uint32_t a2 = y_stage[top * 8 + a_pair + 4];
uint32_t a3 = y_stage[bottom * 8 + a_pair + 4];
uint32_t b0 = diagonal[b_row * 64 + b_pair];
uint32_t b1 = diagonal[b_row * 64 + b_pair + 4];
float x0 = 0.0f;
float x1 = 0.0f;
float x2 = 0.0f;
float x3 = 0.0f;
loco_mma_f16(
a0, a1, a2, a3, b0, b1,
x0, x1, x2, x3);
int x_pair = (q << 3) + (n_half << 2) + rank;
uint32_t top_word = loco_f16x2(x0, x1);
uint32_t bottom_word = loco_f16x2(x2, x3);
x_stage[top * 64 + x_pair] = top_word;
x_stage[bottom * 64 + x_pair] = bottom_word;
int column = (q << 4) + (n_half << 3) + (rank << 1);
int global_top = global_row_base + top;
int global_bottom = global_row_base + bottom;
if (local_row_base + top < rows) {
float* destination = matrix
+ matrix_base
+ (uint64_t)global_top * (uint64_t)n
+ k
+ column;
destination[0] = x0;
destination[1] = x1;
uint64_t packed_destination = (
matrix_base
+ (uint64_t)global_top * (uint64_t)n
+ k
+ (q << 4)) >> 1;
packed_fp16[
packed_destination + (n_half << 2) + rank] = top_word;
}
if (local_row_base + bottom < rows) {
float* destination = matrix
+ matrix_base
+ (uint64_t)global_bottom * (uint64_t)n
+ k
+ column;
destination[0] = x2;
destination[1] = x3;
uint64_t packed_destination = (
matrix_base
+ (uint64_t)global_bottom * (uint64_t)n
+ k
+ (q << 4)) >> 1;
packed_fp16[
packed_destination + (n_half << 2) + rank] = bottom_word;
}
}
if (warp < compute_warps) {
loco_bar_sync_64(warp >> 1);
}
}
}
template <bool PublishHex>
__global__ __launch_bounds__(1024, 1)
void trsm128_kernel(
float* __restrict__ matrix,
uint32_t* __restrict__ packed_fp16,
uint8_t* __restrict__ hex_values,
uint8_t* __restrict__ hex_scales,
int hex_publish_row_begin,
int hex_publish_plane1_end,
int packed_fp16_atoms,
int batch,
int n,
int k,
int bs,
int row_groups,
int active_warps,
int row_owners
) {
extern __shared__ float smem[];
float* diagonal = smem;
float* inverse = diagonal + 128 * 129;
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
if (row_owners < 1
|| gridDim.x % (unsigned int)row_owners)
return;
int owner_group = (int)blockIdx.x % row_owners;
int first_batch = (int)blockIdx.x / row_owners;
int batch_stride = (int)gridDim.x / row_owners;
for (int b = first_batch; b < batch; b += batch_stride) {
uint64_t base = (uint64_t)b * (uint64_t)n * (uint64_t)n;
for (int x = threadIdx.x; x < bs * bs; x += blockDim.x) {
int row = x / bs;
int col = x - row * bs;
diagonal[row * 129 + col] = matrix[
base + (uint64_t)(k + row) * n + k + col];
}
if (threadIdx.x < bs) {
int c = threadIdx.x;
inverse[c] = xcal_rcp_approx(
matrix[base + (uint64_t)(k + c) * n + k + c]);
}
__syncthreads();
for (int row_group = owner_group;
row_group < row_groups;
row_group += row_owners) {
if (warp < active_warps
&& k + bs + row_group * active_warps + warp < n) {
int row = k + bs + row_group * active_warps + warp;
float x0 = row < n && lane < bs
? matrix[base + (uint64_t)row * n + k + lane]
: 0.0f;
float x1 = row < n && lane + 32 < bs
? matrix[base + (uint64_t)row * n + k + lane + 32]
: 0.0f;
float x2 = row < n && lane + 64 < bs
? matrix[base + (uint64_t)row * n + k + lane + 64]
: 0.0f;
float x3 = row < n && lane + 96 < bs
? matrix[base + (uint64_t)row * n + k + lane + 96]
: 0.0f;
for (int c = 0; c < bs; ++c) {
float partial =
lane < c
? x0 * diagonal[c * 129 + lane]
: 0.0f;
if (lane + 32 < c) {
partial += x1 * diagonal[c * 129 + lane + 32];
}
if (lane + 64 < c) {
partial += x2 * diagonal[c * 129 + lane + 64];
}
if (lane + 96 < c) {
partial += x3 * diagonal[c * 129 + lane + 96];
}
partial += __shfl_down_sync(0xffffffffu, partial, 16);
partial += __shfl_down_sync(0xffffffffu, partial, 8);
partial += __shfl_down_sync(0xffffffffu, partial, 4);
partial += __shfl_down_sync(0xffffffffu, partial, 2);
partial += __shfl_down_sync(0xffffffffu, partial, 1);
float sum = __shfl_sync(0xffffffffu, partial, 0);
int owner = c & 31;
int quarter = c >> 5;
float rhs = quarter == 0
? __shfl_sync(0xffffffffu, x0, owner)
: quarter == 1
? __shfl_sync(0xffffffffu, x1, owner)
: quarter == 2
? __shfl_sync(0xffffffffu, x2, owner)
: __shfl_sync(0xffffffffu, x3, owner);
float value = (rhs - sum) * inverse[c];
if (lane == owner) {
if (quarter == 0) x0 = value;
else if (quarter == 1) x1 = value;
else if (quarter == 2) x2 = value;
else x3 = value;
}
}
// FP32 stays canonical. Tight shapes publish descriptor-native
// hi/lo atoms. Exact B60 publishes dense FP16x1 A/B atoms here,
// while every solved value is already live in the warp.
float y0 = __shfl_down_sync(0xffffffffu, x0, 1);
float y1 = __shfl_down_sync(0xffffffffu, x1, 1);
float y2 = __shfl_down_sync(0xffffffffu, x2, 1);
float y3 = __shfl_down_sync(0xffffffffu, x3, 1);
if (row < n) {
if (lane < bs) {
matrix[base + (uint64_t)row * n + k + lane] = x0;
}
if (lane + 32 < bs) {
matrix[base + (uint64_t)row * n + k + lane + 32] = x1;
}
if (lane + 64 < bs) {
matrix[base + (uint64_t)row * n + k + lane + 64] = x2;
}
if (lane + 96 < bs) {
matrix[base + (uint64_t)row * n + k + lane + 96] = x3;
}
uint32_t fp16_0 = 0u;
uint32_t fp16_1 = 0u;
uint32_t fp16_2 = 0u;
uint32_t fp16_3 = 0u;
bool tma_f16_b60 =
XCAL_XRTS_TMA_F16_B60
&& n == 1024
&& batch == 60;
if (bs == 128 && !(lane & 1)) {
int pair = lane >> 1;
if (packed_fp16_atoms && packed_fp16) {
uint32_t h0, h1, h2, h3;
uint32_t l0, l1, l2, l3;
loco_f16x2_hi_lo(x0, y0, h0, l0);
loco_f16x2_hi_lo(x1, y1, h1, l1);
loco_f16x2_hi_lo(x2, y2, h2, l2);
loco_f16x2_hi_lo(x3, y3, h3, l3);
uint64_t w0 = loco_fp16_packed_word(
b, n, row, k + lane, 0);
uint64_t w1 = loco_fp16_packed_word(
b, n, row, k + lane + 32, 0);
uint64_t w2 = loco_fp16_packed_word(
b, n, row, k + lane + 64, 0);
uint64_t w3 = loco_fp16_packed_word(
b, n, row, k + lane + 96, 0);
packed_fp16[w0] = h0;
packed_fp16[w0 + 32] = l0;
packed_fp16[w1] = h1;
packed_fp16[w1 + 32] = l1;
packed_fp16[w2] = h2;
packed_fp16[w2 + 32] = l2;
packed_fp16[w3] = h3;
packed_fp16[w3 + 32] = l3;
fp16_0 = h0;
fp16_1 = h1;
fp16_2 = h2;
fp16_3 = h3;
} else {
fp16_0 = loco_f16x2(x0, y0);
fp16_1 = loco_f16x2(x1, y1);
fp16_2 = loco_f16x2(x2, y2);
fp16_3 = loco_f16x2(x3, y3);
}
if (tma_f16_b60) {
int packet_word =
((row & 63) << 2)
+ (pair & 3)
+ ((pair >> 2) << 8);
uint64_t packet =
(((((uint64_t)b << 4)
+ (uint64_t)(row >> 6))
* 32u
+ (uint64_t)(k >> 5))
<< 10);
packed_fp16[packet + packet_word] = fp16_0;
packed_fp16[packet + 1024 + packet_word] = fp16_1;
packed_fp16[packet + 2048 + packet_word] = fp16_2;
packed_fp16[packet + 3072 + packet_word] = fp16_3;
int upper =
(row & ~63) + (packet_word & 63);
reinterpret_cast<uint32_t*>(matrix)[
base
+ (uint64_t)(
((k >> 5) << 4)
+ (packet_word >> 6)) * n
+ upper] = fp16_0;
reinterpret_cast<uint32_t*>(matrix)[
base
+ (uint64_t)(
(((k >> 5) + 1) << 4)
+ (packet_word >> 6)) * n
+ upper] = fp16_1;
reinterpret_cast<uint32_t*>(matrix)[
base
+ (uint64_t)(
(((k >> 5) + 2) << 4)
+ (packet_word >> 6)) * n
+ upper] = fp16_2;
reinterpret_cast<uint32_t*>(matrix)[
base
+ (uint64_t)(
(((k >> 5) + 3) << 4)
+ (packet_word >> 6)) * n
+ upper] = fp16_3;
}
if (!packed_fp16_atoms
&& packed_fp16
&& !tma_f16_b60) {
uint64_t packed_base =
(base + (uint64_t)row * n + k) >> 1;
packed_fp16[packed_base + pair] = fp16_0;
packed_fp16[packed_base + 16 + pair] = fp16_1;
packed_fp16[packed_base + 32 + pair] = fp16_2;
packed_fp16[packed_base + 48 + pair] = fp16_3;
}
}
if constexpr (PublishHex) {
if (row >= hex_publish_row_begin && bs == 128) {
loco_publish_hexlift_outer_k128_warp(
fp16_0,
fp16_1,
fp16_2,
fp16_3,
hex_values,
hex_scales,
b,
n,
k,
row,
row < hex_publish_plane1_end
? LOCO_HEXLIFT_OUTER_FIRST_PLANE
: LOCO_HEXLIFT_OUTER_FIRST_PLANE + 1);
}
}
}
}
__syncthreads();
}
}
}
struct __align__(128) LocoMaterializeShared {
// Tight rail: eight K16 packets x 6144 u32 of M128/N256 hi/lo atoms.
// Large rail: sixteen K16 packets x 3072 u32 of M128/N256 high halves.
uint32_t packet[16 * 3072]; // 192 KiB in either interpretation
float scratch[8192]; // one M128 x N64 TMEM export slice
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(LocoMaterializeShared) == 229504,
"dense trailing board must retain its 229504-byte envelope");
template <int Q>
struct __align__(128) LocoB60N256Shared {
// Q M128xN256K16 FP16 packets. After K256 completes, the first
// 32 KiB becomes one M128xN64 coalesced TMEM export slice.
uint32_t packet[Q * 3072];
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(LocoB60N256Shared<8>) == 98432
&& sizeof(LocoB60N256Shared<4>) == 49280,
"B60 compact trailing boards must retain their packet envelopes");
struct __align__(128) LocoPersistent512TcgenShared {
// Fixed n512 board:
// tile packed FP32 A00 / A11 128.5 KiB
// inverse reciprocal diagonal 1.0 KiB
// row_packet[8,K16] rotating FP16 TCGEN operand 32.0 KiB
// col_packet[8,K16] rotating FP16 TCGEN operand 32.0 KiB
//
// Before TCGEN, row_packet is the solved FP16 K64 board and col_packet
// is the matching immutable diagonal board. Factor K16 uses the first
// 8 KiB of row_packet. Every hot MMA load is one conflict-free b32.
float tile[256 * 257 / 2];
float inverse[256];
uint32_t row_packet[8 * 1024];
uint32_t col_packet[8 * 1024];
uint32_t trsm_pad[2048];
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(LocoPersistent512TcgenShared) == 206464,
"persistent n512 TCGEN board must retain its 206464-byte envelope");
constexpr int XRSHAPE512_FACTOR_SMEM =
(256 * 257 / 2 + 256) * (int)sizeof(float);
constexpr int XRSHAPE512_TRSM_SMEM =
(256 * 257 / 2 + 256 + 32 * 72) * (int)sizeof(float);
struct __align__(128) XrShape512UpdateShared {
// One quadrant owns one N128 TCGEN D board. The crossed quadrant needs
// both K16 operands; diagonal quadrants alias A/B to the first packet.
uint32_t row_packet[8 * 1024];
uint32_t col_packet[8 * 1024];
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(XrShape512UpdateShared) <= 66 * 1024,
"ranked n512 update board exceeds 66 KiB");
struct __align__(128) LocoNvfp4OuterShared {
// Four descriptor-native K64 packet pairs form one K256 issue group.
// The four UE4M3 bytes per row are the K16 block-absmax scales.
uint8_t packet_a[4][4096];
uint8_t packet_b[4][8192];
uint8_t scale_a[4][512];
uint8_t scale_b[4][1024];
float scratch[8192]; // one M128 x N64 TMEM export slice
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(LocoNvfp4OuterShared) <= 88 * 1024,
"one-rail NVFP4 outer board exceeds 88 KiB");
// Outer FP32 rail:
// x ~= S*q/64
// q = d0 + 4*d1 + 16*d2 + 64*d3
// One s2f6 rail becomes four E2M1 digit planes.
constexpr int LOCO_OUTER_S2F6_RAILS = 1;
constexpr int LOCO_OUTER_DIGITS_PER_RAIL = 4;
constexpr int LOCO_OUTER_ACTIVE_PLANES =
LOCO_OUTER_S2F6_RAILS * LOCO_OUTER_DIGITS_PER_RAIL;
struct __align__(128) LocoFrontierE2M1x4Shared {
// One K64 history keeps all four A and four B planes resident.
// A is M128; B is N128 padded to the native N256 instruction.
uint8_t packet_a[LOCO_OUTER_ACTIVE_PLANES * 4096];
uint8_t packet_b[LOCO_OUTER_ACTIVE_PLANES * 8192];
uint8_t scale_a[LOCO_OUTER_ACTIVE_PLANES * 512];
uint8_t scale_b[LOCO_OUTER_ACTIVE_PLANES * 1024];
float scratch[8192];
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(LocoFrontierE2M1x4Shared) <= 112 * 1024,
"E2M1x4 sibling frontier exceeds 112 KiB");
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200
__device__ __forceinline__ uint64_t loco_operand_desc(uint32_t smem_addr) {
// K-major no-swizzle: TF32 M/N64 x K8 or F16/BF16 M/N64 x K16.
// Encoded byte offsets: LBO=1024 B -> 64, SBO=128 B -> 8.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (64ull << 16)
| (8ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_operand_desc_n256(
uint32_t smem_addr
) {
// K-major no-swizzle: TF32 N256 x K8 or F16/BF16 N256 x K16.
// Encoded byte offsets: LBO=4096 B -> 256, SBO=128 B -> 8.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (256ull << 16)
| (8ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_operand_desc_m128(
uint32_t smem_addr
) {
// K-major no-swizzle: F16 M128 x K16.
// Encoded byte offsets: LBO=2048 B -> 128, SBO=128 B -> 8.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (128ull << 16)
| (8ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_operand_desc_packed_a(
uint32_t smem_addr
) {
// M128: K8 banks are 4096 B apart; 8-row atoms are 256 B apart.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (256ull << 16)
| (16ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_operand_desc_packed_b(
uint32_t smem_addr
) {
// N256: K8 banks are 8192 B apart; 8-row atoms are 256 B apart.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (512ull << 16)
| (16ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_plane_desc_a(
uint32_t smem_addr
) {
// M128 x K64 E2M1: K32 chunks are 2048 B apart.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (128ull << 16)
| (8ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_plane_desc_b(
uint32_t smem_addr
) {
// N256 x K64 E2M1: K32 chunks are 4096 B apart.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| (256ull << 16)
| (8ull << 32)
| (1ull << 46);
}
__device__ __forceinline__ uint64_t loco_inner_plane_desc_b(
uint32_t smem_addr,
int valid_n
) {
// Dense live-N K64 E2M1: K32 chunks are valid_n*16 B apart.
// Descriptor fields are therefore LBO=valid_n and SBO=8.
return ((uint64_t)(smem_addr & 0x3ffffu) >> 4)
| ((uint64_t)valid_n << 16)
| (8ull << 32)
| (1ull << 46);
}
__host__ __device__ __forceinline__ uint64_t loco_plane_scale_offset(
int b,
int history,
int history_panels,
int trailing_rows,
int row
) {
return ((((uint64_t)b * history_panels + history)
* trailing_rows + row) * 4u);
}
__host__ __device__ __forceinline__ uint64_t
loco_outer_plane_value_offset(
int b,
int history,
int history_panels,
int plane,
int chunk,
int trailing_rows,
int row
) {
return ((((((uint64_t)b * history_panels + history)
* LOCO_OUTER_ACTIVE_PLANES + plane) * 2u + chunk)
* trailing_rows + row) * 16u);
}
__host__ __device__ __forceinline__ uint64_t
loco_hexlift_outer_value_offset(
int b,
int n,
int history,
int plane,
int chunk,
int row
) {
uint64_t history_panels = (uint64_t)n >> 6;
return ((((((uint64_t)b * history_panels + (uint64_t)history)
* LOCO_HEXLIFT_OUTER_ACTIVE_PLANES
+ (uint64_t)(plane - LOCO_HEXLIFT_OUTER_FIRST_PLANE))
* 2u + (uint64_t)chunk)
* (uint64_t)n + (uint64_t)row) * 16u);
}
__host__ __device__ __forceinline__ uint64_t
loco_hexlift_outer_scale_offset(
int b,
int n,
int history,
int row
) {
uint64_t history_panels = (uint64_t)n >> 6;
return ((((uint64_t)b * history_panels + (uint64_t)history)
* (uint64_t)n + (uint64_t)row) * 4u);
}
__host__ __device__ __forceinline__ uint64_t
loco_nvfp4_value_offset(
int b,
int n,
int history,
int history_panels,
int chunk,
int trailing_rows,
int row
) {
uint64_t fp16_batch_bytes =
(uint64_t)n * (uint64_t)n * 2u;
return (uint64_t)b * fp16_batch_bytes
+ ((((uint64_t)history * 2u + chunk)
* trailing_rows + row) * 16u);
}
__host__ __device__ __forceinline__ uint64_t
loco_nvfp4_scale_offset(
int b,
int history,
int history_panels,
int trailing_rows,
int row
) {
return ((((uint64_t)b * history_panels + history)
* trailing_rows + row) * 4u);
}
__device__ __forceinline__ uint32_t loco_nvfp4_pack_pair(
float x0,
float x1,
float inverse_scale
) {
float low = x0 * inverse_scale;
float high = x1 * inverse_scale;
uint32_t packed;
asm volatile(
"{\n\t"
".reg .b8 q;\n\t"
"cvt.rn.satfinite.e2m1x2.f32 q, %2, %1;\n\t"
"cvt.u32.u8 %0, q;\n\t"
"}"
: "=r"(packed)
: "f"(low), "f"(high));
return packed;
}
__device__ __forceinline__ float loco_nvfp4_ue4m3(
uint32_t code
) {
uint32_t exponent = (code >> 3) & 15u;
uint32_t mantissa = code & 7u;
if (exponent == 0u) {
return (float)mantissa * 0.001953125f;
}
uint32_t f32_exponent =
exponent == 15u ? 135u : exponent + 120u;
return __uint_as_float(
(f32_exponent << 23) | (mantissa << 20));
}
__device__ __forceinline__ float loco_nvfp4_absmax_scale(
float maximum,
uint8_t& encoded
) {
if (!(maximum > 0.0f)) {
encoded = 0u;
return 0.0f;
}
float target = maximum * 0.1666666716337204f;
float zero = 0.0f;
uint16_t pair;
asm volatile(
"cvt.rn.satfinite.e4m3x2.f32 %0, %1, %2;"
: "=h"(pair)
: "f"(zero), "f"(target));
uint32_t code = (uint32_t)pair & 0xffu;
if (code > 126u) code = 126u;
float scale = loco_nvfp4_ue4m3(code);
if (scale < target && code < 126u) {
++code;
scale = loco_nvfp4_ue4m3(code);
}
encoded = (uint8_t)code;
return scale;
}
__host__ __device__ __forceinline__ int loco_outer_s2f6_digit(
int q,
int digit
) {
uint32_t raw = (uint32_t)(uint8_t)q;
int value = (int)((raw >> (digit << 1)) & 3u);
if (digit == 3 && q < 0) value -= 4;
return value;
}
__device__ __forceinline__ uint32_t loco_outer_pack_digits(
int r0,
int r1
) {
float low = (float)r0;
float high = (float)r1;
uint32_t packed;
asm volatile(
"{\n\t"
".reg .b8 q;\n\t"
"cvt.rn.satfinite.e2m1x2.f32 q, %2, %1;\n\t"
"cvt.u32.u8 %0, q;\n\t"
"}"
: "=r"(packed)
: "f"(low), "f"(high));
return packed;
}
__device__ __forceinline__ float loco_outer_s2f6_scale(
float maximum,
uint8_t& encoded
) {
if (!(maximum > 0.0f)) {
encoded = 0u;
return 0.0f;
}
float target = maximum * (64.0f / 127.0f);
float zero = 0.0f;
uint16_t pair;
asm volatile(
"cvt.rp.satfinite.ue8m0x2.f32 %0, %1, %2;"
: "=h"(pair)
: "f"(zero), "f"(target));
uint32_t code = (uint32_t)pair & 0xffu;
if (code == 255u) code = 254u;
if (code < 13u) code = 13u;
encoded = (uint8_t)code;
return ldexpf(1.0f, (int)code - 127);
}
__device__ __forceinline__ uint16_t loco_outer_s2f6x2(
float x0,
float x1,
uint8_t scale_code
) {
uint16_t packed;
uint16_t scale_pair =
(uint16_t)scale_code | ((uint16_t)scale_code << 8);
asm volatile(
"cvt.rn.satfinite.scaled::n2::ue8m0.s2f6x2.f32 "
"%0, %2, %1, %3;"
: "=h"(packed)
: "f"(x0), "f"(x1), "h"(scale_pair));
return packed;
}
__device__ __forceinline__ uint8_t loco_outer_scale_code(
uint8_t base_code,
int plane
) {
if (base_code == 0u) return 0u;
int rail = plane / LOCO_OUTER_DIGITS_PER_RAIL;
int digit = plane - rail * LOCO_OUTER_DIGITS_PER_RAIL;
int code = (int)base_code - rail * 7 + (digit << 1) - 6;
if (code < 0) code = 0;
if (code > 254) code = 254;
return (uint8_t)code;
}
__device__ __forceinline__ uint32_t loco_outer_scale_word(
uint32_t base_word,
int plane
) {
uint32_t result = 0u;
#pragma unroll
for (int i = 0; i < 4; ++i) {
uint8_t base_code = (uint8_t)(base_word >> (i * 8));
result |= (uint32_t)loco_outer_scale_code(base_code, plane)
<< (i * 8);
}
return result;
}
__device__ __forceinline__ uint32_t loco_tmem_reader_addr(
uint32_t allocation_base,
int reader_warp
) {
return allocation_base | ((uint32_t)reader_warp << 21);
}
__device__ __forceinline__ void loco_mbar_init(uint32_t addr) {
asm volatile(
"mbarrier.init.shared::cta.b64 [%0], 1;"
:: "r"(addr)
: "memory");
}
__device__ __forceinline__ void loco_mbar_wait(
uint32_t addr,
uint32_t phase
) {
uint32_t done;
do {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
".reg .b32 ticks;\n\t"
"mov.b32 ticks, 0x989680;\n\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 "
"p, [%1], %2, ticks;\n\t"
"selp.u32 %0, 1, 0, p;\n\t"
"}"
: "=r"(done)
: "r"(addr), "r"(phase)
: "memory");
} while (!done);
}
__device__ __forceinline__ uint32_t loco_elect_one() {
uint32_t elected = 0;
uint32_t membermask = 0xffffffffu;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync _|p, %1;\n\t"
"@p mov.b32 %0, 1;\n\t"
"}"
: "+r"(elected)
: "r"(membermask));
return elected;
}
// src_addr points at logical [16*w + lane/4, lane%4].
__device__ __forceinline__ void loco_tmem_import(
uint32_t taddr,
uint32_t src_addr
) {
asm volatile(
"{\n\t"
".reg .b32 r<32>;\n\t"
"ld.shared.b32 r0, [%1+0];\n\t"
"ld.shared.b32 r1, [%1+2048];\n\t"
"ld.shared.b32 r2, [%1+16];\n\t"
"ld.shared.b32 r3, [%1+2064];\n\t"
"ld.shared.b32 r4, [%1+32];\n\t"
"ld.shared.b32 r5, [%1+2080];\n\t"
"ld.shared.b32 r6, [%1+48];\n\t"
"ld.shared.b32 r7, [%1+2096];\n\t"
"ld.shared.b32 r8, [%1+64];\n\t"
"ld.shared.b32 r9, [%1+2112];\n\t"
"ld.shared.b32 r10, [%1+80];\n\t"
"ld.shared.b32 r11, [%1+2128];\n\t"
"ld.shared.b32 r12, [%1+96];\n\t"
"ld.shared.b32 r13, [%1+2144];\n\t"
"ld.shared.b32 r14, [%1+112];\n\t"
"ld.shared.b32 r15, [%1+2160];\n\t"
"ld.shared.b32 r16, [%1+128];\n\t"
"ld.shared.b32 r17, [%1+2176];\n\t"
"ld.shared.b32 r18, [%1+144];\n\t"
"ld.shared.b32 r19, [%1+2192];\n\t"
"ld.shared.b32 r20, [%1+160];\n\t"
"ld.shared.b32 r21, [%1+2208];\n\t"
"ld.shared.b32 r22, [%1+176];\n\t"
"ld.shared.b32 r23, [%1+2224];\n\t"
"ld.shared.b32 r24, [%1+192];\n\t"
"ld.shared.b32 r25, [%1+2240];\n\t"
"ld.shared.b32 r26, [%1+208];\n\t"
"ld.shared.b32 r27, [%1+2256];\n\t"
"ld.shared.b32 r28, [%1+224];\n\t"
"ld.shared.b32 r29, [%1+2272];\n\t"
"ld.shared.b32 r30, [%1+240];\n\t"
"ld.shared.b32 r31, [%1+2288];\n\t"
"tcgen05.st.sync.aligned.16x128b.x16.b32 "
"[%0], {r0,r1,r2,r3,r4,r5,r6,r7,"
"r8,r9,r10,r11,r12,r13,r14,r15,"
"r16,r17,r18,r19,r20,r21,r22,r23,"
"r24,r25,r26,r27,r28,r29,r30,r31};\n\t"
"tcgen05.wait::st.sync.aligned;\n\t"
"tcgen05.fence::before_thread_sync;\n\t"
"}"
:
: "r"(taddr), "r"(src_addr)
: "memory");
}
// source map: row = 16*w + lane/4 + 8*(lane&1), col = 4*((lane>>1)&1)
// transpose: four source lanes -> {row+0,row+8} x {col+0,col+4}
// store: .16x128b.x2 x eight N8 slices = one M16 x N64 quarter
__device__ __forceinline__ void loco_tmem_import_global_n8(
uint32_t taddr,
const float* src_addr,
uint32_t lane1,
uint32_t lane2
) {
asm volatile(
"{\n\t"
".reg .pred p1, p2;\n\t"
".reg .b32 x<4>, u<2>, t<2>, y<4>, r<4>;\n\t"
"ld.global.cs.nc.v4.u32 {x0,x1,x2,x3}, [%1];\n\t"
"setp.ne.u32 p1, %2, 0;\n\t"
"setp.ne.u32 p2, %3, 0;\n\t"
"selp.b32 u0, x0, x1, p1;\n\t"
"selp.b32 u1, x2, x3, p1;\n\t"
"shfl.sync.bfly.b32 t0, u0, 1, 0x1c03, 0xffffffff;\n\t"
"shfl.sync.bfly.b32 t1, u1, 1, 0x1c03, 0xffffffff;\n\t"
"selp.b32 y0, t0, x0, p1;\n\t"
"selp.b32 y1, x1, t0, p1;\n\t"
"selp.b32 y2, t1, x2, p1;\n\t"
"selp.b32 y3, x3, t1, p1;\n\t"
"selp.b32 u0, y0, y2, p2;\n\t"
"selp.b32 u1, y1, y3, p2;\n\t"
"shfl.sync.bfly.b32 t0, u0, 2, 0x1c03, 0xffffffff;\n\t"
"shfl.sync.bfly.b32 t1, u1, 2, 0x1c03, 0xffffffff;\n\t"
"selp.b32 r0, t0, y0, p2;\n\t"
"selp.b32 r1, t1, y1, p2;\n\t"
"selp.b32 r2, y2, t0, p2;\n\t"
"selp.b32 r3, y3, t1, p2;\n\t"
"tcgen05.st.sync.aligned.16x128b.x2.b32 "
"[%0], {r0,r1,r2,r3};\n\t"
"}"
:
: "r"(taddr),
"l"(src_addr),
"r"(lane1),
"r"(lane2)
: "memory");
}
__device__ __forceinline__ void loco_tmem_import_global(
uint32_t taddr,
const float* src_addr
) {
const uint32_t lane1 = threadIdx.x & 1;
const uint32_t lane2 = threadIdx.x & 2;
loco_tmem_import_global_n8(
taddr, src_addr, lane1, lane2);
#pragma unroll
for (int n8 = 1; n8 < 8; ++n8) {
loco_tmem_import_global_n8(
taddr + (n8 << 3),
src_addr + (n8 << 3),
lane1,
lane2);
}
asm volatile(
"tcgen05.wait::st.sync.aligned;\n\t"
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
// dst_addr points at logical [16*w + lane/4, lane%4].
__device__ __forceinline__ void loco_tmem_export(
uint32_t taddr,
uint32_t dst_addr
) {
asm volatile(
"{\n\t"
".reg .b32 r<32>;\n\t"
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.ld.sync.aligned.16x128b.x16.b32 "
"{r0,r1,r2,r3,r4,r5,r6,r7,"
"r8,r9,r10,r11,r12,r13,r14,r15,"
"r16,r17,r18,r19,r20,r21,r22,r23,"
"r24,r25,r26,r27,r28,r29,r30,r31}, [%0];\n\t"
"tcgen05.wait::ld.sync.aligned;\n\t"
"st.shared.b32 [%1+0], r0;\n\t"
"st.shared.b32 [%1+2048], r1;\n\t"
"st.shared.b32 [%1+16], r2;\n\t"
"st.shared.b32 [%1+2064], r3;\n\t"
"st.shared.b32 [%1+32], r4;\n\t"
"st.shared.b32 [%1+2080], r5;\n\t"
"st.shared.b32 [%1+48], r6;\n\t"
"st.shared.b32 [%1+2096], r7;\n\t"
"st.shared.b32 [%1+64], r8;\n\t"
"st.shared.b32 [%1+2112], r9;\n\t"
"st.shared.b32 [%1+80], r10;\n\t"
"st.shared.b32 [%1+2128], r11;\n\t"
"st.shared.b32 [%1+96], r12;\n\t"
"st.shared.b32 [%1+2144], r13;\n\t"
"st.shared.b32 [%1+112], r14;\n\t"
"st.shared.b32 [%1+2160], r15;\n\t"
"st.shared.b32 [%1+128], r16;\n\t"
"st.shared.b32 [%1+2176], r17;\n\t"
"st.shared.b32 [%1+144], r18;\n\t"
"st.shared.b32 [%1+2192], r19;\n\t"
"st.shared.b32 [%1+160], r20;\n\t"
"st.shared.b32 [%1+2208], r21;\n\t"
"st.shared.b32 [%1+176], r22;\n\t"
"st.shared.b32 [%1+2224], r23;\n\t"
"st.shared.b32 [%1+192], r24;\n\t"
"st.shared.b32 [%1+2240], r25;\n\t"
"st.shared.b32 [%1+208], r26;\n\t"
"st.shared.b32 [%1+2256], r27;\n\t"
"st.shared.b32 [%1+224], r28;\n\t"
"st.shared.b32 [%1+2272], r29;\n\t"
"st.shared.b32 [%1+240], r30;\n\t"
"st.shared.b32 [%1+2288], r31;\n\t"
"tcgen05.fence::before_thread_sync;\n\t"
"}"
:
: "r"(taddr), "r"(dst_addr)
: "memory");
}
__host__ __device__ __forceinline__ uint64_t
loco_m128n256_task_prefix(int row_pairs) {
// Row-pair r owns floor(r/2)+1 lower-triangular N256 groups.
int pairs = row_pairs >> 1;
int remainder = row_pairs & 1;
return (uint64_t)pairs * (pairs + 1)
+ (uint64_t)remainder * (pairs + 1);
}
__device__ __forceinline__ void loco_decode_m128n256_task(
uint64_t task,
uint64_t grouped_tasks,
int trailing_panels,
int frontier_only,
int frontier_segments,
int& b,
int& row_local,
int& col_local,
int& valid_m_panels,
int& valid_segments
) {
b = (int)(task / grouped_tasks);
uint64_t local = task - (uint64_t)b * grouped_tasks;
int row_pairs = (trailing_panels + 1) >> 1;
int row_pair;
if (frontier_only) {
row_pair = (int)local;
col_local = 0;
} else {
int low = 0;
int high = row_pairs - 1;
while (low < high) {
int middle = (low + high + 1) >> 1;
if (loco_m128n256_task_prefix(middle) <= local) {
low = middle;
} else {
high = middle - 1;
}
}
row_pair = low;
uint64_t row_prefix =
loco_m128n256_task_prefix(row_pair);
int column_group = (int)(local - row_prefix);
col_local = column_group << 2;
}
row_local = row_pair << 1;
valid_m_panels = trailing_panels - row_local;
if (valid_m_panels > 2) valid_m_panels = 2;
// The fused tile carries the union of the two M64 rows. The final lower
// mask suppresses cells above either row's diagonal.
valid_segments =
row_local + valid_m_panels - col_local;
int segment_cap = frontier_only ? frontier_segments : 4;
if (valid_segments > segment_cap) {
valid_segments = segment_cap;
}
if (col_local + valid_segments > trailing_panels) {
valid_segments = trailing_panels - col_local;
}
}
template <bool ResidualN128, bool SplitFrontier = false>
__device__ __forceinline__ void loco_decode_trailing_task(
uint64_t task,
uint64_t grouped_tasks,
int trailing_panels,
int frontier_only,
int frontier_segments,
int& b,
int& row_local,
int& col_local,
int& valid_m_panels,
int& valid_segments
) {
if constexpr (!ResidualN128) {
loco_decode_m128n256_task(
task,
grouped_tasks,
trailing_panels,
frontier_only,
frontier_segments,
b,
row_local,
col_local,
valid_m_panels,
valid_segments);
} else {
b = (int)(task / grouped_tasks);
uint64_t local =
task - (uint64_t)b * grouped_tasks;
if constexpr (SplitFrontier) {
// Split one N256 dependency frontier into exact N128 owners:
//
// row-pair 0 : col-half 0
// row-pair r : col-half 0, 1 for r > 0
//
// The missing first upper half is never scheduled.
if (!local) {
row_local = 0;
col_local = 0;
} else {
--local;
row_local = ((int)(local >> 1) + 1) << 1;
col_local = ((int)local & 1) << 1;
}
valid_m_panels = 2;
valid_segments = 2;
} else {
int low = 0;
int high = (trailing_panels >> 1) - 1;
while (low < high) {
int middle = (low + high + 1) >> 1;
if ((uint64_t)middle * (middle + 1) / 2
<= local) {
low = middle;
} else {
high = middle - 1;
}
}
row_local = low << 1;
col_local = (int)(
local - (uint64_t)low * (low + 1) / 2) << 1;
valid_m_panels = 2;
valid_segments = 2;
}
}
}
__device__ __forceinline__ void loco_fp16_update_m128n256(
uint32_t taddr,
uint32_t packet_base,
uint32_t input_d
) {
// D:F32 | A:F16 | B:F16 | -A | M128 | N256 | K16.
constexpr uint32_t m128n256k16_negate_a = 0x08402010u;
uint64_t a_desc = loco_operand_desc_m128(packet_base);
uint64_t b_desc = loco_operand_desc_n256(packet_base + 4096);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
".reg .b32 zero<4>;\n\t"
"mov.b32 zero0, 0;\n\t"
"mov.b32 zero1, 0;\n\t"
"mov.b32 zero2, 0;\n\t"
"mov.b32 zero3, 0;\n\t"
"setp.ne.u32 accumulate, %4, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %2, %3, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(m128n256k16_negate_a),
"r"(input_d)
: "memory");
}
__device__ __forceinline__ void loco_fp16_update_m128n128(
uint32_t taddr,
uint32_t packet_base,
uint32_t input_d
) {
// D:F32 | A:F16 | B:F16 | -A | M128 | N128 | K16.
constexpr uint32_t m128n128k16_negate_a = 0x08202010u;
uint64_t a_desc = loco_operand_desc_m128(packet_base);
uint64_t b_desc = loco_operand_desc_m128(packet_base + 4096);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
".reg .b32 zero<4>;\n\t"
"mov.b32 zero0, 0;\n\t"
"mov.b32 zero1, 0;\n\t"
"mov.b32 zero2, 0;\n\t"
"mov.b32 zero3, 0;\n\t"
"setp.ne.u32 accumulate, %4, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %2, %3, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(m128n128k16_negate_a),
"r"(input_d)
: "memory");
}
__device__ __forceinline__ void loco_fp16_update_m128n128_desc(
uint32_t taddr,
uint64_t a_desc,
uint64_t b_desc,
uint32_t input_d
) {
// D:F32 | A:F16 | B:F16 | -A | M128 | N128 | K16.
constexpr uint32_t m128n128k16_negate_a = 0x08202010u;
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
".reg .b32 zero<4>;\n\t"
"mov.b32 zero0, 0;\n\t"
"mov.b32 zero1, 0;\n\t"
"mov.b32 zero2, 0;\n\t"
"mov.b32 zero3, 0;\n\t"
"setp.ne.u32 accumulate, %4, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %2, %3, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(m128n128k16_negate_a),
"r"(input_d)
: "memory");
}
__device__ __forceinline__ void loco_tmem_load16(
uint32_t taddr,
float& r0,
float& r1,
float& r2,
float& r3,
float& r4,
float& r5,
float& r6,
float& r7,
float& r8,
float& r9,
float& r10,
float& r11,
float& r12,
float& r13,
float& r14,
float& r15
) {
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x16.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,"
"%8,%9,%10,%11,%12,%13,%14,%15}, [%16];\n\t"
"tcgen05.wait::ld.sync.aligned;"
: "=f"(r0), "=f"(r1), "=f"(r2), "=f"(r3),
"=f"(r4), "=f"(r5), "=f"(r6), "=f"(r7),
"=f"(r8), "=f"(r9), "=f"(r10), "=f"(r11),
"=f"(r12), "=f"(r13), "=f"(r14), "=f"(r15)
: "r"(taddr)
: "memory");
}
__device__ __forceinline__ void loco_fp16x3_update_m128n256(
uint32_t taddr,
uint32_t packet_base
) {
// D -= Ah*Bh + Ah*Bl + Al*Bh. The omitted Al*Bl term is second order.
constexpr uint32_t m128n256k16_negate_a = 0x08402010u;
uint64_t ah_desc =
loco_operand_desc_packed_a(packet_base);
uint64_t al_desc =
loco_operand_desc_packed_a(packet_base + 128);
uint64_t bh_desc =
loco_operand_desc_packed_b(packet_base + 8192);
uint64_t bl_desc =
loco_operand_desc_packed_b(packet_base + 8320);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
".reg .b32 zero<4>;\n\t"
"mov.b32 zero0, 0;\n\t"
"mov.b32 zero1, 0;\n\t"
"mov.b32 zero2, 0;\n\t"
"mov.b32 zero3, 0;\n\t"
"setp.ne.u32 accumulate, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %2, %5, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %4, %5, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %3, %2, %5, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(ah_desc),
"l"(bh_desc),
"l"(al_desc),
"l"(bl_desc),
"r"(m128n256k16_negate_a),
"r"(1u)
: "memory");
}
__device__ __forceinline__ void loco_fp16x3_update_m128n128(
uint32_t taddr,
uint32_t packet_base
) {
// Both M128 and N128 use the same packed K8-bank geometry.
constexpr uint32_t m128n128k16_negate_a = 0x08202010u;
uint64_t ah_desc =
loco_operand_desc_packed_a(packet_base);
uint64_t al_desc =
loco_operand_desc_packed_a(packet_base + 128);
uint64_t bh_desc =
loco_operand_desc_packed_a(packet_base + 8192);
uint64_t bl_desc =
loco_operand_desc_packed_a(packet_base + 8320);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
".reg .b32 zero<4>;\n\t"
"mov.b32 zero0, 0;\n\t"
"mov.b32 zero1, 0;\n\t"
"mov.b32 zero2, 0;\n\t"
"mov.b32 zero3, 0;\n\t"
"setp.ne.u32 accumulate, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %2, %5, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %4, %5, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %3, %2, %5, "
"{zero0,zero1,zero2,zero3}, accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(ah_desc),
"l"(bh_desc),
"l"(al_desc),
"l"(bl_desc),
"r"(m128n128k16_negate_a),
"r"(1u)
: "memory");
}
struct __align__(128) XrtsF16Shared {
// A is stationary for the whole lower-row sweep. Each dense M64 packet
// occupies the top half of an M128 descriptor; the bottom half is zero.
// B is one eight-issue TMA stage and is recycled only after mma_done.
// D is one row-major 16 KiB residual tile.
uint32_t a[16][1024];
uint32_t b[64 * 64];
uint32_t d[64 * 64];
uint64_t tma_done;
uint64_t mma_done;
uint32_t taddr_slot;
};
static_assert(
sizeof(XrtsF16Shared) == 98432,
"xRTS FP16x1 owner must retain its conflict-free 96.125 KiB board");
__device__ __forceinline__ void xrts_v8nc_load_a(
const uint32_t* factor,
uint32_t* destination,
int matrix_y,
int row_x,
int history,
int waves,
int tid
) {
// Two dense M64xK16 packets occupy each 16x64 upper-mirror rectangle.
// Expand each one into the top half of a dense M128 descriptor and zero
// the unused bottom half once for the whole lower-row sweep.
for (int vector = tid;
vector < (waves << 7);
vector += 256) {
int issue = vector >> 6;
int atom = vector & 63;
int source_word = (issue << 9) + (atom << 3);
int bank = atom >> 5;
int bank_atom = atom & 31;
uint4 first;
uint4 second;
loco_ld_global_u8(
factor
+ (uint64_t)(
matrix_y
+ ((history >> 5) << 4)
+ (source_word >> 6)) * 1024u
+ (row_x & ~63)
+ (source_word & 63),
first,
second);
int top_word =
(issue << 10)
+ (bank << 9)
+ (bank_atom << 3);
loco_st_shared_u8(
(uint32_t)__cvta_generic_to_shared(
destination + top_word),
first,
second);
uint4 zero = make_uint4(0u, 0u, 0u, 0u);
loco_st_shared_u8(
(uint32_t)__cvta_generic_to_shared(
destination + top_word + 256),
zero,
zero);
}
}
__device__ __forceinline__ void xrts_tma_load_residual16k(
const CUtensorMap& map,
uint32_t destination,
uint32_t barrier,
int row_y,
int matrix_row
) {
asm volatile(
"mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 "
"_, [%0], 16384;\n\t"
"cp.async.bulk.tensor.2d.shared::cta.global."
"mbarrier::complete_tx::bytes "
"[%1], [%2, {%3,%4}], [%0];"
:
: "r"(barrier),
"r"(destination),
"l"(&map),
"r"(row_y),
"r"(matrix_row)
: "memory");
}
__device__ __forceinline__ void xrts_tma_load_packed16k(
const CUtensorMap& map,
uint32_t destination,
uint32_t barrier,
int y
) {
asm volatile(
"mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 "
"_, [%0], 16384;\n\t"
"cp.async.bulk.tensor.2d.shared::cta.global."
"mbarrier::complete_tx::bytes "
"[%1], [%2, {0,%3}], [%0];"
:
: "r"(barrier),
"r"(destination),
"l"(&map),
"r"(y)
: "memory");
}
__device__ __forceinline__ void xrts_tma_store_residual(
const CUtensorMap& map,
uint32_t source,
int row_y,
int matrix_row
) {
asm volatile(
"cp.async.bulk.tensor.2d.global.shared::cta.bulk_group "
"[%0, {%1,%2}], [%3];\n\t"
"cp.async.bulk.commit_group;\n\t"
"cp.async.bulk.wait_group 0;"
:
: "l"(&map),
"r"(row_y),
"r"(matrix_row),
"r"(source)
: "memory");
}
__device__ __forceinline__ void xrts_f16_issue(
uint32_t taddr,
uint32_t a,
uint32_t b,
uint32_t accumulate
) {
// D:F32 | A/B:F16 | -A | dense M128 N64 K16.
constexpr uint32_t m128n64k16_f16_negate_a = 0x08102010u;
uint64_t adesc = loco_operand_desc_m128(a);
uint64_t bdesc = loco_operand_desc(b);
asm volatile(
"{\n\t"
".reg .pred acc;\n\t"
".reg .b32 z<4>;\n\t"
"setp.ne.u32 acc, %4, 0;\n\t"
"mov.b32 z0, 0;\n\t"
"mov.b32 z1, 0;\n\t"
"mov.b32 z2, 0;\n\t"
"mov.b32 z3, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 "
"[%0], %1, %2, %3, {z0,z1,z2,z3}, acc;\n\t"
"}"
:
: "r"(taddr),
"l"(adesc),
"l"(bdesc),
"r"(m128n64k16_f16_negate_a),
"r"(accumulate)
: "memory");
}
/*
One lower-row owner:
packed A(X) -> V8.nc once, retained through the column sweep
packed B(Y) -> 16 KiB TMA, eight K16 issues per stage
residual D -> 16 KiB TMA, retained row-major in shared
dense FP16x1 -> -X*Y'
final -> D + delta
K3/TRSM publishes both descriptor-native views while the solved values are
live: dense A packets in conflict-free upper-mirror rectangles and dense B
packets in the FP16 shadow. K4 performs no format conversion.
sibling=1 owns lower(H1,H1) plus R[I,H1], I>H1 after H0.
sibling=0 owns one inclusive-lower M64 row and sweeps all columns <= row.
*/
__global__ __launch_bounds__(256, 2)
void xrts_f16_m128n64_b60(
float* __restrict__ factor,
const __grid_constant__ CUtensorMap residual_map,
const __grid_constant__ CUtensorMap packed_map,
int history,
int first,
int waves,
int sibling
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
XrtsF16Shared& shared =
*reinterpret_cast<XrtsF16Shared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int panels = (8 - first) << 1;
int per_batch = sibling
? 2 + ((7 - first) << 1)
: panels;
if (tid < 2) {
loco_mbar_init((uint32_t)__cvta_generic_to_shared(
tid ? &shared.mma_done : &shared.tma_done));
}
if (tid < 2) {
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 7) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 64;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
uint32_t taddr = shared.taddr_slot;
uint32_t tma_bar = (uint32_t)__cvta_generic_to_shared(
&shared.tma_done);
uint32_t mma_bar = (uint32_t)__cvta_generic_to_shared(
&shared.mma_done);
uint32_t tma_phase = 0u;
uint32_t mma_phase = 0u;
for (int task = blockIdx.x;
task < 60 * per_batch;
task += gridDim.x) {
int b = task / per_batch;
int local = task - b * per_batch;
int row_x = (first << 7) + (local << 6);
int columns =
sibling && local > 1 ? 2 : local + 1;
int column_begin = first << 7;
int matrix_y = b << 10;
xrts_v8nc_load_a(
reinterpret_cast<const uint32_t*>(factor),
shared.a[0],
matrix_y,
row_x,
history,
waves,
tid);
__syncthreads();
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
for (int column = 0; column < columns; ++column) {
int row_y = column_begin + (column << 6);
if (tid == 0) {
xrts_tma_load_residual16k(
residual_map,
(uint32_t)__cvta_generic_to_shared(shared.d),
tma_bar,
row_y,
matrix_y + row_x);
}
loco_mbar_wait(tma_bar, tma_phase);
tma_phase ^= 1u;
__syncthreads();
#pragma unroll
for (int phase = 0; phase < (waves >> 2); ++phase) {
if (tid == 0) {
xrts_tma_load_packed16k(
packed_map,
(uint32_t)__cvta_generic_to_shared(shared.b),
tma_bar,
(((((b << 4) + (row_y >> 6)) * 32)
+ (history >> 5)
+ (phase << 2))
<< 2));
}
loco_mbar_wait(tma_bar, tma_phase);
tma_phase ^= 1u;
__syncthreads();
if (tid == 0) {
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
xrts_f16_issue(
taddr,
(uint32_t)__cvta_generic_to_shared(
shared.a[0]
+ (((phase << 3) + issue) << 10)),
(uint32_t)__cvta_generic_to_shared(
shared.b + (issue << 9)),
(uint32_t)(phase || issue));
}
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:
: "r"(mma_bar)
: "memory");
}
if (warp < 2) {
loco_mbar_wait(mma_bar, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
if (warp < 2) {
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int q = 0; q < 4; ++q) {
float value[16];
loco_tmem_load16(
(taddr | ((uint32_t)warp << 21)) + (q << 4),
value[0],value[1],value[2],value[3],
value[4],value[5],value[6],value[7],
value[8],value[9],value[10],value[11],
value[12],value[13],value[14],value[15]);
int row = (warp << 5) + lane;
int first_col = q << 4;
float* residual =
reinterpret_cast<float*>(shared.d)
+ row * 64 + first_col;
#pragma unroll
for (int v = 0; v < 16; ++v) {
if (row_x != row_y || row >= first_col + v) {
residual[v] += value[v];
}
}
}
asm volatile(
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
__syncthreads();
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
__syncthreads();
if (tid == 0) {
xrts_tma_store_residual(
residual_map,
(uint32_t)__cvta_generic_to_shared(shared.d),
row_y,
matrix_y + row_x);
}
__syncthreads();
}
}
if (warp == 7) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 64;"
:
: "r"(taddr)
: "memory");
}
__syncthreads();
#endif
}
__global__ __launch_bounds__(256, 8)
void xrts_zero_upper_b60(
float* __restrict__ factor
) {
for (int task = blockIdx.x; task < 60 * 36; task += gridDim.x) {
int b = task / 36;
int local = task - b * 36;
int row_block;
int col_block;
if (local < 28) {
col_block = 1;
while (local >= col_block) {
local -= col_block;
++col_block;
}
row_block = local;
} else {
row_block = local - 28;
col_block = row_block;
}
uint64_t base =
((uint64_t)b << 20)
+ ((uint64_t)(row_block << 7) << 10)
+ (col_block << 7);
for (int vector = threadIdx.x;
vector < 128 * 32;
vector += 256) {
int row = vector >> 5;
int col = (vector & 31) << 2;
if (row_block != col_block) {
*reinterpret_cast<float4*>(
factor + base + (uint64_t)row * 1024u + col) =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
} else {
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row < col + e) {
factor[
base
+ (uint64_t)row * 1024u
+ col + e] = 0.0f;
}
}
}
}
}
}
__device__ __forceinline__ void loco_plane_store_scale_a(
uint32_t taddr,
const uint8_t* __restrict__ scale
) {
int lane = threadIdx.x & 31;
const uint32_t* words = reinterpret_cast<const uint32_t*>(scale);
uint32_t r0 = words[lane];
uint32_t r1 = words[lane + 32];
uint32_t r2 = words[lane + 64];
uint32_t r3 = words[lane + 96];
asm volatile(
"tcgen05.st.sync.aligned.32x32b.x4.b32 "
"[%0], {%1,%2,%3,%4};"
:
: "r"(taddr), "r"(r0), "r"(r1), "r"(r2), "r"(r3)
: "memory");
}
__device__ __forceinline__ void loco_plane_store_scale_b(
uint32_t taddr,
const uint8_t* __restrict__ scale
) {
int lane = threadIdx.x & 31;
const uint32_t* words = reinterpret_cast<const uint32_t*>(scale);
uint32_t r0 = words[lane];
uint32_t r1 = words[lane + 32];
uint32_t r2 = words[lane + 64];
uint32_t r3 = words[lane + 96];
uint32_t r4 = words[lane + 128];
uint32_t r5 = words[lane + 160];
uint32_t r6 = words[lane + 192];
uint32_t r7 = words[lane + 224];
asm volatile(
"tcgen05.st.sync.aligned.32x32b.x8.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8};"
:
: "r"(taddr),
"r"(r0), "r"(r1), "r"(r2), "r"(r3),
"r"(r4), "r"(r5), "r"(r6), "r"(r7)
: "memory");
}
__device__ __forceinline__ void loco_persistent_plane_commit(
uint32_t done_bar
) {
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done_bar)
: "memory");
}
__device__ __forceinline__ void loco_nvfp4_store_scales(
uint32_t partition_base,
const uint8_t* __restrict__ scale_a,
const uint8_t* __restrict__ scale_b
) {
#pragma unroll
for (int issue = 0; issue < 4; ++issue) {
loco_plane_store_scale_a(
partition_base + 256u + (uint32_t)(issue << 2),
scale_a + issue * 512);
loco_plane_store_scale_b(
partition_base + 272u + (uint32_t)(issue << 3),
scale_b + issue * 1024);
}
asm volatile(
"tcgen05.wait::st.sync.aligned;\n\t"
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
__device__ __forceinline__ void
loco_nvfp4_update_m128n256k64(
uint32_t taddr,
uint32_t packet_a,
uint32_t packet_b,
int issue
) {
// UE4M3 block16 scales, dense E2M1, FP32 D, and descriptor negate-A.
constexpr uint32_t instruction = 0x08402480u;
uint64_t a_desc = loco_plane_desc_a(packet_a);
uint64_t b_desc = loco_plane_desc_b(packet_b);
uint32_t sa = taddr + 256u + (uint32_t)(issue << 2);
uint32_t sb = taddr + 272u + (uint32_t)(issue << 3);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
"setp.ne.u32 accumulate, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4."
"block_scale.scale_vec::4X "
"[%0], %1, %2, %3, [%4], [%5], accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(instruction),
"r"(sa),
"r"(sb),
"r"(1u)
: "memory");
}
__device__ __forceinline__ void loco_e2m1x4_store_all_scales(
uint32_t partition_base,
const uint8_t* __restrict__ scale_a,
const uint8_t* __restrict__ scale_b
) {
// D uses [0,256). Four SFA boards consume 16 columns; four SFB boards
// consume 32 more. All scales remain resident for the 16 cross-products.
#pragma unroll
for (int plane = 0; plane < LOCO_OUTER_ACTIVE_PLANES; ++plane) {
loco_plane_store_scale_a(
partition_base + 256u + (uint32_t)(plane << 2),
scale_a + plane * 512);
}
#pragma unroll
for (int plane = 0; plane < LOCO_OUTER_ACTIVE_PLANES; ++plane) {
loco_plane_store_scale_b(
partition_base + 272u + (uint32_t)(plane << 3),
scale_b + plane * 1024);
}
asm volatile(
"tcgen05.wait::st.sync.aligned;\n\t"
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
__device__ __forceinline__ void loco_inner_hexlift_store_scales(
uint32_t partition_base,
const uint8_t* __restrict__ scale_a,
const uint8_t* __restrict__ scale_b,
int first_plane
) {
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
loco_plane_store_scale_a(
partition_base + 256u + (uint32_t)(plane << 2),
scale_a + plane * 512);
}
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
loco_plane_store_scale_b(
partition_base + 272u + (uint32_t)(plane << 3),
scale_b + plane * 1024);
}
asm volatile(
"tcgen05.wait::st.sync.aligned;\n\t"
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
__device__ __forceinline__ void loco_e2m1x4_issue_m128n256k64(
uint32_t taddr,
uint32_t packet_a,
uint32_t packet_b,
int plane_a,
int plane_b
) {
constexpr uint32_t instruction = 0x08c02480u;
uint64_t a_desc = loco_plane_desc_a(packet_a);
uint64_t b_desc = loco_plane_desc_b(packet_b);
uint32_t sa = taddr + 256u + (uint32_t)(plane_a << 2);
uint32_t sb = taddr + 272u + (uint32_t)(plane_b << 3);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
"setp.ne.u32 accumulate, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4."
"block_scale.scale_vec::2X "
"[%0], %1, %2, %3, [%4], [%5], accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(instruction),
"r"(sa),
"r"(sb),
"r"(1u)
: "memory");
}
__device__ __forceinline__ void
loco_hexlift_outer_issue_m128n256k64(
uint32_t taddr,
uint32_t packet_a,
uint32_t packet_b,
int plane_a,
int plane_b
) {
// UE8M0 stays in idesc; 4X consumes four K16 scale bytes per K64.
constexpr uint32_t instruction = 0x08c02480u;
uint64_t a_desc = loco_plane_desc_a(packet_a);
uint64_t b_desc = loco_plane_desc_b(packet_b);
uint32_t sa = taddr + 256u + (uint32_t)(plane_a << 2);
uint32_t sb = taddr + 272u + (uint32_t)(plane_b << 3);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
"setp.ne.u32 accumulate, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4."
"block_scale.scale_vec::4X "
"[%0], %1, %2, %3, [%4], [%5], accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(instruction),
"r"(sa),
"r"(sb),
"r"(1u)
: "memory");
}
__device__ __forceinline__ void loco_inner_hexlift_issue_m128nk64(
uint32_t taddr,
uint32_t packet_a,
uint32_t packet_b,
int valid_n,
int plane_a,
int plane_b
) {
// valid_n={192,128,64} -> {0x08b02480,0x08a02480,0x08902480}.
constexpr uint32_t instruction_base = 0x08802480u;
uint32_t instruction =
instruction_base | ((uint32_t)(valid_n >> 3) << 17);
uint64_t a_desc = loco_plane_desc_a(packet_a);
uint64_t b_desc = loco_inner_plane_desc_b(packet_b, valid_n);
uint32_t sa = taddr + 256u + (uint32_t)(plane_a << 2);
uint32_t sb = taddr + 272u + (uint32_t)(plane_b << 3);
asm volatile(
"{\n\t"
".reg .pred accumulate;\n\t"
"setp.ne.u32 accumulate, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4."
"block_scale.scale_vec::2X "
"[%0], %1, %2, %3, [%4], [%5], accumulate;\n\t"
"}"
:
: "r"(taddr),
"l"(a_desc),
"l"(b_desc),
"r"(instruction),
"r"(sa),
"r"(sb),
"r"(1u)
: "memory");
}
__device__ __forceinline__ int loco_inner_hexlift_nearest_digit(
int value,
int weight
) {
// Voronoi cells of |2*E2M1| = {0,1,2,3,4,6,8,12}.
int magnitude = value < 0 ? -value : value;
int twice = magnitude << 1;
int digit = twice <= weight * 1 ? 0
: twice <= weight * 3 ? 1
: twice <= weight * 5 ? 2
: twice <= weight * 7 ? 3
: twice <= weight * 10 ? 4
: twice <= weight * 14 ? 6
: twice <= weight * 20 ? 8
: 12;
return value < 0 ? -digit : digit;
}
__device__ __forceinline__ int loco_inner_fp16_exponent(uint16_t bits) {
if (!(bits & 0x7fffu)) return 0;
int exponent = (bits >> 10) & 31;
if (exponent == 31) exponent = 30;
return exponent ? exponent : 1;
}
__device__ __forceinline__ uint32_t loco_inner_hexlift_digits(
uint16_t bits,
int block_exponent
) {
int exponent = (bits >> 10) & 31;
int mantissa = bits & 1023;
if (exponent == 31) {
exponent = 30;
mantissa = 1023;
}
int magnitude = exponent ? 1024 + mantissa : mantissa;
if (!magnitude) return 0u;
int value_exponent = exponent ? exponent : 1;
int shift = block_exponent - value_exponent;
if (shift >= 12) {
magnitude = 0;
} else if (shift > 0) {
int retained = magnitude >> shift;
int remainder = magnitude & ((1 << shift) - 1);
int halfway = 1 << (shift - 1);
if (remainder > halfway
|| (remainder == halfway && (retained & 1))) {
++retained;
}
magnitude = retained;
}
int r3 = loco_inner_hexlift_nearest_digit(magnitude, 256);
magnitude -= r3 * 256;
int r2 = loco_inner_hexlift_nearest_digit(magnitude, 32);
magnitude -= r2 * 32;
int r1 = loco_inner_hexlift_nearest_digit(magnitude, 4);
int r0 = magnitude - r1 * 4;
int sign = bits & 0x8000u ? -1 : 1;
return (uint32_t)(uint8_t)(int8_t)(sign * r0)
| ((uint32_t)(uint8_t)(int8_t)(sign * r1) << 8)
| ((uint32_t)(uint8_t)(int8_t)(sign * r2) << 16)
| ((uint32_t)(uint8_t)(int8_t)(sign * r3) << 24);
}
__device__ __forceinline__ uint8_t loco_inner_hexlift_pack(
int low_digit,
int high_digit
) {
float low = (float)low_digit * 0.5f;
float high = (float)high_digit * 0.5f;
uint32_t packed;
asm volatile(
"{\n\t"
".reg .b8 q;\n\t"
"cvt.rn.satfinite.e2m1x2.f32 q, %2, %1;\n\t"
"cvt.u32.u8 %0, q;\n\t"
"}"
: "=r"(packed)
: "f"(low), "f"(high));
return (uint8_t)packed;
}
__device__ __forceinline__ uint8_t loco_inner_hexlift_scale(
int block_exponent,
int plane
) {
if (!block_exponent) return 0u;
int shift = plane == 0 ? 0 : plane == 1 ? 2 : plane == 2 ? 5 : 8;
int code = block_exponent + 103 + shift;
if (code > 254) code = 254;
return (uint8_t)code;
}
__device__ __forceinline__ uint32_t loco_inner_hexlift_scale_word(
uint32_t base_word,
int plane
) {
uint32_t result = 0u;
#pragma unroll
for (int chunk = 0; chunk < 2; ++chunk) {
uint8_t base_code = (uint8_t)(base_word >> (chunk << 3));
if (base_code) {
int shift =
plane == 0 ? 0 : plane == 1 ? 2 : plane == 2 ? 5 : 8;
int code = (int)base_code + shift;
if (code > 254) code = 254;
result |= (uint32_t)code << (chunk << 3);
}
}
return result;
}
__device__ __forceinline__ uint32_t
loco_hexlift_outer_k16_scale_word(
uint32_t base_word,
int plane
) {
uint32_t result = 0u;
int shift = plane == 1 ? 2 : plane == 2 ? 5 : 8;
#pragma unroll
for (int block = 0; block < 4; ++block) {
uint8_t base_code =
(uint8_t)(base_word >> (block << 3));
if (base_code) {
int code = (int)base_code + shift;
if (code > 254) code = 254;
result |= (uint32_t)code << (block << 3);
}
}
return result;
}
__device__ __forceinline__ int loco_inner_hexlift_row_prefix(
int stage
) {
// Prefix of live B rows for N={192,128,64}: {0,192,320}.
return 32 * stage * (7 - stage);
}
__device__ __forceinline__ int loco_inner_hexlift_value_bytes(
int first_plane
) {
return (LOCO_INNER_HEXLIFT_PLANES - first_plane)
* 2 * LOCO_INNER_HEXLIFT_STAGE_ROWS * 16;
}
__device__ __forceinline__ int loco_inner_hexlift_block_exponent(
uint32_t fp16,
uint32_t mask
) {
int exponent = loco_inner_fp16_exponent((uint16_t)fp16);
int exponent_1 =
loco_inner_fp16_exponent((uint16_t)(fp16 >> 16));
if (exponent_1 > exponent) exponent = exponent_1;
int other = __shfl_down_sync(mask, exponent, 8, 16);
if (other > exponent) exponent = other;
other = __shfl_down_sync(mask, exponent, 4, 16);
if (other > exponent) exponent = other;
other = __shfl_down_sync(mask, exponent, 2, 16);
if (other > exponent) exponent = other;
other = __shfl_down_sync(mask, exponent, 1, 16);
if (other > exponent) exponent = other;
return __shfl_sync(mask, exponent, 0, 16);
}
__device__ __forceinline__ uint8_t
loco_hexlift_outer_nibble(int digit) {
int magnitude = digit < 0 ? -digit : digit;
uint8_t code = magnitude <= 4 ? (uint8_t)magnitude
: magnitude == 6 ? 5u
: magnitude == 8 ? 6u
: 7u;
return code | (digit < 0 ? 8u : 0u);
}
__device__ __forceinline__ uint8_t
loco_hexlift_outer_pack(int low_digit, int high_digit) {
return loco_hexlift_outer_nibble(low_digit)
| (uint8_t)(loco_hexlift_outer_nibble(high_digit) << 4);
}
__device__ __forceinline__ void
loco_publish_hexlift_outer_k64_warp(
uint32_t fp16_0,
uint32_t fp16_1,
uint8_t* __restrict__ value_row,
uint32_t* __restrict__ scale_row,
uint64_t chunk_stride,
int first_plane
) {
int lane = threadIdx.x & 31;
uint32_t block_mask = 0xffffu << (lane & 16);
int exponent_0 =
loco_inner_hexlift_block_exponent(fp16_0, block_mask);
int exponent_1 =
loco_inner_hexlift_block_exponent(fp16_1, block_mask);
int upper_exponent_0 =
__shfl_sync(0xffffffffu, exponent_0, 16);
int upper_exponent_1 =
__shfl_sync(0xffffffffu, exponent_1, 16);
if (!(lane & 1)) {
int byte = lane >> 1;
uint32_t digits_00 =
loco_inner_hexlift_digits((uint16_t)fp16_0, exponent_0);
uint32_t digits_01 = loco_inner_hexlift_digits(
(uint16_t)(fp16_0 >> 16), exponent_0);
uint32_t digits_10 =
loco_inner_hexlift_digits((uint16_t)fp16_1, exponent_1);
uint32_t digits_11 = loco_inner_hexlift_digits(
(uint16_t)(fp16_1 >> 16), exponent_1);
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
uint64_t plane_stride =
(uint64_t)(plane - LOCO_HEXLIFT_OUTER_FIRST_PLANE)
* (chunk_stride << 1);
int shift = plane << 3;
int d00 = (int)(int8_t)(digits_00 >> shift);
int d01 = (int)(int8_t)(digits_01 >> shift);
int d10 = (int)(int8_t)(digits_10 >> shift);
int d11 = (int)(int8_t)(digits_11 >> shift);
value_row[plane_stride + byte] =
loco_hexlift_outer_pack(d00, d01);
value_row[plane_stride + chunk_stride + byte] =
loco_hexlift_outer_pack(d10, d11);
}
}
if (lane == 0) {
*scale_row =
(uint32_t)loco_inner_hexlift_scale(exponent_0, 0)
| ((uint32_t)loco_inner_hexlift_scale(
upper_exponent_0, 0) << 8)
| ((uint32_t)loco_inner_hexlift_scale(
exponent_1, 0) << 16)
| ((uint32_t)loco_inner_hexlift_scale(
upper_exponent_1, 0) << 24);
}
}
__device__ __forceinline__ void
loco_publish_hexlift_outer_k128_warp(
uint32_t fp16_0,
uint32_t fp16_1,
uint32_t fp16_2,
uint32_t fp16_3,
uint8_t* __restrict__ packed_values,
uint8_t* __restrict__ packed_scales,
int b,
int n,
int k,
int row,
int first_plane
) {
uint64_t chunk_stride = (uint64_t)n << 4;
int history = k >> 6;
uint8_t* value_row =
packed_values + loco_hexlift_outer_value_offset(
b,
n,
history,
LOCO_HEXLIFT_OUTER_FIRST_PLANE,
0,
row);
uint32_t* scale_row = reinterpret_cast<uint32_t*>(
packed_scales + loco_hexlift_outer_scale_offset(
b, n, history, row));
loco_publish_hexlift_outer_k64_warp(
fp16_0,
fp16_1,
value_row,
scale_row,
chunk_stride,
first_plane);
value_row = packed_values + loco_hexlift_outer_value_offset(
b,
n,
history + 1,
LOCO_HEXLIFT_OUTER_FIRST_PLANE,
0,
row);
scale_row = reinterpret_cast<uint32_t*>(
packed_scales + loco_hexlift_outer_scale_offset(
b, n, history + 1, row));
loco_publish_hexlift_outer_k64_warp(
fp16_2,
fp16_3,
value_row,
scale_row,
chunk_stride,
first_plane);
}
__device__ __forceinline__ void loco_inner_hexlift_pack_factor(
LocoHexLiftInnerShared& shared,
uint8_t* __restrict__ diagonal_shadow,
int panel,
int start,
int row_base,
int valid_m,
bool pack_b,
int first_plane
) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int chunk = lane >> 4;
int rank = lane & 15;
uint32_t block_mask = 0xffffu << (chunk << 4);
for (int row = warp; row < 128; row += 32) {
int source_row = row_base + row;
uint32_t fp16 = 0u;
if (row < valid_m) {
int k = panel + (chunk << 5) + (rank << 1);
int source = loco_lower256(source_row, k);
fp16 = loco_f16x2(
shared.tile[source],
shared.tile[source + 1]);
}
int block_exponent =
loco_inner_hexlift_block_exponent(fp16, block_mask);
int other_exponent =
__shfl_sync(0xffffffffu, block_exponent, 16);
uint32_t digits_0 =
loco_inner_hexlift_digits((uint16_t)fp16, block_exponent);
uint32_t digits_1 = loco_inner_hexlift_digits(
(uint16_t)(fp16 >> 16), block_exponent);
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
int r0 = (int)(int8_t)(digits_0 >> (plane << 3));
int r1 = (int)(int8_t)(digits_1 >> (plane << 3));
shared.packet_a[plane][
(chunk << 11) + (row << 4) + rank
] = loco_inner_hexlift_pack(r0, r1);
if (lane == 0) {
uint32_t scales =
(uint32_t)loco_inner_hexlift_scale(
block_exponent, plane)
| ((uint32_t)loco_inner_hexlift_scale(
other_exponent, plane) << 8);
reinterpret_cast<uint32_t*>(
shared.scale_a[plane])[row] = scales;
}
}
}
if (pack_b) {
int valid_n = 256 - start;
int stage = panel >> 6;
int row_prefix = loco_inner_hexlift_row_prefix(stage);
int plane_count = LOCO_INNER_HEXLIFT_PLANES - first_plane;
int stage_values = plane_count * 32 * row_prefix;
int plane_values = valid_n * 32;
int scale_base =
loco_inner_hexlift_value_bytes(first_plane)
+ (row_prefix << 2);
// Bounding the dense loop prevents chunk-zero padding from aliasing
// the live chunk-one rows when N is below 256.
for (int row = warp; row < valid_n; row += 32) {
int source_row = start + row;
int k = panel + (chunk << 5) + (rank << 1);
int source = loco_lower256(source_row, k);
uint32_t fp16 = loco_f16x2(
shared.tile[source],
shared.tile[source + 1]);
int block_exponent =
loco_inner_hexlift_block_exponent(fp16, block_mask);
int other_exponent =
__shfl_sync(0xffffffffu, block_exponent, 16);
uint32_t digits_0 =
loco_inner_hexlift_digits(
(uint16_t)fp16, block_exponent);
uint32_t digits_1 = loco_inner_hexlift_digits(
(uint16_t)(fp16 >> 16), block_exponent);
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
int r0 = (int)(int8_t)(digits_0 >> (plane << 3));
int r1 = (int)(int8_t)(digits_1 >> (plane << 3));
int packet_offset =
chunk * valid_n * 16 + (row << 4) + rank;
uint8_t packed = loco_inner_hexlift_pack(r0, r1);
shared.packet_b[plane][packet_offset] = packed;
diagonal_shadow[
stage_values
+ (plane - first_plane) * plane_values
+ packet_offset
] = packed;
if (lane == 0) {
uint32_t scales =
(uint32_t)loco_inner_hexlift_scale(
block_exponent, plane)
| ((uint32_t)loco_inner_hexlift_scale(
other_exponent, plane) << 8);
reinterpret_cast<uint32_t*>(
shared.scale_b[plane])[row] = scales;
}
}
if (lane == 0) {
uint32_t base_scales =
(uint32_t)loco_inner_hexlift_scale(
block_exponent, 0)
| ((uint32_t)loco_inner_hexlift_scale(
other_exponent, 0) << 8);
reinterpret_cast<uint32_t*>(
diagonal_shadow + scale_base)[row] = base_scales;
}
}
}
__syncthreads();
}
__device__ __forceinline__ void loco_inner_hexlift_pack_trsm(
LocoHexLiftInnerShared& shared,
const uint8_t* __restrict__ diagonal_shadow,
const float* __restrict__ rows,
int row_stride,
int row_count,
int row0,
int panel,
int start,
bool load_b,
int first_plane
) {
int tid = threadIdx.x;
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int chunk = lane >> 4;
int rank = lane & 15;
uint32_t block_mask = 0xffffu << (chunk << 4);
int valid_m = row_count - row0;
if (valid_m > 128) valid_m = 128;
int a_warps = load_b ? 16 : 32;
for (int row = warp; warp < a_warps && row < 128; row += a_warps) {
uint32_t fp16 = 0u;
if (row < valid_m) {
int k = panel + (chunk << 5) + (rank << 1);
const float* source =
rows + (uint64_t)(row0 + row) * row_stride + k;
fp16 = loco_f16x2(source[0], source[1]);
}
int block_exponent =
loco_inner_hexlift_block_exponent(fp16, block_mask);
int other_exponent =
__shfl_sync(0xffffffffu, block_exponent, 16);
uint32_t digits_0 =
loco_inner_hexlift_digits((uint16_t)fp16, block_exponent);
uint32_t digits_1 = loco_inner_hexlift_digits(
(uint16_t)(fp16 >> 16), block_exponent);
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
int r0 = (int)(int8_t)(digits_0 >> (plane << 3));
int r1 = (int)(int8_t)(digits_1 >> (plane << 3));
shared.packet_a[plane][
(chunk << 11) + (row << 4) + rank
] = loco_inner_hexlift_pack(r0, r1);
if (lane == 0) {
uint32_t scales =
(uint32_t)loco_inner_hexlift_scale(
block_exponent, plane)
| ((uint32_t)loco_inner_hexlift_scale(
other_exponent, plane) << 8);
reinterpret_cast<uint32_t*>(
shared.scale_a[plane])[row] = scales;
}
}
}
if (load_b && warp >= 16) {
int valid_n = 256 - start;
int local_tid = tid - 512;
int stage = panel >> 6;
int row_prefix = loco_inner_hexlift_row_prefix(stage);
int plane_count = LOCO_INNER_HEXLIFT_PLANES - first_plane;
int stage_values = plane_count * 32 * row_prefix;
int plane_values = valid_n * 32;
int row_pairs = valid_n >> 1;
int value_tasks = plane_count * 2 * row_pairs;
// All row-group CTAs reread this immutable B packet. One retained
// 32-byte transaction lands two adjacent packed rows directly.
for (int x = local_tid; x < value_tasks; x += 512) {
int pair = x % row_pairs;
int q = x / row_pairs;
int packed_plane = q >> 1;
int b_chunk = q & 1;
int plane = first_plane + packed_plane;
int packet_offset =
b_chunk * valid_n * 16 + (pair << 5);
const uint8_t* src =
diagonal_shadow
+ stage_values
+ packed_plane * plane_values
+ packet_offset;
uint4 first;
uint4 second;
loco_ld_global_u8_ca(
reinterpret_cast<const uint32_t*>(src),
first,
second);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
&shared.packet_b[plane][packet_offset]);
loco_st_shared_u8(dst, first, second);
}
int scale_base =
loco_inner_hexlift_value_bytes(first_plane)
+ (row_prefix << 2);
for (int row = local_tid; row < valid_n; row += 512) {
uint32_t base_word = loco_ld_global_u32_ca(
reinterpret_cast<const uint32_t*>(
diagonal_shadow + scale_base) + row);
#pragma unroll
for (int plane = first_plane;
plane < LOCO_INNER_HEXLIFT_PLANES;
++plane) {
reinterpret_cast<uint32_t*>(
shared.scale_b[plane])[row] =
loco_inner_hexlift_scale_word(base_word, plane);
}
}
}
__syncthreads();
}
__device__ __forceinline__ void loco_inner_hexlift_import_factor(
LocoHexLiftInnerShared& shared,
int start,
int row_base,
int valid_m
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int valid_n = 256 - start;
int valid_segments = valid_n >> 6;
#pragma unroll
for (int segment = 0; segment < valid_segments; ++segment) {
for (int vector = tid; vector < 2048; vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
int global_row = row_base + row;
int global_col = start + (segment << 6) + col4;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* words = reinterpret_cast<float*>(&value);
if (row < valid_m && global_col < start + valid_n) {
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (global_col + e < 256
&& global_row >= global_col + e) {
words[e] = shared.tile[
loco_lower256(global_row, global_col + e)];
}
}
}
reinterpret_cast<float4*>(shared.scratch)[vector] = value;
}
__syncthreads();
if (warp < 8) {
int lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int local_row = lane_base + (lane >> 2);
int local_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
uint32_t src = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (local_row << 6) + local_col);
loco_tmem_import(taddr, src);
}
__syncthreads();
}
}
__device__ __forceinline__ void loco_inner_hexlift_import_trsm(
LocoHexLiftInnerShared& shared,
const float* __restrict__ rows,
int row_stride,
int row_count,
int row0,
int start
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int valid_m = row_count - row0;
if (valid_m > 128) valid_m = 128;
int valid_n = 256 - start;
int valid_segments = valid_n >> 6;
#pragma unroll
for (int segment = 0; segment < valid_segments; ++segment) {
for (int vector = tid; vector < 2048; vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
int col = start + (segment << 6) + col4;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (row < valid_m && col < start + valid_n) {
value = loco_ld_global_v4(
rows + (uint64_t)(row0 + row) * row_stride + col);
}
reinterpret_cast<float4*>(shared.scratch)[vector] = value;
}
__syncthreads();
if (warp < 8) {
int lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int local_row = lane_base + (lane >> 2);
int local_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
uint32_t src = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (local_row << 6) + local_col);
loco_tmem_import(taddr, src);
}
__syncthreads();
}
}
__device__ __forceinline__ void loco_inner_hexlift_issue(
LocoHexLiftInnerShared& shared,
int valid_n,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase
) {
int warp = threadIdx.x >> 5;
int first_plane = hex10 ? 0 : 1;
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
if (warp < 4) {
uint32_t partition =
loco_tmem_reader_addr(shared.taddr_slot, warp);
loco_inner_hexlift_store_scales(
partition,
&shared.scale_a[0][0],
&shared.scale_b[0][0],
first_plane);
}
__syncthreads();
uint32_t packet_a =
(uint32_t)__cvta_generic_to_shared(&shared.packet_a[0][0]);
uint32_t packet_b =
(uint32_t)__cvta_generic_to_shared(&shared.packet_b[0][0]);
if (issuer) {
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b + 3 * 8192,
valid_n,
3, 3);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b + 2 * 8192,
valid_n,
3, 2);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 2 * 4096,
packet_b + 3 * 8192,
valid_n,
2, 3);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b + 1 * 8192,
valid_n,
3, 1);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 1 * 4096,
packet_b + 3 * 8192,
valid_n,
1, 3);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 2 * 4096,
packet_b + 2 * 8192,
valid_n,
2, 2);
if (hex10) {
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b,
valid_n,
3, 0);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a,
packet_b + 3 * 8192,
valid_n,
0, 3);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 2 * 4096,
packet_b + 1 * 8192,
valid_n,
2, 1);
loco_inner_hexlift_issue_m128nk64(
shared.taddr_slot,
packet_a + 1 * 4096,
packet_b + 2 * 8192,
valid_n,
1, 2);
}
loco_persistent_plane_commit(done);
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
__device__ __forceinline__ void loco_inner_hexlift_export_factor(
LocoHexLiftInnerShared& shared,
int start,
int row_base,
int valid_m
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int valid_n = 256 - start;
int valid_segments = valid_n >> 6;
#pragma unroll
for (int segment = 0; segment < valid_segments; ++segment) {
if (warp < 8) {
int lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int local_row = lane_base + (lane >> 2);
int local_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (local_row << 6) + local_col);
loco_tmem_export(taddr, dst);
}
__syncthreads();
for (int vector = tid; vector < 2048; vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
int global_row = row_base + row;
int global_col = start + (segment << 6) + col4;
if (row < valid_m && global_col < start + valid_n) {
float4 value =
reinterpret_cast<float4*>(shared.scratch)[vector];
float* words = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (global_col + e < 256
&& global_row >= global_col + e) {
shared.tile[
loco_lower256(global_row, global_col + e)
] = words[e];
}
}
}
}
__syncthreads();
}
}
__device__ __forceinline__ void loco_inner_hexlift_export_trsm(
LocoHexLiftInnerShared& shared,
float* __restrict__ rows,
int row_stride,
int row_count,
int row0,
int start
) {
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int valid_m = row_count - row0;
if (valid_m > 128) valid_m = 128;
int valid_n = 256 - start;
int valid_segments = valid_n >> 6;
#pragma unroll
for (int segment = 0; segment < valid_segments; ++segment) {
if (warp < 8) {
int lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int local_row = lane_base + (lane >> 2);
int local_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (local_row << 6) + local_col);
loco_tmem_export(taddr, dst);
}
__syncthreads();
int segment_col = start + (segment << 6);
if (segment_col < start + valid_n) {
for (int vector = tid; vector < 2048; vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
if (row < valid_m) {
*reinterpret_cast<float4*>(
rows
+ (uint64_t)(row0 + row) * row_stride
+ segment_col
+ col4
) = reinterpret_cast<float4*>(
shared.scratch)[vector];
}
}
}
__syncthreads();
}
}
__device__ __forceinline__ void loco_hexlift_factor_update(
LocoHexLiftInnerShared& shared,
uint8_t* __restrict__ diagonal_shadow,
int panel,
int start,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase
) {
int first_plane = hex10 ? 0 : 1;
for (int row_base = start; row_base < 256; row_base += 128) {
int valid_m = 256 - row_base;
if (valid_m > 128) valid_m = 128;
loco_inner_hexlift_pack_factor(
shared,
diagonal_shadow,
panel,
start,
row_base,
valid_m,
row_base == start,
first_plane);
loco_inner_hexlift_import_factor(
shared,
start,
row_base,
valid_m);
loco_inner_hexlift_issue(
shared,
256 - start,
hex10,
issuer,
mma_phase);
loco_inner_hexlift_export_factor(
shared,
start,
row_base,
valid_m);
}
}
__device__ __forceinline__ void loco_hexlift_trsm_update(
LocoHexLiftInnerShared& shared,
const uint8_t* __restrict__ diagonal_shadow,
float* __restrict__ rows,
int row_stride,
int row_count,
int row0,
int panel,
int start,
bool load_b,
bool hex10,
uint32_t issuer,
uint32_t& mma_phase
) {
int first_plane = hex10 ? 0 : 1;
loco_inner_hexlift_pack_trsm(
shared,
diagonal_shadow,
rows,
row_stride,
row_count,
row0,
panel,
start,
load_b,
first_plane);
loco_inner_hexlift_import_trsm(
shared,
rows,
row_stride,
row_count,
row0,
start);
loco_inner_hexlift_issue(
shared,
256 - start,
hex10,
issuer,
mma_phase);
loco_inner_hexlift_export_trsm(
shared,
rows,
row_stride,
row_count,
row0,
start);
}
#endif
/*
Ranked shapes 5/6: B16/B640 x N512.
A = [A00 * ] L = [L00 0 ]
[A10 A11] [L10 L11]
Four ordered launches expose a different board at each dependency edge:
B CTAs factor L00
B*{8,1} CTAs solve L10
B*3 CTA512 form the three lower A11 quadrants
B CTAs factor L11
The update owns only N128 TMEM. Three CTA512s therefore replace the old
full-TMEM CTA1024 matrix owner without changing the FP16 TCGEN product.
*/
__global__ __launch_bounds__(1024, 1)
void xrshape512_factor_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch,
int diagonal
) {
extern __shared__ __align__(128) float smem[];
float* tile = smem;
float* inverse = tile + 256 * 257 / 2;
int tid = threadIdx.x;
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
uint64_t base = (uint64_t)b << 18;
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int column = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
input
+ base
+ (uint64_t)(diagonal + row) * 512u
+ diagonal
+ column);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
tile[loco_lower256(row, column + e)] = word[e];
}
}
}
__syncthreads();
loco_factor_k16<256, 32>(tile, inverse);
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int column = (vector & 63) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
word[e] = tile[loco_lower256(row, column + e)];
}
}
*reinterpret_cast<float4*>(
output
+ base
+ (uint64_t)(diagonal + row) * 512u
+ diagonal
+ column) = value;
if (!diagonal) {
*reinterpret_cast<float4*>(
output
+ base
+ (uint64_t)row * 512u
+ 256u
+ column) =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
}
__syncthreads();
}
}
__global__ __launch_bounds__(1024, 1)
void xrshape512_trsm_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch,
int row_groups,
int rows_per_group
) {
extern __shared__ __align__(128) float smem[];
float* tile = smem;
float* inverse = tile + 256 * 257 / 2;
float* a_smem = inverse + 256;
int tid = threadIdx.x;
for (int task = blockIdx.x;
task < batch * row_groups;
task += gridDim.x) {
int b = task / row_groups;
int group = task - b * row_groups;
uint64_t base = (uint64_t)b << 18;
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int column = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
output
+ base
+ (uint64_t)row * 512u
+ column);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= column + e) {
tile[loco_lower256(row, column + e)] = word[e];
}
}
}
__syncthreads();
if (tid < 256) {
inverse[tid] = xcal_rcp_approx(
tile[loco_lower256(tid, tid)]);
}
int first = group * rows_per_group;
int last = first + rows_per_group;
if (last > 256) last = 256;
for (int vector = tid;
vector < (last - first) * 64;
vector += 1024) {
int row = vector >> 6;
int column = (vector & 63) << 2;
*reinterpret_cast<float4*>(
output
+ base
+ (uint64_t)(256 + first + row) * 512u
+ column) = loco_ld_global_v4_cs(
input
+ base
+ (uint64_t)(256 + first + row) * 512u
+ column);
}
__syncthreads();
loco_persistent_trsm256<false>(
output
+ base
+ (uint64_t)(256 + first) * 512u,
512,
last - first,
tile,
inverse,
a_smem);
__syncthreads();
}
}
__global__ __launch_bounds__(512, 3)
void xrshape512_update_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1000
extern __shared__ __align__(128) unsigned char dynamic_smem[];
XrShape512UpdateShared& shared =
*reinterpret_cast<XrShape512UpdateShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 8) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 128;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t row_packet =
(uint32_t)__cvta_generic_to_shared(shared.row_packet);
uint32_t col_packet =
(uint32_t)__cvta_generic_to_shared(shared.col_packet);
uint32_t mma_phase = 0;
for (int task = blockIdx.x; task < batch * 3; task += gridDim.x) {
int b = task / 3;
int quadrant = task - b * 3;
uint64_t base = (uint64_t)b << 18;
#pragma unroll
for (int phase = 0; phase < 2; ++phase) {
int packet_row_index = tid >> 2;
int quad = tid & 3;
int k_quad = quad << 2;
int packet_word =
(packet_row_index << 2)
+ ((quad & 1) << 1)
+ ((quad >> 1) << 9);
#pragma unroll
for (int wave = 0; wave < 4; ++wave) {
#pragma unroll
for (int issue_half = 0; issue_half < 2; ++issue_half) {
int local_issue = (wave << 1) + issue_half;
int issue = (phase << 3) + local_issue;
int k = (issue << 4) + k_quad;
int word = (local_issue << 10) + packet_word;
if (quadrant != 2) {
float4 l0 = loco_ld_global_v4_cs(
output
+ base
+ (uint64_t)(256 + packet_row_index) * 512u
+ k);
if (!quadrant) {
shared.row_packet[word] =
loco_f16x2(l0.x, l0.y);
shared.row_packet[word + 1] =
loco_f16x2(l0.z, l0.w);
} else {
shared.col_packet[word] =
loco_f16x2(l0.x, l0.y);
shared.col_packet[word + 1] =
loco_f16x2(l0.z, l0.w);
}
}
if (quadrant) {
float4 l1 = loco_ld_global_v4_cs(
output
+ base
+ (uint64_t)(384 + packet_row_index) * 512u
+ k);
shared.row_packet[word] =
loco_f16x2(l1.x, l1.y);
shared.row_packet[word + 1] =
loco_f16x2(l1.z, l1.w);
}
}
}
__syncthreads();
if (issuer) {
asm volatile(
"fence.proxy.async.shared::cta;\n\t"
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16_update_m128n128_desc(
shared.taddr_slot,
loco_operand_desc_m128(
row_packet + (issue << 12)),
loco_operand_desc_m128(
(quadrant == 1 ? col_packet : row_packet)
+ (issue << 12)),
phase != 0 || issue != 0);
}
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
if (warp < 4) {
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
int row =
(quadrant != 0 ? 128 : 0)
+ (warp << 5)
+ lane;
#pragma unroll
for (int column = 0; column < 128; column += 16) {
float result[16];
loco_tmem_load16(
(shared.taddr_slot + column)
| ((uint32_t)(warp << 5) << 16),
result[0], result[1], result[2], result[3],
result[4], result[5], result[6], result[7],
result[8], result[9], result[10], result[11],
result[12], result[13], result[14], result[15]);
#pragma unroll
for (int e = 0; e < 16; ++e) {
int col =
(quadrant == 2 ? 128 : 0)
+ column
+ e;
if (row >= col) {
output[
base
+ (uint64_t)(256 + row) * 512u
+ 256u
+ col] =
input[
base
+ (uint64_t)(256 + row) * 512u
+ 256u
+ col]
+ result[e];
}
}
}
asm volatile(
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
__syncthreads();
}
if (warp == 8) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 128;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
__global__ __launch_bounds__(1024, 1)
void loco_persistent512_tcgen_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoPersistent512TcgenShared& shared =
*reinterpret_cast<LocoPersistent512TcgenShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
float* tile = shared.tile;
float* inverse = shared.inverse;
float* a_smem = reinterpret_cast<float*>(shared.row_packet);
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t row_packet =
(uint32_t)__cvta_generic_to_shared(shared.row_packet);
uint32_t col_packet =
(uint32_t)__cvta_generic_to_shared(shared.col_packet);
uint32_t mma_phase = 0;
// Resident-sized matrix queue. Both local K256 factors use K64 fronts;
// L10 retains its measured K64 solve and the A11 Gram cone stays TCGEN.
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
uint64_t base = (uint64_t)b * 512u * 512u;
// A00 FP32 dense -> packed lower. Evict-first loads do not displace the
// rotating L10 packet board from L1 while the batch queue advances.
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
input + base + (uint64_t)row * 512 + col);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
tile[loco_lower256(row, col + e)] = word[e];
}
}
}
__syncthreads();
// Exact B640 control: scalar fronts remain FP32; each completed
// K16 cone is one rounded-FP16 MMA.
loco_factor_k16<256, 32, true>(
tile,
inverse,
shared.row_packet);
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
word[e] = tile[loco_lower256(row, col + e)];
}
}
*reinterpret_cast<float4*>(
output + base + (uint64_t)row * 512 + col) = value;
*reinterpret_cast<float4*>(
output + base + (uint64_t)row * 512 + 256 + col) =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
__syncthreads();
// L10: exact K64 solve fronts, one rounded-FP16 tensor drain.
// One coalesced preamble makes the RHS mutable for the four fronts.
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
*reinterpret_cast<float4*>(
output
+ base
+ (uint64_t)(256 + row) * 512
+ col) = loco_ld_global_v4_cs(
input
+ base
+ (uint64_t)(256 + row) * 512
+ col);
}
__syncthreads();
loco_persistent_trsm256<true, true>(
output + base + 256u * 512u,
512,
256,
tile,
inverse,
a_smem,
shared.col_packet);
// A11 is loaded once, then stays packed and resident through all
// three lower M128xN128 quadrants.
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
input
+ base
+ (uint64_t)(256 + row) * 512
+ 256
+ col);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
tile[loco_lower256(row, col + e)] = word[e];
}
}
}
__syncthreads();
// Lower A11 board lives in three adjacent TMEM M128xN128 regions:
//
// TMEM[ 0:128] = -L0 L0^T q0 = C[ 0:128, 0:128]
// TMEM[128:256] = -L1 L0^T q1 = C[128:256, 0:128]
// TMEM[256:384] = -L1 L1^T q2 = C[128:256,128:256]
//
// L0 and L1 are each packed once per K128 phase. All three products
// share one commit, so K256 costs two completion edges instead of six.
#pragma unroll
for (int phase = 0; phase < 2; ++phase) {
int packet_row_index = (tid & 511) >> 2;
int quad = tid & 3;
int issue_half = tid >> 9;
int k_quad = quad << 2;
int packet_word =
(packet_row_index << 2)
+ ((quad & 1) << 1)
+ ((quad >> 1) << 9);
#pragma unroll
for (int wave = 0; wave < 4; ++wave) {
int local_issue = (wave << 1) + issue_half;
int issue = (phase << 3) + local_issue;
int k = (issue << 4) + k_quad;
float4 l0 = loco_ld_global_v4_cs(
output
+ base
+ (uint64_t)(256 + packet_row_index) * 512
+ k);
float4 l1 = loco_ld_global_v4_cs(
output
+ base
+ (uint64_t)(384 + packet_row_index) * 512
+ k);
int word = (local_issue << 10) + packet_word;
shared.row_packet[word] = loco_f16x2(l0.x, l0.y);
shared.row_packet[word + 1] = loco_f16x2(l0.z, l0.w);
shared.col_packet[word] = loco_f16x2(l1.x, l1.y);
shared.col_packet[word + 1] = loco_f16x2(l1.z, l1.w);
}
__syncthreads();
if (issuer) {
asm volatile(
"fence.proxy.async.shared::cta;\n\t"
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int quadrant = 0; quadrant < 3; ++quadrant) {
uint32_t q_taddr =
shared.taddr_slot + (quadrant << 7);
uint32_t a_packet =
quadrant == 0 ? row_packet : col_packet;
uint32_t b_packet =
quadrant == 2 ? col_packet : row_packet;
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16_update_m128n128_desc(
q_taddr,
loco_operand_desc_m128(
a_packet + (issue << 12)),
loco_operand_desc_m128(
b_packet + (issue << 12)),
phase != 0 || issue != 0);
}
}
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
// Four warps drain all three disjoint quadrants directly into packed
// FP32 A11. There is one CTA edge after the complete TMEM export.
if (warp < 4) {
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int quadrant = 0; quadrant < 3; ++quadrant) {
int row_block = quadrant != 0;
int col_block = quadrant == 2;
int row = (row_block << 7) + (warp << 5) + lane;
#pragma unroll
for (int column = 0; column < 128; column += 16) {
float result[16];
uint32_t taddr =
(shared.taddr_slot
+ (quadrant << 7)
+ column)
| ((uint32_t)(warp << 5) << 16);
loco_tmem_load16(
taddr,
result[0], result[1], result[2], result[3],
result[4], result[5], result[6], result[7],
result[8], result[9], result[10], result[11],
result[12], result[13], result[14], result[15]);
#pragma unroll
for (int e = 0; e < 16; ++e) {
int col = (col_block << 7) + column + e;
if (row >= col) {
tile[loco_lower256(row, col)] += result[e];
}
}
}
}
asm volatile(
"tcgen05.fence::before_thread_sync;"
::: "memory");
}
__syncthreads();
loco_factor_k16<256, 32, true>(
tile,
inverse,
shared.row_packet);
// L11 and its zero upper triangle finish the complete output matrix.
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
word[e] = tile[loco_lower256(row, col + e)];
}
}
*reinterpret_cast<float4*>(
output
+ base
+ (uint64_t)(256 + row) * 512
+ 256
+ col) = value;
}
__syncthreads();
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
/*
Ranked N1024/B4 board:
copy
[ factor K256 -> solve/publish K256 -> update every lower residual ] x 3
factor K256
Four exact K256 fronts replace eight K128 fronts. The TRSM board owns at
most 32 rows, so its stride-72 K64 feeder remains the measured N512 board.
*/
__global__ __launch_bounds__(1024, 1)
void xrshape1024_factor256_b4_kernel(
float* __restrict__ factor,
int panel
) {
extern __shared__ __align__(128) float smem[];
float* tile = smem;
float* inverse = tile + 256 * 257 / 2;
int tid = threadIdx.x;
for (int b = blockIdx.x; b < 4; b += gridDim.x) {
uint64_t base = (uint64_t)b << 20;
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
factor
+ base
+ (uint64_t)(panel + row) * 1024u
+ panel
+ col);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
tile[loco_lower256(row, col + e)] = word[e];
}
}
}
__syncthreads();
loco_factor_k16<256, 32>(tile, inverse);
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
word[e] =
tile[loco_lower256(row, col + e)];
}
}
*reinterpret_cast<float4*>(
factor
+ base
+ (uint64_t)(panel + row) * 1024u
+ panel
+ col) = value;
}
__syncthreads();
}
}
__global__ __launch_bounds__(1024, 1)
void xrshape1024_trsm256_b4_kernel(
float* __restrict__ factor,
uint32_t* __restrict__ packed_fp16,
int panel,
int row_groups
) {
extern __shared__ __align__(128) float smem[];
float* tile = smem;
float* inverse = tile + 256 * 257 / 2;
float* a_smem = inverse + 256;
int tid = threadIdx.x;
for (int task = blockIdx.x;
task < 4 * row_groups;
task += gridDim.x) {
int b = task / row_groups;
int group = task - b * row_groups;
uint64_t base = (uint64_t)b << 20;
for (int vector = tid; vector < 256 * 64; vector += 1024) {
int row = vector >> 6;
int col = (vector & 63) << 2;
float4 value = loco_ld_global_v4_cs(
factor
+ base
+ (uint64_t)(panel + row) * 1024u
+ panel
+ col);
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (row >= col + e) {
tile[loco_lower256(row, col + e)] = word[e];
}
}
}
__syncthreads();
if (tid < 256) {
inverse[tid] = xcal_rcp_approx(
tile[loco_lower256(tid, tid)]);
}
int first = panel + 256 + (group << 5);
int last = first + 32;
if (last > 1024) last = 1024;
__syncthreads();
loco_persistent_trsm256<false>(
factor + base + (uint64_t)first * 1024u + panel,
1024,
last - first,
tile,
inverse,
a_smem);
__syncthreads();
// Publish descriptor-native high/residual-low atoms once. The K256
// residual wave consumes two K128 groups without re-reading FP32.
for (int pair = tid;
pair < (last - first) * 128;
pair += 1024) {
int row = pair >> 7;
int col = (pair & 127) << 1;
float x0 = factor[
base
+ (uint64_t)(first + row) * 1024u
+ panel
+ col];
float x1 = factor[
base
+ (uint64_t)(first + row) * 1024u
+ panel
+ col
+ 1];
uint32_t high;
uint32_t low;
loco_f16x2_hi_lo(x0, x1, high, low);
packed_fp16[loco_fp16_packed_word(
b, 1024, first + row, panel + col, 0)] = high;
packed_fp16[loco_fp16_packed_word(
b, 1024, first + row, panel + col, 0) + 32] = low;
}
__syncthreads();
}
}
/*
N1024/B60 three-resident residual board:
P12 : 60 x 21 M128N128 owners
P8 : 60 x 10 M128N128 owners
P4 : 60 x 3 M128N128 owners
CTA256 reuses the measured N512 descriptor board: 64 KiB of eight-issue
A/B packets, 128 TMEM columns, two K128 phases, then the dead A packet is
the coalesced export slice. Three CTAs fit both RMEM and TMEM budgets.
*/
template <int P>
__global__ __launch_bounds__(256, 3)
void loco_trailing_tcgen_b60_m128n128(
float* __restrict__ factor,
const uint32_t* __restrict__ packed_fp16,
int batch
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
XrShape512UpdateShared& shared =
*reinterpret_cast<XrShape512UpdateShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 4) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 128;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t row_packets =
(uint32_t)__cvta_generic_to_shared(shared.row_packet);
uint32_t col_packets =
(uint32_t)__cvta_generic_to_shared(shared.col_packet);
uint32_t mma_phase = 0;
for (int task = blockIdx.x;
task < batch * (P == 12 ? 21 : P == 8 ? 10 : 3);
task += gridDim.x) {
int b = task / (P == 12 ? 21 : P == 8 ? 10 : 3);
int local =
task - b * (P == 12 ? 21 : P == 8 ? 10 : 3);
int row;
if constexpr (P == 12) {
row = local < 1 ? 0
: local < 3 ? 1
: local < 6 ? 2
: local < 10 ? 3
: local < 15 ? 4
: 5;
} else if constexpr (P == 8) {
row = local < 1 ? 0
: local < 3 ? 1
: local < 6 ? 2
: 3;
} else {
row = local < 1 ? 0 : 1;
}
int column = local - (row * (row + 1) >> 1);
int row_base = (8 - (P >> 1) + row) << 7;
int col_base = (8 - (P >> 1) + column) << 7;
uint64_t base = (uint64_t)b << 20;
#pragma unroll
for (int segment = 0; segment < 2; ++segment) {
int lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
loco_tmem_import_global(
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16),
factor
+ base
+ ((uint64_t)(
row_base
+ lane_base
+ (lane >> 2)
+ ((lane & 1) << 3)) << 10)
+ col_base
+ (segment << 6)
+ ((lane & 2) << 1));
__syncthreads();
}
#pragma unroll
for (int phase = 0; phase < 2; ++phase) {
for (int vector = tid;
vector < 2048;
vector += 256) {
uint4 first;
uint4 second;
if (vector < 1024) {
loco_ld_global_u8(
packed_fp16
+ ((base
+ ((uint64_t)(
row_base + (vector & 127)) << 10)
+ ((12 - P) << 6)
+ (phase << 7)
+ ((vector >> 7) << 4)) >> 1),
first,
second);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.row_packet
+ (vector >> 7) * 1024
+ ((vector & 127) << 2)),
first);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.row_packet
+ (vector >> 7) * 1024
+ 512
+ ((vector & 127) << 2)),
second);
} else {
int bvector = vector - 1024;
loco_ld_global_u8(
packed_fp16
+ ((base
+ ((uint64_t)(
col_base + (bvector & 127)) << 10)
+ ((12 - P) << 6)
+ (phase << 7)
+ ((bvector >> 7) << 4)) >> 1),
first,
second);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.col_packet
+ (bvector >> 7) * 1024
+ ((bvector & 127) << 2)),
first);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.col_packet
+ (bvector >> 7) * 1024
+ 512
+ ((bvector & 127) << 2)),
second);
}
}
__syncthreads();
if (issuer) {
asm volatile(
"fence.proxy.async.shared::cta;\n\t"
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16_update_m128n128_desc(
shared.taddr_slot,
loco_operand_desc_m128(
row_packets + (issue << 12)),
loco_operand_desc_m128(
col_packets + (issue << 12)),
1u);
}
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
#pragma unroll
for (int segment = 0; segment < 2; ++segment) {
int export_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
loco_tmem_export(
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)export_base << 16),
(uint32_t)__cvta_generic_to_shared(
shared.row_packet
+ ((export_base + (lane >> 2)) << 6)
+ (lane & 3)));
__syncthreads();
for (int vector = tid;
vector < 2048;
vector += 256) {
int r = vector >> 4;
int col4 = (vector & 15) << 2;
float4 value =
reinterpret_cast<float4*>(
shared.row_packet)[vector];
if (row != column || r >= (segment << 6) + col4 + 3) {
*reinterpret_cast<float4*>(
factor
+ base
+ ((uint64_t)(row_base + r) << 10)
+ col_base
+ (segment << 6)
+ col4) = value;
} else if (r >= (segment << 6) + col4) {
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (r >= (segment << 6) + col4 + e) {
factor[
base
+ ((uint64_t)(row_base + r) << 10)
+ col_base
+ (segment << 6)
+ col4
+ e] = word[e];
}
}
}
}
__syncthreads();
}
}
if (warp == 4) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 128;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
__host__ __device__ __forceinline__ uint64_t
loco_scaled_n_task_prefix(
int tile_rows,
int owner128
) {
// Before scaled-N row r:
// r full off-diagonal owners + r full diagonal owners.
return (
(uint64_t)tile_rows
* owner128
* ((uint64_t)owner128 * tile_rows + 1)
) >> 1;
}
// Scaled low-B outer board:
//
// one sealed P panel
// -> every lower scaled-N residual owner in parallel
// -> native M128N128 owner, CTA256, 128 TMEM columns
//
// One persistent owner is launched per SM. D remains in TMEM across the full
// P history; the FP16 shadow is loaded in K128 packets, then D exports once.
template <bool OverlapHistory>
__global__ __launch_bounds__(256, 1)
void loco_outer_scaled_n_parallel(
float* __restrict__ factor,
const uint32_t* __restrict__ packed_fp16,
int batch,
int n,
int start_panel,
int history_begin,
int history_end,
int trailing128,
int owner128
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
XrShape512UpdateShared& shared =
*reinterpret_cast<XrShape512UpdateShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 4) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 128;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t row_packets =
(uint32_t)__cvta_generic_to_shared(shared.row_packet);
uint32_t col_packets =
(uint32_t)__cvta_generic_to_shared(shared.col_packet);
uint32_t mma_phase = 0u;
uint64_t grouped_tasks =
(uint64_t)trailing128 * (trailing128 + 1) >> 1;
uint64_t tasks = (uint64_t)batch * grouped_tasks;
for (uint64_t task = blockIdx.x;
task < tasks;
task += gridDim.x) {
int b = (int)(task / grouped_tasks);
uint64_t local =
task - (uint64_t)b * grouped_tasks;
int owner_rows =
(trailing128 + owner128 - 1) / owner128;
int low = 0;
int high = owner_rows - 1;
while (low < high) {
int middle = (low + high + 1) >> 1;
if (loco_scaled_n_task_prefix(
middle,
owner128) <= local) {
low = middle;
} else {
high = middle - 1;
}
}
int tile_row = low;
int in_row =
(int)(local - loco_scaled_n_task_prefix(
tile_row,
owner128));
int rows_in_owner =
trailing128 - tile_row * owner128;
if (rows_in_owner > owner128) {
rows_in_owner = owner128;
}
int off_diagonal =
rows_in_owner * tile_row * owner128;
int row128;
int col128;
if (in_row < off_diagonal) {
int owner_cells = rows_in_owner * owner128;
int tile_col = in_row / owner_cells;
int cell = in_row - tile_col * owner_cells;
int row_in_owner = cell / owner128;
row128 = tile_row * owner128 + row_in_owner;
col128 =
tile_col * owner128
+ cell
- row_in_owner * owner128;
} else {
int cell = in_row - off_diagonal;
int diagonal_row = 0;
int diagonal_high = rows_in_owner - 1;
while (diagonal_row < diagonal_high) {
int middle =
(diagonal_row + diagonal_high + 1) >> 1;
if ((uint64_t)middle * (middle + 1) / 2
<= (uint64_t)cell) {
diagonal_row = middle;
} else {
diagonal_high = middle - 1;
}
}
row128 = tile_row * owner128 + diagonal_row;
col128 =
tile_row * owner128
+ cell
- (diagonal_row * (diagonal_row + 1) >> 1);
}
int row_base = (start_panel << 6) + (row128 << 7);
int col_base = (start_panel << 6) + (col128 << 7);
uint64_t base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
#pragma unroll
for (int segment = 0; segment < 2; ++segment) {
int lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
loco_tmem_import_global(
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16),
factor
+ base
+ (uint64_t)(
row_base
+ lane_base
+ (lane >> 2)
+ ((lane & 1) << 3)) * n
+ col_base
+ (segment << 6)
+ ((lane & 2) << 1));
__syncthreads();
}
for (int history = history_begin;
history < history_end;
history += 2) {
// Eight independent v8 loads are issued before any shared store.
// The one-CTA/SM board spends its free RMEM on memory-level
// parallelism. Every carrier and destination is explicit.
int packet_row = tid & 127;
int issue_rank = tid >> 7;
uint64_t row_word =
(base
+ (uint64_t)(row_base + packet_row) * n
+ ((uint64_t)history << 6)) >> 1;
uint64_t col_word =
(base
+ (uint64_t)(col_base + packet_row) * n
+ ((uint64_t)history << 6)) >> 1;
int packet_word =
(issue_rank << 10) + (packet_row << 2);
uint4 a00, a01, a20, a21;
uint4 a40, a41, a60, a61;
uint4 b00, b01, b20, b21;
uint4 b40, b41, b60, b61;
loco_ld_global_u8(
packed_fp16 + row_word + (issue_rank << 3),
a00, a01);
loco_ld_global_u8(
packed_fp16 + row_word + (issue_rank << 3) + 16,
a20, a21);
loco_ld_global_u8(
packed_fp16 + row_word + (issue_rank << 3) + 32,
a40, a41);
loco_ld_global_u8(
packed_fp16 + row_word + (issue_rank << 3) + 48,
a60, a61);
loco_ld_global_u8(
packed_fp16 + col_word + (issue_rank << 3),
b00, b01);
loco_ld_global_u8(
packed_fp16 + col_word + (issue_rank << 3) + 16,
b20, b21);
loco_ld_global_u8(
packed_fp16 + col_word + (issue_rank << 3) + 32,
b40, b41);
loco_ld_global_u8(
packed_fp16 + col_word + (issue_rank << 3) + 48,
b60, b61);
// N>=4096 keeps the next K128's complete 64 KiB carrier bundle in
// RMEM while TCGEN consumes the previous shared packet.
if constexpr (OverlapHistory) {
if (history != history_begin) {
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
}
uint32_t row_packet =
(uint32_t)__cvta_generic_to_shared(
shared.row_packet + packet_word);
uint32_t col_packet =
(uint32_t)__cvta_generic_to_shared(
shared.col_packet + packet_word);
loco_st_shared_u4(row_packet, a00);
loco_st_shared_u4(row_packet + 8192, a20);
loco_st_shared_u4(row_packet + 16384, a40);
loco_st_shared_u4(row_packet + 24576, a60);
loco_st_shared_u4(row_packet + 2048, a01);
loco_st_shared_u4(row_packet + 10240, a21);
loco_st_shared_u4(row_packet + 18432, a41);
loco_st_shared_u4(row_packet + 26624, a61);
loco_st_shared_u4(col_packet, b00);
loco_st_shared_u4(col_packet + 8192, b20);
loco_st_shared_u4(col_packet + 16384, b40);
loco_st_shared_u4(col_packet + 24576, b60);
loco_st_shared_u4(col_packet + 2048, b01);
loco_st_shared_u4(col_packet + 10240, b21);
loco_st_shared_u4(col_packet + 18432, b41);
loco_st_shared_u4(col_packet + 26624, b61);
__syncthreads();
if (issuer) {
asm volatile(
"fence.proxy.async.shared::cta;\n\t"
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16_update_m128n128_desc(
shared.taddr_slot,
loco_operand_desc_m128(
row_packets + (issue << 12)),
loco_operand_desc_m128(
col_packets + (issue << 12)),
1u);
}
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
if constexpr (!OverlapHistory) {
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
}
if constexpr (OverlapHistory) {
if (history_begin < history_end) {
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
}
#pragma unroll
for (int segment = 0; segment < 2; ++segment) {
int export_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
loco_tmem_export(
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)export_base << 16),
(uint32_t)__cvta_generic_to_shared(
shared.row_packet
+ ((export_base + (lane >> 2)) << 6)
+ (lane & 3)));
__syncthreads();
for (int vector = tid;
vector < 2048;
vector += 256) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
float4 value =
reinterpret_cast<float4*>(
shared.row_packet)[vector];
int global_row = row_base + row;
int global_col4 =
col_base + (segment << 6) + col4;
if (global_row >= global_col4 + 3) {
*reinterpret_cast<float4*>(
factor
+ base
+ (uint64_t)global_row * n
+ global_col4) = value;
} else if (global_row >= global_col4) {
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (global_row >= global_col4 + e) {
factor[
base
+ (uint64_t)global_row * n
+ global_col4
+ e] = word[e];
}
}
}
}
__syncthreads();
}
}
if (warp == 4) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 128;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
/*
N1024/B60 compact residual board:
P12 : 60 x 12 M128N256 owners
P8 : 60 x 6 M128N256 owners
P4 : 60 x 2 M128N256 owners
CTA512 keeps the control's four N64 packet. Q-issue phases reuse one compact
packet, D stays in 256 TMEM columns for the full K256, and the dead packet
becomes the coalesced export slice. R CTAs fit per SM.
*/
template <int P, int Q, int R>
__global__ __launch_bounds__(512, R)
void loco_trailing_tcgen_b60_m128n256(
float* __restrict__ factor,
const uint32_t* __restrict__ packed_fp16,
int batch
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoB60N256Shared<Q>& shared =
*reinterpret_cast<LocoB60N256Shared<Q>*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 8) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 256;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
uint32_t packets =
(uint32_t)__cvta_generic_to_shared(shared.packet);
uint32_t mma_phase = 0;
for (int task = blockIdx.x;
task < batch * (P == 12 ? 12 : P == 8 ? 6 : 2);
task += gridDim.x) {
int b = task / (P == 12 ? 12 : P == 8 ? 6 : 2);
int local =
task - b * (P == 12 ? 12 : P == 8 ? 6 : 2);
int row_pair;
if constexpr (P == 12) {
row_pair = local < 1 ? 0
: local < 2 ? 1
: local < 4 ? 2
: local < 6 ? 3
: local < 9 ? 4
: 5;
} else if constexpr (P == 8) {
row_pair = local < 1 ? 0
: local < 2 ? 1
: local < 4 ? 2
: 3;
} else {
row_pair = local < 1 ? 0 : 1;
}
int column_group = local
- (int)loco_m128n256_task_prefix(row_pair);
int valid_segments =
(row_pair << 1) + 2 - (column_group << 2);
if (valid_segments > 4) valid_segments = 4;
int row_base =
(16 - P + (row_pair << 1)) << 6;
int col_base =
(16 - P + (column_group << 2)) << 6;
uint64_t base = (uint64_t)b << 20;
// Seed the full N256 board. Diagonal tasks export only their live
// N128 half, but seeding all four segments keeps every accumulate
// bit defined without a separate zero board.
#pragma unroll
for (int segment = 0; segment < 4; segment += 2) {
if (warp < 16) {
int local_segment = segment + (warp >> 3);
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
loco_tmem_import_global(
(shared.taddr_slot
+ (uint32_t)(local_segment << 6))
| ((uint32_t)lane_base << 16),
factor
+ base
+ ((uint64_t)(
row_base
+ lane_base
+ (lane >> 2)
+ ((lane & 1) << 3)) << 10)
+ col_base
+ (local_segment << 6)
+ ((lane & 2) << 1));
}
__syncthreads();
}
#pragma unroll
for (int phase = 0; phase < 16 / Q; ++phase) {
static_assert(!(Q & 1), "B60 issue group must split evenly");
constexpr int half_q = Q >> 1;
#pragma unroll
for (int half = 0; half < 2; ++half) {
int issue_begin = half * half_q;
// Private A packets for this half. Their shared addresses are
// disjoint from the half already executing on TCGEN.
for (int local = tid;
local < half_q * 128;
local += 512) {
int issue = issue_begin + (local >> 7);
int row = local & 127;
uint4 first;
uint4 second;
loco_ld_global_u8(
packed_fp16
+ ((base
+ ((uint64_t)(row_base + row) << 10)
+ ((12 - P) << 6)
+ phase * (Q << 4)
+ (issue << 4)) >> 1),
first,
second);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.packet
+ issue * 3072
+ (row << 2)),
first);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.packet
+ issue * 3072
+ 512
+ (row << 2)),
second);
}
// Matching B packets. Missing diagonal segments stay zero.
for (int local = tid;
local < half_q * 256;
local += 512) {
int issue = issue_begin + (local >> 8);
int row = local & 255;
uint4 first = make_uint4(0u, 0u, 0u, 0u);
uint4 second = make_uint4(0u, 0u, 0u, 0u);
if (row < (valid_segments << 6)) {
loco_ld_global_u8(
packed_fp16
+ ((base
+ ((uint64_t)(col_base + row) << 10)
+ ((12 - P) << 6)
+ phase * (Q << 4)
+ (issue << 4)) >> 1),
first,
second);
}
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.packet
+ issue * 3072
+ 1024
+ (row << 2)),
first);
loco_st_shared_u4(
(uint32_t)__cvta_generic_to_shared(
shared.packet
+ issue * 3072
+ 2048
+ (row << 2)),
second);
}
__syncthreads();
if (issuer) {
asm volatile(
"fence.proxy.async.shared::cta;\n\t"
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int local_issue = 0;
local_issue < half_q;
++local_issue) {
int issue = issue_begin + local_issue;
loco_fp16_update_m128n256(
shared.taddr_slot,
packets + issue * 12288,
1u);
}
// One completion tracks both ordered halves. During the
// first half, non-issuer warps immediately build half 1.
if (half == 1) {
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
}
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
// Packet lifetime ended at the second wait. Reuse its first 32 KiB
// as the exact control export map, one coalesced N64 slice at a time.
#pragma unroll
for (int segment = 0;
segment < valid_segments;
++segment) {
if (warp < 8) {
int export_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
loco_tmem_export(
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)export_base << 16),
(uint32_t)__cvta_generic_to_shared(
shared.packet
+ ((export_base + (lane >> 2)) << 6)
+ (lane & 3)));
}
__syncthreads();
for (int vector = tid;
vector < 2048;
vector += 512) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
float4 value =
reinterpret_cast<float4*>(shared.packet)[vector];
int global_row = row_base + row;
int global_col4 =
col_base + (segment << 6) + col4;
if (global_row >= global_col4 + 3) {
*reinterpret_cast<float4*>(
factor
+ base
+ ((uint64_t)global_row << 10)
+ global_col4) = value;
} else if (global_row >= global_col4) {
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (global_row >= global_col4 + e) {
factor[
base
+ ((uint64_t)global_row << 10)
+ global_col4
+ e] = word[e];
}
}
}
}
__syncthreads();
}
}
if (warp == 8) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 256;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
template <
bool ResidualN128,
bool SplitFrontier = false,
bool OverlapHistory = false>
__global__ __launch_bounds__(1024, 1)
void loco_trailing_tcgen_m128n256(
float* __restrict__ factor,
const uint32_t* __restrict__ packed_fp16,
int packed_fp16_atoms,
int batch,
int n,
int start_panel,
int history_begin,
int history_end,
int trailing_panels,
int frontier_only,
int frontier_segments,
int clusters_per_matrix,
int gang
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoMaterializeShared& shared =
*reinterpret_cast<LocoMaterializeShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
// A clustered tail rounds its final worker team up to G CTAs. Cull a
// rank with no N256 owner before it reserves TMEM or initializes state.
if (clusters_per_matrix > 0) {
int row_pairs = (trailing_panels + 1) >> 1;
uint64_t useful = frontier_only
? (uint64_t)row_pairs
: loco_m128n256_task_prefix(row_pairs);
int team = clusters_per_matrix * gang;
int block = (int)blockIdx.x;
int worker = block - (block / team) * team;
if ((uint64_t)worker >= useful) return;
}
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
// The control owns M128xN256. ResidualN128 exposes true M128xN128
// owners; SplitFrontier limits that split to the live N256 dependency.
uint32_t mma_phase = 0;
int row_pairs = (trailing_panels + 1) >> 1;
uint64_t grouped_tasks;
if constexpr (ResidualN128) {
if constexpr (SplitFrontier) {
grouped_tasks = ((uint64_t)row_pairs << 1) - 1;
} else {
grouped_tasks =
(uint64_t)row_pairs * (row_pairs + 1) >> 1;
}
} else {
grouped_tasks = frontier_only
? (uint64_t)row_pairs
: loco_m128n256_task_prefix(row_pairs);
}
uint64_t tasks = (uint64_t)batch * grouped_tasks;
const bool sibling_n128 = ResidualN128
|| (frontier_only && frontier_segments == 2);
const int seed_segments = sibling_n128 ? 2 : 4;
// Clustered launch ownership:
//
// cluster_linear = block / G
// b = cluster_linear / K_b
// c = cluster_linear % K_b
// r = block % G
// worker = c*G + r
//
// K_b and G are launch-selected from measured cluster residency and the
// useful tile count. G=1 / K_b=0 retains the original global task board.
int assigned_batch = -1;
uint64_t task_begin = blockIdx.x;
uint64_t task_step = gridDim.x;
uint64_t task_limit = tasks;
if (clusters_per_matrix > 0 && gang > 0) {
int block = (int)blockIdx.x;
int cluster_linear = block / gang;
assigned_batch = cluster_linear / clusters_per_matrix;
int cluster = cluster_linear
- assigned_batch * clusters_per_matrix;
int rank = block - (block / gang) * gang;
task_begin = (uint64_t)(cluster * gang + rank);
task_step = (uint64_t)(clusters_per_matrix * gang);
task_limit = grouped_tasks;
if (assigned_batch >= batch) task_limit = 0;
}
int tmem_stage = 0;
bool seed_ready = false;
for (uint64_t scheduled = task_begin;
scheduled < task_limit;
scheduled += task_step) {
uint64_t task = assigned_batch >= 0
? (uint64_t)assigned_batch * grouped_tasks + scheduled
: scheduled;
int b;
int row_local;
int col_local;
int valid_m_panels;
int valid_segments;
loco_decode_trailing_task<ResidualN128, SplitFrontier>(
task,
grouped_tasks,
trailing_panels,
frontier_only,
frontier_segments,
b,
row_local,
col_local,
valid_m_panels,
valid_segments);
int row_base = (start_panel + row_local) << 6;
int col_base = (start_panel + col_local) << 6;
int valid_m_rows = valid_m_panels << 6;
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
uint32_t d_tmem =
shared.taddr_slot + (uint32_t)(tmem_stage << 8);
#if XCAL_XRCHOLV18_PAYLOAD_PIPELINE
uint64_t next_scheduled = scheduled + task_step;
bool has_next = next_scheduled < task_limit;
uint64_t next_task = assigned_batch >= 0
? (uint64_t)assigned_batch * grouped_tasks + next_scheduled
: next_scheduled;
int next_b = 0;
int next_row_local = 0;
int next_col_local = 0;
int next_valid_m_panels = 0;
int next_valid_segments = 0;
if (has_next) {
loco_decode_trailing_task<ResidualN128, SplitFrontier>(
next_task,
grouped_tasks,
trailing_panels,
frontier_only,
frontier_segments,
next_b,
next_row_local,
next_col_local,
next_valid_m_panels,
next_valid_segments);
}
#endif
// Scratch is only a zero source for missing seed rows/segments.
// Join current and pipelined-next requirements before any export can
// overwrite it. A fully valid task skips the old 32 KiB clear.
int zero_begin = 2048;
if (!seed_ready) {
if (valid_segments < seed_segments) {
zero_begin = 0;
} else if (valid_m_rows < 128) {
zero_begin = 1024;
}
}
#if XCAL_XRCHOLV18_PAYLOAD_PIPELINE
if (has_next) {
if (next_valid_segments < seed_segments) {
zero_begin = 0;
} else if (next_valid_m_panels < 2
&& zero_begin > 1024) {
zero_begin = 1024;
}
}
#endif
for (int vector = zero_begin + tid;
vector < 2048;
vector += 1024) {
reinterpret_cast<float4*>(shared.scratch)[vector] =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
__syncthreads();
// N256 uses four N64 segments; the sibling N128 issue uses two.
if (!seed_ready && warp < seed_segments * 8) {
int segment = warp >> 3;
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
uint32_t taddr =
(d_tmem + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
if (lane_base < valid_m_rows
&& segment < valid_segments) {
int local_row =
lane_base
+ (lane >> 2)
+ ((lane & 1) << 3);
int local_column = (lane & 2) << 1;
const float* src =
factor
+ matrix_base
+ (uint64_t)(row_base + local_row) * n
+ col_base
+ (segment << 6)
+ local_column;
loco_tmem_import_global(taddr, src);
} else {
int local_row =
lane_base + (lane >> 2);
int local_column = lane & 3;
uint32_t src =
(uint32_t)__cvta_generic_to_shared(
shared.scratch
+ (local_row << 6)
+ local_column);
loco_tmem_import(taddr, src);
}
}
__syncthreads();
// Rounded-FP16 N>=4096 splits the 192 KiB packet envelope into
// descriptor-native K128 stages:
//
// N128 : 64 KiB + 64 KiB + 64 KiB
// N256 : 96 KiB + 96 KiB
//
// The N128 issuer carries three ordered MMA groups before one commit;
// that commit tracks the complete group. N256 retains the measured
// two-stage load/compute overlap. D stays resident in FP32 TMEM.
int history_step = OverlapHistory
? 2
: packed_fp16_atoms || SplitFrontier ? 2 : 4;
int history_slot = 0;
for (int history = history_begin;
history < history_end;
history += history_step) {
int history_col = history << 6;
int history_depth = (history_end - history) << 6;
int history_cap = history_step << 6;
if (history_depth > history_cap) {
history_depth = history_cap;
}
int issue_count = history_depth >> 4;
int packet_stage = history_slot
* (sibling_n128 ? 16384 : 24576);
// One N128 commit closes three same-issuer, same-accumulator
// tcgen05 pipelines. Its completion frees all three packet slots.
if constexpr (ResidualN128 && OverlapHistory) {
if (!history_slot && history != history_begin) {
if (warp < 4) {
loco_mbar_wait(
(uint32_t)__cvta_generic_to_shared(
&shared.mma_done),
mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
}
if (packed_fp16_atoms && !OverlapHistory) {
// The global shadow already matches the consumer K-major
// atoms. Each v8 moves two adjacent rows within one rail.
int issue_words = sibling_n128 ? 4096 : 6144;
int a_vectors = issue_count << 8;
#pragma unroll
for (int vector = tid;
vector < a_vectors;
vector += 1024) {
int a_issue = vector >> 8;
int a_local = vector & 255;
int a_half = a_local >> 7;
int a_group = (a_local >> 3) & 15;
int a_rail = (a_local >> 2) & 1;
int a_row_pair = a_local & 3;
int a_row =
(a_group << 3) + (a_row_pair << 1);
uint4 first =
make_uint4(0u, 0u, 0u, 0u);
uint4 second =
make_uint4(0u, 0u, 0u, 0u);
if (a_row < valid_m_rows) {
int a_col = history_col
+ (a_issue << 4)
+ (a_half << 3);
uint64_t word = loco_fp16_packed_word(
b,
n,
row_base + a_row,
a_col,
a_rail);
loco_ld_global_u8(
packed_fp16 + word,
first,
second);
}
int a_word = a_issue * issue_words
+ (a_half << 10)
+ (a_group << 6)
+ (a_rail << 5)
+ (a_row_pair << 3);
loco_st_shared_u8(
(uint32_t)__cvta_generic_to_shared(
shared.packet + packet_stage + a_word),
first,
second);
}
int b_vectors = sibling_n128
? issue_count << 8
: issue_count << 9;
#pragma unroll 4
for (int vector = tid;
vector < b_vectors;
vector += 1024) {
int b_issue = sibling_n128
? vector >> 8
: vector >> 9;
int b_local = sibling_n128
? vector & 255
: vector & 511;
int b_half = sibling_n128
? b_local >> 7
: b_local >> 8;
int b_group = (b_local >> 3)
& (sibling_n128 ? 15 : 31);
int b_rail = (b_local >> 2) & 1;
int b_row_pair = b_local & 3;
int b_row =
(b_group << 3) + (b_row_pair << 1);
bool b_valid =
b_row < valid_segments * 64;
uint4 first =
make_uint4(0u, 0u, 0u, 0u);
uint4 second =
make_uint4(0u, 0u, 0u, 0u);
if (b_valid) {
int b_col = history_col
+ (b_issue << 4)
+ (b_half << 3);
uint64_t word =
loco_fp16_packed_word(
b,
n,
col_base + b_row,
b_col,
b_rail);
loco_ld_global_u8(
packed_fp16 + word,
first,
second);
}
int b_word = b_issue * issue_words
+ 2048
+ (b_half << (sibling_n128 ? 10 : 11))
+ (b_group << 6)
+ (b_rail << 5)
+ (b_row_pair << 3);
loco_st_shared_u8(
(uint32_t)__cvta_generic_to_shared(
shared.packet + packet_stage + b_word),
first,
second);
}
} else {
// Large shapes retain the row-major high shadow. One v8
// load is split across the descriptor's two K8 banks.
int issue_words = sibling_n128 ? 2048 : 3072;
int a_vectors = issue_count << 7;
#pragma unroll
for (int vector = tid;
vector < a_vectors;
vector += 1024) {
int a_issue = vector >> 7;
int a_row = vector & 127;
uint4 ah0 =
make_uint4(0u, 0u, 0u, 0u);
uint4 ah1 =
make_uint4(0u, 0u, 0u, 0u);
if (a_row < valid_m_rows) {
uint64_t element = matrix_base
+ (uint64_t)(row_base + a_row) * n
+ history_col
+ (a_issue << 4);
loco_ld_global_u8(
packed_fp16 + (element >> 1),
ah0,
ah1);
}
int a_word =
a_issue * issue_words + (a_row << 2);
*reinterpret_cast<uint4*>(
shared.packet + packet_stage + a_word) = ah0;
*reinterpret_cast<uint4*>(
shared.packet
+ packet_stage
+ a_word
+ 512) = ah1;
}
int b_vectors = sibling_n128
? issue_count << 7
: issue_count << 8;
#pragma unroll 4
for (int vector = tid;
vector < b_vectors;
vector += 1024) {
int b_issue = sibling_n128
? vector >> 7
: vector >> 8;
int b_row = vector
& (sibling_n128 ? 127 : 255);
bool b_valid =
b_row < valid_segments * 64;
uint4 bh0 =
make_uint4(0u, 0u, 0u, 0u);
uint4 bh1 =
make_uint4(0u, 0u, 0u, 0u);
if (b_valid) {
uint64_t element = matrix_base
+ (uint64_t)(col_base + b_row) * n
+ history_col
+ (b_issue << 4);
loco_ld_global_u8(
packed_fp16 + (element >> 1),
bh0,
bh1);
}
int b_word =
b_issue * issue_words
+ 1024
+ (b_row << 2);
*reinterpret_cast<uint4*>(
shared.packet + packet_stage + b_word) = bh0;
*reinterpret_cast<uint4*>(
shared.packet
+ packet_stage
+ b_word
+ (sibling_n128 ? 512 : 1024)) = bh1;
}
}
// N256 waits after filling its alternate stage. N128 owns three
// disjoint stages and waits only before wrapping back to stage 0.
if constexpr (OverlapHistory && !ResidualN128) {
if (history != history_begin) {
if (warp < 4) {
loco_mbar_wait(
(uint32_t)__cvta_generic_to_shared(
&shared.mma_done),
mma_phase);
}
}
}
__syncthreads();
if constexpr (OverlapHistory && !ResidualN128) {
if (history != history_begin) {
mma_phase ^= 1u;
}
}
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
uint32_t done =
(uint32_t)__cvta_generic_to_shared(
&shared.mma_done);
if (issuer) {
uint32_t packets =
(uint32_t)__cvta_generic_to_shared(
shared.packet + packet_stage);
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
if (packed_fp16_atoms && !OverlapHistory) {
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
if (sibling_n128) {
loco_fp16x3_update_m128n128(
d_tmem,
packets + issue * 16384);
} else {
loco_fp16x3_update_m128n256(
d_tmem,
packets + issue * 24576);
}
}
} else if (sibling_n128) {
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16_update_m128n128(
d_tmem,
packets + issue * 8192,
1u);
}
} else if (history_depth == 128) {
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16_update_m128n256(
d_tmem,
packets + issue * 12288,
1u);
}
} else {
#pragma unroll
for (int issue = 0; issue < 16; ++issue) {
loco_fp16_update_m128n256(
d_tmem,
packets + issue * 12288,
1u);
}
}
if constexpr (ResidualN128 && OverlapHistory) {
if (history_slot == 2
|| history + history_step >= history_end) {
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
} else {
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
}
#if XCAL_XRCHOLV18_PAYLOAD_PIPELINE
// Park the next seed at the disjoint 256-column TMEM stage.
if (has_next
&& history == history_begin
&& warp < seed_segments * 8) {
int segment = warp >> 3;
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
uint32_t next_tmem =
shared.taddr_slot
+ (uint32_t)((tmem_stage ^ 1) << 8);
uint32_t taddr =
(next_tmem + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
int next_valid_m_rows =
next_valid_m_panels << 6;
if (lane_base < next_valid_m_rows
&& segment < next_valid_segments) {
int local_row =
lane_base
+ (lane >> 2)
+ ((lane & 1) << 3);
int local_column = (lane & 2) << 1;
int next_row_base =
(start_panel + next_row_local) << 6;
int next_col_base =
(start_panel + next_col_local) << 6;
uint64_t next_matrix_base =
(uint64_t)next_b
* (uint64_t)n
* (uint64_t)n;
const float* src =
factor
+ next_matrix_base
+ (uint64_t)(next_row_base + local_row) * n
+ next_col_base
+ (segment << 6)
+ local_column;
loco_tmem_import_global(taddr, src);
} else {
int local_row =
lane_base + (lane >> 2);
int local_column = lane & 3;
uint32_t src =
(uint32_t)__cvta_generic_to_shared(
shared.scratch
+ (local_row << 6)
+ local_column);
loco_tmem_import(taddr, src);
}
}
#endif
if constexpr (!OverlapHistory) {
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
// The completion wait and CTA sync protect the reusable packet board.
__syncthreads();
mma_phase ^= 1u;
}
if constexpr (OverlapHistory) {
if constexpr (ResidualN128) {
if (++history_slot == 3) {
history_slot = 0;
}
} else {
history_slot ^= 1;
}
}
}
if constexpr (OverlapHistory) {
if (history_begin < history_end) {
if (warp < 4) {
loco_mbar_wait(
(uint32_t)__cvta_generic_to_shared(
&shared.mma_done),
mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
}
// Export one N64 slice at a time through the 32 KiB scratch board.
#pragma unroll
for (int segment = 0; segment < seed_segments; ++segment) {
if (warp < 8 && segment < valid_segments) {
int export_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int export_row =
export_base + (lane >> 2);
int export_col = lane & 3;
uint32_t taddr =
(d_tmem + (uint32_t)(segment << 6))
| ((uint32_t)export_base << 16);
uint32_t dst =
(uint32_t)__cvta_generic_to_shared(
shared.scratch
+ (export_row << 6)
+ export_col);
loco_tmem_export(taddr, dst);
}
__syncthreads();
if (segment < valid_segments) {
for (int vector = tid;
vector < 2048;
vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
if (row < valid_m_rows) {
float4 value =
reinterpret_cast<float4*>(
shared.scratch)[vector];
int global_row = row_base + row;
int global_col4 =
col_base + (segment << 6) + col4;
if (global_row >= global_col4 + 3) {
*reinterpret_cast<float4*>(
factor
+ matrix_base
+ (uint64_t)global_row * n
+ global_col4) = value;
} else if (global_row >= global_col4) {
float* words =
reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
int global_col = global_col4 + e;
if (global_row >= global_col) {
factor[
matrix_base
+ (uint64_t)global_row * n
+ global_col] = words[e];
}
}
}
}
}
}
__syncthreads();
}
#if XCAL_XRCHOLV18_PAYLOAD_PIPELINE
seed_ready = has_next;
tmem_stage ^= 1;
#else
seed_ready = false;
#endif
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
__global__ __launch_bounds__(256, 2)
void loco_quantize_outer_nvfp4(
const float* __restrict__ factor,
uint8_t* __restrict__ packed_values,
uint8_t* __restrict__ packed_scales,
int batch,
int n,
int macro_end,
int history_begin,
int history_panels,
int trailing_rows
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int group = lane >> 3;
int rank = lane & 7;
uint32_t octet_mask = 0xffu << (group << 3);
uint64_t tasks =
(uint64_t)batch * history_panels * trailing_rows;
// One task publishes one solved K64 row. Each octet owns one K16
// absmax, one UE4M3 scale, and sixteen E2M1 values.
for (uint64_t task = (uint64_t)blockIdx.x * 8u + warp;
task < tasks;
task += (uint64_t)gridDim.x * 8u) {
int row = (int)(task % trailing_rows);
uint64_t q = task / trailing_rows;
int history = (int)(q % history_panels);
int b = (int)(q / history_panels);
int history_col = (history_begin + history) << 6;
int k0 = (group << 4) + (rank << 1);
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
const float* src = factor
+ matrix_base
+ (uint64_t)(macro_end + row) * n
+ history_col
+ k0;
float x0 = src[0];
float x1 = src[1];
float maximum = fmaxf(fabsf(x0), fabsf(x1));
maximum = fmaxf(
maximum,
__shfl_down_sync(octet_mask, maximum, 4, 8));
maximum = fmaxf(
maximum,
__shfl_down_sync(octet_mask, maximum, 2, 8));
maximum = fmaxf(
maximum,
__shfl_down_sync(octet_mask, maximum, 1, 8));
maximum = __shfl_sync(
octet_mask, maximum, 0, 8);
uint8_t scale_code = 0u;
float scale = rank == 0
? loco_nvfp4_absmax_scale(maximum, scale_code)
: 0.0f;
scale = __shfl_sync(
octet_mask, scale, 0, 8);
uint32_t code_word = __shfl_sync(
octet_mask, (uint32_t)scale_code, 0, 8);
float inverse_scale =
scale > 0.0f ? xcal_rcp_approx(scale) : 0.0f;
uint32_t packed =
loco_nvfp4_pack_pair(x0, x1, inverse_scale);
int chunk = lane >> 4;
int byte = lane & 15;
uint64_t value_offset = loco_nvfp4_value_offset(
b,
n,
history,
history_panels,
chunk,
trailing_rows,
row);
packed_values[value_offset + byte] = (uint8_t)packed;
if (rank == 0) {
uint64_t scale_offset = loco_nvfp4_scale_offset(
b,
history,
history_panels,
trailing_rows,
row);
packed_scales[scale_offset + group] =
(uint8_t)code_word;
}
}
#endif
}
__global__ __launch_bounds__(1024, 1)
void loco_outer_tcgen_nvfp4_m128n256(
float* __restrict__ factor,
const uint8_t* __restrict__ packed_values,
const uint8_t* __restrict__ packed_scales,
int batch,
int n,
int start_panel,
int history_begin,
int history_end,
int trailing_panels,
int trailing_rows,
int history_panels
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoNvfp4OuterShared& shared =
*reinterpret_cast<LocoNvfp4OuterShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
uint32_t done =
(uint32_t)__cvta_generic_to_shared(&shared.mma_done);
if (tid == 0) {
loco_mbar_init(done);
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
uint32_t mma_phase = 0u;
int row_pairs = (trailing_panels + 1) >> 1;
uint64_t grouped_tasks =
loco_m128n256_task_prefix(row_pairs);
uint64_t tasks = (uint64_t)batch * grouped_tasks;
for (uint64_t task = blockIdx.x;
task < tasks;
task += gridDim.x) {
int b = (int)(task / grouped_tasks);
uint64_t local =
task - (uint64_t)b * grouped_tasks;
int low = 0;
int high = row_pairs - 1;
while (low < high) {
int middle = (low + high + 1) >> 1;
if (loco_m128n256_task_prefix(middle) <= local) {
low = middle;
} else {
high = middle - 1;
}
}
int row_pair = low;
uint64_t row_prefix =
loco_m128n256_task_prefix(row_pair);
int column_group = (int)(local - row_prefix);
int row_local = row_pair << 1;
int col_local = column_group << 2;
int valid_m_panels = trailing_panels - row_local;
if (valid_m_panels > 2) valid_m_panels = 2;
int valid_n_panels = trailing_panels - col_local;
if (valid_n_panels > 4) valid_n_panels = 4;
int row_base = (start_panel + row_local) << 6;
int col_base = (start_panel + col_local) << 6;
int valid_m_rows = valid_m_panels << 6;
int valid_n_rows = valid_n_panels << 6;
int row_workspace = row_local << 6;
int col_workspace = col_local << 6;
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
for (int vector = tid; vector < 2048; vector += 1024) {
reinterpret_cast<float4*>(shared.scratch)[vector] =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
__syncthreads();
// Four N64 slices x eight warps cover the full M128 TMEM tile.
int n_segment = warp >> 3;
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
uint32_t d_taddr =
(shared.taddr_slot + (uint32_t)(n_segment << 6))
| ((uint32_t)lane_base << 16);
if (lane_base < valid_m_rows
&& n_segment < valid_n_panels) {
int local_row =
lane_base + (lane >> 2) + ((lane & 1) << 3);
int local_column = (lane & 2) << 1;
const float* src = factor
+ matrix_base
+ (uint64_t)(row_base + local_row) * n
+ col_base
+ (n_segment << 6)
+ local_column;
loco_tmem_import_global(d_taddr, src);
} else {
int local_row = lane_base + (lane >> 2);
int local_column = lane & 3;
uint32_t src = (uint32_t)__cvta_generic_to_shared(
shared.scratch
+ (local_row << 6)
+ local_column);
loco_tmem_import(d_taddr, src);
}
__syncthreads();
// Four K64 packets share one TCGEN completion edge.
for (int history = history_begin;
history < history_end;
history += 4) {
int issue_count = history_end - history;
if (issue_count > 4) issue_count = 4;
// A: four issues x two K32 chunks x M128.
{
int issue = tid >> 8;
int local_a = tid & 255;
int chunk = local_a >> 7;
int row = local_a & 127;
uint4 value = make_uint4(0u, 0u, 0u, 0u);
if (issue < issue_count && row < valid_m_rows) {
uint64_t offset = loco_nvfp4_value_offset(
b,
n,
history - history_begin + issue,
history_panels,
chunk,
trailing_rows,
row_workspace + row);
value = *reinterpret_cast<const uint4*>(
packed_values + offset);
}
*reinterpret_cast<uint4*>(
shared.packet_a[issue]
+ chunk * 2048
+ row * 16) = value;
}
// B: four issues x two K32 chunks x N256.
for (int vector = tid; vector < 2048; vector += 1024) {
int issue = vector >> 9;
int local_b = vector & 511;
int chunk = local_b >> 8;
int row = local_b & 255;
uint4 value = make_uint4(0u, 0u, 0u, 0u);
if (issue < issue_count && row < valid_n_rows) {
uint64_t offset = loco_nvfp4_value_offset(
b,
n,
history - history_begin + issue,
history_panels,
chunk,
trailing_rows,
col_workspace + row);
value = *reinterpret_cast<const uint4*>(
packed_values + offset);
}
*reinterpret_cast<uint4*>(
shared.packet_b[issue]
+ chunk * 4096
+ row * 16) = value;
}
if (tid < 512) {
int issue = tid >> 7;
int row = tid & 127;
uint32_t word = 0u;
if (issue < issue_count && row < valid_m_rows) {
uint64_t offset = loco_nvfp4_scale_offset(
b,
history - history_begin + issue,
history_panels,
trailing_rows,
row_workspace + row);
word = *reinterpret_cast<const uint32_t*>(
packed_scales + offset);
}
reinterpret_cast<uint32_t*>(
&shared.scale_a[0][0])[tid] = word;
}
{
int issue = tid >> 8;
int row = tid & 255;
uint32_t word = 0u;
if (issue < issue_count && row < valid_n_rows) {
uint64_t offset = loco_nvfp4_scale_offset(
b,
history - history_begin + issue,
history_panels,
trailing_rows,
col_workspace + row);
word = *reinterpret_cast<const uint32_t*>(
packed_scales + offset);
}
reinterpret_cast<uint32_t*>(
&shared.scale_b[0][0])[tid] = word;
}
__syncthreads();
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
if (warp < 4) {
uint32_t partition = loco_tmem_reader_addr(
shared.taddr_slot, warp);
loco_nvfp4_store_scales(
partition,
&shared.scale_a[0][0],
&shared.scale_b[0][0]);
}
__syncthreads();
if (issuer) {
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int issue = 0; issue < 4; ++issue) {
if (issue < issue_count) {
loco_nvfp4_update_m128n256k64(
shared.taddr_slot,
(uint32_t)__cvta_generic_to_shared(
shared.packet_a[issue]),
(uint32_t)__cvta_generic_to_shared(
shared.packet_b[issue]),
issue);
}
}
loco_persistent_plane_commit(done);
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
// Export one N64 slice at a time through the 32 KiB scratch board.
#pragma unroll
for (int segment = 0; segment < 4; ++segment) {
if (warp < 8) {
int export_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int export_row =
export_base + (lane >> 2);
int export_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)export_base << 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.scratch
+ (export_row << 6)
+ export_col);
loco_tmem_export(taddr, dst);
}
__syncthreads();
if (segment < valid_n_panels) {
for (int vector = tid;
vector < 2048;
vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
if (row < valid_m_rows) {
float4 value =
reinterpret_cast<float4*>(
shared.scratch)[vector];
int global_row = row_base + row;
int global_col4 =
col_base + (segment << 6) + col4;
if (global_row >= global_col4 + 3) {
*reinterpret_cast<float4*>(
factor
+ matrix_base
+ (uint64_t)global_row * n
+ global_col4) = value;
} else if (global_row >= global_col4) {
float* words =
reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
int global_col = global_col4 + e;
if (global_row >= global_col) {
factor[
matrix_base
+ (uint64_t)global_row * n
+ global_col] = words[e];
}
}
}
}
}
}
__syncthreads();
}
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
__global__ __launch_bounds__(32, 1)
void xrshape1024_b4_barrier_init(
uint32_t* __restrict__ packed_fp16
) {
if (threadIdx.x == 0) {
packed_fp16[
((uint64_t)(blockIdx.x + 1) << 20) - 2] = 0u;
packed_fp16[
((uint64_t)(blockIdx.x + 1) << 20) - 1] = 0u;
}
}
__device__ __forceinline__ void xrshape1024_b4_barrier(
uint32_t* control,
uint32_t phase
) {
__syncthreads();
if (threadIdx.x == 0) {
if (atomicAdd(control, 1u) == 35u) {
atomicExch(control, 0u);
__threadfence();
atomicExch(control + 1, phase);
} else {
while (atomicAdd(control + 1, 0u) < phase) {
asm volatile("nanosleep.u32 64;");
}
}
}
__syncthreads();
}
__device__ __forceinline__ void xrshape1024_b4_tmem_to_tile(
uint32_t d_tmem,
LocoMaterializeShared& shared,
float* tile
) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
#pragma unroll
for (int segment = 0; segment < 2; ++segment) {
if (warp < 8) {
int export_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int export_row = export_base + (lane >> 2);
int export_col = lane & 3;
loco_tmem_export(
(d_tmem + (uint32_t)(segment << 6))
| ((uint32_t)export_base << 16),
(uint32_t)__cvta_generic_to_shared(
shared.scratch
+ (export_row << 6)
+ export_col));
}
__syncthreads();
for (int vector = threadIdx.x;
vector < 2048;
vector += 1024) {
float4 value =
reinterpret_cast<float4*>(shared.scratch)[vector];
float* word = reinterpret_cast<float*>(&value);
int row = vector >> 4;
int col = (segment << 6) + ((vector & 15) << 2);
tile[row * 129 + col] = word[0];
tile[row * 129 + col + 1] = word[1];
tile[row * 129 + col + 2] = word[2];
tile[row * 129 + col + 3] = word[3];
}
__syncthreads();
}
}
__device__ __forceinline__ void xrshape1024_b4_factor(
float* __restrict__ factor,
uint64_t matrix_base,
int row_base,
float* tile
) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
#pragma unroll 1
for (int p = 0; p < 128; ++p) {
if (warp == 0) {
float reciprocal = 0.0f;
if (lane == 0) {
float diagonal = xcal_sqrt_approx(
fmaxf(tile[p * 129 + p], 1.0e-20f));
tile[p * 129 + p] = diagonal;
reciprocal = xcal_rcp_approx(diagonal);
}
reciprocal =
__shfl_sync(0xffffffffu, reciprocal, 0);
for (int row = p + 1 + lane;
row < 128;
row += 32) {
tile[row * 129 + p] *= reciprocal;
}
}
__syncthreads();
for (int row = p + 1 + warp;
row < 128;
row += 32) {
float lip = tile[row * 129 + p];
for (int col = p + 1 + lane;
col <= row;
col += 32) {
tile[row * 129 + col] = fmaf(
-lip,
tile[col * 129 + p],
tile[row * 129 + col]);
}
}
__syncthreads();
}
for (int x = threadIdx.x; x < 128 * 128; x += 1024) {
int row = x >> 7;
int col = x & 127;
if (row >= col) {
factor[
matrix_base
+ (uint64_t)(row_base + row) * 1024
+ row_base
+ col] = tile[row * 129 + col];
}
}
}
__device__ __forceinline__ void xrshape1024_b4_trsm(
float* __restrict__ factor,
uint32_t* __restrict__ packed_fp16,
uint64_t matrix_base,
int b,
int row_base,
int col_base,
float* tile
) {
float* diagonal = tile + 128 * 129;
float* inverse = diagonal + 128 * 129;
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
for (int x = threadIdx.x; x < 128 * 128; x += 1024) {
diagonal[x + (x >> 7)] = factor[
matrix_base
+ (uint64_t)(col_base + (x >> 7)) * 1024
+ col_base
+ (x & 127)];
}
if (threadIdx.x < 128) {
inverse[threadIdx.x] = xcal_rcp_approx(
factor[
matrix_base
+ (uint64_t)(col_base + threadIdx.x) * 1024
+ col_base
+ threadIdx.x]);
}
__syncthreads();
for (int row = warp; row < 128; row += 32) {
float x0 = tile[row * 129 + lane];
float x1 = tile[row * 129 + lane + 32];
float x2 = tile[row * 129 + lane + 64];
float x3 = tile[row * 129 + lane + 96];
#pragma unroll 1
for (int c = 0; c < 128; ++c) {
float partial = lane < c
? x0 * diagonal[c * 129 + lane]
: 0.0f;
if (lane + 32 < c) {
partial = fmaf(
x1,
diagonal[c * 129 + lane + 32],
partial);
}
if (lane + 64 < c) {
partial = fmaf(
x2,
diagonal[c * 129 + lane + 64],
partial);
}
if (lane + 96 < c) {
partial = fmaf(
x3,
diagonal[c * 129 + lane + 96],
partial);
}
partial +=
__shfl_down_sync(0xffffffffu, partial, 16);
partial +=
__shfl_down_sync(0xffffffffu, partial, 8);
partial +=
__shfl_down_sync(0xffffffffu, partial, 4);
partial +=
__shfl_down_sync(0xffffffffu, partial, 2);
partial +=
__shfl_down_sync(0xffffffffu, partial, 1);
float sum =
__shfl_sync(0xffffffffu, partial, 0);
int owner = c & 31;
float rhs = c < 32
? __shfl_sync(0xffffffffu, x0, owner)
: c < 64
? __shfl_sync(0xffffffffu, x1, owner)
: c < 96
? __shfl_sync(
0xffffffffu, x2, owner)
: __shfl_sync(
0xffffffffu, x3, owner);
float value = (rhs - sum) * inverse[c];
if (lane == owner) {
if (c < 32) x0 = value;
else if (c < 64) x1 = value;
else if (c < 96) x2 = value;
else x3 = value;
}
}
factor[
matrix_base
+ (uint64_t)(row_base + row) * 1024
+ col_base
+ lane] = x0;
factor[
matrix_base
+ (uint64_t)(row_base + row) * 1024
+ col_base
+ lane
+ 32] = x1;
factor[
matrix_base
+ (uint64_t)(row_base + row) * 1024
+ col_base
+ lane
+ 64] = x2;
factor[
matrix_base
+ (uint64_t)(row_base + row) * 1024
+ col_base
+ lane
+ 96] = x3;
float y0 = __shfl_down_sync(0xffffffffu, x0, 1);
float y1 = __shfl_down_sync(0xffffffffu, x1, 1);
float y2 = __shfl_down_sync(0xffffffffu, x2, 1);
float y3 = __shfl_down_sync(0xffffffffu, x3, 1);
if (!(lane & 1)) {
uint32_t h0;
uint32_t h1;
uint32_t h2;
uint32_t h3;
uint32_t l0;
uint32_t l1;
uint32_t l2;
uint32_t l3;
loco_f16x2_hi_lo(x0, y0, h0, l0);
loco_f16x2_hi_lo(x1, y1, h1, l1);
loco_f16x2_hi_lo(x2, y2, h2, l2);
loco_f16x2_hi_lo(x3, y3, h3, l3);
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane, 0)] = h0;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane, 0) + 32] = l0;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane + 32, 0)] = h1;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane + 32, 0) + 32] = l1;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane + 64, 0)] = h2;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane + 64, 0) + 32] = l2;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane + 96, 0)] = h3;
packed_fp16[loco_fp16_packed_word(
b, 1024, row_base + row, col_base + lane + 96, 0) + 32] = l3;
}
}
__syncthreads();
}
__device__ __forceinline__ void xrshape1024_b4_update(
LocoMaterializeShared& shared,
uint32_t d_tmem,
uint32_t* __restrict__ packed_fp16,
int b,
int row_base,
int col_base,
int history_col,
uint32_t issuer,
uint32_t& mma_phase
) {
for (int vector = threadIdx.x;
vector < 2048;
vector += 1024) {
int issue = vector >> 8;
int local = vector & 255;
int half = local >> 7;
int group = (local >> 3) & 15;
int rail = (local >> 2) & 1;
int row_pair = local & 3;
int row = (group << 3) + (row_pair << 1);
int k = history_col
+ (issue << 4)
+ (half << 3);
uint4 a0;
uint4 a1;
uint4 b0;
uint4 b1;
loco_ld_global_u8(
packed_fp16 + loco_fp16_packed_word(
b, 1024, row_base + row, k, rail),
a0,
a1);
loco_ld_global_u8(
packed_fp16 + loco_fp16_packed_word(
b, 1024, col_base + row, k, rail),
b0,
b1);
int packet = issue * 4096
+ (half << 10)
+ (group << 6)
+ (rail << 5)
+ (row_pair << 3);
loco_st_shared_u8(
(uint32_t)__cvta_generic_to_shared(
shared.packet + packet),
a0,
a1);
loco_st_shared_u8(
(uint32_t)__cvta_generic_to_shared(
shared.packet + packet + 2048),
b0,
b1);
}
__syncthreads();
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
uint32_t done = (uint32_t)__cvta_generic_to_shared(
&shared.mma_done);
if (issuer) {
uint32_t packets =
(uint32_t)__cvta_generic_to_shared(shared.packet);
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int issue = 0; issue < 8; ++issue) {
loco_fp16x3_update_m128n128(
d_tmem,
packets + issue * 16384);
}
asm volatile(
"tcgen05.commit.cta_group::1."
"mbarrier::arrive::one.b64 [%0];"
:: "r"(done)
: "memory");
}
if ((threadIdx.x >> 5) < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
__global__ __launch_bounds__(1024, 1)
void xrshape1024_residual_owner_b4_kernel(
const float* __restrict__ input,
float* __restrict__ factor,
uint32_t* __restrict__ packed_fp16
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoMaterializeShared& shared =
*reinterpret_cast<LocoMaterializeShared*>(dynamic_smem);
int tid = threadIdx.x;
int warp = tid >> 5;
int lane = tid & 31;
int b = blockIdx.x / 36;
int tile = blockIdx.x - b * 36;
int tile_col = tile >= 35 ? 7
: tile >= 33 ? 6
: tile >= 30 ? 5
: tile >= 26 ? 4
: tile >= 21 ? 3
: tile >= 15 ? 2
: tile >= 8 ? 1
: 0;
int tile_row = tile_col == 0 ? tile
: tile_col == 1 ? tile - 7
: tile_col == 2 ? tile - 13
: tile_col == 3 ? tile - 18
: tile_col == 4 ? tile - 22
: tile_col == 5 ? tile - 25
: tile_col == 6 ? tile - 27
: 7;
uint64_t matrix_base = (uint64_t)b << 20;
uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
if (tid == 0) {
loco_mbar_init(
(uint32_t)__cvta_generic_to_shared(
&shared.mma_done));
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
// One lower N128 tile remains D-resident for its entire lifetime:
//
// 36 owners/matrix x B4 = 144 CTA1024
// D(i,j) -= L(i,k) L(j,k)^T, k < j
//
// A tile leaves TMEM only when its column becomes the live frontier.
if (warp < 16) {
int segment = warp >> 3;
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
int local_row =
lane_base
+ (lane >> 2)
+ ((lane & 1) << 3);
int local_column = (lane & 2) << 1;
loco_tmem_import_global(
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16),
input
+ matrix_base
+ (uint64_t)((tile_row << 7) + local_row) * 1024
+ (tile_col << 7)
+ (segment << 6)
+ local_column);
}
__syncthreads();
// Every owner also clears its unique reflected upper tile. Diagonal
// owners clear only the strict upper half of their N128 block.
for (int vector = tid; vector < 4096; vector += 1024) {
int row = vector >> 5;
int col = (vector & 31) << 2;
if (tile_row != tile_col) {
*reinterpret_cast<float4*>(
factor
+ matrix_base
+ (uint64_t)((tile_col << 7) + row) * 1024
+ (tile_row << 7)
+ col) = make_float4(
0.0f, 0.0f, 0.0f, 0.0f);
} else if (col > row) {
*reinterpret_cast<float4*>(
factor
+ matrix_base
+ (uint64_t)((tile_row << 7) + row) * 1024
+ (tile_col << 7)
+ col) = make_float4(
0.0f, 0.0f, 0.0f, 0.0f);
} else if (col + 3 > row) {
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (col + e > row) {
factor[
matrix_base
+ (uint64_t)((tile_row << 7) + row) * 1024
+ (tile_col << 7)
+ col
+ e] = 0.0f;
}
}
}
}
uint32_t* control =
packed_fp16 + ((uint64_t)(b + 1) << 20) - 2;
uint32_t grid_phase = 0u;
uint32_t mma_phase = 0u;
float* tile_smem =
reinterpret_cast<float*>(shared.packet);
#pragma unroll 1
for (int k = 0; k < 8; ++k) {
if (tile_row == k && tile_col == k) {
xrshape1024_b4_tmem_to_tile(
shared.taddr_slot,
shared,
tile_smem);
xrshape1024_b4_factor(
factor,
matrix_base,
k << 7,
tile_smem);
// Every factor publisher releases its own Lkk word before this
// CTA contributes the phase arrival.
__threadfence();
}
xrshape1024_b4_barrier(
control, ++grid_phase);
if (tile_row > k && tile_col == k) {
xrshape1024_b4_tmem_to_tile(
shared.taddr_slot,
shared,
tile_smem);
xrshape1024_b4_trsm(
factor,
packed_fp16,
matrix_base,
b,
tile_row << 7,
k << 7,
tile_smem);
// Only panel owners publish cross-CTA FP32/BF16 state.
__threadfence();
}
xrshape1024_b4_barrier(
control, ++grid_phase);
if (tile_col > k) {
xrshape1024_b4_update(
shared,
shared.taddr_slot,
packed_fp16,
b,
tile_row << 7,
tile_col << 7,
k << 7,
issuer,
mma_phase);
}
xrshape1024_b4_barrier(
control, ++grid_phase);
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 "
"%0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
struct LocoTrailingClusterBoard {
// active[x] is the measured device-wide resident cluster count for
// G={1,2,4,8,16}. G1 is the ordinary one-CTA cluster.
int active[5];
};
static void loco_measure_trailing_cluster_board(
LocoTrailingClusterBoard& board,
int resident
) {
const int gangs[5] = {1, 2, 4, 8, 16};
board.active[0] = resident;
for (int x = 1; x < 5; ++x) {
int gang = gangs[x];
board.active[x] = 0;
if (gang > XCAL_XRCHOLV18_MAX_G) continue;
cudaLaunchAttribute attribute[1]{};
attribute[0].id = cudaLaunchAttributeClusterDimension;
attribute[0].val.clusterDim.x = gang;
attribute[0].val.clusterDim.y = 1;
attribute[0].val.clusterDim.z = 1;
cudaLaunchConfig_t occupancy{};
// A deliberately long candidate grid prevents the occupancy query
// from reporting only the one cluster present in a short grid.
occupancy.gridDim = dim3((unsigned int)(gang * 1024), 1, 1);
occupancy.blockDim = dim3(1024, 1, 1);
occupancy.dynamicSmemBytes =
(int)sizeof(LocoMaterializeShared);
occupancy.attrs = attribute;
occupancy.numAttrs = 1;
int active = 0;
cudaError_t status = cudaOccupancyMaxActiveClusters(
&active,
(const void*)loco_trailing_tcgen_m128n256<false>,
&occupancy);
if (status == cudaSuccess && active > 0) {
board.active[x] = active;
} else if (status != cudaSuccess) {
// An unsupported candidate is not a launch failure; it simply
// cannot participate in the selector.
(void)cudaGetLastError();
}
}
}
static cudaError_t loco_launch_outer_scaled_n_parallel(
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int start_panel,
int history_begin,
int history_end,
int trailing_panels,
int owner_width,
int resident
) {
int history_panels = history_end - history_begin;
if (history_panels < 2
|| (history_panels & 1)
|| (trailing_panels & 1)
|| owner_width < 128
|| (owner_width & (owner_width - 1))) {
return cudaErrorInvalidValue;
}
int trailing128 = trailing_panels >> 1;
int owner128 = owner_width >> 7;
uint64_t useful =
(uint64_t)trailing128 * (trailing128 + 1) >> 1;
uint64_t total = (uint64_t)batch * useful;
int grid = total < (uint64_t)resident
? (int)total
: resident;
if (n >= 4096) {
loco_outer_scaled_n_parallel<true>
<<<grid,
256,
(int)sizeof(XrShape512UpdateShared)>>>(
factor,
packed_fp16,
batch,
n,
start_panel,
history_begin,
history_end,
trailing128,
owner128);
} else {
loco_outer_scaled_n_parallel<false>
<<<grid,
256,
(int)sizeof(XrShape512UpdateShared)>>>(
factor,
packed_fp16,
batch,
n,
start_panel,
history_begin,
history_end,
trailing128,
owner128);
}
return cudaGetLastError();
}
static cudaError_t loco_launch_trailing_tcgen_m128n256(
const LocoTrailingClusterBoard& board,
float* factor,
const uint32_t* packed_fp16,
int packed_fp16_atoms,
int batch,
int n,
int start_panel,
int history_begin,
int history_end,
int trailing_panels,
int frontier_only,
int frontier_segments,
int resident
) {
int row_pairs = (trailing_panels + 1) >> 1;
bool overlap_history =
n >= 4096 && !packed_fp16_atoms;
if (frontier_only == 2 || frontier_only == 3) {
uint64_t useful = frontier_only == 3
? ((uint64_t)row_pairs << 1) - 1
: (uint64_t)row_pairs * (row_pairs + 1) >> 1;
uint64_t total = (uint64_t)batch * useful;
int grid = total < (uint64_t)resident
? (int)total
: resident;
if (frontier_only == 3) {
if (overlap_history) {
loco_trailing_tcgen_m128n256<true, true, true>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
0,
1);
} else {
loco_trailing_tcgen_m128n256<true, true>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
0,
1);
}
return cudaGetLastError();
}
if (overlap_history) {
loco_trailing_tcgen_m128n256<true, false, true>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
0,
1);
} else {
loco_trailing_tcgen_m128n256<true>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
0,
1);
}
return cudaGetLastError();
}
uint64_t useful = frontier_only
? (uint64_t)row_pairs
: loco_m128n256_task_prefix(row_pairs);
const int gangs[5] = {1, 2, 4, 8, 16};
int candidate_k[5] = {0, 0, 0, 0, 0};
int candidate_effective[5] = {0, 0, 0, 0, 0};
int max_effective = 0;
// K_b(G) = min(resident_clusters(G)/B,
// ceil(useful_tiles_per_matrix/G)).
//
// The effective worker score clips padded tail ranks. Among candidates
// within 95% of the best score, prefer larger G: it preserves the
// common-B cluster geometry without sacrificing a useful machine wave.
for (int x = 0; x < 5; ++x) {
int gang = gangs[x];
if (gang > XCAL_XRCHOLV18_MAX_G
|| (uint64_t)gang > useful
|| board.active[x] < batch) {
continue;
}
int clusters = board.active[x] / batch;
uint64_t useful_clusters =
(useful + (uint64_t)gang - 1u) / (uint64_t)gang;
if ((uint64_t)clusters > useful_clusters) {
clusters = (int)useful_clusters;
}
if (clusters < 1) continue;
int workers = clusters * gang;
int effective = (uint64_t)workers < useful
? workers
: (int)useful;
candidate_k[x] = clusters;
candidate_effective[x] = effective;
if (effective > max_effective) {
max_effective = effective;
}
}
int selected = -1;
for (int x = 0; x < 5; ++x) {
if (candidate_k[x] > 0
&& candidate_effective[x] * 20 >= max_effective * 19) {
selected = x;
}
}
// If no cluster candidate can give every live matrix a resident team,
// retain v21's unrestricted global task board.
if (selected < 0) {
uint64_t total = (uint64_t)batch * useful;
int grid = total < (uint64_t)resident
? (int)total
: resident;
if (overlap_history) {
loco_trailing_tcgen_m128n256<false, false, true>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
0,
1);
} else {
loco_trailing_tcgen_m128n256<false>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
0,
1);
}
return cudaGetLastError();
}
int gang = gangs[selected];
int clusters = candidate_k[selected];
int grid = batch * clusters * gang;
if (gang == 1) {
if (overlap_history) {
loco_trailing_tcgen_m128n256<false, false, true>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
clusters,
gang);
} else {
loco_trailing_tcgen_m128n256<false>
<<<grid, 1024, (int)sizeof(LocoMaterializeShared)>>>(
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
frontier_only,
frontier_segments,
clusters,
gang);
}
return cudaGetLastError();
}
cudaLaunchAttribute attribute[1]{};
attribute[0].id = cudaLaunchAttributeClusterDimension;
attribute[0].val.clusterDim.x = gang;
attribute[0].val.clusterDim.y = 1;
attribute[0].val.clusterDim.z = 1;
cudaLaunchConfig_t launch{};
launch.gridDim = dim3((unsigned int)grid, 1, 1);
launch.blockDim = dim3(1024, 1, 1);
launch.dynamicSmemBytes = (int)sizeof(LocoMaterializeShared);
launch.attrs = attribute;
launch.numAttrs = 1;
void* args[] = {
&factor,
&packed_fp16,
&packed_fp16_atoms,
&batch,
&n,
&start_panel,
&history_begin,
&history_end,
&trailing_panels,
&frontier_only,
&frontier_segments,
&clusters,
&gang,
};
if (overlap_history) {
return cudaLaunchKernelExC(
&launch,
(const void*)loco_trailing_tcgen_m128n256<
false,
false,
true>,
args);
}
return cudaLaunchKernelExC(
&launch,
(const void*)loco_trailing_tcgen_m128n256<false>,
args);
}
__global__ __launch_bounds__(256, 2)
void loco_quantize_outer_planes(
const float* __restrict__ factor,
uint8_t* __restrict__ packed_values,
uint8_t* __restrict__ packed_scales,
int batch,
int n,
int row_begin,
int history_begin,
int history_panels,
int trailing_rows
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int group = lane >> 4;
int rank = lane & 15;
uint32_t block_mask = 0xffffu << (group << 4);
uint64_t tasks =
(uint64_t)batch * history_panels * trailing_rows;
// _____________________________________________________________________
// FP32 BLOCK32 -> ONE S2F6 RAIL -> FOUR E2M1 PLANES
//
// q0 = RN_s2f6(x / S)
// x' = S*q0/64
// q = d0 + 4*d1 + 16*d2 + 64*d3
//
// This preamble replaces TF32x3 only on the K128 sibling dependency edge.
// Every lower frontier tile consumes the same packed E2M1 bytes.
// _____________________________________________________________________
for (uint64_t task = (uint64_t)blockIdx.x * 8u + warp;
task < tasks;
task += (uint64_t)gridDim.x * 8u) {
int row = (int)(task % trailing_rows);
uint64_t q = task / trailing_rows;
int history = (int)(q % history_panels);
int b = (int)(q / history_panels);
int k0 = (group << 5) + (rank << 1);
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
const float* src = factor
+ matrix_base
+ (uint64_t)(row_begin + row) * n
+ ((history_begin + history) << 6)
+ k0;
float x0 = src[0];
float x1 = src[1];
float maximum = fmaxf(fabsf(x0), fabsf(x1));
maximum = fmaxf(
maximum,
__shfl_down_sync(block_mask, maximum, 8, 16));
maximum = fmaxf(
maximum,
__shfl_down_sync(block_mask, maximum, 4, 16));
maximum = fmaxf(
maximum,
__shfl_down_sync(block_mask, maximum, 2, 16));
maximum = fmaxf(
maximum,
__shfl_down_sync(block_mask, maximum, 1, 16));
maximum = __shfl_sync(block_mask, maximum, 0, 16);
uint8_t scale_code = 0u;
float scale = 0.0f;
if (rank == 0) {
scale = loco_outer_s2f6_scale(maximum, scale_code);
}
scale_code = (uint8_t)__shfl_sync(
block_mask, (uint32_t)scale_code, 0, 16);
scale = __shfl_sync(block_mask, scale, 0, 16);
uint16_t rail0 = loco_outer_s2f6x2(x0, x1, scale_code);
int q0 = (int)(int8_t)(rail0 & 0xffu);
int q1 = (int)(int8_t)(rail0 >> 8);
int chunk = lane >> 4;
int byte = lane & 15;
#pragma unroll
for (int plane = 0;
plane < LOCO_OUTER_ACTIVE_PLANES;
++plane) {
int r0 = loco_outer_s2f6_digit(q0, plane);
int r1 = loco_outer_s2f6_digit(q1, plane);
uint32_t packed = loco_outer_pack_digits(r0, r1);
uint64_t value_offset = loco_outer_plane_value_offset(
b,
history,
history_panels,
plane,
chunk,
trailing_rows,
row);
packed_values[value_offset + byte] = (uint8_t)packed;
}
if (rank == 0) {
uint64_t scale_offset = loco_plane_scale_offset(
b,
history,
history_panels,
trailing_rows,
row);
packed_scales[scale_offset + group] = scale_code;
if (group == 0) {
packed_scales[scale_offset + 2] = 0u;
packed_scales[scale_offset + 3] = 0u;
}
}
}
#endif
}
__global__ __launch_bounds__(1024, 1)
void loco_frontier_e2m1x4_m128n256(
float* __restrict__ factor,
const uint8_t* __restrict__ packed_values,
const uint8_t* __restrict__ packed_scales,
int batch,
int n,
int row_begin,
int history_panels,
int trailing_rows
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoFrontierE2M1x4Shared& shared =
*reinterpret_cast<LocoFrontierE2M1x4Shared*>(dynamic_smem);
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
const uint32_t done = (uint32_t)__cvta_generic_to_shared(
&shared.mma_done);
uint32_t mma_phase = 0u;
if (tid == 0) {
loco_mbar_init(done);
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
int row_tiles = trailing_rows >> 7;
uint64_t tasks = (uint64_t)batch * row_tiles;
for (uint64_t task = blockIdx.x; task < tasks; task += gridDim.x) {
int b = (int)(task / row_tiles);
int row_tile = (int)(task - (uint64_t)b * row_tiles);
int row_local_base = row_tile << 7;
int row_base = row_begin + row_local_base;
int col_base = row_begin;
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
// One immutable zero M128xN64 tile seeds both padded N64 quarters.
for (int vector = tid; vector < 2048; vector += 1024) {
reinterpret_cast<float4*>(shared.scratch)[vector] =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
__syncthreads();
// Current FP32 sibling residual -> TMEM M128xN128; N128:255 is zero.
int segment = warp >> 3;
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
uint32_t d_taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
if (segment < 2) {
int local_row =
lane_base + (lane >> 2) + ((lane & 1) << 3);
int local_col = (lane & 2) << 1;
const float* src = factor
+ matrix_base
+ (uint64_t)(row_base + local_row) * n
+ col_base
+ (segment << 6)
+ local_col;
loco_tmem_import_global(d_taddr, src);
} else {
int local_row = lane_base + (lane >> 2);
int local_col = lane & 3;
uint32_t src = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (local_row << 6) + local_col);
loco_tmem_import(d_taddr, src);
}
__syncthreads();
#pragma unroll
for (int history = 0; history < 2; ++history) {
// All four A planes: M128 x K64, two K32 chunks.
for (int vector = tid;
vector < LOCO_OUTER_ACTIVE_PLANES * 256;
vector += 1024) {
int plane = vector >> 8;
int local = vector & 255;
int chunk = local >> 7;
int row = local & 127;
uint64_t offset = loco_outer_plane_value_offset(
b,
history,
history_panels,
plane,
chunk,
trailing_rows,
row_local_base + row);
reinterpret_cast<uint4*>(shared.packet_a)[vector] =
loco_ld_global_u4(packed_values + offset);
}
// All four B planes: live N128 plus a zero-padded N128 half.
for (int vector = tid;
vector < LOCO_OUTER_ACTIVE_PLANES * 512;
vector += 1024) {
int plane = vector >> 9;
int local = vector & 511;
int chunk = local >> 8;
int row = local & 255;
uint4 value = make_uint4(0u, 0u, 0u, 0u);
if (row < 128) {
uint64_t offset = loco_outer_plane_value_offset(
b,
history,
history_panels,
plane,
chunk,
trailing_rows,
row);
value = loco_ld_global_u4(packed_values + offset);
}
reinterpret_cast<uint4*>(shared.packet_b)[vector] = value;
}
// Derive every plane's UE8M0 board once from the two base scales.
for (int x = tid;
x < LOCO_OUTER_ACTIVE_PLANES * 128;
x += 1024) {
int plane = x >> 7;
int row = x & 127;
uint64_t offset = loco_plane_scale_offset(
b,
history,
history_panels,
trailing_rows,
row_local_base + row);
uint32_t base_word =
*reinterpret_cast<const uint32_t*>(
packed_scales + offset);
reinterpret_cast<uint32_t*>(shared.scale_a)[x] =
loco_outer_scale_word(base_word, plane);
}
for (int x = tid;
x < LOCO_OUTER_ACTIVE_PLANES * 256;
x += 1024) {
int plane = x >> 8;
int row = x & 255;
uint32_t word = 0u;
if (row < 128) {
uint64_t offset = loco_plane_scale_offset(
b,
history,
history_panels,
trailing_rows,
row);
uint32_t base_word =
*reinterpret_cast<const uint32_t*>(
packed_scales + offset);
word = loco_outer_scale_word(base_word, plane);
}
reinterpret_cast<uint32_t*>(shared.scale_b)[x] = word;
}
__syncthreads();
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
if (warp < 4) {
uint32_t partition = loco_tmem_reader_addr(
shared.taddr_slot, warp);
loco_e2m1x4_store_all_scales(
partition,
shared.scale_a,
shared.scale_b);
}
__syncthreads();
// The complete 4x4 rail is sixteen MXFP4 products. Keep it in
// one asynchronous group: one completion edge per K64 history.
if (issuer) {
uint32_t packet_a =
(uint32_t)__cvta_generic_to_shared(shared.packet_a);
uint32_t packet_b =
(uint32_t)__cvta_generic_to_shared(shared.packet_b);
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
#pragma unroll
for (int plane_a = 0;
plane_a < LOCO_OUTER_ACTIVE_PLANES;
++plane_a) {
#pragma unroll
for (int plane_b = 0;
plane_b < LOCO_OUTER_ACTIVE_PLANES;
++plane_b) {
loco_e2m1x4_issue_m128n256k64(
shared.taddr_slot,
packet_a + plane_a * 4096,
packet_b + plane_b * 8192,
plane_a,
plane_b);
}
}
loco_persistent_plane_commit(done);
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
// Publish only the live N128 sibling. Preserve lower-only output.
#pragma unroll
for (int n64 = 0; n64 < 2; ++n64) {
if (warp < 8) {
int export_lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int export_row = export_lane_base + (lane >> 2);
int export_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(n64 << 6))
| ((uint32_t)export_lane_base << 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (export_row << 6) + export_col);
loco_tmem_export(taddr, dst);
}
__syncthreads();
for (int vector = tid; vector < 2048; vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
int global_row = row_base + row;
int global_col = col_base + (n64 << 6) + col4;
float4 value =
reinterpret_cast<float4*>(shared.scratch)[vector];
if (global_row >= global_col + 3) {
*reinterpret_cast<float4*>(
factor
+ matrix_base
+ (uint64_t)global_row * n
+ global_col) = value;
} else if (global_row >= global_col) {
float* word = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (global_row >= global_col + e) {
factor[
matrix_base
+ (uint64_t)global_row * n
+ global_col + e] = word[e];
}
}
}
}
__syncthreads();
}
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
__global__ __launch_bounds__(1024, 1)
void loco_outer_hexlift_m128n256(
float* __restrict__ factor,
const uint8_t* __restrict__ packed_values,
const uint8_t* __restrict__ packed_scales,
int batch,
int n,
int row_begin,
int column_begin,
int column_width,
int critical_end,
int history_begin,
int history_panels
) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000 || __CUDA_ARCH__ == 1200)
extern __shared__ __align__(128) unsigned char dynamic_smem[];
LocoFrontierE2M1x4Shared& shared =
*reinterpret_cast<LocoFrontierE2M1x4Shared*>(dynamic_smem);
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const uint32_t issuer = warp == 0 ? loco_elect_one() : 0u;
const uint32_t done = (uint32_t)__cvta_generic_to_shared(
&shared.mma_done);
uint32_t mma_phase = 0u;
if (tid == 0) {
loco_mbar_init(done);
asm volatile(
"fence.mbarrier_init.release.cluster;"
::: "memory");
}
if (warp == 16) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 "
"[%0], 512;"
:
: "r"((uint32_t)__cvta_generic_to_shared(
&shared.taddr_slot))
: "memory");
}
__syncthreads();
int trailing_rows = n - row_begin;
int row_tiles = trailing_rows >> 7;
int live_segments = column_width >> 6;
uint64_t tasks = (uint64_t)batch * (uint64_t)row_tiles;
for (uint64_t task = blockIdx.x; task < tasks; task += gridDim.x) {
int b = (int)(task / row_tiles);
int row_tile = (int)(task - (uint64_t)b * row_tiles);
int row_base = row_begin + (row_tile << 7);
bool hex6 =
!XCAL_XRCHOLV18_HEX4_GE8192
&& (XCAL_XRCHOLV18_HEX6_ALL
|| row_base < critical_end);
int first_plane = hex6
? LOCO_HEXLIFT_OUTER_FIRST_PLANE
: LOCO_HEXLIFT_OUTER_FIRST_PLANE + 1;
int plane_count =
LOCO_INNER_HEXLIFT_PLANES - first_plane;
uint64_t matrix_base =
(uint64_t)b * (uint64_t)n * (uint64_t)n;
for (int vector = tid; vector < 2048; vector += 1024) {
reinterpret_cast<float4*>(shared.scratch)[vector] =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
__syncthreads();
// Current FP32 residual enters TMEM once. A final N128 frontier uses
// the same native N256 instruction with two zero-padded segments.
int segment = warp >> 3;
int reader = warp & 7;
int lane_base =
((reader & 3) << 5) + ((reader >> 2) << 4);
uint32_t d_taddr =
(shared.taddr_slot + (uint32_t)(segment << 6))
| ((uint32_t)lane_base << 16);
if (segment < live_segments) {
int local_row =
lane_base + (lane >> 2) + ((lane & 1) << 3);
int local_col = (lane & 2) << 1;
const float* src = factor
+ matrix_base
+ (uint64_t)(row_base + local_row) * n
+ column_begin
+ (segment << 6)
+ local_col;
loco_tmem_import_global(d_taddr, src);
} else {
int local_row = lane_base + (lane >> 2);
int local_col = lane & 3;
uint32_t src = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (local_row << 6) + local_col);
loco_tmem_import(d_taddr, src);
}
__syncthreads();
for (int local_history = 0;
local_history < history_panels;
++local_history) {
int history = history_begin + local_history;
// Three retained A planes, two adjacent rows per 32-byte load.
constexpr int a_row_pairs = 64;
int a_value_tasks =
plane_count * 2 * a_row_pairs;
for (int x = tid; x < a_value_tasks; x += 1024) {
int pair = x % a_row_pairs;
int q = x / a_row_pairs;
int plane = first_plane + (q >> 1);
int chunk = q & 1;
int row = pair << 1;
const uint8_t* src =
packed_values + loco_hexlift_outer_value_offset(
b, n, history, plane, chunk, row_base + row);
uint4 first;
uint4 second;
loco_ld_global_u8_ca(
reinterpret_cast<const uint32_t*>(src),
first,
second);
uint32_t packet_offset =
(uint32_t)(plane * 4096
+ chunk * 2048
+ row * 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.packet_a + packet_offset);
loco_st_shared_u8(dst, first, second);
}
// B is N256. Width128 leaves its upper two segments zero.
constexpr int b_row_pairs = 128;
int b_value_tasks =
plane_count * 2 * b_row_pairs;
int live_b_pairs = column_width >> 1;
for (int x = tid; x < b_value_tasks; x += 1024) {
int pair = x % b_row_pairs;
int q = x / b_row_pairs;
int plane = first_plane + (q >> 1);
int chunk = q & 1;
int row = pair << 1;
uint4 first = make_uint4(0u, 0u, 0u, 0u);
uint4 second = make_uint4(0u, 0u, 0u, 0u);
if (pair < live_b_pairs) {
const uint8_t* src =
packed_values + loco_hexlift_outer_value_offset(
b,
n,
history,
plane,
chunk,
column_begin + row);
loco_ld_global_u8_ca(
reinterpret_cast<const uint32_t*>(src),
first,
second);
}
uint32_t packet_offset =
(uint32_t)(plane * 8192
+ chunk * 4096
+ row * 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.packet_b + packet_offset);
loco_st_shared_u8(dst, first, second);
}
for (int x = tid;
x < plane_count * 128;
x += 1024) {
int plane = first_plane + (x >> 7);
int row = x & 127;
uint32_t base_word = loco_ld_global_u32_ca(
reinterpret_cast<const uint32_t*>(
packed_scales
+ loco_hexlift_outer_scale_offset(
b, n, history, row_base + row)));
reinterpret_cast<uint32_t*>(
shared.scale_a + plane * 512)[row] =
loco_hexlift_outer_k16_scale_word(
base_word, plane);
}
for (int x = tid;
x < plane_count * 256;
x += 1024) {
int plane = first_plane + (x >> 8);
int row = x & 255;
uint32_t word = 0u;
if (row < column_width) {
uint32_t base_word = loco_ld_global_u32_ca(
reinterpret_cast<const uint32_t*>(
packed_scales
+ loco_hexlift_outer_scale_offset(
b,
n,
history,
column_begin + row)));
word =
loco_hexlift_outer_k16_scale_word(
base_word, plane);
}
reinterpret_cast<uint32_t*>(
shared.scale_b + plane * 1024)[row] = word;
}
__syncthreads();
asm volatile(
"fence.proxy.async.shared::cta;"
::: "memory");
if (warp < 4) {
uint32_t partition = loco_tmem_reader_addr(
shared.taddr_slot, warp);
loco_inner_hexlift_store_scales(
partition,
shared.scale_a,
shared.scale_b,
first_plane);
}
__syncthreads();
if (issuer) {
uint32_t packet_a =
(uint32_t)__cvta_generic_to_shared(shared.packet_a);
uint32_t packet_b =
(uint32_t)__cvta_generic_to_shared(shared.packet_b);
asm volatile(
"tcgen05.fence::after_thread_sync;"
::: "memory");
loco_hexlift_outer_issue_m128n256k64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b + 3 * 8192,
3,
3);
loco_hexlift_outer_issue_m128n256k64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b + 2 * 8192,
3,
2);
loco_hexlift_outer_issue_m128n256k64(
shared.taddr_slot,
packet_a + 2 * 4096,
packet_b + 3 * 8192,
2,
3);
if (hex6) {
loco_hexlift_outer_issue_m128n256k64(
shared.taddr_slot,
packet_a + 3 * 4096,
packet_b + 1 * 8192,
3,
1);
loco_hexlift_outer_issue_m128n256k64(
shared.taddr_slot,
packet_a + 1 * 4096,
packet_b + 3 * 8192,
1,
3);
}
loco_hexlift_outer_issue_m128n256k64(
shared.taddr_slot,
packet_a + 2 * 4096,
packet_b + 2 * 8192,
2,
2);
loco_persistent_plane_commit(done);
}
if (warp < 4) {
loco_mbar_wait(done, mma_phase);
}
__syncthreads();
mma_phase ^= 1u;
}
for (int n64 = 0; n64 < live_segments; ++n64) {
if (warp < 8) {
int export_lane_base =
((warp & 3) << 5) + ((warp >> 2) << 4);
int export_row = export_lane_base + (lane >> 2);
int export_col = lane & 3;
uint32_t taddr =
(shared.taddr_slot + (uint32_t)(n64 << 6))
| ((uint32_t)export_lane_base << 16);
uint32_t dst = (uint32_t)__cvta_generic_to_shared(
shared.scratch + (export_row << 6) + export_col);
loco_tmem_export(taddr, dst);
}
__syncthreads();
for (int vector = tid; vector < 2048; vector += 1024) {
int row = vector >> 4;
int col4 = (vector & 15) << 2;
int global_row = row_base + row;
int global_col = column_begin + (n64 << 6) + col4;
float4 value =
reinterpret_cast<float4*>(shared.scratch)[vector];
if (global_row >= global_col + 3) {
*reinterpret_cast<float4*>(
factor
+ matrix_base
+ (uint64_t)global_row * n
+ global_col) = value;
} else if (global_row >= global_col) {
float* words = reinterpret_cast<float*>(&value);
#pragma unroll
for (int e = 0; e < 4; ++e) {
if (global_row >= global_col + e) {
factor[
matrix_base
+ (uint64_t)global_row * n
+ global_col + e] = words[e];
}
}
}
}
__syncthreads();
}
}
if (warp == 16) {
asm volatile(
"tcgen05.fence::after_thread_sync;\n\t"
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, 512;\n\t"
"tcgen05.relinquish_alloc_permit."
"cta_group::1.sync.aligned;"
:
: "r"(shared.taddr_slot)
: "memory");
}
__syncthreads();
#endif
}
static bool loco_run_frontier_e2m1(
float* __restrict__ factor,
uint8_t* __restrict__ plane_values,
uint8_t* __restrict__ plane_scales,
int batch,
int n,
int history_begin,
int history_panels,
int row_begin,
int column_end,
int resident
) {
if (!plane_values
|| !plane_scales
|| history_panels != 2
|| row_begin >= n
|| column_end - row_begin != 128
|| (n & 127)
|| ((history_begin + history_panels) > (row_begin >> 6))) {
return false;
}
int trailing_rows = n - row_begin;
if (trailing_rows & 127) return false;
// One s2f6 rail becomes four E2M1 planes. Quantize the solved K128 edge
// once, then reuse its descriptor-native 4x4 products for every sibling.
uint64_t quant_tasks =
(uint64_t)batch * history_panels * trailing_rows;
uint64_t quant_blocks = (quant_tasks + 7u) >> 3;
int quant_grid = quant_blocks < (uint64_t)(resident * 4)
? (int)quant_blocks
: resident * 4;
if (quant_grid < 1) quant_grid = 1;
loco_quantize_outer_planes<<<quant_grid, 256>>>(
factor,
plane_values,
plane_scales,
batch,
n,
row_begin,
history_begin,
history_panels,
trailing_rows);
uint64_t tile_tasks =
(uint64_t)batch * (uint64_t)(trailing_rows >> 7);
int grid = tile_tasks < (uint64_t)resident
? (int)tile_tasks
: resident;
if (grid < 1) return false;
loco_frontier_e2m1x4_m128n256
<<<grid,
1024,
(int)sizeof(LocoFrontierE2M1x4Shared)>>>(
factor,
plane_values,
plane_scales,
batch,
n,
row_begin,
history_panels,
trailing_rows);
return true;
}
static bool loco_run_outer_hexlift_frontier(
float* __restrict__ factor,
const uint8_t* __restrict__ packed_values,
const uint8_t* __restrict__ packed_scales,
int batch,
int n,
int row_begin,
int column_begin,
int column_width,
int critical_end,
int history_begin,
int history_panels,
int resident
) {
if (!factor
|| !packed_values
|| !packed_scales
|| batch < 1
|| row_begin >= n
|| row_begin != column_begin
|| (column_width != 128 && column_width != 256)
|| critical_end < column_begin + column_width
|| critical_end > n
|| history_panels < 1
|| history_begin < 0
|| history_begin + history_panels > (column_begin >> 6)
|| (n & 127)
|| (critical_end & 127)
|| ((n - row_begin) & 127)) {
return false;
}
uint64_t tasks =
(uint64_t)batch * (uint64_t)((n - row_begin) >> 7);
int grid = tasks < (uint64_t)resident
? (int)tasks
: resident;
if (grid < 1) return false;
loco_outer_hexlift_m128n256
<<<grid,
1024,
(int)sizeof(LocoFrontierE2M1x4Shared)>>>(
factor,
packed_values,
packed_scales,
batch,
n,
row_begin,
column_begin,
column_width,
critical_end,
history_begin,
history_panels);
return true;
}
static cudaError_t loco_blas_batched_left_frontier(
cublasHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int panel,
int width);
static cudaError_t loco_lt_left_frontier(
cublasLtHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int panel,
int width,
void* workspace,
uint64_t workspace_bytes);
static cudaError_t xrshape1024_run_b4_k256(
const float* __restrict__ input,
float* __restrict__ factor,
uint32_t* __restrict__ packed_fp16,
int resident
) {
lower_copy_kernel<<<resident, 1024>>>(
input,
factor,
4ull << 20,
1024);
cudaError_t status = cudaGetLastError();
if (status != cudaSuccess) return status;
LocoTrailingClusterBoard board{};
board.active[0] = resident;
for (int panel = 0; panel < 1024; panel += 256) {
xrshape1024_factor256_b4_kernel
<<<4, 1024, XRSHAPE512_FACTOR_SMEM>>>(
factor,
panel);
status = cudaGetLastError();
if (status != cudaSuccess) return status;
int rows = 768 - panel;
if (rows <= 0) continue;
int row_groups = rows >> 5;
xrshape1024_trsm256_b4_kernel
<<<4 * row_groups,
1024,
XRSHAPE512_TRSM_SMEM>>>(
factor,
packed_fp16,
panel,
row_groups);
status = cudaGetLastError();
if (status != cudaSuccess) return status;
status = loco_launch_trailing_tcgen_m128n256(
board,
factor,
packed_fp16,
1,
4,
1024,
(panel + 256) >> 6,
panel >> 6,
(panel + 256) >> 6,
rows >> 6,
2,
2,
resident);
if (status != cudaSuccess) return status;
}
return cudaSuccess;
}
static cudaError_t xrshape1024_run_b4_frontier(
const float* __restrict__ input,
float* __restrict__ factor,
uint32_t* __restrict__ packed_fp16,
cublasLtHandle_t lt_handle,
cublasHandle_t blas_handle,
void* lt_workspace,
uint64_t lt_workspace_bytes,
int resident
) {
lower_copy_kernel<<<resident << 1, 1024>>>(
input,
factor,
4ull << 20,
1024);
cudaError_t status = cudaGetLastError();
if (status != cudaSuccess) return status;
LocoTrailingClusterBoard board{};
board.active[0] = resident;
// P128 dependency rail. Before a panel becomes active, one strided-
// batched GEMM consumes every previously solved FP16 column exactly once.
// The panel then stays native: register POTRF followed by MMA TRSM.
for (int panel = 0; panel < 1024; panel += 128) {
if (panel > 0) {
status = loco_lt_left_frontier(
lt_handle,
factor,
packed_fp16,
4,
1024,
panel,
128,
lt_workspace,
lt_workspace_bytes);
if (status != cudaSuccess) {
status = loco_blas_batched_left_frontier(
blas_handle,
factor,
packed_fp16,
4,
1024,
panel,
128);
}
if (status != cudaSuccess) return status;
}
xrshape128_diag_mid_kernel<true, true, true>
<<<4, 1024>>>(
factor,
packed_fp16,
4,
1024,
panel);
status = cudaGetLastError();
if (status != cudaSuccess) return status;
int rows = 896 - panel;
if (rows <= 0) continue;
// B4 has only four matrices, so choose M by total live rows rather
// than by per-matrix extent. M32 amortizes the 32 KiB diagonal
// sidecar while the board still exposes roughly one B200 wave;
// M16 takes over once M32 would leave most SMs idle.
if (rows >= 768) {
int row_tiles = (rows + 31) >> 5;
trsm128_mma_f16_kernel<32>
<<<4 * row_tiles,
256,
(8192 + 32 * 64 + 32 * 8) * 4>>>(
factor,
packed_fp16,
4,
1024,
panel,
rows);
} else {
int row_tiles = (rows + 15) >> 4;
trsm128_mma_f16_kernel<16>
<<<4 * row_tiles,
256,
(8192 + 16 * 64 + 16 * 8) * 4>>>(
factor,
packed_fp16,
4,
1024,
panel,
rows);
}
status = cudaGetLastError();
if (status != cudaSuccess) return status;
}
return cudaSuccess;
}
static cudaError_t xrts_encode_map_b60(
CUtensorMap* map,
float* factor
) {
uint64_t global_dim[2] = {1024u, 60u * 1024u};
uint64_t global_stride[1] = {1024u * sizeof(float)};
uint32_t box_dim[2] = {64u, 64u};
uint32_t element_stride[2] = {1u, 1u};
CUresult result = cuTensorMapEncodeTiled(
map,
CU_TENSOR_MAP_DATA_TYPE_UINT32,
2,
factor,
global_dim,
global_stride,
box_dim,
element_stride,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_NONE,
CU_TENSOR_MAP_L2_PROMOTION_L2_256B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
return result == CUDA_SUCCESS
? cudaSuccess
: cudaErrorInvalidValue;
}
static cudaError_t xrts_encode_packed_map_b60(
CUtensorMap* map,
uint32_t* packed_fp16
) {
// Four adjacent 4 KiB K32 packets form one 16 KiB TMA stage.
uint64_t global_dim[2] = {
256u,
60u * 16u * 32u * 4u
};
uint64_t global_stride[1] = {256u * sizeof(uint32_t)};
uint32_t box_dim[2] = {256u, 16u};
uint32_t element_stride[2] = {1u, 1u};
CUresult result = cuTensorMapEncodeTiled(
map,
CU_TENSOR_MAP_DATA_TYPE_UINT32,
2,
packed_fp16,
global_dim,
global_stride,
box_dim,
element_stride,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_NONE,
CU_TENSOR_MAP_L2_PROMOTION_L2_256B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
return result == CUDA_SUCCESS
? cudaSuccess
: cudaErrorInvalidValue;
}
static cudaError_t loco_cublas_to_cuda(cublasStatus_t status) {
switch (status) {
case CUBLAS_STATUS_SUCCESS:
return cudaSuccess;
case CUBLAS_STATUS_ALLOC_FAILED:
return cudaErrorMemoryAllocation;
case CUBLAS_STATUS_INVALID_VALUE:
return cudaErrorInvalidValue;
case CUBLAS_STATUS_ARCH_MISMATCH:
case CUBLAS_STATUS_NOT_SUPPORTED:
return cudaErrorNotSupported;
case CUBLAS_STATUS_EXECUTION_FAILED:
return cudaErrorLaunchFailure;
default:
return cudaErrorUnknown;
}
}
static cudaError_t loco_blas_left_frontier(
cublasHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int n,
int panel,
int width
) {
int history = panel;
int rows = n - panel;
int columns = rows < width ? rows : width;
if (history <= 0 || rows <= 0 || columns <= 0) {
return cudaSuccess;
}
if (!handle || !factor || !packed_fp16) {
return cudaErrorInvalidValue;
}
// Row-major:
// C[rows,columns] -= A[rows,history] * B[columns,history]^T
// The same storage is column-major C^T[columns,rows], so issue
// C^T -= B * A^T
// without a transform or a second carrier.
const uint16_t* shadow =
reinterpret_cast<const uint16_t*>(packed_fp16);
const void* b_transposed = shadow + (uint64_t)panel * n;
const void* a_transposed = b_transposed;
float* c_transposed = factor + (uint64_t)panel * n + panel;
float alpha = -1.0f;
float beta = 1.0f;
cublasStatus_t status = cublasGemmEx(
handle,
CUBLAS_OP_T,
CUBLAS_OP_N,
columns,
rows,
history,
&alpha,
b_transposed,
CUDA_R_16F,
n,
a_transposed,
CUDA_R_16F,
n,
&beta,
c_transposed,
CUDA_R_32F,
n,
CUBLAS_COMPUTE_32F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP);
return loco_cublas_to_cuda(status);
}
static cudaError_t loco_blas_panel_frontier(
cublasHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int n,
int destination,
int width,
int history_begin,
int history_end
) {
int history = history_end - history_begin;
int rows = n - destination;
int columns = rows < width ? rows : width;
if (history <= 0 || rows <= 0 || columns <= 0) {
return cudaSuccess;
}
if (!handle || !factor || !packed_fp16
|| history_begin < 0
|| history_end > destination) {
return cudaErrorInvalidValue;
}
// One active lower column slab:
//
// C[d:n, d:d+w]
// -= L[d:n, h0:h1] * L[d:d+w, h0:h1]^T.
//
// Solved L is already published once in the canonical FP16 shadow.
// Reinterpret row-major storage as its column-major transpose, exactly
// like the old-history P2048 frontier, with no repack or conversion.
const uint16_t* shadow =
reinterpret_cast<const uint16_t*>(packed_fp16);
const void* b_transposed =
shadow + (uint64_t)destination * n + history_begin;
const void* a_transposed = b_transposed;
float* c_transposed =
factor + (uint64_t)destination * n + destination;
float alpha = -1.0f;
float beta = 1.0f;
cublasStatus_t status = cublasGemmEx(
handle,
CUBLAS_OP_T,
CUBLAS_OP_N,
columns,
rows,
history,
&alpha,
b_transposed,
CUDA_R_16F,
n,
a_transposed,
CUDA_R_16F,
n,
&beta,
c_transposed,
CUDA_R_32F,
n,
CUBLAS_COMPUTE_32F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP);
return loco_cublas_to_cuda(status);
}
// Exact ranked mid-shape left-looking slab.
//
// One call updates every matrix in the batch:
//
// C[p:n, p:p+w] -= L[p:n, 0:p] * L[p:p+w, 0:p]^T.
//
// Solved L is already row-major FP16. Reinterpret each row-major matrix as
// its column-major transpose and use a strided-batched GEMM; no repack, no
// transform, and no intermediate right-looking residual writeback.
static cudaError_t loco_blas_batched_left_frontier(
cublasHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int panel,
int width
) {
int history = panel;
int rows = n - panel;
int columns = rows < width ? rows : width;
if (history <= 0 || rows <= 0 || columns <= 0 || batch <= 0) {
return cudaSuccess;
}
if (!handle || !factor || !packed_fp16) {
return cudaErrorInvalidValue;
}
const uint16_t* shadow =
reinterpret_cast<const uint16_t*>(packed_fp16);
const void* b_transposed = shadow + (uint64_t)panel * n;
const void* a_transposed = b_transposed;
float* c_transposed = factor + (uint64_t)panel * n + panel;
int64_t stride = (int64_t)n * (int64_t)n;
float alpha = -1.0f;
float beta = 1.0f;
cublasStatus_t status = cublasGemmStridedBatchedEx(
handle,
CUBLAS_OP_T,
CUBLAS_OP_N,
columns,
rows,
history,
&alpha,
b_transposed,
CUDA_R_16F,
n,
stride,
a_transposed,
CUDA_R_16F,
n,
stride,
&beta,
c_transposed,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP);
return loco_cublas_to_cuda(status);
}
static cublasStatus_t loco_lt_create_col_layout(
cublasLtMatrixLayout_t* layout,
cudaDataType_t type,
int rows,
int columns,
int leading_dimension,
int batch,
int64_t batch_stride
) {
*layout = nullptr;
cublasStatus_t status = cublasLtMatrixLayoutCreate(
layout,
type,
rows,
columns,
leading_dimension);
if (status != CUBLAS_STATUS_SUCCESS) return status;
cublasLtOrder_t order = CUBLASLT_ORDER_COL;
int32_t batch_count = batch;
status = cublasLtMatrixLayoutSetAttribute(
*layout,
CUBLASLT_MATRIX_LAYOUT_ORDER,
&order,
sizeof(order));
if (status == CUBLAS_STATUS_SUCCESS) {
status = cublasLtMatrixLayoutSetAttribute(
*layout,
CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT,
&batch_count,
sizeof(batch_count));
}
if (status == CUBLAS_STATUS_SUCCESS) {
status = cublasLtMatrixLayoutSetAttribute(
*layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&batch_stride,
sizeof(batch_stride));
}
if (status != CUBLAS_STATUS_SUCCESS) {
cublasLtMatrixLayoutDestroy(*layout);
*layout = nullptr;
}
return status;
}
struct LocoLtFrontierPlan {
int valid;
int device;
int batch;
int n;
int panel;
int width;
int history;
int has_algo;
size_t workspace_size;
cublasLtMatmulAlgo_t algo;
cublasLtMatmulDesc_t op_desc;
cublasLtMatrixLayout_t a_desc;
cublasLtMatrixLayout_t b_desc;
cublasLtMatrixLayout_t c_desc;
};
static void loco_lt_destroy_frontier_plan(LocoLtFrontierPlan& plan) {
if (plan.c_desc) cublasLtMatrixLayoutDestroy(plan.c_desc);
if (plan.b_desc) cublasLtMatrixLayoutDestroy(plan.b_desc);
if (plan.a_desc) cublasLtMatrixLayoutDestroy(plan.a_desc);
if (plan.op_desc) cublasLtMatmulDescDestroy(plan.op_desc);
plan = LocoLtFrontierPlan{};
}
static cublasStatus_t loco_lt_get_frontier_plan(
LocoLtFrontierPlan*& result,
cublasLtHandle_t handle,
int batch,
int n,
int panel,
int width,
int history,
uint64_t workspace_bytes
) {
constexpr int cache_capacity = 256;
static thread_local LocoLtFrontierPlan cache[cache_capacity]{};
static thread_local int replacement = 0;
static thread_local int device = -1;
if (!handle) return CUBLAS_STATUS_INVALID_VALUE;
if (device < 0) {
cudaError_t cuda_status = cudaGetDevice(&device);
if (cuda_status != cudaSuccess) {
return CUBLAS_STATUS_NOT_INITIALIZED;
}
}
for (int i = 0; i < cache_capacity; ++i) {
LocoLtFrontierPlan& plan = cache[i];
if (plan.valid
&& plan.device == device
&& plan.batch == batch
&& plan.n == n
&& plan.panel == panel
&& plan.width == width
&& plan.history == history) {
result = &plan;
return CUBLAS_STATUS_SUCCESS;
}
}
LocoLtFrontierPlan& plan = cache[replacement];
replacement = replacement + 1 == cache_capacity ? 0 : replacement + 1;
if (plan.valid) loco_lt_destroy_frontier_plan(plan);
int rows = n - panel;
int columns = rows < width ? rows : width;
int64_t batch_stride = (int64_t)n * (int64_t)n;
// Match the measured cublasGemmEx TN formulation directly. Row-major
// history X is viewed as column-major X^T; C is viewed as C^T.
cublasOperation_t op_a = CUBLAS_OP_T;
cublasOperation_t op_b = CUBLAS_OP_N;
cublasStatus_t status = cublasLtMatmulDescCreate(
&plan.op_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F);
if (status == CUBLAS_STATUS_SUCCESS) {
status = cublasLtMatmulDescSetAttribute(
plan.op_desc, CUBLASLT_MATMUL_DESC_TRANSA,
&op_a, sizeof(op_a));
}
if (status == CUBLAS_STATUS_SUCCESS) {
status = cublasLtMatmulDescSetAttribute(
plan.op_desc, CUBLASLT_MATMUL_DESC_TRANSB,
&op_b, sizeof(op_b));
}
if (status == CUBLAS_STATUS_SUCCESS) {
status = loco_lt_create_col_layout(
&plan.a_desc, CUDA_R_16F, history, columns, n,
batch, batch_stride);
}
if (status == CUBLAS_STATUS_SUCCESS) {
status = loco_lt_create_col_layout(
&plan.b_desc, CUDA_R_16F, history, rows, n,
batch, batch_stride);
}
if (status == CUBLAS_STATUS_SUCCESS) {
status = loco_lt_create_col_layout(
&plan.c_desc, CUDA_R_32F, columns, rows, n,
batch, batch_stride);
}
if (status != CUBLAS_STATUS_SUCCESS) {
loco_lt_destroy_frontier_plan(plan);
return status;
}
cublasLtMatmulPreference_t preference = nullptr;
status = cublasLtMatmulPreferenceCreate(&preference);
if (status == CUBLAS_STATUS_SUCCESS) {
status = cublasLtMatmulPreferenceSetAttribute(
preference,
CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_bytes,
sizeof(workspace_bytes));
}
if (status == CUBLAS_STATUS_SUCCESS) {
cublasLtMatmulHeuristicResult_t candidates[16]{};
int returned = 0;
status = cublasLtMatmulAlgoGetHeuristic(
handle,
plan.op_desc,
plan.a_desc,
plan.b_desc,
plan.c_desc,
plan.c_desc,
preference,
16,
candidates,
&returned);
if (status == CUBLAS_STATUS_SUCCESS) {
for (int i = 0; i < returned; ++i) {
if (candidates[i].state == CUBLAS_STATUS_SUCCESS
&& candidates[i].workspaceSize <= workspace_bytes) {
plan.algo = candidates[i].algo;
plan.workspace_size = candidates[i].workspaceSize;
plan.has_algo = 1;
break;
}
}
} else if (status == CUBLAS_STATUS_NOT_SUPPORTED) {
status = CUBLAS_STATUS_SUCCESS;
}
}
if (preference) cublasLtMatmulPreferenceDestroy(preference);
if (status != CUBLAS_STATUS_SUCCESS) {
loco_lt_destroy_frontier_plan(plan);
return status;
}
plan.valid = 1;
plan.device = device;
plan.batch = batch;
plan.n = n;
plan.panel = panel;
plan.width = width;
plan.history = history;
result = &plan;
return CUBLAS_STATUS_SUCCESS;
}
static cudaError_t loco_lt_frontier(
cublasLtHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int destination,
int width,
int history_begin,
int history_end,
void* workspace,
uint64_t workspace_bytes
) {
int history = history_end - history_begin;
int rows = n - destination;
int columns = rows < width ? rows : width;
if (history <= 0 || rows <= 0 || columns <= 0 || batch <= 0) {
return cudaSuccess;
}
if (!handle || !factor || !packed_fp16
|| !workspace || !workspace_bytes
|| history_begin < 0 || history_end > destination) {
return cudaErrorInvalidValue;
}
LocoLtFrontierPlan* plan = nullptr;
cublasStatus_t status = loco_lt_get_frontier_plan(
plan, handle, batch, n, destination, width, history,
workspace_bytes);
if (status != CUBLAS_STATUS_SUCCESS) {
return loco_cublas_to_cuda(status);
}
float alpha = -1.0f;
float beta = 1.0f;
const uint16_t* shadow = reinterpret_cast<const uint16_t*>(packed_fp16);
const void* a = shadow + (uint64_t)destination * n + history_begin;
const void* b = a;
float* c = factor + (uint64_t)destination * n + destination;
const cublasLtMatmulAlgo_t* algo = plan->has_algo ? &plan->algo : nullptr;
void* selected_workspace =
plan->has_algo && plan->workspace_size ? workspace : nullptr;
size_t selected_workspace_size =
plan->has_algo ? plan->workspace_size : 0;
status = cublasLtMatmul(
handle, plan->op_desc,
&alpha, a, plan->a_desc,
b, plan->b_desc,
&beta, c, plan->c_desc,
c, plan->c_desc,
algo, selected_workspace, selected_workspace_size, 0);
return loco_cublas_to_cuda(status);
}
static cudaError_t loco_lt_left_frontier(
cublasLtHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int panel,
int width,
void* workspace,
uint64_t workspace_bytes
) {
return loco_lt_frontier(
handle, factor, packed_fp16, batch, n,
panel, width, 0, panel, workspace, workspace_bytes);
}
static cudaError_t loco_lt_panel_frontier(
cublasLtHandle_t handle,
float* factor,
const uint32_t* packed_fp16,
int batch,
int n,
int destination,
int width,
int history_begin,
int history_end,
void* workspace,
uint64_t workspace_bytes
) {
return loco_lt_frontier(
handle, factor, packed_fp16, batch, n,
destination, width, history_begin, history_end,
workspace, workspace_bytes);
}
// V18 macro-left scheduler with bounded in-macro updates.
static cudaError_t loco_run_tcgen_macro(
const float* __restrict__ input,
float* __restrict__ factor,
uint8_t* __restrict__ plane_values,
uint8_t* __restrict__ plane_scales,
uint32_t* __restrict__ packed_fp16,
int packed_fp16_atoms,
cublasLtHandle_t lt_handle,
int use_lt_large3,
cublasHandle_t blas_handle,
int use_blas32,
int use_blas_mid,
void* lt_workspace,
uint64_t lt_workspace_bytes,
int batch,
int input_n,
int n,
int resident,
int use_e2m1
) {
uint64_t total =
(uint64_t)batch * (uint64_t)n * (uint64_t)n;
int copy_threads = 1024;
int copy_grid =
(int)((total + copy_threads - 1) / copy_threads);
int copy_cap = resident << 1;
if (copy_grid > copy_cap) copy_grid = copy_cap;
if (input_n == n) {
lower_copy_kernel<<<copy_grid, copy_threads>>>(
input,
factor,
total,
n);
} else {
lower_pad_copy_kernel<<<copy_grid, copy_threads>>>(
input,
factor,
total,
input_n,
n);
}
bool incremental_n1024 =
!use_blas_mid
&& !packed_fp16_atoms
&& input_n == 1024
&& n == 1024
&& batch >= 16;
bool xrts_f16_b60 =
XCAL_XRTS_TMA_F16_B60
&& incremental_n1024
&& batch == 60;
CUtensorMap xrts_map;
CUtensorMap xrts_packed_map;
if (xrts_f16_b60) {
cudaError_t map_status = xrts_encode_map_b60(
&xrts_map,
factor);
if (map_status != cudaSuccess) return map_status;
map_status = xrts_encode_packed_map_b60(
&xrts_packed_map,
packed_fp16);
if (map_status != cudaSuccess) return map_status;
}
// K256 remains an opt-in experiment. The measured control keeps the
// K128 dependency front while the outer board is retiled independently.
bool throughput_k256 =
XCAL_XRCHOLV18_K256_FRONT
&& !packed_fp16_atoms
&& (incremental_n1024
|| (n == 2048 && batch >= 8)
|| (n == 4096 && batch >= 2)
|| n >= 8192);
// Normalized low-B ladder:
// N1024/B4: P128 -> J128 exact frontier route above.
// N2048/B2: P256 -> J256 scaled M128N128 outer owners.
// N2048/B8: P1024, split-N128 in-J dependency owners.
// N4096/B1: P512 -> J512 scaled outer + split-N128 in-J.
//
// Retained measured boards:
// N4096/B2: P2048 -> J2048, split-N128 in-J.
// N8192/B1: P2048 -> J2048.
// N16384/B1 and N32768/B1 keep their vendor frontier schedules.
bool scaled_n_outer =
(n == 2048 && batch == 2)
|| (n >= 4096 && !(n & 511));
int outer_width =
n == 2048 && batch == 2
? 256
: n == 4096 && batch == 1
? 512
: 2048;
int macro_width =
n == 2048 && batch == 2
? 256
: n == 4096 && batch == 1
? 512
: scaled_n_outer
? 2048
: n <= 4096
? 1024
: (1 << (30 - __builtin_clz((unsigned)n)));
if (macro_width > n) macro_width = n;
if (use_blas_mid) {
macro_width = n;
scaled_n_outer = false;
}
bool split_ranked_frontier =
(n == 2048 && batch == 8)
|| (n == 4096 && (batch == 1 || batch == 2));
bool use_hexlift_outer = false;
uint8_t* sibling_plane_values = plane_values;
uint8_t* sibling_plane_scales = plane_scales;
if (use_hexlift_outer) {
uint64_t elements =
(uint64_t)batch * (uint64_t)n * (uint64_t)n;
sibling_plane_values += elements * 3u / 2u;
sibling_plane_scales += elements / 16u;
}
int hexlift_first_macro = n;
int hexlift_last_macro = 0;
int hexlift_min_ctas = XCAL_XRCHOLV18_HEXLIFT_MIN_CTAS;
if (XCAL_XRCHOLV18_HEX4_GE8192
&& input_n >= 8192) {
int size_floor = n >> 8;
if (size_floor < hexlift_min_ctas) {
hexlift_min_ctas = size_floor;
}
}
if (use_hexlift_outer) {
for (int candidate = macro_width;
candidate < n;
candidate += macro_width) {
int second_row_tiles =
(n - candidate - 256) >> 7;
int history_panels = candidate >> 6;
if (batch * second_row_tiles
>= hexlift_min_ctas
&& history_panels
>= XCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY) {
if (hexlift_first_macro == n) {
hexlift_first_macro = candidate;
}
hexlift_last_macro = candidate;
}
}
}
LocoTrailingClusterBoard trailing_board{};
trailing_board.active[0] = resident;
// Cluster geometry only pays its occupancy-query cost on the low-batch
// shapes where one matrix can actually consume a device-wide gang.
if (n >= 8192 && batch < 8) {
loco_measure_trailing_cluster_board(
trailing_board,
resident);
}
for (int macro = 0; macro < n; macro += macro_width) {
int macro_end =
macro + macro_width < n
? macro + macro_width
: n;
int second_row_tiles =
macro_end - macro == 512
? (n - macro - 256) >> 7
: 0;
int history_panels = macro >> 6;
bool hexlift_front =
use_hexlift_outer
&& macro > 0
&& batch * second_row_tiles
>= hexlift_min_ctas
&& history_panels
>= XCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY;
bool publish_hexlift =
use_hexlift_outer
&& macro < hexlift_last_macro;
int hex_publish_row_begin =
macro_end > hexlift_first_macro
? macro_end
: hexlift_first_macro;
int hex_publish_plane1_end =
XCAL_XRCHOLV18_HEX4_GE8192
? 0
: XCAL_XRCHOLV18_HEX6_ALL
? n
: hexlift_last_macro + macro_width;
if (hex_publish_plane1_end > n) {
hex_publish_plane1_end = n;
}
// Workspace-enabled cuBLASLt owns every exact large old-history
// frontier. N32768 keeps the measured cublasGemmEx fallback.
if (use_lt_large3 && macro > 0) {
cudaError_t lt_status = loco_lt_left_frontier(
lt_handle,
factor,
packed_fp16,
batch,
n,
macro,
macro_end - macro,
lt_workspace,
lt_workspace_bytes);
if (lt_status != cudaSuccess) {
lt_status = loco_blas_left_frontier(
blas_handle,
factor,
packed_fp16,
n,
macro,
macro_end - macro);
}
if (lt_status != cudaSuccess) return lt_status;
} else if (use_blas32 && macro > 0) {
cudaError_t blas_status = loco_blas_left_frontier(
blas_handle,
factor,
packed_fp16,
n,
macro,
macro_end - macro);
if (blas_status != cudaSuccess) return blas_status;
}
// One panel is N256. The x3/NVFP4 route replays bounded macro history.
// Exact N1024 high-batch enters with every completed N256 contribution
// already accumulated in the lower FP32 state.
for (int panel = macro; panel < macro_end; panel += 256) {
int panel_end =
panel + 256 < macro_end
? panel + 256
: macro_end;
if (use_blas_mid && panel > 0) {
cudaError_t mid_status = loco_lt_left_frontier(
lt_handle,
factor,
packed_fp16,
batch,
n,
panel,
panel_end - panel,
lt_workspace,
lt_workspace_bytes);
if (mid_status != cudaSuccess) {
mid_status = loco_blas_batched_left_frontier(
blas_handle,
factor,
packed_fp16,
batch,
n,
panel,
panel_end - panel);
}
if (mid_status != cudaSuccess) return mid_status;
} else if (!incremental_n1024
&& panel > macro) {
bool lt_panel = use_lt_large3
&& (use_blas32 || panel - macro >= 512);
if (lt_panel) {
cudaError_t lt_status = loco_lt_panel_frontier(
lt_handle,
factor,
packed_fp16,
batch,
n,
panel,
panel_end - panel,
macro,
panel,
lt_workspace,
lt_workspace_bytes);
if (lt_status != cudaSuccess) {
lt_status = loco_blas_panel_frontier(
blas_handle,
factor,
packed_fp16,
n,
panel,
panel_end - panel,
macro,
panel);
}
if (lt_status != cudaSuccess) return lt_status;
} else if (use_blas32) {
cudaError_t blas_status = loco_blas_panel_frontier(
blas_handle,
factor,
packed_fp16,
n,
panel,
panel_end - panel,
macro,
panel);
if (blas_status != cudaSuccess) return blas_status;
} else {
int start_panel = panel >> 6;
int history_begin = macro >> 6;
int history_end = panel >> 6;
int trailing_panels = (n - panel) >> 6;
cudaError_t launch_status =
loco_launch_trailing_tcgen_m128n256(
trailing_board,
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
split_ranked_frontier ? 3 : 1,
split_ranked_frontier ? 2 : 4,
resident);
if (launch_status != cudaSuccess) {
return launch_status;
}
}
}
bool panel_k256 =
throughput_k256 && panel_end - panel == 256;
if (panel_k256) {
int factor_grid =
batch < resident ? batch : resident;
loco_factor256_kernel
<<<factor_grid,
1024,
LOCO_HEXLIFT256_SMEM>>>(
factor,
packed_fp16,
batch,
n,
panel);
int rows = n - panel_end;
if (rows) {
int row_groups =
(resident + batch - 1) / batch;
int max_groups = (rows + 127) / 128;
if (row_groups > max_groups) {
row_groups = max_groups;
}
if (row_groups < 1) row_groups = 1;
int rows_per_group =
(rows + row_groups - 1) / row_groups;
int trsm_tasks = batch * row_groups;
int trsm_grid =
trsm_tasks < resident
? trsm_tasks
: resident;
loco_trsm256_kernel
<<<trsm_grid,
1024,
LOCO_HEXLIFT256_SMEM>>>(
factor,
packed_fp16,
batch,
n,
panel,
row_groups,
rows_per_group);
}
} else {
for (int k = panel; k < panel_end; k += 128) {
int bs =
panel_end - k < 128
? panel_end - k
: 128;
bool firstprinciples_front = bs == 128;
bool exact_b640 =
n == 512 && batch == 640;
int potrf_grid =
firstprinciples_front && exact_b640
? batch
: (batch < resident ? batch : resident);
int mma_trsm_rows = n - k - bs;
int64_t mma_trsm_work =
(int64_t)batch * (int64_t)mma_trsm_rows;
// B16/B640 publish the solved K128 panel once into the FP16
// sidecar. B640 then exposes every matrix/frontier owner.
bool mma_trsm_front =
firstprinciples_front
&& !packed_fp16_atoms
&& (n >= 1024
|| (n == 512
&& (batch == 16 || batch == 640)))
&& mma_trsm_work >= 1536;
if (firstprinciples_front) {
if (mma_trsm_front) {
xrshape128_diag_mid_kernel<true, true, true>
<<<potrf_grid, 1024>>>(
factor,
packed_fp16,
batch,
n,
k);
} else if (n >= 512) {
xrshape128_diag_mid_kernel<true, false, false>
<<<potrf_grid, 1024>>>(
factor,
nullptr,
batch,
n,
k);
} else {
xrshape128_diag_mid_kernel<false, false, false>
<<<potrf_grid, 1024>>>(
factor,
nullptr,
batch,
n,
k);
}
} else {
potrf128_kernel
<<<potrf_grid, 1024, LOCO_FRONTIER128_SMEM>>>(
factor,
batch,
n,
k,
bs);
}
int start = k + bs;
int rows = n - start;
if (!rows) continue;
int target_groups =
(resident + batch - 1) / batch;
if (target_groups > rows) {
target_groups = rows;
}
int active_warps =
target_groups > 1
? (rows - 1) / (target_groups - 1)
: rows;
if (active_warps < 1) active_warps = 1;
if (active_warps > 32) active_warps = 32;
int row_groups =
(rows + active_warps - 1) / active_warps;
// A batch-owner keeps this diagonal resident and walks its
// disjoint row-groups; no group reloads the same 128x128 L.
int row_owners = 1;
if (batch < resident) {
row_owners = resident / batch;
if (row_owners > row_groups)
row_owners = row_groups;
}
int trsm_grid =
(batch < resident ? batch : resident)
* row_owners;
int trsm_warps =
active_warps < 4 ? 4 : active_warps;
int trsm_threads = trsm_warps << 5;
bool use_mma_trsm = mma_trsm_front;
if (use_mma_trsm) {
if (n == 512 && batch == 16 && rows == 384) {
int row_tiles = 8;
trsm128_mma_f16_kernel<48>
<<<batch * row_tiles,
256,
(8192 + 48 * 64 + 48 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
} else if (exact_b640 && rows == 384) {
int row_tiles = 2;
trsm128_mma_f16_kernel<192>
<<<batch * row_tiles,
768,
(8192 + 192 * 64 + 192 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
} else if (mma_trsm_work >= 24576
&& (!exact_b640 || rows >= 256)) {
int row_tiles = (rows + 255) >> 8;
trsm128_mma_f16_kernel<256>
<<<batch * row_tiles,
1024,
(8192 + 256 * 64 + 256 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
} else if (mma_trsm_work >= 12288
&& (!exact_b640 || rows >= 128)) {
int row_tiles = (rows + 127) >> 7;
trsm128_mma_f16_kernel<128>
<<<batch * row_tiles,
512,
(8192 + 128 * 64 + 128 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
} else if (mma_trsm_work >= 6144
&& (!exact_b640 || rows >= 64)) {
int row_tiles = (rows + 63) >> 6;
trsm128_mma_f16_kernel<64>
<<<batch * row_tiles,
256,
(8192 + 64 * 64 + 64 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
} else if (mma_trsm_work >= 3072
&& (!exact_b640 || rows >= 32)) {
int row_tiles = (rows + 31) >> 5;
trsm128_mma_f16_kernel<32>
<<<batch * row_tiles,
256,
(8192 + 32 * 64 + 32 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
} else {
int row_tiles = (rows + 15) >> 4;
trsm128_mma_f16_kernel<16>
<<<batch * row_tiles,
256,
(8192 + 16 * 64 + 16 * 8) * 4>>>(
factor, packed_fp16, batch, n, k, rows);
}
cudaError_t trsm_status = cudaGetLastError();
if (trsm_status != cudaSuccess) {
return trsm_status;
}
} else if (publish_hexlift) {
trsm128_kernel<true>
<<<trsm_grid,
trsm_threads,
LOCO_FRONTIER128_SMEM>>>(
factor,
packed_fp16,
plane_values,
plane_scales,
hex_publish_row_begin,
hex_publish_plane1_end,
packed_fp16_atoms,
batch,
n,
k,
bs,
row_groups,
active_warps,
row_owners);
} else {
trsm128_kernel<false>
<<<trsm_grid,
trsm_threads,
LOCO_FRONTIER128_SMEM>>>(
factor,
packed_fp16,
nullptr,
nullptr,
n,
n,
packed_fp16_atoms,
batch,
n,
k,
bs,
row_groups,
active_warps,
row_owners);
}
// K128_1 consumes K128_0 before its POTRF. Tight checker
// shapes use the packed-FP16 TCGEN rail; large shapes use
// the reusable E2M1 rail.
if (k == panel && start < panel_end) {
if (xrts_f16_b60) {
int first = start >> 7;
int owners =
2 + ((7 - first) << 1);
if (owners) {
int grid = batch * owners;
if (grid > resident * 2) {
grid = resident * 2;
}
xrts_f16_m128n64_b60
<<<grid,
256,
sizeof(XrtsF16Shared)>>>(
factor,
xrts_map,
xrts_packed_map,
k,
first,
4,
1);
cudaError_t sibling_status =
cudaGetLastError();
if (sibling_status != cudaSuccess) {
return sibling_status;
}
}
} else if (use_e2m1) {
bool e2m1_frontier = loco_run_frontier_e2m1(
factor,
sibling_plane_values,
sibling_plane_scales,
batch,
n,
k >> 6,
bs >> 6,
start,
panel_end,
resident);
if (!e2m1_frontier) {
return cudaErrorInvalidValue;
}
} else {
int sibling_panels = (n - start) >> 6;
cudaError_t sibling_status =
loco_launch_trailing_tcgen_m128n256(
trailing_board,
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start >> 6,
k >> 6,
start >> 6,
sibling_panels,
1,
2,
resident);
if (sibling_status != cudaSuccess) {
return sibling_status;
}
}
}
}
}
// State invariant after panel j:
// C_future = A_future - sum_{p=0..j} L_p * L_p^T.
// The completed panel is four row-major FP16 K64 histories. Seed D
// from the current FP32 lower state and visit every future column.
if (incremental_n1024 && panel_end < n) {
int start_panel = panel_end >> 6;
int history_begin = panel >> 6;
int history_end = panel_end >> 6;
int trailing_panels = (n - panel_end) >> 6;
cudaError_t launch_status;
if (xrts_f16_b60) {
int first = panel_end >> 7;
int panels = (8 - first) << 1;
int owners = panels;
int grid = batch * owners;
if (grid > resident * 2) {
grid = resident * 2;
}
xrts_f16_m128n64_b60
<<<grid,
256,
sizeof(XrtsF16Shared)>>>(
factor,
xrts_map,
xrts_packed_map,
panel,
first,
8,
0);
launch_status = cudaGetLastError();
} else if (batch == 60) {
#if XCAL_XRCHOLV18_B60_N128
int grid = batch
* (trailing_panels == 12
? 21
: trailing_panels == 8 ? 10 : 3);
if (grid > resident * 3) {
grid = resident * 3;
}
if (trailing_panels == 12) {
loco_trailing_tcgen_b60_m128n128<12>
<<<grid,
256,
sizeof(XrShape512UpdateShared)>>>(
factor,
packed_fp16,
batch);
} else if (trailing_panels == 8) {
loco_trailing_tcgen_b60_m128n128<8>
<<<grid,
256,
sizeof(XrShape512UpdateShared)>>>(
factor,
packed_fp16,
batch);
} else {
loco_trailing_tcgen_b60_m128n128<4>
<<<grid,
256,
sizeof(XrShape512UpdateShared)>>>(
factor,
packed_fp16,
batch);
}
#else
int grid = batch
* (trailing_panels == 12
? 12
: trailing_panels == 8 ? 6 : 2);
if (grid > resident
* XCAL_XRCHOLV18_B60_RESIDENT) {
grid = resident
* XCAL_XRCHOLV18_B60_RESIDENT;
}
if (trailing_panels == 12) {
loco_trailing_tcgen_b60_m128n256<
12,
XCAL_XRCHOLV18_B60_ISSUES,
XCAL_XRCHOLV18_B60_RESIDENT>
<<<grid,
512,
sizeof(LocoB60N256Shared<
XCAL_XRCHOLV18_B60_ISSUES>)>>>(
factor,
packed_fp16,
batch);
} else if (trailing_panels == 8) {
loco_trailing_tcgen_b60_m128n256<
8,
XCAL_XRCHOLV18_B60_ISSUES,
XCAL_XRCHOLV18_B60_RESIDENT>
<<<grid,
512,
sizeof(LocoB60N256Shared<
XCAL_XRCHOLV18_B60_ISSUES>)>>>(
factor,
packed_fp16,
batch);
} else {
loco_trailing_tcgen_b60_m128n256<
4,
XCAL_XRCHOLV18_B60_ISSUES,
XCAL_XRCHOLV18_B60_RESIDENT>
<<<grid,
512,
sizeof(LocoB60N256Shared<
XCAL_XRCHOLV18_B60_ISSUES>)>>>(
factor,
packed_fp16,
batch);
}
#endif
launch_status = cudaGetLastError();
} else {
launch_status =
loco_launch_trailing_tcgen_m128n256(
trailing_board,
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
0,
4,
resident);
}
if (launch_status != cudaSuccess) {
return launch_status;
}
}
}
// A sealed macro publishes once through its lower residual board.
if (!incremental_n1024
&& !use_lt_large3
&& !use_blas32
&& macro_end < n) {
int start_panel = macro_end >> 6;
int history_begin = macro >> 6;
int history_end = macro_end >> 6;
int trailing_panels =
(n - macro_end) >> 6;
if (scaled_n_outer) {
cudaError_t launch_status =
loco_launch_outer_scaled_n_parallel(
factor,
packed_fp16,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
outer_width,
resident);
if (launch_status != cudaSuccess) {
return launch_status;
}
} else if (XCAL_XRCHOLV18_NVFP4_OUTER && n > 8192) {
int trailing_rows = n - macro_end;
int history_panels =
history_end - history_begin;
uint8_t* nvfp4_values =
reinterpret_cast<uint8_t*>(packed_fp16);
uint64_t quantize_tasks =
(uint64_t)batch
* history_panels
* trailing_rows;
uint64_t quantize_blocks =
(quantize_tasks + 7u) >> 3;
int quantize_grid =
quantize_blocks < (uint64_t)(resident * 4)
? (int)quantize_blocks
: resident * 4;
loco_quantize_outer_nvfp4
<<<quantize_grid, 256>>>(
factor,
nvfp4_values,
plane_scales,
batch,
n,
macro_end,
history_begin,
history_panels,
trailing_rows);
int row_pairs =
(trailing_panels + 1) >> 1;
uint64_t grouped_tasks =
loco_m128n256_task_prefix(row_pairs);
uint64_t total_tasks =
(uint64_t)batch * grouped_tasks;
int grid =
total_tasks < (uint64_t)resident
? (int)total_tasks
: resident;
loco_outer_tcgen_nvfp4_m128n256
<<<grid,
1024,
(int)sizeof(LocoNvfp4OuterShared)>>>(
factor,
nvfp4_values,
plane_scales,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
trailing_rows,
history_panels);
} else {
cudaError_t launch_status =
loco_launch_trailing_tcgen_m128n256(
trailing_board,
factor,
packed_fp16,
packed_fp16_atoms,
batch,
n,
start_panel,
history_begin,
history_end,
trailing_panels,
0,
4,
resident);
if (launch_status != cudaSuccess) {
return launch_status;
}
}
}
}
if (use_blas_mid) {
int dirty_panels = (n >> 8) - 1;
if (dirty_panels > 0) {
uint64_t vectors =
(uint64_t)batch * (uint64_t)dirty_panels * 128u * 32u;
int zero_grid = (int)((vectors + 255u) >> 8);
int zero_cap = resident << 2;
if (zero_grid > zero_cap) zero_grid = zero_cap;
zero_mid_blas_cross_upper_kernel
<<<zero_grid, 256>>>(factor, batch, n);
cudaError_t zero_status = cudaGetLastError();
if (zero_status != cudaSuccess) return zero_status;
}
}
if (xrts_f16_b60) {
int zero_grid = 60 * 36;
if (zero_grid > resident * 8) {
zero_grid = resident * 8;
}
xrts_zero_upper_b60<<<zero_grid, 256>>>(factor);
}
return cudaGetLastError();
}
/* xR1 terminal lower-only contract. */
template<int R, int C>
static __device__ __forceinline__ void xr1_zupper1(
float* F,
int N
) {
if (((blockIdx.z << 7) + (threadIdx.x >> 5) + R < (unsigned int)N)
&& ((blockIdx.x << 7) + ((threadIdx.x & 31u) << 2) + C
< (unsigned int)N)
&& ((blockIdx.x << 7) + ((threadIdx.x & 31u) << 2) + C
> (blockIdx.z << 7) + (threadIdx.x >> 5) + R))
asm volatile(
"{\n\t"
".reg .b32 z;\n\t"
"mov.b32 z, 0;\n\t"
"st.global.b32 [%0], z;\n\t"
"}"
:
: "l"(F
+ (unsigned long long)blockIdx.y
* (unsigned int)N * (unsigned int)N
+ (unsigned long long)(
(blockIdx.z << 7) + (threadIdx.x >> 5) + R)
* (unsigned int)N
+ (blockIdx.x << 7)
+ ((threadIdx.x & 31u) << 2) + C)
: "memory");
}
template<int R>
static __device__ __forceinline__ void xr1_zupper4(
float* F,
int N
) {
if (blockIdx.x > blockIdx.z
&& !(N & 3)
&& ((blockIdx.z << 7) + (threadIdx.x >> 5) + R
< (unsigned int)N)
&& ((blockIdx.x << 7) + ((threadIdx.x & 31u) << 2) + 3u
< (unsigned int)N)) {
asm volatile(
"{\n\t"
".reg .b32 z;\n\t"
"mov.b32 z, 0;\n\t"
"st.global.v4.b32 [%0], {z,z,z,z};\n\t"
"}"
:
: "l"(F
+ (unsigned long long)blockIdx.y
* (unsigned int)N * (unsigned int)N
+ (unsigned long long)(
(blockIdx.z << 7) + (threadIdx.x >> 5) + R)
* (unsigned int)N
+ (blockIdx.x << 7)
+ ((threadIdx.x & 31u) << 2))
: "memory");
return;
}
if (blockIdx.x < blockIdx.z)
return;
xr1_zupper1<R,0>(F,N);
xr1_zupper1<R,1>(F,N);
xr1_zupper1<R,2>(F,N);
xr1_zupper1<R,3>(F,N);
}
__global__ __launch_bounds__(256, 4) void xr1_terminal_upper128(
float* F,
int N
) {
if (blockDim.x != 256)
return;
if (blockIdx.x < blockIdx.z)
return;
xr1_zupper4< 0>(F,N); xr1_zupper4< 8>(F,N);
xr1_zupper4< 16>(F,N); xr1_zupper4< 24>(F,N);
xr1_zupper4< 32>(F,N); xr1_zupper4< 40>(F,N);
xr1_zupper4< 48>(F,N); xr1_zupper4< 56>(F,N);
xr1_zupper4< 64>(F,N); xr1_zupper4< 72>(F,N);
xr1_zupper4< 80>(F,N); xr1_zupper4< 88>(F,N);
xr1_zupper4< 96>(F,N); xr1_zupper4<104>(F,N);
xr1_zupper4<112>(F,N); xr1_zupper4<120>(F,N);
}
static cudaError_t loco_configure_v18_once() {
static thread_local int configured_device = -1;
// The ranked harness is single-device. Keep the hot invocation entirely
// host-API free after the first configuration pass.
if (configured_device >= 0) return cudaSuccess;
int device = -1;
cudaError_t status = cudaGetDevice(&device);
if (status != cudaSuccess) return status;
#define LOCO_SET_MAX_SMEM(kernel, bytes) do { \
status = cudaFuncSetAttribute( \
(const void*)(kernel), \
cudaFuncAttributeMaxDynamicSharedMemorySize, \
(int)(bytes)); \
if (status != cudaSuccess) return status; \
status = cudaFuncSetAttribute( \
(const void*)(kernel), \
cudaFuncAttributePreferredSharedMemoryCarveout, \
100); \
if (status != cudaSuccess) return status; \
} while (0)
LOCO_SET_MAX_SMEM(potrf_small_kernel, LOCO_FRONTIER128_SMEM);
LOCO_SET_MAX_SMEM(
loco_potrf64_warp_kernel,
LOCO_N64_SMEM);
LOCO_SET_MAX_SMEM(
loco_potrf64_bf16x3_kernel,
LOCO_N64_X3_SMEM);
LOCO_SET_MAX_SMEM(
loco_potrf128_k16x3_kernel,
LOCO_N128_K16_SMEM);
LOCO_SET_MAX_SMEM(
loco_persistent256_x3_kernel,
LOCO_FACTOR256_K16_SMEM);
LOCO_SET_MAX_SMEM(
xrshape256_trsm_b64_kernel,
XRSHAPE256_TRSM_SMEM);
LOCO_SET_MAX_SMEM(
xrshape256_update_b64_kernel,
XRSHAPE256_UPDATE_SMEM);
LOCO_SET_MAX_SMEM(
xrshape512_factor_kernel,
XRSHAPE512_FACTOR_SMEM);
LOCO_SET_MAX_SMEM(
xrshape512_trsm_kernel,
XRSHAPE512_TRSM_SMEM);
LOCO_SET_MAX_SMEM(
xrshape512_update_kernel,
sizeof(XrShape512UpdateShared));
LOCO_SET_MAX_SMEM(
loco_outer_scaled_n_parallel<false>,
sizeof(XrShape512UpdateShared));
LOCO_SET_MAX_SMEM(
loco_outer_scaled_n_parallel<true>,
sizeof(XrShape512UpdateShared));
LOCO_SET_MAX_SMEM(
loco_trailing_tcgen_b60_m128n128<12>,
sizeof(XrShape512UpdateShared));
LOCO_SET_MAX_SMEM(
loco_trailing_tcgen_b60_m128n128<8>,
sizeof(XrShape512UpdateShared));
LOCO_SET_MAX_SMEM(
loco_trailing_tcgen_b60_m128n128<4>,
sizeof(XrShape512UpdateShared));
LOCO_SET_MAX_SMEM(
loco_persistent512_tcgen_kernel,
sizeof(LocoPersistent512TcgenShared));
LOCO_SET_MAX_SMEM(loco_factor256_kernel, LOCO_HEXLIFT256_SMEM);
LOCO_SET_MAX_SMEM(loco_trsm256_kernel, LOCO_HEXLIFT256_SMEM);
LOCO_SET_MAX_SMEM(potrf128_kernel, LOCO_FRONTIER128_SMEM);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<16>),
(8192 + 16 * 64 + 16 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<32>),
(8192 + 32 * 64 + 32 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<48>),
(8192 + 48 * 64 + 48 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<64>),
(8192 + 64 * 64 + 64 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<128>),
(8192 + 128 * 64 + 128 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<128, 256>),
(8192 + 128 * 64 + 128 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<192>),
(8192 + 192 * 64 + 192 * 8) * 4);
LOCO_SET_MAX_SMEM(
(trsm128_mma_f16_kernel<256>),
(8192 + 256 * 64 + 256 * 8) * 4);
LOCO_SET_MAX_SMEM(
trsm128_kernel<false>,
LOCO_FRONTIER128_SMEM);
LOCO_SET_MAX_SMEM(
trsm128_kernel<true>,
LOCO_FRONTIER128_SMEM);
LOCO_SET_MAX_SMEM(
xrts_f16_m128n64_b60,
sizeof(XrtsF16Shared));
LOCO_SET_MAX_SMEM(
loco_frontier_e2m1x4_m128n256,
sizeof(LocoFrontierE2M1x4Shared));
LOCO_SET_MAX_SMEM(
loco_outer_hexlift_m128n256,
sizeof(LocoFrontierE2M1x4Shared));
LOCO_SET_MAX_SMEM(
loco_trailing_tcgen_m128n256<false>,
sizeof(LocoMaterializeShared));
LOCO_SET_MAX_SMEM(
loco_trailing_tcgen_m128n256<true>,
sizeof(LocoMaterializeShared));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_m128n256<true, true>),
sizeof(LocoMaterializeShared));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_m128n256<false, false, true>),
sizeof(LocoMaterializeShared));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_m128n256<true, false, true>),
sizeof(LocoMaterializeShared));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_m128n256<true, true, true>),
sizeof(LocoMaterializeShared));
#if XCAL_XRCHOLV18_B60_ISSUES == 4
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_b60_m128n256<12, 4, 3>),
sizeof(LocoB60N256Shared<4>));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_b60_m128n256<8, 4, 3>),
sizeof(LocoB60N256Shared<4>));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_b60_m128n256<4, 4, 3>),
sizeof(LocoB60N256Shared<4>));
#else
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_b60_m128n256<12, 8, 2>),
sizeof(LocoB60N256Shared<8>));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_b60_m128n256<8, 8, 2>),
sizeof(LocoB60N256Shared<8>));
LOCO_SET_MAX_SMEM(
(loco_trailing_tcgen_b60_m128n256<4, 8, 2>),
sizeof(LocoB60N256Shared<8>));
#endif
LOCO_SET_MAX_SMEM(
loco_outer_tcgen_nvfp4_m128n256,
sizeof(LocoNvfp4OuterShared));
LOCO_SET_MAX_SMEM(
xrshape1024_factor256_b4_kernel,
XRSHAPE512_FACTOR_SMEM);
LOCO_SET_MAX_SMEM(
xrshape1024_trsm256_b4_kernel,
XRSHAPE512_TRSM_SMEM);
status = cudaFuncSetAttribute(
(const void*)loco_trailing_tcgen_m128n256<false>,
cudaFuncAttributeNonPortableClusterSizeAllowed,
1);
if (status != cudaSuccess) return status;
status = cudaFuncSetAttribute(
(const void*)loco_trailing_tcgen_m128n256<
false,
false,
true>,
cudaFuncAttributeNonPortableClusterSizeAllowed,
1);
if (status != cudaSuccess) return status;
#undef LOCO_SET_MAX_SMEM
configured_device = device;
return cudaSuccess;
}
static __forceinline__ int loco_batch4_wave_grid(
int batch,
int resident
) {
int grid = resident << 2;
return batch < grid ? batch : grid;
}
extern "C" int xrcholv18_cholesky_launch(
const float* input,
float* factor,
float* output,
uint8_t* plane_values,
uint8_t* plane_scales,
void* lt_handle,
void* blas_handle,
void* lt_workspace,
uint64_t lt_workspace_bytes,
int batch,
int input_n,
int work_n,
int resident,
int use_tcgen,
int use_e2m1
) {
if (!batch || !input_n) return (int)cudaSuccess;
cudaError_t status = loco_configure_v18_once();
if (status != cudaSuccess) return (int)status;
// N32 uses one warp per matrix. N64/N128 retain the proven multi-warp
// scalar board so they keep the original SM warp occupancy.
if (!use_tcgen) {
if (input_n > 128 || work_n != input_n) {
return (int)cudaErrorInvalidValue;
}
if (input_n == 32) {
int grid = (batch + LOCO_N32_WARPS - 1) / LOCO_N32_WARPS;
loco_potrf32_warp_kernel
<<<grid, LOCO_N32_THREADS>>>(
input, output, batch);
return (int)cudaGetLastError();
}
if (input_n == 64) {
if (batch == 1024) {
int grid =
(batch + LOCO_N64_X3_MATRICES - 1)
/ LOCO_N64_X3_MATRICES;
loco_potrf64_bf16x3_kernel
<<<grid, LOCO_N64_X3_THREADS, LOCO_N64_X3_SMEM>>>(
input, output, batch);
} else {
int grid =
(batch + LOCO_N64_WARPS - 1) / LOCO_N64_WARPS;
loco_potrf64_warp_kernel
<<<grid, LOCO_N64_THREADS, LOCO_N64_SMEM>>>(
input, output, batch);
}
return (int)cudaGetLastError();
}
if (input_n == 128) {
if (batch == 256) {
xrshape128_b256_kernel<<<256, 1024>>>(input, output);
} else {
loco_potrf128_k16x3_kernel
<<<batch,
LOCO_N128_K16_THREADS,
LOCO_N128_K16_SMEM>>>(
input, output, batch);
}
return (int)cudaGetLastError();
}
int warps = (input_n + 3) >> 2;
if (warps < 4) warps = 4;
if (warps > 32) warps = 32;
int threads = warps << 5;
int blocks_per_sm = 64 / warps;
if (blocks_per_sm < 1) blocks_per_sm = 1;
int grid = batch < resident * blocks_per_sm
? batch
: resident * blocks_per_sm;
int smem =
(input_n * (input_n + 1) + input_n) * (int)sizeof(float);
potrf_small_kernel<<<grid, threads, smem>>>(
input, output, batch, input_n);
return (int)cudaGetLastError();
}
// B64/N256:
// lower F32 -> K128 FP16 diag + FP16 sidecar
// -> M128xN128 MMA solve
// -> FP16 A11 -> K128 FP16 diag.
if (input_n == 256 && work_n == 256) {
if (batch == 64) {
uint32_t* packed_fp16 =
reinterpret_cast<uint32_t*>(plane_values);
if (!packed_fp16) return (int)cudaErrorInvalidValue;
lower_copy_kernel<<<resident, 1024>>>(
input,
factor,
64ull << 16,
256);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
xrshape128_diag_mid_kernel<true, true, true>
<<<64, 1024>>>(
factor,
packed_fp16,
64,
256,
0);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
trsm128_mma_f16_kernel<64>
<<<128,
256,
(8192 + 64 * 64 + 64 * 8) * 4>>>(
factor,
packed_fp16,
64,
256,
0,
128);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
xrshape256_update_b64_kernel
<<<128, 1024, XRSHAPE256_UPDATE_SMEM>>>(
input, output);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
xrshape128_diag_mid_kernel<true, false, true>
<<<64, 1024>>>(
factor,
nullptr,
64,
256,
128);
return (int)cudaGetLastError();
}
int grid = batch < resident ? batch : resident;
loco_persistent256_x3_kernel
<<<grid, 1024, LOCO_FACTOR256_K16_SMEM>>>(
input, output, batch);
return (int)cudaGetLastError();
}
// Ranked N512 board:
// L00 -> L10 -> three independent A11 quadrants -> L11.
// Exact B16 and B640 fall through to the K128 + block-MMA front below.
// Other N512 batches retain the staged four-launch board.
if (input_n == 512
&& work_n == 512
&& batch != 16
&& batch != 640) {
int factor_grid = batch < resident ? batch : resident;
xrshape512_factor_kernel
<<<factor_grid, 1024, XRSHAPE512_FACTOR_SMEM>>>(
input,
output,
batch,
0);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
int row_groups = (resident + batch - 1) / batch;
if (row_groups > 8) row_groups = 8;
if (row_groups < 1) row_groups = 1;
int rows_per_group =
(256 + row_groups - 1) / row_groups;
int trsm_grid = batch * row_groups;
if (trsm_grid > resident) trsm_grid = resident;
xrshape512_trsm_kernel
<<<trsm_grid, 1024, XRSHAPE512_TRSM_SMEM>>>(
input,
output,
batch,
row_groups,
rows_per_group);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
int update_grid = batch * 3;
if (update_grid > resident * 3) {
update_grid = resident * 3;
}
xrshape512_update_kernel
<<<update_grid,
512,
(int)sizeof(XrShape512UpdateShared)>>>(
input,
output,
batch);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
xrshape512_factor_kernel
<<<factor_grid, 1024, XRSHAPE512_FACTOR_SMEM>>>(
output,
output,
batch,
256);
return (int)cudaGetLastError();
}
if (input_n == 1024
&& work_n == 1024
&& batch == 4) {
if (!lt_handle || !blas_handle
|| !lt_workspace || !lt_workspace_bytes) {
return (int)cudaErrorInvalidValue;
}
return (int)xrshape1024_run_b4_frontier(
input,
factor,
reinterpret_cast<uint32_t*>(plane_values),
reinterpret_cast<cublasLtHandle_t>(lt_handle),
reinterpret_cast<cublasHandle_t>(blas_handle),
lt_workspace,
lt_workspace_bytes,
resident);
}
if (input_n <= 128
|| work_n < 256
|| work_n < input_n
|| (work_n & 127)) {
return (int)cudaErrorInvalidValue;
}
// N>=256 keeps one exact K128 frontier machine. The dormant K256 kernels
// remain available for a separately validated board.
//
// n < 2048 && B < 16 : FP16 hi/lo x3 M128xN128 TCGEN05
// non-E2M1 fast sibling: rounded-FP16 M128xN128 TCGEN05
// large sibling K128 : four-plane E2M1 TCGEN05
// n < 2048 && B < 16 : FP16 hi/lo x3 history, FP32 D
// n1024 && B >= 16 : rounded-FP16 incremental N256 outer, FP32 D
// otherwise history : rounded-FP16 blocked-lower update, FP32 D
// n > 8192 outer : one-rail NVFP4 M128xN256K64, FP32 TMEM D
//
uint32_t* packed_fp16 =
reinterpret_cast<uint32_t*>(plane_values);
int packed_fp16_atoms =
input_n < 2048 && batch < 16 ? 1 : 0;
int use_incremental1024 =
input_n == 1024
&& work_n == 1024
&& batch >= 16;
int use_lt_large3 =
batch == 1
&& input_n == work_n
&& (work_n == 8192
|| work_n == 16384
|| work_n == 32768);
int use_blas32 =
batch == 1
&& input_n == work_n
&& work_n == 32768;
int use_blas_mid =
input_n == work_n
&& ((work_n == 512 && (batch == 16 || batch == 640))
|| (work_n == 1024 && (batch == 4 || batch == 60))
|| (work_n == 2048 && (batch == 2 || batch == 8))
|| (work_n == 4096 && (batch == 1 || batch == 2)));
uint8_t* e2m1_values = nullptr;
if (packed_fp16) {
uint64_t elements =
(uint64_t)batch * (uint64_t)work_n * (uint64_t)work_n;
uint64_t fp16_words = (elements + 1u) >> 1;
if (use_e2m1) {
e2m1_values =
plane_values + fp16_words * sizeof(uint32_t);
}
}
if (!packed_fp16
|| (use_e2m1 && (!e2m1_values || !plane_scales))
|| ((use_lt_large3 || use_blas_mid)
&& (!lt_handle || !lt_workspace || !lt_workspace_bytes))
|| ((use_blas32 || use_blas_mid || use_lt_large3)
&& !blas_handle)) {
return (int)cudaErrorInvalidValue;
}
status = loco_run_tcgen_macro(
input,
factor,
e2m1_values,
plane_scales,
packed_fp16,
packed_fp16_atoms,
reinterpret_cast<cublasLtHandle_t>(lt_handle),
use_lt_large3,
reinterpret_cast<cublasHandle_t>(blas_handle),
use_blas32,
use_blas_mid,
lt_workspace,
lt_workspace_bytes,
batch,
input_n,
work_n,
resident,
use_e2m1);
if (status != cudaSuccess) return (int)status;
if (input_n >= 8192) {
// Initial lower-copy already zeroed the untouched upper triangle.
// Library writes dirty only P2048 diagonal macro squares plus each
// local P256 cross-half, so clear those exact regions instead of
// rewriting the complete O(N^2/2) strict upper triangle.
int macro_width = 2048;
int dirty_blocks = (work_n + macro_width - 1) / macro_width - 1;
if (dirty_blocks > 0) {
uint64_t vectors = (uint64_t)batch
* (uint64_t)dirty_blocks
* (uint64_t)macro_width
* (uint64_t)(macro_width >> 2);
int grid = (int)((vectors + 1023u) >> 10);
int cap = resident << 2;
if (grid > cap) grid = cap;
zero_macro_upper_kernel<<<grid, 1024>>>(
factor, batch, work_n, macro_width);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
}
int dirty_panels = (work_n >> 8) - 1;
if (dirty_panels > 0) {
uint64_t vectors = (uint64_t)batch
* (uint64_t)dirty_panels * 128u * 32u;
int grid = (int)((vectors + 255u) >> 8);
int cap = resident << 2;
if (grid > cap) grid = cap;
zero_mid_blas_cross_upper_kernel<<<grid, 256>>>(
factor, batch, work_n);
status = cudaGetLastError();
if (status != cudaSuccess) return (int)status;
}
}
if (work_n != input_n) {
uint64_t total =
(uint64_t)batch
* (uint64_t)input_n
* (uint64_t)input_n;
int threads = 1024;
int grid = (int)((total + threads - 1) / threads);
if (grid > resident) grid = resident;
lower_crop_kernel<<<grid, threads>>>(
factor,
output,
total,
input_n,
work_n);
return (int)cudaGetLastError();
}
return (int)cudaSuccess;
}
"""
_BUILD_KEY = "x38-sprint-ltall-copy2-m48m192-targetupper-r1"
_SOURCE_TAG = hashlib.sha256(
(CPP + "\0" + CUDA + "\0" + _BUILD_KEY).encode("utf-8")
).hexdigest()[:16]
_FILE_DIR = Path(
globals().get("__file__", "/tmp/x38_sprint_lt.py")
).resolve().parent
_BUILD_DIR = _FILE_DIR / ".build_x38_sprint_lt"
_BUILD_DIR.mkdir(exist_ok=True)
_CU13_ROOT = Path(torch.__file__).resolve().parent.parent / "nvidia" / "cu13"
_CU13_LIB = _CU13_ROOT / "lib"
load_inline(
name=f"xcaliber_x38_sprint_lt_{_SM_ARCH}_{_SOURCE_TAG}",
cpp_sources=CPP,
cuda_sources=CUDA,
is_python_module=False,
no_implicit_headers=True,
extra_include_paths=[str(_CU13_ROOT / "include")],
extra_cflags=[
"-O3",
f"-DXCAL_XRCHOLV18_NVFP4_OUTER={_NVFP4_OUTER}",
f"-DXCAL_XRCHOLV18_LEFT_MACRO={_LEFT_MACRO}",
f"-DXCAL_XRCHOLV18_K256_FRONT={_K256_FRONT}",
f"-DXCAL_XRCHOLV18_HEXLIFT_OUTER={_HEXLIFT_OUTER}",
f"-DXCAL_XRCHOLV18_HEXLIFT_MIN_CTAS={_HEXLIFT_MIN_CTAS}",
f"-DXCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY={_HEXLIFT_MIN_HISTORY}",
f"-DXCAL_XRCHOLV18_HEX6_ALL={_HEX6_ALL}",
f"-DXCAL_XRCHOLV18_HEX4_GE8192={_HEX4_GE8192}",
],
extra_cuda_cflags=[
"-O3",
"--extra-device-vectorization",
f"-gencode=arch=compute_{_SM_ARCH},code=sm_{_SM_ARCH}",
f"-DXCAL_XRCHOLV18_MAX_G={_MAX_G}",
f"-DXCAL_XRCHOLV18_PAYLOAD_PIPELINE={_PAYLOAD_PIPELINE}",
f"-DXCAL_XRCHOLV18_B60_ISSUES={_B60_ISSUES}",
f"-DXCAL_XRCHOLV18_B60_N128={_B60_N128}",
f"-DXCAL_XRTS_TMA_F16_B60={_TMA_F16_B60}",
f"-DXCAL_XRCHOLV18_NVFP4_OUTER={_NVFP4_OUTER}",
f"-DXCAL_XRCHOLV18_LEFT_MACRO={_LEFT_MACRO}",
f"-DXCAL_XRCHOLV18_K256_FRONT={_K256_FRONT}",
f"-DXCAL_XRCHOLV18_HEXLIFT_OUTER={_HEXLIFT_OUTER}",
f"-DXCAL_XRCHOLV18_HEXLIFT_MIN_CTAS={_HEXLIFT_MIN_CTAS}",
f"-DXCAL_XRCHOLV18_HEXLIFT_MIN_HISTORY={_HEXLIFT_MIN_HISTORY}",
f"-DXCAL_XRCHOLV18_HEX6_ALL={_HEX6_ALL}",
f"-DXCAL_XRCHOLV18_HEX4_GE8192={_HEX4_GE8192}",
],
extra_ldflags=[
f"-L{_CU13_LIB}",
f"-Wl,-rpath,{_CU13_LIB}",
"-l:libcublas.so.13",
"-l:libcublasLt.so.13",
"-lcuda",
],
build_directory=str(_BUILD_DIR),
with_cuda=True,
verbose=False,
)
def custom_kernel(data: input_t) -> output_t:
return torch.ops.xcaliber_x38_sprint_lt.chol(data)
scrolls · 14951 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