import statistics import time import warnings import torch import triton import triton.language as tl BLOCK_TOKENS: int = 512 BLOCK_NONZERO: int = 16 BLOCK_VALUES: int = 32 BLOCK_REDUCE: int = 256 IN_FEATURES: int = 3072 OUT_FEATURES: int = 3072 TOKENS: int = 1024 DENSITY: float = 0.005 DTYPE: torch.dtype = torch.bfloat16 WARMUP: int = 5 STEPS: int = 30 SEED: int = 20260723 @triton.jit def csr_spmm_kernel( x_ptr, out_ptr, crow_ptr, column_ptr, value_ptr, tokens, stride_xi, stride_xt, stride_ot, stride_oo, BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, ): row = tl.program_id(0) token_index = tl.program_id(1) * BLOCK_T + tl.arange(0, BLOCK_T) token_mask = token_index < tokens token_offset = token_index.to(tl.int64) start = tl.load(crow_ptr + row).to(tl.int32) end = tl.load(crow_ptr + row + 1).to(tl.int32) accumulator = tl.zeros((BLOCK_T,), dtype=tl.float32) for block in range(start, end, BLOCK_K): nonzero_index = block + tl.arange(0, BLOCK_K) nonzero_mask = nonzero_index < end columns = tl.load( column_ptr + nonzero_index, mask=nonzero_mask, other=0, ).to(tl.int64) values = tl.load( value_ptr + nonzero_index, mask=nonzero_mask, other=0.0, ) activations = tl.load( x_ptr + columns[:, None] * stride_xi + token_offset[None, :] * stride_xt, mask=nonzero_mask[:, None] & token_mask[None, :], other=0.0, ).to(tl.float32) accumulator += tl.sum(activations * values[:, None], axis=0) tl.store( out_ptr + token_offset * stride_ot + row.to(tl.int64) * stride_oo, accumulator.to(out_ptr.dtype.element_ty), mask=token_mask, ) @triton.jit def csc_spmm_kernel( grad_ptr, out_ptr, ccol_ptr, row_ptr, permutation_ptr, value_ptr, tokens, stride_go, stride_gt, stride_ot, stride_oi, BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, ): column = tl.program_id(0) token_index = tl.program_id(1) * BLOCK_T + tl.arange(0, BLOCK_T) token_mask = token_index < tokens token_offset = token_index.to(tl.int64) start = tl.load(ccol_ptr + column).to(tl.int32) end = tl.load(ccol_ptr + column + 1).to(tl.int32) accumulator = tl.zeros((BLOCK_T,), dtype=tl.float32) for block in range(start, end, BLOCK_K): nonzero_index = block + tl.arange(0, BLOCK_K) nonzero_mask = nonzero_index < end rows = tl.load( row_ptr + nonzero_index, mask=nonzero_mask, other=0, ).to(tl.int64) slots = tl.load( permutation_ptr + nonzero_index, mask=nonzero_mask, other=0, ) values = tl.load( value_ptr + slots, mask=nonzero_mask, other=0.0, ) gradients = tl.load( grad_ptr + rows[:, None] * stride_go + token_offset[None, :] * stride_gt, mask=nonzero_mask[:, None] & token_mask[None, :], other=0.0, ).to(tl.float32) accumulator += tl.sum(gradients * values[:, None], axis=0) tl.store( out_ptr + token_offset * stride_ot + column.to(tl.int64) * stride_oi, accumulator.to(out_ptr.dtype.element_ty), mask=token_mask, ) @triton.jit def csr_sddmm_kernel( grad_ptr, x_ptr, row_ptr, column_ptr, out_ptr, nonzeros, tokens, stride_go, stride_gt, stride_xi, stride_xt, BLOCK_N: tl.constexpr, BLOCK_T: tl.constexpr, ): nonzero_index = tl.program_id(0) * BLOCK_N + tl.arange(0, BLOCK_N) nonzero_mask = nonzero_index < nonzeros rows = tl.load(row_ptr + nonzero_index, mask=nonzero_mask, other=0).to(tl.int64) columns = tl.load(column_ptr + nonzero_index, mask=nonzero_mask, other=0).to(tl.int64) accumulator = tl.zeros((BLOCK_N,), dtype=tl.float32) for block in range(0, tokens, BLOCK_T): token_index = block + tl.arange(0, BLOCK_T) token_mask = token_index < tokens token_offset = token_index.to(tl.int64) active = nonzero_mask[:, None] & token_mask[None, :] gradients = tl.load( grad_ptr + rows[:, None] * stride_go + token_offset[None, :] * stride_gt, mask=active, other=0.0, ).to(tl.float32) activations = tl.load( x_ptr + columns[:, None] * stride_xi + token_offset[None, :] * stride_xt, mask=active, other=0.0, ).to(tl.float32) accumulator += tl.sum(gradients * activations, axis=1) tl.store(out_ptr + nonzero_index, accumulator, mask=nonzero_mask) def build_sparse_layout( out_features: int, in_features: int, density: float, seed: int, device: torch.device, ): total = out_features * in_features nonzeros = min(total, max(1, round(total * density))) generator = torch.Generator(device="cpu").manual_seed(seed) chosen = torch.randperm(total, generator=generator, dtype=torch.int64)[:nonzeros] chosen = chosen.sort().values rows = torch.div(chosen, in_features, rounding_mode="floor") columns = chosen - rows * in_features crow = torch.zeros(out_features + 1, dtype=torch.int64) crow[1:] = torch.bincount(rows, minlength=out_features).cumsum(0) order = torch.argsort(columns, stable=True) transposed_crow = torch.zeros(in_features + 1, dtype=torch.int64) transposed_crow[1:] = torch.bincount(columns, minlength=in_features).cumsum(0) return ( crow.to(device), columns.to(device), rows.to(device), transposed_crow.to(device), rows.index_select(0, order).to(device), order.to(device), torch.zeros(nonzeros, dtype=torch.float32, device=device), ) def torch_forward(x, values, layout): crow, columns, _, _, _, _, _ = layout sparse = torch.sparse_csr_tensor( crow, columns, values, size=(OUT_FEATURES, IN_FEATURES), check_invariants=False ) return torch.sparse.mm(sparse, x.float().t()).t().to(DTYPE) def torch_backward(x, values, grad, layout): crow, columns, _, transposed_crow, transposed_columns, permutation, zeros = layout grad = grad.float() transposed = torch.sparse_csr_tensor( transposed_crow, transposed_columns, values.index_select(0, permutation), size=(IN_FEATURES, OUT_FEATURES), check_invariants=False, ) pattern = torch.sparse_csr_tensor( crow, columns, zeros, size=(OUT_FEATURES, IN_FEATURES), check_invariants=False ) grad_input = torch.sparse.mm(transposed, grad.t()).t().to(DTYPE) grad_values = torch.sparse.sampled_addmm(pattern, grad.t(), x.float(), beta=0.0).values() return grad_input, grad_values def triton_forward(x, values, layout): crow, columns, _, _, _, _, _ = layout transposed_x = x.t().contiguous() tokens = int(transposed_x.shape[1]) result = torch.zeros((tokens, OUT_FEATURES), device=x.device, dtype=DTYPE) grid = (OUT_FEATURES, triton.cdiv(tokens, BLOCK_TOKENS)) csr_spmm_kernel[grid]( transposed_x, result, crow, columns, values, tokens, transposed_x.stride(0), transposed_x.stride(1), result.stride(0), result.stride(1), BLOCK_T=BLOCK_TOKENS, BLOCK_K=BLOCK_NONZERO, ) return result def triton_backward(x, values, grad, layout): _, columns, rows, transposed_crow, transposed_columns, permutation, _ = layout transposed_grad = grad.t().contiguous() transposed_x = x.t().contiguous() tokens = int(transposed_grad.shape[1]) grad_input = torch.zeros((tokens, IN_FEATURES), device=x.device, dtype=DTYPE) grid = (IN_FEATURES, triton.cdiv(tokens, BLOCK_TOKENS)) csc_spmm_kernel[grid]( transposed_grad, grad_input, transposed_crow, transposed_columns, permutation, values, tokens, transposed_grad.stride(0), transposed_grad.stride(1), grad_input.stride(0), grad_input.stride(1), BLOCK_T=BLOCK_TOKENS, BLOCK_K=BLOCK_NONZERO, ) nonzeros = int(columns.numel()) grad_values = torch.zeros(nonzeros, device=x.device, dtype=torch.float32) grid = (triton.cdiv(nonzeros, BLOCK_VALUES),) csr_sddmm_kernel[grid]( transposed_grad, transposed_x, rows, columns, grad_values, nonzeros, tokens, transposed_grad.stride(0), transposed_grad.stride(1), transposed_x.stride(0), transposed_x.stride(1), BLOCK_N=BLOCK_VALUES, BLOCK_T=BLOCK_REDUCE, ) return grad_input, grad_values def benchmark_backend( forward, backward, x, values, grad, layout, ): forward_times = [] backward_times = [] for step in range(WARMUP + STEPS): torch.cuda.synchronize() begin = time.perf_counter() forward(x, values, layout) torch.cuda.synchronize() middle = time.perf_counter() backward(x, values, grad, layout) torch.cuda.synchronize() end = time.perf_counter() if step >= WARMUP: forward_times.append((middle - begin) * 1000.0) backward_times.append((end - middle) * 1000.0) return statistics.median(forward_times), statistics.median(backward_times) def main(): warnings.filterwarnings("ignore", message="Sparse .*", category=UserWarning) torch.sparse.check_sparse_tensor_invariants.disable() device = torch.device("cuda") layout = build_sparse_layout(OUT_FEATURES, IN_FEATURES, DENSITY, SEED, device) torch.manual_seed(SEED + 100) x = torch.randn(TOKENS, IN_FEATURES, device=device, dtype=DTYPE) values = torch.randn(layout[1].numel(), device=device, dtype=torch.float32) grad = torch.randn(TOKENS, OUT_FEATURES, device=device, dtype=DTYPE) torch_forward_ms, torch_backward_ms = benchmark_backend( torch_forward, torch_backward, x, values, grad, layout ) triton_forward_ms, triton_backward_ms = benchmark_backend( triton_forward, triton_backward, x, values, grad, layout ) torch_total = torch_forward_ms + torch_backward_ms triton_total = triton_forward_ms + triton_backward_ms print( f"torch: forward={torch_forward_ms:.3f} ms backward={torch_backward_ms:.3f} ms " f"total={torch_total:.3f} ms" ) print( f"triton: forward={triton_forward_ms:.3f} ms backward={triton_backward_ms:.3f} ms " f"total={triton_total:.3f} ms speedup={torch_total / triton_total:.2f}x" ) if __name__ == "__main__": main()