Compare commits
7
Commits
main
...
vsa_bwd_fix
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2d58c9bfde | ||
|
|
62f0bc1daa | ||
|
|
60165c7b65 | ||
|
|
d2c9d61240 | ||
|
|
57bb232432 | ||
|
|
88319b75d5 | ||
|
|
749a9c5d97 |
@@ -411,6 +411,12 @@ if(BUILD_CXX_KERNELS)
|
||||
set_source_files_properties(csrc/attention/block_sparse_sm100a.cu
|
||||
csrc/attention/block_sparse_blk128_sm100a.cu PROPERTIES
|
||||
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a;-gencode;arch=compute_103a,code=sm_103a;-DVSA_BHSD=true")
|
||||
# VSA block-sparse attention BACKWARD, 64-token blocks. sm_100a only for now: validated
|
||||
# on GB200, not yet on B300/GB300, so no sm_103a image is built and the Python side keeps
|
||||
# the Triton backward for sm_103a devices.
|
||||
list(APPEND EXTENSION_SOURCES csrc/attention/block_sparse_bwd_sm100a.cu)
|
||||
set_source_files_properties(csrc/attention/block_sparse_bwd_sm100a.cu PROPERTIES
|
||||
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a;-DVSA_BHSD=true")
|
||||
endif()
|
||||
|
||||
Python_add_library(fastvideo_kernel_ops MODULE USE_SABI ${SKBUILD_SABI_VERSION} WITH_SOABI
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,342 @@
|
||||
// block_sparse_bwd_launch_sm100a.cuh -- host surface of the VSA block-sparse backward drop:
|
||||
// argument struct, workspace sizes, the support predicate and the stream-chained launch
|
||||
// (preprocess -> order -> main -> postprocess). Tensor maps are encoded per call (no static
|
||||
// cache: a torch caller hands us fresh pointers every time).
|
||||
#ifndef BLOCK_SPARSE_VSA_BWD_LAUNCH_SM100A_CUH
|
||||
#define BLOCK_SPARSE_VSA_BWD_LAUNCH_SM100A_CUH
|
||||
|
||||
#include <algorithm>
|
||||
#include <cmath>
|
||||
#include "block_sparse_bwd_kernel_sm100a.cuh"
|
||||
|
||||
#ifndef VSA_BHSD
|
||||
#define VSA_BHSD false
|
||||
#endif
|
||||
#ifndef VSA_BWD_DQ_F16
|
||||
#define VSA_BWD_DQ_F16 false
|
||||
#endif
|
||||
#ifndef VSA_BWD_USE_CLC
|
||||
#define VSA_BWD_USE_CLC true
|
||||
#endif
|
||||
|
||||
namespace vsa_bwd_blk64 {
|
||||
|
||||
#if VSA_BWD_DQ_F16
|
||||
using dq_accum_t = uint16_t;
|
||||
#else
|
||||
using dq_accum_t = float;
|
||||
#endif
|
||||
|
||||
struct BlockSparseVsaBwdArgs {
|
||||
// Activations are bf16, contiguous, [B, H, S, 128] under VSA_BHSD, else [B*S, H, 128].
|
||||
// nb below = num_kv_blocks_per_seq = S / 64.
|
||||
|
||||
// Forward operands and results.
|
||||
const __nv_bfloat16* q;
|
||||
const __nv_bfloat16* k;
|
||||
const __nv_bfloat16* v;
|
||||
const __nv_bfloat16* o;
|
||||
// Gradient of the forward output.
|
||||
const __nv_bfloat16* dout;
|
||||
// [B, H, S] fp32 log-sum-exp in Triton's M form: max(qk * sm_scale * log2e) + log2(l).
|
||||
const float* lse;
|
||||
|
||||
// Sparsity metadata, FastVideo's invert_indices layout.
|
||||
// [B*H*nb, max_q_blocks] int32: q blocks selecting each kv block; entries past the count unread.
|
||||
const int* k2q_idx;
|
||||
// [B*H*nb] int32: valid entries per k2q_idx row (0 allowed).
|
||||
const int* k2q_num;
|
||||
// [nb] int32: valid kv tokens per block (<= 64); kv rows at or past the count are masked.
|
||||
const int* variable_block_sizes;
|
||||
|
||||
// Work order: which (batch, head, kv block) item each CTA processes.
|
||||
// [B*H*nb] int32 work id -> item ((b*H + h)*nb + kv). nullptr: identity order below
|
||||
// ORDER_MIN_KV_BLOCKS, else the launch computes the length-binned order into order_workspace.
|
||||
const int* workitem_remap;
|
||||
// [B*H*nb] int32; required when workitem_remap is nullptr and nb >= ORDER_MIN_KV_BLOCKS.
|
||||
int* order_workspace;
|
||||
|
||||
// Outputs, inputs' layout; dk/dv rows of unselected kv blocks are zeroed by the preprocess.
|
||||
__nv_bfloat16* dq;
|
||||
__nv_bfloat16* dk;
|
||||
__nv_bfloat16* dv;
|
||||
|
||||
// Scratch, caller-allocated; byte sizes from the block_sparse_bwd_*_bytes helpers below.
|
||||
// [B*H*S*128] dq_accum_t, drain-native; preprocess zeroes, main reduce-adds, postprocess reads.
|
||||
dq_accum_t* dqaccum;
|
||||
// [H*128, B*S] Q^T, written by the preprocess.
|
||||
__nv_bfloat16* qt;
|
||||
// [H*128, B*S] dO^T, written by the preprocess.
|
||||
__nv_bfloat16* dot;
|
||||
// [B*H*S] fp32 rowsum(bf16(o) * dout), written by the preprocess.
|
||||
float* delta;
|
||||
|
||||
int batch;
|
||||
int num_heads;
|
||||
// S; a multiple of 128 (the preprocess works in 128-token blocks).
|
||||
int seqlen;
|
||||
// Must be 128.
|
||||
int head_dim;
|
||||
// nb = seqlen / 64.
|
||||
int num_kv_blocks_per_seq;
|
||||
// k2q_idx row stride (FastVideo passes nb).
|
||||
int max_q_blocks;
|
||||
// Softmax scale; dq and dk carry it, dv does not.
|
||||
float sm_scale;
|
||||
};
|
||||
|
||||
__host__ inline size_t block_sparse_bwd_dqaccum_bytes(int batch, int num_heads, int seqlen) {
|
||||
return (size_t)batch * num_heads * seqlen * HEAD_DIM * sizeof(dq_accum_t);
|
||||
}
|
||||
__host__ inline size_t block_sparse_bwd_order_bytes(int batch, int num_heads,
|
||||
int num_kv_blocks_per_seq) {
|
||||
return (size_t)batch * num_heads * num_kv_blocks_per_seq * sizeof(int);
|
||||
}
|
||||
__host__ inline size_t block_sparse_bwd_transposed_bytes(int batch, int num_heads, int seqlen) {
|
||||
return (size_t)num_heads * HEAD_DIM * (size_t)batch * seqlen * sizeof(__nv_bfloat16);
|
||||
}
|
||||
__host__ inline size_t block_sparse_bwd_delta_bytes(int batch, int num_heads, int seqlen) {
|
||||
return (size_t)batch * num_heads * seqlen * sizeof(float);
|
||||
}
|
||||
|
||||
// Below this many kv blocks per sequence (S < 65536) the identity order is as fast as the
|
||||
// length-binned one and the order kernel's own time is not (fv_perf_log.md 2026-09-04: -9% at
|
||||
// 4k, -2.8% at 16k), so the main kernel runs the identity order there (workitem_remap == nullptr)
|
||||
// and the order kernel is left to the larger shapes.
|
||||
constexpr int ORDER_MIN_KV_BLOCKS = 1024;
|
||||
|
||||
__host__ inline cudaError_t block_sparse_bwd_supported(const BlockSparseVsaBwdArgs& args) {
|
||||
if (args.head_dim != HEAD_DIM) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (args.num_kv_blocks_per_seq < 1 || args.num_kv_blocks_per_seq % PRE_QBLOCKS != 0) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (args.seqlen != args.num_kv_blocks_per_seq * BLOCK) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (args.batch < 1 || args.num_heads < 1 || args.max_q_blocks < 1) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (!std::isfinite(args.sm_scale)) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (!args.q || !args.k || !args.v || !args.o || !args.dout || !args.lse) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (!args.dq || !args.dk || !args.dv) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (!args.k2q_idx || !args.k2q_num || !args.variable_block_sizes) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (!args.dqaccum || !args.qt || !args.dot || !args.delta) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
// No explicit order and large enough for the order kernel: it needs the workspace and two ints
|
||||
// of SMEM per kv block.
|
||||
if (!args.workitem_remap && args.num_kv_blocks_per_seq >= ORDER_MIN_KV_BLOCKS &&
|
||||
(!args.order_workspace || args.num_kv_blocks_per_seq > ORDER_MAX_BLOCKS)) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
return cudaSuccess;
|
||||
}
|
||||
|
||||
// K, V, dK, dV tensor maps, one 64-token x 64-hd box per TMA (two per tile):
|
||||
// BSHD: 3D [64 hd, B*S tokens, H*2 hd units], strides {H*128*2, 128} bytes.
|
||||
// BHSD: 4D [64 hd, S tokens, 2 hd units, B*H], strides {128*2, 128, S*128*2} bytes.
|
||||
__host__ inline cudaError_t make_tma_kv_units(CUtensorMap* map, const __nv_bfloat16* ptr, int B,
|
||||
int H, int S) {
|
||||
CUresult r;
|
||||
if (VSA_BHSD) {
|
||||
uint64_t gd[4] = {(uint64_t)SUB_COLS_BF16, (uint64_t)S, (uint64_t)KV_SUBTILES, (uint64_t)B * H};
|
||||
uint64_t gs[3] = {(uint64_t)HEAD_DIM * 2, (uint64_t)SUB_COLS_BYTES, (uint64_t)S * HEAD_DIM * 2};
|
||||
uint32_t bd[4] = {(uint32_t)SUB_COLS_BF16, (uint32_t)BLOCK, 1u, 1u};
|
||||
uint32_t es[4] = {1u, 1u, 1u, 1u};
|
||||
r = cuTensorMapEncodeTiled(
|
||||
map, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, const_cast<__nv_bfloat16*>(ptr), gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
|
||||
} else {
|
||||
uint64_t gd[3] = {(uint64_t)SUB_COLS_BF16, (uint64_t)B * S, (uint64_t)H * KV_SUBTILES};
|
||||
uint64_t gs[2] = {(uint64_t)H * HEAD_DIM * 2, (uint64_t)SUB_COLS_BYTES};
|
||||
uint32_t bd[3] = {(uint32_t)SUB_COLS_BF16, (uint32_t)BLOCK, 1u};
|
||||
uint32_t es[3] = {1u, 1u, 1u};
|
||||
r = cuTensorMapEncodeTiled(
|
||||
map, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 3, const_cast<__nv_bfloat16*>(ptr), gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
|
||||
}
|
||||
return (r == CUDA_SUCCESS) ? cudaSuccess : cudaErrorInvalidValue;
|
||||
}
|
||||
|
||||
// Above the L2-capacity transition the main kernel is DRAM-bound and keeping a fractional
|
||||
// subset of the repeatedly reduced dQ lines resident pays (~0.6% at 524K/1M tokens); below it the
|
||||
// policy register costs more than it saves. Threshold scales with accumulator BYTES.
|
||||
constexpr int CACHE_WAVE_MIN_SEQ_LEN = 524288;
|
||||
// The in-kernel identity order (workitem_remap == nullptr below ORDER_MIN_KV_BLOCKS) has no
|
||||
// sub-array to offset a chunk into, so the chunked launches that DQ_L2_KEEP enables must only
|
||||
// ever run with a device-computed order: the L2 transition has to sit at or above the
|
||||
// order-kernel threshold. keep_dq_l2 <=> S * sizeof(dq_accum_t) >= 2 * CACHE_WAVE_MIN_SEQ_LEN.
|
||||
static_assert((size_t)CACHE_WAVE_MIN_SEQ_LEN * 2 / sizeof(dq_accum_t) >=
|
||||
(size_t)ORDER_MIN_KV_BLOCKS * BLOCK,
|
||||
"DQ_L2_KEEP chunked launches need the device-computed work order");
|
||||
|
||||
template <bool DQ_L2_KEEP, bool USE_CLC, bool BHSD>
|
||||
__host__ inline cudaError_t launch_main(const BlockSparseVsaBwdArgs& args, const int* work_remap,
|
||||
const CUtensorMap& tk, const CUtensorMap& tv,
|
||||
const CUtensorMap& tqt, const CUtensorMap& tdot,
|
||||
const CUtensorMap& tdk, const CUtensorMap& tdv, int sms,
|
||||
cudaStream_t stream) {
|
||||
auto kernel = vsa_bwd_main_kernel<DQ_L2_KEEP, USE_CLC, BHSD, dq_accum_t>;
|
||||
cudaError_t e =
|
||||
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_TOTAL);
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
const float scale_log2 = args.sm_scale * 1.4426950408889634f;
|
||||
const int B = args.batch, H = args.num_heads, S = args.seqlen;
|
||||
const int total = B * H * args.num_kv_blocks_per_seq;
|
||||
if constexpr (USE_CLC) {
|
||||
// Above the L2 transition, one SM-wide launch at a time keeps each list neighbourhood
|
||||
// resident; below it one launch of every item lets CLC steal freely.
|
||||
const int chunk = DQ_L2_KEEP ? std::min(sms, total) : total;
|
||||
cudaLaunchConfig_t cfg = {};
|
||||
cfg.blockDim = dim3(N_WARPS * 32, 1, 1);
|
||||
cfg.dynamicSmemBytes = SMEM_TOTAL;
|
||||
cfg.stream = stream;
|
||||
cudaLaunchAttribute at[1];
|
||||
at[0].id = cudaLaunchAttributeClusterDimension;
|
||||
at[0].val.clusterDim.x = 1;
|
||||
at[0].val.clusterDim.y = 1;
|
||||
at[0].val.clusterDim.z = 1;
|
||||
cfg.attrs = at;
|
||||
cfg.numAttrs = 1;
|
||||
// Guarded by the static_assert above; kept as a runtime check for other callers of the
|
||||
// launch API that pass their own thresholds.
|
||||
if (work_remap == nullptr && chunk != total) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
for (int base = 0; base < total; base += chunk) {
|
||||
const int count = std::min(chunk, total - base);
|
||||
cfg.gridDim = dim3((unsigned)count, 1, 1);
|
||||
// A chunk starts at work id `base`: it gets the order's sub-array.
|
||||
e = cudaLaunchKernelEx(&cfg, kernel, tk, tv, tqt, tdot, tdk, tdv, args.dqaccum, args.lse,
|
||||
args.delta, args.k2q_idx, args.k2q_num,
|
||||
work_remap ? work_remap + base : nullptr,
|
||||
args.variable_block_sizes, args.max_q_blocks, B, H, S, scale_log2,
|
||||
args.sm_scale);
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
}
|
||||
return cudaSuccess;
|
||||
} else {
|
||||
const int grid = std::min(total, sms);
|
||||
kernel<<<dim3((unsigned)grid, 1, 1), dim3(N_WARPS * 32, 1, 1), SMEM_TOTAL, stream>>>(
|
||||
tk, tv, tqt, tdot, tdk, tdv, args.dqaccum, args.lse, args.delta, args.k2q_idx, args.k2q_num,
|
||||
work_remap, args.variable_block_sizes, args.max_q_blocks, B, H, S, scale_log2,
|
||||
args.sm_scale);
|
||||
return cudaGetLastError();
|
||||
}
|
||||
}
|
||||
|
||||
__host__ inline cudaError_t launch_block_sparse_bwd_sm100a(const BlockSparseVsaBwdArgs& args,
|
||||
cudaStream_t stream) {
|
||||
const cudaError_t supported = block_sparse_bwd_supported(args);
|
||||
if (supported != cudaSuccess) {
|
||||
return supported;
|
||||
}
|
||||
const int B = args.batch, H = args.num_heads, S = args.seqlen;
|
||||
const long n_tokens = (long)B * S;
|
||||
|
||||
CUtensorMap tk, tv, tqt, tdot, tdk, tdv;
|
||||
if (make_tma_kv_units(&tk, args.k, B, H, S) != cudaSuccess) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (make_tma_kv_units(&tv, args.v, B, H, S) != cudaSuccess) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (make_tma_kv_units(&tdk, args.dk, B, H, S) != cudaSuccess) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (make_tma_kv_units(&tdv, args.dv, B, H, S) != cudaSuccess) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
// Q^T, dO^T ([H*hd rows, B*S cols], token contiguous): box [hd rows, BLOCK cols] = one q64
|
||||
// block per TMA.
|
||||
if (make_tma_2d_tiled(&tqt, args.qt, H * HEAD_DIM, (int)n_tokens, HEAD_DIM, BLOCK, 2,
|
||||
CU_TENSOR_MAP_DATA_TYPE_BFLOAT16) != cudaSuccess) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
if (make_tma_2d_tiled(&tdot, args.dot, H * HEAD_DIM, (int)n_tokens, HEAD_DIM, BLOCK, 2,
|
||||
CU_TENSOR_MAP_DATA_TYPE_BFLOAT16) != cudaSuccess) {
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
|
||||
int dev = 0, sms = 0;
|
||||
cudaError_t e = cudaGetDevice(&dev);
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
e = cudaDeviceGetAttribute(&sms, cudaDevAttrMultiProcessorCount, dev);
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
|
||||
vsa_bwd_preprocess_kernel<VSA_BHSD, dq_accum_t>
|
||||
<<<dim3((unsigned)(S / PRE_TOKENS), (unsigned)(B * H), 1), dim3(256, 1, 1), 0, stream>>>(
|
||||
args.q, args.o, args.dout, args.delta, args.dqaccum, args.qt, args.dot, args.dk, args.dv,
|
||||
args.k2q_num, B, H, S);
|
||||
e = cudaGetLastError();
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
|
||||
const bool keep_dq_l2 = (size_t)S * sizeof(dq_accum_t) >= (size_t)CACHE_WAVE_MIN_SEQ_LEN * 2;
|
||||
|
||||
// Work-item order: explicit if the caller passed one; else identity (nullptr) below
|
||||
// ORDER_MIN_KV_BLOCKS, else computed on device into order_workspace (length bins; the same L2
|
||||
// transition that selects DQ_L2_KEEP selects the wider bins plus the midpoint snake).
|
||||
const int* work_remap = args.workitem_remap;
|
||||
if (work_remap == nullptr && args.num_kv_blocks_per_seq >= ORDER_MIN_KV_BLOCKS) {
|
||||
const int order_smem = 2 * args.num_kv_blocks_per_seq * (int)sizeof(int);
|
||||
if (order_smem > 48 * 1024) {
|
||||
e = cudaFuncSetAttribute(vsa_bwd_order_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
order_smem);
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
}
|
||||
const unsigned item_chunks =
|
||||
(unsigned)((args.num_kv_blocks_per_seq + ORDER_THREADS - 1) / ORDER_THREADS);
|
||||
vsa_bwd_order_kernel<<<dim3((unsigned)(B * H), item_chunks, 1), dim3(ORDER_THREADS, 1, 1),
|
||||
order_smem, stream>>>(args.k2q_idx, args.k2q_num, args.max_q_blocks,
|
||||
args.num_kv_blocks_per_seq, keep_dq_l2 ? 12 : 8,
|
||||
keep_dq_l2, args.order_workspace);
|
||||
e = cudaGetLastError();
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
work_remap = args.order_workspace;
|
||||
}
|
||||
|
||||
e = keep_dq_l2 ? launch_main<true, VSA_BWD_USE_CLC, VSA_BHSD>(args, work_remap, tk, tv, tqt, tdot,
|
||||
tdk, tdv, sms, stream)
|
||||
: launch_main<false, VSA_BWD_USE_CLC, VSA_BHSD>(args, work_remap, tk, tv, tqt,
|
||||
tdot, tdk, tdv, sms, stream);
|
||||
if (e != cudaSuccess) {
|
||||
return e;
|
||||
}
|
||||
|
||||
vsa_bwd_postprocess_kernel<VSA_BHSD, dq_accum_t>
|
||||
<<<dim3((unsigned)(S / BLOCK), (unsigned)(B * H), 1), dim3(128, 1, 1), 0, stream>>>(
|
||||
args.dqaccum, args.dq, H, S, args.sm_scale);
|
||||
return cudaGetLastError();
|
||||
}
|
||||
|
||||
} // namespace vsa_bwd_blk64
|
||||
|
||||
using namespace vsa_bwd_blk64;
|
||||
|
||||
#endif // BLOCK_SPARSE_VSA_BWD_LAUNCH_SM100A_CUH
|
||||
@@ -0,0 +1,166 @@
|
||||
// block_sparse_bwd_sm100a.cu -- torch binding for the sm_100a VSA block-sparse FMHA backward.
|
||||
//
|
||||
// Pairs with block_sparse_sm100a_fwd. Inputs are the forward's operands plus its output o and
|
||||
// its lse; lse is the Triton "M format" tensor the forward returns -- [B, H, S] fp32,
|
||||
// M = max(qk * sm_scale * log2e) + log2(l) -- and is consumed as-is. Sparsity arrives as
|
||||
// FastVideo's k2q metadata (fastvideo_kernel.triton_kernels.index.invert_indices): for every
|
||||
// (batch, head, kv block) the LOCAL q64 block ids that selected it, padded to max_q_blocks,
|
||||
// plus a count; entries past the count are never read. Returns {dq, dk, dv} in bf16 with the
|
||||
// inputs' layout and the Triton backward's scaling: dk and dq carry sm_scale, dv does not.
|
||||
//
|
||||
// The layout is fixed at compile time: VSA_BHSD true -> [B, H, S, 128] (FastVideo's build),
|
||||
// false -> [B, S, H, 128] (repo native; the kernel addresses it as [B*S tokens, H, 128]).
|
||||
#include <torch/extension.h>
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <vector>
|
||||
|
||||
#include "block_sparse_bwd_launch_sm100a.cuh"
|
||||
|
||||
namespace {
|
||||
|
||||
void check_activation(const torch::Tensor& t, const char* name, int64_t B, int64_t H, int64_t S,
|
||||
int64_t D) {
|
||||
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
|
||||
TORCH_CHECK(t.scalar_type() == at::kBFloat16, name, " must be bfloat16, got ", t.scalar_type());
|
||||
TORCH_CHECK(t.dim() == 4, name, " must be 4-D, got ", t.dim(), " dims");
|
||||
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
|
||||
if (VSA_BHSD) {
|
||||
TORCH_CHECK(t.size(0) == B && t.size(1) == H && t.size(2) == S && t.size(3) == D, name,
|
||||
" has shape ", t.sizes(), ", expected [", B, ",", H, ",", S, ",", D, "]");
|
||||
} else {
|
||||
TORCH_CHECK(t.size(0) == B && t.size(1) == S && t.size(2) == H && t.size(3) == D, name,
|
||||
" has shape ", t.sizes(), ", expected [", B, ",", S, ",", H, ",", D, "]");
|
||||
}
|
||||
}
|
||||
|
||||
void check_index(const torch::Tensor& t, const char* name) {
|
||||
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
|
||||
TORCH_CHECK(t.scalar_type() == at::kInt, name, " must be int32, got ", t.scalar_type());
|
||||
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
|
||||
}
|
||||
|
||||
__nv_bfloat16* bf16_ptr(const torch::Tensor& t) {
|
||||
return reinterpret_cast<__nv_bfloat16*>(t.data_ptr());
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
// Returns {dq, dk, dv}: bf16, each with the shape and layout of q, k, v respectively.
|
||||
std::vector<torch::Tensor> block_sparse_sm100a_bwd(torch::Tensor grad_o, torch::Tensor q,
|
||||
torch::Tensor k, torch::Tensor v,
|
||||
torch::Tensor o, torch::Tensor lse,
|
||||
torch::Tensor k2q_idx, torch::Tensor k2q_num,
|
||||
torch::Tensor variable_block_sizes,
|
||||
double sm_scale) {
|
||||
const c10::cuda::OptionalCUDAGuard guard(device_of(q));
|
||||
|
||||
TORCH_CHECK(q.dim() == 4, "q must be 4-D, got ", q.dim(), " dims");
|
||||
const int64_t B = q.size(0);
|
||||
const int64_t H = VSA_BHSD ? q.size(1) : q.size(2);
|
||||
const int64_t S = VSA_BHSD ? q.size(2) : q.size(1);
|
||||
const int64_t D = q.size(3);
|
||||
|
||||
check_activation(q, "q", B, H, S, D);
|
||||
check_activation(k, "k", B, H, S, D);
|
||||
check_activation(v, "v", B, H, S, D);
|
||||
check_activation(o, "o", B, H, S, D);
|
||||
check_activation(grad_o, "grad_o", B, H, S, D);
|
||||
|
||||
TORCH_CHECK(lse.is_cuda(), "lse must be a CUDA tensor");
|
||||
TORCH_CHECK(lse.scalar_type() == at::kFloat, "lse must be float32, got ", lse.scalar_type());
|
||||
TORCH_CHECK(lse.is_contiguous(), "lse must be contiguous");
|
||||
TORCH_CHECK(lse.numel() == B * H * S, "lse must hold [B, H, S] = ", B * H * S,
|
||||
" values (Triton M format), got ", lse.numel());
|
||||
|
||||
check_index(k2q_idx, "k2q_idx");
|
||||
check_index(k2q_num, "k2q_num");
|
||||
check_index(variable_block_sizes, "variable_block_sizes");
|
||||
|
||||
const int64_t num_kv_blocks_per_seq = variable_block_sizes.numel();
|
||||
TORCH_CHECK(S == num_kv_blocks_per_seq * BLOCK, "seqlen ", S,
|
||||
" must equal num_kv_blocks_per_seq (", num_kv_blocks_per_seq, ") * ", BLOCK,
|
||||
"; FastVideo pads the sequence up to whole blocks");
|
||||
|
||||
const int64_t num_items = B * H * num_kv_blocks_per_seq;
|
||||
TORCH_CHECK(k2q_idx.dim() == 4 || k2q_idx.dim() == 2,
|
||||
"k2q_idx must be [B, H, num_kv_blocks_per_seq, max_q_blocks] or "
|
||||
"[B*H*num_kv_blocks_per_seq, max_q_blocks], got shape ",
|
||||
k2q_idx.sizes());
|
||||
const int64_t max_q_blocks = k2q_idx.size(-1);
|
||||
const bool k2q_dims_ok = k2q_idx.dim() == 2 || (k2q_idx.size(0) == B && k2q_idx.size(1) == H &&
|
||||
k2q_idx.size(2) == num_kv_blocks_per_seq);
|
||||
TORCH_CHECK(k2q_dims_ok && k2q_idx.numel() == num_items * max_q_blocks, "k2q_idx has shape ",
|
||||
k2q_idx.sizes(), ", expected [", B, ",", H, ",", num_kv_blocks_per_seq,
|
||||
",max_q_blocks] or [", num_items, ",max_q_blocks]");
|
||||
TORCH_CHECK(k2q_num.numel() == num_items,
|
||||
"k2q_num must hold one count per (batch, head, kv "
|
||||
"block) = ",
|
||||
num_items, " values, got ", k2q_num.numel());
|
||||
|
||||
auto dq = torch::empty_like(q);
|
||||
auto dk = torch::empty_like(k);
|
||||
auto dv = torch::empty_like(v);
|
||||
|
||||
// Workspace: torch::empty is enough. The preprocess kernel zeroes dqaccum and fully writes
|
||||
// qt, dot and delta before the main kernel reads them; it also zeroes the dk/dv rows of kv
|
||||
// blocks that no q block selects, so the empty_like outputs above come back fully defined.
|
||||
const auto bytes = q.options().dtype(at::kByte);
|
||||
const int b = (int)B, h = (int)H, s = (int)S;
|
||||
auto dqaccum = torch::empty({(int64_t)block_sparse_bwd_dqaccum_bytes(b, h, s)}, bytes);
|
||||
auto qt = torch::empty({(int64_t)block_sparse_bwd_transposed_bytes(b, h, s)}, bytes);
|
||||
auto dot = torch::empty({(int64_t)block_sparse_bwd_transposed_bytes(b, h, s)}, bytes);
|
||||
auto delta = torch::empty({(int64_t)block_sparse_bwd_delta_bytes(b, h, s)}, bytes);
|
||||
// Work-item order: from ORDER_MIN_KV_BLOCKS on, the launch computes the length-binned order
|
||||
// into this workspace; below, the main kernel runs the identity order (no array).
|
||||
torch::Tensor order;
|
||||
const bool device_order = num_kv_blocks_per_seq >= ORDER_MIN_KV_BLOCKS;
|
||||
if (device_order) {
|
||||
order = torch::empty({(int64_t)block_sparse_bwd_order_bytes(b, h, (int)num_kv_blocks_per_seq)},
|
||||
bytes);
|
||||
}
|
||||
|
||||
BlockSparseVsaBwdArgs a{};
|
||||
a.q = bf16_ptr(q);
|
||||
a.k = bf16_ptr(k);
|
||||
a.v = bf16_ptr(v);
|
||||
a.o = bf16_ptr(o);
|
||||
a.dout = bf16_ptr(grad_o);
|
||||
a.dq = bf16_ptr(dq);
|
||||
a.dk = bf16_ptr(dk);
|
||||
a.dv = bf16_ptr(dv);
|
||||
a.lse = lse.data_ptr<float>();
|
||||
a.k2q_idx = k2q_idx.data_ptr<int>();
|
||||
a.k2q_num = k2q_num.data_ptr<int>();
|
||||
a.variable_block_sizes = variable_block_sizes.data_ptr<int>();
|
||||
a.workitem_remap = nullptr;
|
||||
a.order_workspace = device_order ? reinterpret_cast<int*>(order.data_ptr()) : nullptr;
|
||||
a.dqaccum = reinterpret_cast<dq_accum_t*>(dqaccum.data_ptr());
|
||||
a.qt = bf16_ptr(qt);
|
||||
a.dot = bf16_ptr(dot);
|
||||
a.delta = reinterpret_cast<float*>(delta.data_ptr());
|
||||
a.batch = b;
|
||||
a.num_heads = h;
|
||||
a.seqlen = s;
|
||||
a.head_dim = (int)D;
|
||||
a.num_kv_blocks_per_seq = (int)num_kv_blocks_per_seq;
|
||||
a.max_q_blocks = (int)max_q_blocks;
|
||||
a.sm_scale = (float)sm_scale;
|
||||
|
||||
// Report an unsupported regime loudly rather than returning plausible-looking wrong values.
|
||||
TORCH_CHECK(block_sparse_bwd_supported(a) == cudaSuccess,
|
||||
"block_sparse_sm100a_bwd: unsupported configuration -- requires head_dim==", HEAD_DIM,
|
||||
", seqlen == num_kv_blocks_per_seq*", BLOCK,
|
||||
" with seqlen % 128 == 0, "
|
||||
"max_q_blocks >= 1 and a finite sm_scale. Got head_dim=",
|
||||
D, " num_kv_blocks_per_seq=", num_kv_blocks_per_seq, " seqlen=", S,
|
||||
" max_q_blocks=", max_q_blocks, " sm_scale=", sm_scale);
|
||||
|
||||
const cudaError_t err = launch_block_sparse_bwd_sm100a(a, at::cuda::getCurrentCUDAStream());
|
||||
TORCH_CHECK(err == cudaSuccess,
|
||||
"block_sparse_sm100a_bwd launch failed: ", cudaGetErrorString(err));
|
||||
|
||||
return {dq, dk, dv};
|
||||
}
|
||||
@@ -874,3 +874,114 @@ __device__ __forceinline__
|
||||
void sts_f32(uint32_t smem_addr, float val) {
|
||||
asm volatile("st.shared.f32 [%0], %1;" :: "r"(smem_addr), "f"(val) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_mma_ws_f16_ss_1sm_predicated(
|
||||
uint32_t issue, uint32_t tmem_d, uint64_t desc_a, uint64_t desc_b,
|
||||
uint32_t idesc, bool enable_input_d) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p, q;\n\t"
|
||||
"setp.ne.b32 q, %0, 0;\n\t"
|
||||
"setp.ne.b32 p, %5, 0;\n\t"
|
||||
"@q tcgen05.mma.ws.cta_group::1.kind::f16 "
|
||||
"[%1], %2, %3, %4, p, 0;\n\t"
|
||||
"}\n"
|
||||
:: "r"(issue), "r"(tmem_d), "l"(desc_a), "l"(desc_b),
|
||||
"r"(idesc), "r"(enable_input_d ? 1u : 0u));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_mma_ws_f16_ts_1sm_predicated(
|
||||
uint32_t issue, uint32_t tmem_d, uint32_t tmem_a, uint64_t desc_b,
|
||||
uint32_t idesc, bool enable_input_d) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p, q;\n\t"
|
||||
"setp.ne.b32 q, %0, 0;\n\t"
|
||||
"setp.ne.b32 p, %5, 0;\n\t"
|
||||
"@q tcgen05.mma.ws.cta_group::1.kind::f16 "
|
||||
"[%1], [%2], %3, %4, p, 0;\n\t"
|
||||
"}\n"
|
||||
:: "r"(issue), "r"(tmem_d), "r"(tmem_a), "l"(desc_b),
|
||||
"r"(idesc), "r"(enable_input_d ? 1u : 0u));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_wait_ld() {
|
||||
asm volatile("tcgen05.wait::ld.sync.aligned;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void cpasync_bulk_load_mbarrier(uint32_t smem_dst, const void* gmem_src,
|
||||
uint32_t bytes, uint32_t mbar_smem) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1], %2, [%3];\n"
|
||||
:: "r"(smem_dst), "l"(gmem_src), "r"(bytes), "r"(mbar_smem)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void cpasync_reduce_bulk_add_f32(float* global_dst, uint32_t smem_src,
|
||||
uint32_t bytes) {
|
||||
asm volatile(
|
||||
"cp.reduce.async.bulk.global.shared::cta.bulk_group.add.f32"
|
||||
" [%0], [%1], %2;\n"
|
||||
:: "l"(global_dst), "r"(smem_src), "r"(bytes)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void cpasync_reduce_bulk_add_f16(uint16_t* global_dst, uint32_t smem_src,
|
||||
uint32_t bytes) {
|
||||
asm volatile(
|
||||
"cp.reduce.async.bulk.global.shared::cta.bulk_group.add.noftz.f16"
|
||||
" [%0], [%1], %2;\n"
|
||||
:: "l"(global_dst), "r"(smem_src), "r"(bytes)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void cpasync_reduce_bulk_add_f32_l2hint(float* global_dst, uint32_t smem_src,
|
||||
uint32_t bytes, uint64_t cache_policy) {
|
||||
asm volatile(
|
||||
"cp.reduce.async.bulk.global.shared::cta.bulk_group.L2::cache_hint.add.f32"
|
||||
" [%0], [%1], %2, %3;\n"
|
||||
:: "l"(global_dst), "r"(smem_src), "r"(bytes), "l"(cache_policy)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void cpasync_reduce_bulk_add_f16_l2hint(uint16_t* global_dst, uint32_t smem_src,
|
||||
uint32_t bytes, uint64_t cache_policy) {
|
||||
asm volatile(
|
||||
"cp.reduce.async.bulk.global.shared::cta.bulk_group.L2::cache_hint.add.noftz.f16"
|
||||
" [%0], [%1], %2, %3;\n"
|
||||
:: "l"(global_dst), "r"(smem_src), "r"(bytes), "l"(cache_policy)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void smem_desc_add_lo(uint64_t& d, uint32_t inc) {
|
||||
asm volatile("{\n\t"
|
||||
".reg .b32 lo, hi;\n\t"
|
||||
"mov.b64 {lo, hi}, %0;\n\t"
|
||||
"add.u32 lo, lo, %1;\n\t"
|
||||
"mov.b64 %0, {lo, hi};\n\t"
|
||||
"}" : "+l"(d) : "r"(inc));
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t cvt_f32x2_to_f16x2(float a, float b) {
|
||||
uint32_t r;
|
||||
asm volatile("cvt.rn.f16x2.f32 %0, %2, %1;\n"
|
||||
: "=r"(r) : "f"(a), "f"(b));
|
||||
return r;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint64_t make_l2cache_policy_fractional_evict_last_unchanged(float fraction_f32) {
|
||||
uint64_t policy;
|
||||
asm("createpolicy.fractional.L2::evict_last.L2::evict_unchanged.b64"
|
||||
" %0, %1;\n"
|
||||
: "=l"(policy)
|
||||
: "f"(fraction_f32));
|
||||
return policy;
|
||||
}
|
||||
|
||||
@@ -42,6 +42,10 @@ extern std::vector<torch::Tensor> block_sparse_sm100a_blk128_fwd(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
|
||||
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
|
||||
double sm_scale, bool need_lse);
|
||||
extern std::vector<torch::Tensor> block_sparse_sm100a_bwd(
|
||||
torch::Tensor grad_o, torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
torch::Tensor lse, torch::Tensor k2q_idx, torch::Tensor k2q_num,
|
||||
torch::Tensor variable_block_sizes, double sm_scale);
|
||||
#endif
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
@@ -54,6 +58,9 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("block_sparse_sm100a_blk128_fwd",
|
||||
torch::wrap_pybind_function(block_sparse_sm100a_blk128_fwd),
|
||||
"VSA block-sparse attention forward, 128-token blocks (Blackwell sm100a/sm103a)");
|
||||
m.def("block_sparse_sm100a_bwd",
|
||||
torch::wrap_pybind_function(block_sparse_sm100a_bwd),
|
||||
"VSA block-sparse attention backward, 64-token blocks (Blackwell sm100a)");
|
||||
#endif
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
|
||||
@@ -372,13 +372,65 @@ block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_co
|
||||
# ---------------------------------------------------------------------------
|
||||
# Data-center Blackwell backend custom op (index-native; legacy sm100a API name)
|
||||
#
|
||||
# Forward runs the sm_100a/sm_103a CUDA extension; backward reuses the Triton kernels.
|
||||
# Forward runs the sm_100a/sm_103a CUDA extension. Backward runs the sm_100a CUDA
|
||||
# backward when block_sparse_attn_bwd_sm100a.is_supported passes (64-token blocks,
|
||||
# sm_100a device, extension built with the op) and the Triton kernels otherwise.
|
||||
# The native forward emits lse in exactly Triton's M format (max*log2e +
|
||||
# log2(l)), so the pairing needs no conversion. The Triton backward is
|
||||
# log2(l)), so either pairing needs no conversion. Both backwards are
|
||||
# hardcoded to 64-token blocks, hence the block-size assert below.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_backward_sm100a",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def block_sparse_attn_backward_sm100a(
|
||||
grad_o: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
lse: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
from fastvideo_kernel.block_sparse_attn_bwd_sm100a import (
|
||||
block_sparse_attn_backward_sm100a_from_k2q, )
|
||||
|
||||
num_kv_blocks = variable_block_sizes.numel()
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(q2k_idx, q2k_num, num_kv_blocks)
|
||||
dq, dk, dv = block_sparse_attn_backward_sm100a_from_k2q(
|
||||
grad_o.contiguous(), q.contiguous(), k.contiguous(), v.contiguous(), o.contiguous(),
|
||||
lse.contiguous(), k2q_idx, k2q_num, variable_block_sizes)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm100a")
|
||||
def _block_sparse_attn_backward_sm100a_fake(
|
||||
grad_o: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
lse: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
return torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
|
||||
|
||||
|
||||
def _sm100a_backward_is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
|
||||
try:
|
||||
from fastvideo_kernel import block_sparse_attn_bwd_sm100a as vsa_bwd_sm100a
|
||||
except ImportError: # pragma: no cover - extension not built
|
||||
return False
|
||||
return vsa_bwd_sm100a.is_supported(q, variable_block_sizes)
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_sm100a",
|
||||
mutates_args=(),
|
||||
@@ -428,11 +480,15 @@ def _backward_sm100a(ctx, grad_o, grad_M):
|
||||
block = q.shape[2] // variable_block_sizes.numel()
|
||||
if block != 64:
|
||||
raise RuntimeError(
|
||||
"block_sparse_attn_sm100a backward pairs the sm_100a/sm_103a forward with the "
|
||||
f"Triton backward, which is hardcoded to 64-token blocks; got {block}. "
|
||||
"block_sparse_attn_sm100a backward pairs the sm_100a/sm_103a forward with a "
|
||||
f"backward that is hardcoded to 64-token blocks; got {block}. "
|
||||
"Run 128-token-block metadata without grad, or use the Triton forward.")
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, q2k_idx,
|
||||
q2k_num, variable_block_sizes)
|
||||
if _sm100a_backward_is_supported(q, variable_block_sizes):
|
||||
dq, dk, dv = block_sparse_attn_backward_sm100a(grad_o, q, k, v, o, M, q2k_idx,
|
||||
q2k_num, variable_block_sizes)
|
||||
else:
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, q2k_idx,
|
||||
q2k_num, variable_block_sizes)
|
||||
return dq, dk, dv, None, None, None
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""sm_100a (Blackwell) CUDA block-sparse VSA backward.
|
||||
|
||||
Companion of ``block_sparse_attn_sm100a`` (the forward): consumes the forward's ``lse`` in the
|
||||
Triton "M format" (``max(qk * sm_scale * log2e) + log2(l)``, ``[B, H, S]`` fp32) unchanged and
|
||||
FastVideo's k2q index metadata, returns ``(dq, dk, dv)`` in bf16 with the inputs' layout and the
|
||||
Triton backward's scaling (dq and dk carry sm_scale, dv does not). 64-token blocks only; any
|
||||
other configuration falls back to Triton via ``is_supported``.
|
||||
"""
|
||||
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
# The pybind symbols live on fastvideo_kernel_ops, NOT on the _C package that contains it
|
||||
# (its __init__ is empty, so hasattr on the package fails with the kernel built and present).
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops as _C
|
||||
_BWD = getattr(_C, "block_sparse_sm100a_bwd", None)
|
||||
_HAS_VSA_BWD_SM100A = _BWD is not None
|
||||
except ImportError: # pragma: no cover - extension not built
|
||||
_C = None
|
||||
_BWD = None
|
||||
_HAS_VSA_BWD_SM100A = False
|
||||
|
||||
_SM100 = (10, 0)
|
||||
HEAD_DIM = 128
|
||||
BLOCK = 64
|
||||
# Must match the -DVSA_BHSD the extension was compiled with (FastVideo builds with true).
|
||||
BHSD = True
|
||||
|
||||
|
||||
def set_extension(module) -> None:
|
||||
"""Use an already-loaded extension module exposing ``block_sparse_sm100a_bwd``.
|
||||
|
||||
A standalone build of ``block_sparse_bwd_sm100a.cu`` (for example through
|
||||
``torch.utils.cpp_extension.load`` with a ten-line pybind wrapper) can be injected here, so
|
||||
the backend can be exercised without rebuilding the fastvideo_kernel wheel.
|
||||
"""
|
||||
global _C, _BWD, _HAS_VSA_BWD_SM100A
|
||||
_C = module
|
||||
_BWD = getattr(module, "block_sparse_sm100a_bwd", None)
|
||||
_HAS_VSA_BWD_SM100A = _BWD is not None
|
||||
|
||||
|
||||
def _seqlen(q: torch.Tensor) -> int:
|
||||
return q.shape[2] if BHSD else q.shape[1]
|
||||
|
||||
|
||||
def is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
|
||||
"""True iff this build can run these tensors; otherwise the caller uses Triton.
|
||||
|
||||
Static facts only (shapes, dtypes, arch, layout), never tensor contents, so it is cheap
|
||||
enough for a per-layer dispatch path. The kernel is fixed at 64-token blocks with
|
||||
head_dim 128 and needs seqlen == 64 * num_blocks with an even num_blocks (its preprocess
|
||||
works in 128-token blocks). Per-row k2q counts may be anything in [0, num_q_blocks],
|
||||
including 0: unselected kv blocks get exactly-zero dk/dv rows.
|
||||
"""
|
||||
if not _HAS_VSA_BWD_SM100A or not q.is_cuda:
|
||||
return False
|
||||
if torch.cuda.get_device_capability(q.device) != _SM100:
|
||||
return False
|
||||
if q.dtype != torch.bfloat16 or q.dim() != 4 or q.shape[-1] != HEAD_DIM:
|
||||
return False
|
||||
if not q.is_contiguous():
|
||||
return False
|
||||
# Metadata must be integer-typed so the wrapper's int32 conversion is value-preserving.
|
||||
if not variable_block_sizes.is_cuda or variable_block_sizes.dtype not in (torch.int32,
|
||||
torch.int64):
|
||||
return False
|
||||
num_blocks = variable_block_sizes.numel()
|
||||
if num_blocks == 0 or num_blocks % 2 != 0:
|
||||
return False
|
||||
if _seqlen(q) != BLOCK * num_blocks:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def block_sparse_attn_backward_sm100a_from_k2q(
|
||||
grad_o: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
lse: torch.Tensor,
|
||||
k2q_idx: torch.Tensor,
|
||||
k2q_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Backward from k2q metadata already in hand (``invert_indices`` layout).
|
||||
|
||||
``k2q_idx`` is ``[B, H, num_kv_blocks, max_q_blocks]`` (or the flat 2-D view) of LOCAL q64
|
||||
block ids, ``k2q_num`` ``[B, H, num_kv_blocks]``; entries past a row's count are never read.
|
||||
"""
|
||||
sm_scale = 1.0 / (q.shape[-1]**0.5)
|
||||
idx = k2q_idx.to(torch.int32).contiguous()
|
||||
num = k2q_num.to(torch.int32).contiguous()
|
||||
vbs = variable_block_sizes.to(torch.int32).contiguous()
|
||||
res = _BWD(grad_o.contiguous(), q.contiguous(), k.contiguous(), v.contiguous(),
|
||||
o.contiguous(), lse.contiguous(), idx, num, vbs, sm_scale)
|
||||
return res[0], res[1], res[2]
|
||||
|
||||
|
||||
def block_sparse_attn_backward_sm100a(
|
||||
grad_o: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
lse: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Backward pass from the forward's q2k metadata. Returns ``(dq, dk, dv)``.
|
||||
|
||||
Mirrors ``block_sparse_attn_backward_triton``: the k2q inversion is recomputed here with
|
||||
FastVideo's Triton ``invert_indices`` rather than saved by the forward.
|
||||
"""
|
||||
from fastvideo_kernel.triton_kernels.index import invert_indices
|
||||
|
||||
num_kv_blocks = variable_block_sizes.numel()
|
||||
batch = q.shape[0]
|
||||
heads = q.shape[1] if BHSD else q.shape[2]
|
||||
idx = q2k_idx.to(torch.int32).contiguous()
|
||||
num = q2k_num.to(torch.int32).contiguous()
|
||||
if idx.dim() != 4:
|
||||
idx = idx.view(batch, heads, -1, idx.shape[-1])
|
||||
if num.dim() != 3:
|
||||
num = num.view(batch, heads, -1)
|
||||
k2q_idx, k2q_num = invert_indices(idx, num, num_kv_blocks)
|
||||
return block_sparse_attn_backward_sm100a_from_k2q(grad_o, q, k, v, o, lse, k2q_idx, k2q_num,
|
||||
variable_block_sizes)
|
||||
@@ -4,10 +4,12 @@
|
||||
The historical ``sm100a`` module and symbol names are retained for compatibility, but the
|
||||
extension carries native sm_100a and sm_103a images and supports both device generations.
|
||||
|
||||
A third backend behind the same VSA op as the Triton and CuTe-DSL paths. Forward only: it
|
||||
returns ``(out, lse)`` with ``lse`` in exactly the form ``triton_block_sparse_attn_forward``
|
||||
writes -- ``max(qk * qk_scale) + log2(l)``, ``[B, H, S]`` fp32 -- so
|
||||
``block_sparse_attn_backward_triton`` runs against it unchanged.
|
||||
A third backend behind the same VSA op as the Triton and CuTe-DSL paths. This module is the
|
||||
forward: it returns ``(out, lse)`` with ``lse`` in exactly the form
|
||||
``triton_block_sparse_attn_forward`` writes -- ``max(qk * qk_scale) + log2(l)``, ``[B, H, S]``
|
||||
fp32 -- so both ``block_sparse_attn_backward_triton`` and the sm_100a CUDA backward
|
||||
(``block_sparse_attn_bwd_sm100a``, 64-token blocks, sm_100a devices only) run against it
|
||||
unchanged.
|
||||
|
||||
The extension carries TWO instantiations of the kernel, for 64- and 128-token sparse blocks
|
||||
(tile volumes 64 and 128 in ``build_vsa_metadata``); the block size is inferred from the
|
||||
|
||||
@@ -0,0 +1,260 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Correctness tests for the sm_100a CUDA block-sparse VSA backward.
|
||||
|
||||
Reference: fp32 dense attention restricted to the selected blocks, keys past
|
||||
variable_block_sizes masked to -inf, differentiated with torch autograd on fp32 copies of the
|
||||
bf16 inputs with loss = (out * grad_o).sum(). The kernel is fed the reference's lse in Triton's
|
||||
M format (logsumexp * log2e) and the reference output rounded to bf16 as ``o`` -- exactly what
|
||||
the sm_100a forward hands it in FastVideo. The k2q inversion is done here in torch (a stable
|
||||
sort) so the tests do not need Triton.
|
||||
|
||||
Run with: python -m pytest tests/test_block_sparse_bwd_sm100a.py -v
|
||||
"""
|
||||
|
||||
import itertools
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel import block_sparse_attn_bwd_sm100a as bwd
|
||||
|
||||
HEAD_DIM = 128
|
||||
BLOCK = 64
|
||||
LOG2E = 1.4426950408889634
|
||||
|
||||
# Per-tensor tolerances on the bf16 outputs against the fp32 reference. Measured over all 15
|
||||
# cases on GB200 (2026-09-03, VSA_BWD_TEST_VERBOSE=1): max|diff|/max|ref| up to 3.0e-3 (dq) and
|
||||
# 5.3e-3 (dk, dv); mean|diff| up to 2.8e-4 against mean|ref| of 0.06-0.12. The bounds below leave
|
||||
# about 2x (rel max) and 3.5x (mean abs) headroom.
|
||||
REL_MAX_TOL = 1e-2 # max|got - ref| / max|ref|
|
||||
MEAN_ABS_TOL = 1e-3 # mean|got - ref|
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0)
|
||||
or not bwd._HAS_VSA_BWD_SM100A,
|
||||
reason="requires a compute-capability (10, 0) GPU (sm_100a) and a fastvideo_kernel "
|
||||
"extension built with block_sparse_sm100a_bwd",
|
||||
)
|
||||
|
||||
|
||||
def make_case(num_blocks=8, topk=4, heads=4, batch=1, ragged=False, seed=0, kv_pool=None):
|
||||
"""Random bf16 q/k/v/grad_o plus q2k metadata [B, H, Nq, topk] / [B, H, Nq] and vbs.
|
||||
|
||||
``kv_pool`` restricts the kv blocks a q block may select (default: all), which is how the
|
||||
zero-count case leaves some kv blocks unselected by every q block.
|
||||
"""
|
||||
torch.manual_seed(seed)
|
||||
S = num_blocks * BLOCK
|
||||
shape = (batch, heads, S, HEAD_DIM) if bwd.BHSD else (batch, S, heads, HEAD_DIM)
|
||||
q, k, v, grad_o = (torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(4))
|
||||
|
||||
pool = torch.arange(num_blocks) if kv_pool is None else torch.as_tensor(list(kv_pool))
|
||||
assert topk <= pool.numel()
|
||||
rows = batch * heads * num_blocks
|
||||
idx = torch.empty((rows, topk), dtype=torch.int32)
|
||||
for r in range(rows):
|
||||
idx[r] = pool[torch.randperm(pool.numel())[:topk]].sort().values.to(torch.int32)
|
||||
idx = idx.view(batch, heads, num_blocks, topk).cuda()
|
||||
num = torch.full((batch, heads, num_blocks), topk, dtype=torch.int32, device="cuda")
|
||||
|
||||
if ragged:
|
||||
vbs = torch.randint(BLOCK // 2, BLOCK + 1, (num_blocks, ), dtype=torch.int32,
|
||||
device="cuda")
|
||||
else:
|
||||
vbs = torch.full((num_blocks, ), BLOCK, dtype=torch.int32, device="cuda")
|
||||
return q, k, v, grad_o, idx, num, vbs
|
||||
|
||||
|
||||
def invert_indices_torch(q2k_idx, q2k_num, num_kv_blocks, pad_value=0):
|
||||
"""k2q from q2k without Triton: a stable sort, so each row lists its q blocks ascending.
|
||||
|
||||
Returns (k2q_idx [B, H, num_kv_blocks, Nq] int32, k2q_num [B, H, num_kv_blocks] int32), the
|
||||
layout fastvideo_kernel.triton_kernels.index.invert_indices produces. Entries past a row's
|
||||
count hold ``pad_value`` -- a VALID block id, so an over-read fails by wrong values rather
|
||||
than by luck.
|
||||
"""
|
||||
B, H, Nq, Mk = q2k_idx.shape
|
||||
device = q2k_idx.device
|
||||
valid = torch.arange(Mk, device=device).view(1, 1, 1, Mk) < q2k_num.view(B, H, Nq, 1)
|
||||
row = (torch.arange(B * H, device=device).view(B, H, 1, 1) * num_kv_blocks
|
||||
+ q2k_idx.long())[valid]
|
||||
qblock = torch.arange(Nq, device=device).view(1, 1, Nq, 1).expand(B, H, Nq, Mk)[valid]
|
||||
order = torch.sort(row * Nq + qblock).indices
|
||||
row, qblock = row[order], qblock[order]
|
||||
counts = torch.bincount(row, minlength=B * H * num_kv_blocks)
|
||||
starts = torch.cumsum(counts, 0) - counts
|
||||
slot = torch.arange(row.numel(), device=device) - starts[row]
|
||||
k2q_idx = torch.full((B * H * num_kv_blocks, Nq), pad_value, dtype=torch.int32,
|
||||
device=device)
|
||||
k2q_idx[row, slot] = qblock.to(torch.int32)
|
||||
return (k2q_idx.view(B, H, num_kv_blocks, Nq),
|
||||
counts.to(torch.int32).view(B, H, num_kv_blocks))
|
||||
|
||||
|
||||
def reference(q, k, v, grad_o, idx, num, vbs):
|
||||
"""fp32 masked-dense autograd reference.
|
||||
|
||||
Returns (o bf16, lse fp32 [B, H, S] in M format, dq, dk, dv fp32) -- o and the grads in
|
||||
q's layout.
|
||||
"""
|
||||
if not bwd.BHSD:
|
||||
q, k, v, grad_o = (t.transpose(1, 2) for t in (q, k, v, grad_o)) # -> [B, H, S, D]
|
||||
B, H, S, D = q.shape
|
||||
num_blocks = vbs.numel()
|
||||
scale = 1.0 / (D**0.5)
|
||||
|
||||
idx, num, vbs = idx.cpu(), num.cpu(), vbs.cpu()
|
||||
keep = torch.zeros((B, H, S, S), dtype=torch.bool, device=q.device)
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for qb in range(num_blocks):
|
||||
for j in range(int(num[b, h, qb])):
|
||||
kb = int(idx[b, h, qb, j])
|
||||
valid = int(vbs[kb])
|
||||
keep[b, h, qb * BLOCK:(qb + 1) * BLOCK, kb * BLOCK:kb * BLOCK + valid] = True
|
||||
|
||||
q32, k32, v32 = (t.detach().float().requires_grad_(True) for t in (q, k, v))
|
||||
scores = (q32 @ k32.transpose(-1, -2)) * scale
|
||||
scores = scores.masked_fill(~keep, float("-inf"))
|
||||
p = torch.softmax(scores, dim=-1)
|
||||
out = p @ v32
|
||||
loss = (out * grad_o.float()).sum()
|
||||
dq, dk, dv = torch.autograd.grad(loss, (q32, k32, v32))
|
||||
lse = (torch.logsumexp(scores, dim=-1) * LOG2E).detach().contiguous()
|
||||
o = out.detach().to(torch.bfloat16)
|
||||
|
||||
if not bwd.BHSD:
|
||||
o, dq, dk, dv = (t.transpose(1, 2).contiguous() for t in (o, dq, dk, dv))
|
||||
return o, lse, dq.detach(), dk.detach(), dv.detach()
|
||||
|
||||
|
||||
def check_close(name, got, ref):
|
||||
got, ref = got.float(), ref.float()
|
||||
assert got.shape == ref.shape, f"{name}: shape {tuple(got.shape)} vs {tuple(ref.shape)}"
|
||||
assert torch.isfinite(got).all(), f"{name}: non-finite values"
|
||||
diff = (got - ref).abs()
|
||||
rel_max = diff.max().item() / max(ref.abs().max().item(), 1e-6)
|
||||
mean_abs = diff.mean().item()
|
||||
if os.environ.get("VSA_BWD_TEST_VERBOSE"):
|
||||
print(f"{name}: rel_max={rel_max:.3e} mean_abs={mean_abs:.3e} "
|
||||
f"mean|ref|={ref.abs().mean().item():.3e}")
|
||||
assert rel_max <= REL_MAX_TOL and mean_abs <= MEAN_ABS_TOL, (
|
||||
f"{name}: max|diff|/max|ref| = {rel_max:.3e} (tol {REL_MAX_TOL:.1e}), "
|
||||
f"mean|diff| = {mean_abs:.3e} (tol {MEAN_ABS_TOL:.1e}), "
|
||||
f"max|ref| = {ref.abs().max().item():.3e}, mean|ref| = {ref.abs().mean().item():.3e}")
|
||||
|
||||
|
||||
def run_and_compare(num_blocks=8, topk=4, heads=4, batch=1, ragged=False, seed=0, kv_pool=None):
|
||||
q, k, v, grad_o, idx, num, vbs = make_case(num_blocks=num_blocks, topk=topk, heads=heads,
|
||||
batch=batch, ragged=ragged, seed=seed,
|
||||
kv_pool=kv_pool)
|
||||
assert bwd.is_supported(q, vbs)
|
||||
o, lse, ref_dq, ref_dk, ref_dv = reference(q, k, v, grad_o, idx, num, vbs)
|
||||
k2q_idx, k2q_num = invert_indices_torch(idx, num, num_blocks)
|
||||
dq, dk, dv = bwd.block_sparse_attn_backward_sm100a_from_k2q(grad_o, q, k, v, o, lse,
|
||||
k2q_idx, k2q_num, vbs)
|
||||
torch.cuda.synchronize()
|
||||
for name, got, ref in (("dq", dq, ref_dq), ("dk", dk, ref_dk), ("dv", dv, ref_dv)):
|
||||
assert got.dtype == torch.bfloat16, f"{name}: dtype {got.dtype}"
|
||||
check_close(name, got, ref)
|
||||
return (dq, dk, dv), (ref_dq, ref_dk, ref_dv), vbs
|
||||
|
||||
|
||||
def _to_bhsd(t):
|
||||
return t if bwd.BHSD else t.transpose(1, 2)
|
||||
|
||||
|
||||
def test_backward_matches_reference():
|
||||
run_and_compare()
|
||||
|
||||
|
||||
def test_batch_two():
|
||||
run_and_compare(batch=2)
|
||||
|
||||
|
||||
def test_ragged_block_sizes():
|
||||
"""variable_block_sizes is what FastVideo always passes; padded keys must be masked."""
|
||||
run_and_compare(ragged=True)
|
||||
|
||||
|
||||
def test_ragged_padded_key_rows_are_zero():
|
||||
"""Keys at or past a block's count get P^T = 0 in-kernel, so their dk/dv rows are exactly 0
|
||||
(Triton's backward stores zeros there as well)."""
|
||||
(dq, dk, dv), _, vbs = run_and_compare(ragged=True, seed=1)
|
||||
dk, dv = _to_bhsd(dk).float(), _to_bhsd(dv).float()
|
||||
checked = 0
|
||||
for kb in range(vbs.numel()):
|
||||
if int(vbs[kb]) == BLOCK:
|
||||
continue # a full block has no padded rows
|
||||
rows = slice(kb * BLOCK + int(vbs[kb]), (kb + 1) * BLOCK)
|
||||
assert dk[:, :, rows].abs().max().item() == 0.0, f"dk: padded rows of kv block {kb}"
|
||||
assert dv[:, :, rows].abs().max().item() == 0.0, f"dv: padded rows of kv block {kb}"
|
||||
checked += 1
|
||||
assert checked > 0, "the ragged draw produced no padded kv block; change the seed"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("topk", [1, 2, 3, 5, 7])
|
||||
def test_topk_not_a_multiple_of_the_quad(topk):
|
||||
"""The kernel walks each kv block's q list in quads of 4; a ragged tail quad must be exact."""
|
||||
run_and_compare(ragged=True, num_blocks=8, topk=topk)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("num_blocks", [4, 8, 16])
|
||||
def test_sequence_lengths(num_blocks):
|
||||
run_and_compare(ragged=True, num_blocks=num_blocks, topk=3)
|
||||
|
||||
|
||||
def test_zero_count_kv_blocks():
|
||||
"""kv blocks no q block selects: the main kernel skips them and the preprocess must write
|
||||
their dk/dv rows as exact zeros (the outputs are empty_like), every other block exact."""
|
||||
num_blocks = 8
|
||||
excluded = (2, 5)
|
||||
pool = [kb for kb in range(num_blocks) if kb not in excluded]
|
||||
(dq, dk, dv), _, _ = run_and_compare(num_blocks=num_blocks, topk=4, ragged=True,
|
||||
kv_pool=pool)
|
||||
dk, dv = _to_bhsd(dk).float(), _to_bhsd(dv).float()
|
||||
for kb in excluded:
|
||||
rows = slice(kb * BLOCK, (kb + 1) * BLOCK)
|
||||
assert dk[:, :, rows].abs().max().item() == 0.0, f"dk: unselected kv block {kb}"
|
||||
assert dv[:, :, rows].abs().max().item() == 0.0, f"dv: unselected kv block {kb}"
|
||||
|
||||
|
||||
def test_invert_indices_torch_matches_q2k():
|
||||
"""Guards the test's own k2q builder: every (q, kv) pair appears once, counts add up."""
|
||||
num_blocks = 8
|
||||
_, _, _, _, idx, num, _ = make_case(num_blocks=num_blocks, topk=3, heads=2, batch=2)
|
||||
k2q_idx, k2q_num = invert_indices_torch(idx, num, num_blocks)
|
||||
B, H, Nq, Mk = idx.shape
|
||||
assert k2q_idx.shape == (B, H, num_blocks, Nq) and k2q_num.shape == (B, H, num_blocks)
|
||||
assert int(k2q_num.sum()) == B * H * Nq * Mk
|
||||
for b, h, qb, j in itertools.product(range(B), range(H), range(Nq), range(Mk)):
|
||||
kb = int(idx[b, h, qb, j])
|
||||
listed = k2q_idx[b, h, kb, :int(k2q_num[b, h, kb])].tolist()
|
||||
assert listed.count(qb) == 1
|
||||
assert listed == sorted(listed)
|
||||
|
||||
|
||||
def test_unsupported_is_rejected():
|
||||
q, k, v, grad_o, idx, num, vbs = make_case()
|
||||
assert not bwd.is_supported(q.float(), vbs) # wrong dtype
|
||||
assert not bwd.is_supported(q[..., :64].contiguous(), vbs) # wrong head_dim
|
||||
seven = torch.full((7, ), 64, dtype=torch.int32, device="cuda")
|
||||
assert not bwd.is_supported(q, seven) # seqlen != 64 * num_blocks
|
||||
q500 = (q[:, :, :500] if bwd.BHSD else q[:, :500]).contiguous()
|
||||
assert not bwd.is_supported(q500, vbs) # S not a multiple of 64
|
||||
# An odd block count (S % 128 != 0) is refused statically too, so FastVideo falls back to
|
||||
# Triton instead of tripping the binding's check.
|
||||
q448 = (q[:, :, :448] if bwd.BHSD else q[:, :448]).contiguous()
|
||||
assert not bwd.is_supported(q448, seven)
|
||||
|
||||
# The binding itself refuses bad dtypes before touching the GPU.
|
||||
k2q_idx, k2q_num = invert_indices_torch(idx, num, vbs.numel())
|
||||
o = torch.zeros_like(q)
|
||||
heads = q.shape[1] if bwd.BHSD else q.shape[2]
|
||||
lse = torch.zeros((q.shape[0], heads, vbs.numel() * BLOCK), dtype=torch.float32,
|
||||
device="cuda")
|
||||
with pytest.raises(RuntimeError):
|
||||
bwd.block_sparse_attn_backward_sm100a_from_k2q(grad_o.float(), q.float(), k.float(),
|
||||
v.float(), o.float(), lse, k2q_idx,
|
||||
k2q_num, vbs)
|
||||
@@ -2,14 +2,18 @@
|
||||
"""Routing tests for the opt-in sm_100a/sm_103a dispatch in ``block_sparse_attn_from_indices``.
|
||||
|
||||
``FASTVIDEO_VSA_SM100A=1`` routes the forward to the sm_100a extension when
|
||||
``block_sparse_attn_sm100a.is_supported`` passes, pairing it with the Triton
|
||||
backward (the sm_100a lse is already in Triton's M format). Everything else --
|
||||
env unset, unsupported input, ``FASTVIDEO_VSA_TRITON`` override -- must keep
|
||||
the pre-existing selection, bit-for-bit.
|
||||
``block_sparse_attn_sm100a.is_supported`` passes. Its backward is the sm_100a
|
||||
CUDA backward where that op is built and ``block_sparse_attn_bwd_sm100a.is_supported``
|
||||
passes (64-token blocks on an sm_100a device), the Triton backward otherwise; the
|
||||
sm_100a lse is already in Triton's M format, so either pairing needs no conversion.
|
||||
Everything else -- env unset, unsupported input, ``FASTVIDEO_VSA_TRITON`` override --
|
||||
must keep the pre-existing selection, bit-for-bit.
|
||||
|
||||
Run with: python -m pytest tests/test_block_sparse_sm100a_dispatch.py -v
|
||||
"""
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
@@ -107,27 +111,125 @@ def test_force_triton_overrides_sm100a(monkeypatch):
|
||||
assert not torch.equal(got[0], sm100a_o)
|
||||
|
||||
|
||||
def test_backward_runs_triton_and_matches(monkeypatch):
|
||||
"""sm_100a forward + Triton backward: grads match the all-Triton path."""
|
||||
monkeypatch.setenv(ENV, "1")
|
||||
q, k, v, idx, num, vbs = make_case(64, requires_grad=True)
|
||||
out, _ = block_sparse_attn_from_indices(q, k, v, idx, num, vbs)
|
||||
out.float().square().sum().backward()
|
||||
got = [t.grad.float().clone() for t in (q, k, v)]
|
||||
def _grads_sm100a_route_vs_triton(monkeypatch, vbs=None):
|
||||
"""Grads of (out**2).sum() through the sm_100a route vs the all-Triton route.
|
||||
|
||||
On a device where the sm_100a backward is built and supported, the route must not enter
|
||||
the Triton backward at all (it is monkeypatched to fail); elsewhere the route pairs the
|
||||
sm_100a forward with the Triton backward, as before.
|
||||
"""
|
||||
# The package exports a FUNCTION named block_sparse_attn; the module needs importlib.
|
||||
dispatch = importlib.import_module("fastvideo_kernel.block_sparse_attn")
|
||||
from fastvideo_kernel import block_sparse_attn_bwd_sm100a as vsa_bwd
|
||||
|
||||
q, k, v, idx, num, default_vbs = make_case(64, requires_grad=True)
|
||||
vbs = default_vbs if vbs is None else vbs
|
||||
|
||||
q2, k2, v2 = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
|
||||
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
|
||||
out2, _ = block_sparse_attn_from_indices(q2, k2, v2, idx, num, vbs)
|
||||
out2.float().square().sum().backward()
|
||||
ref = [t.grad.float() for t in (q2, k2, v2)]
|
||||
monkeypatch.delenv("FASTVIDEO_VSA_TRITON")
|
||||
|
||||
monkeypatch.setenv(ENV, "1")
|
||||
sm100a_backward = vsa_bwd.is_supported(q, vbs)
|
||||
if sm100a_backward:
|
||||
|
||||
def no_triton_backward(*args, **kwargs):
|
||||
raise AssertionError("Triton backward entered on a supported sm_100a input")
|
||||
|
||||
monkeypatch.setattr(dispatch, "block_sparse_attn_backward_triton", no_triton_backward)
|
||||
out, _ = block_sparse_attn_from_indices(q, k, v, idx, num, vbs)
|
||||
out.float().square().sum().backward()
|
||||
got = [t.grad.float() for t in (q, k, v)]
|
||||
return got, ref, sm100a_backward
|
||||
|
||||
|
||||
def _assert_grads_close(got, ref):
|
||||
# Two bf16 kernels against each other (not an fp32 reference): twice the tolerances the
|
||||
# sm_100a backward holds against fp32 in tests/test_block_sparse_bwd_sm100a.py.
|
||||
for g, r, name in zip(got, ref, "qkv"):
|
||||
assert torch.allclose(g, r, atol=5e-2, rtol=5e-2), \
|
||||
f"d{name} max|diff|={(g - r).abs().max().item()}"
|
||||
diff = (g - r).abs()
|
||||
rel_max = diff.max().item() / max(r.abs().max().item(), 1e-6)
|
||||
mean_abs = diff.mean().item()
|
||||
assert rel_max <= 2e-2 and mean_abs <= 2e-3, \
|
||||
f"d{name}: max|diff|/max|ref|={rel_max:.3e} mean|diff|={mean_abs:.3e}"
|
||||
|
||||
|
||||
def test_backward_matches_triton(monkeypatch):
|
||||
"""sm_100a route (CUDA backward where built, Triton backward otherwise) vs all-Triton."""
|
||||
got, ref, _ = _grads_sm100a_route_vs_triton(monkeypatch)
|
||||
_assert_grads_close(got, ref)
|
||||
|
||||
|
||||
def test_backward_ragged_block_sizes_matches_triton(monkeypatch):
|
||||
"""variable_block_sizes below 64 (padded kv rows) through the same two routes."""
|
||||
torch.manual_seed(1)
|
||||
vbs = torch.randint(32, 65, (8, ), dtype=torch.int32, device="cuda")
|
||||
got, ref, _ = _grads_sm100a_route_vs_triton(monkeypatch, vbs=vbs)
|
||||
_assert_grads_close(got, ref)
|
||||
|
||||
|
||||
def test_backward_uses_sm100a_kernel_when_built(monkeypatch):
|
||||
"""Guards against a silent Triton fallback: with the op built on an sm_100a device the
|
||||
CUDA backward must be the one that runs."""
|
||||
from fastvideo_kernel import block_sparse_attn_bwd_sm100a as vsa_bwd
|
||||
if not vsa_bwd._HAS_VSA_BWD_SM100A or torch.cuda.get_device_capability() != (10, 0):
|
||||
pytest.skip("sm_100a backward not built for this device")
|
||||
_, _, sm100a_backward = _grads_sm100a_route_vs_triton(monkeypatch)
|
||||
assert sm100a_backward
|
||||
|
||||
|
||||
def test_backward_large_seq_matches_triton_and_is_deterministic(monkeypatch):
|
||||
"""65536 tokens (1024 kv blocks, 12.5% density): the device-computed work order (nb >= 1024)
|
||||
and the sequence regime where a TMEM ordering bug corrupted dk/dv while every <= 16-block
|
||||
test stayed green. sm_100a route vs all-Triton, and two sm_100a runs must agree to within
|
||||
summation-order noise: invert_indices compacts each kv row's q list with atomics, so the
|
||||
kernel sums quads in a different order per call (measured run-to-run delta on GB200:
|
||||
rel_max 5e-3, mean|diff| 7e-6). The TMEM race this guards against gave rel_max 0.6-1.0
|
||||
and mean|diff| 4e-2, an order of magnitude past the bounds below on both metrics."""
|
||||
from fastvideo_kernel import block_sparse_attn_bwd_sm100a as vsa_bwd
|
||||
|
||||
torch.manual_seed(0)
|
||||
block, num_blocks, heads, topk = 64, 1024, 4, 128
|
||||
S = num_blocks * block
|
||||
shape = (1, heads, S, HEAD_DIM) if vsa.BHSD else (1, S, heads, HEAD_DIM)
|
||||
q, k, v = (torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3))
|
||||
scores = torch.rand(1, heads, num_blocks, num_blocks, device="cuda")
|
||||
idx = scores.topk(topk, dim=-1).indices.sort(dim=-1).values.to(torch.int32)
|
||||
del scores
|
||||
num = torch.full((1, heads, num_blocks), topk, dtype=torch.int32, device="cuda")
|
||||
vbs = torch.full((num_blocks, ), block, dtype=torch.int32, device="cuda")
|
||||
if not vsa_bwd.is_supported(q, vbs):
|
||||
pytest.skip("sm_100a backward not built for this device")
|
||||
|
||||
def grads(route):
|
||||
qq, kk, vv = (t.clone().requires_grad_(True) for t in (q, k, v))
|
||||
if route == "triton":
|
||||
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
|
||||
else:
|
||||
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
|
||||
monkeypatch.setenv(ENV, "1")
|
||||
out, _ = block_sparse_attn_from_indices(qq, kk, vv, idx, num, vbs)
|
||||
out.float().square().sum().backward()
|
||||
return [t.grad.float() for t in (qq, kk, vv)]
|
||||
|
||||
ref = grads("triton")
|
||||
got = grads("sm100a")
|
||||
again = grads("sm100a")
|
||||
_assert_grads_close(got, ref)
|
||||
for g1, g2, name in zip(got, again, "qkv"):
|
||||
diff = (g1 - g2).abs()
|
||||
rel_max = diff.max().item() / max(g1.abs().max().item(), 1e-6)
|
||||
mean_abs = diff.mean().item()
|
||||
assert rel_max <= 2e-2 and mean_abs <= 1e-4, \
|
||||
f"d{name}: two sm_100a runs differ beyond summation-order noise: " \
|
||||
f"max|diff|/max|ref|={rel_max:.3e} mean|diff|={mean_abs:.3e}"
|
||||
|
||||
|
||||
def test_blk128_backward_raises(monkeypatch):
|
||||
"""128-token blocks: forward runs, backward refuses (Triton bwd is 64-block only)."""
|
||||
"""128-token blocks: forward runs, backward refuses (both backwards are 64-block only)."""
|
||||
monkeypatch.setenv(ENV, "1")
|
||||
q, k, v, idx, num, vbs = make_case(128, requires_grad=True)
|
||||
out, _ = block_sparse_attn_from_indices(q, k, v, idx, num, vbs)
|
||||
|
||||
Reference in New Issue
Block a user