#pragma once #include #include #include #include // Pure POD/traits header — no .cuh/CUDA-kernel includes; raw __nv_* type // spellings only. namespace astrai { namespace fp8 { // Compile-time FP8 format: E4M3 (forward / high precision, max 448) or // E5M2 (gradient / large dynamic range, max 57344). enum class FP8Format : int { E4M3 = 0, E5M2 = 1, }; // Operand memory layouts as types (CUTLASS-style tags). The tag names the // storage order of the raw buffer relative to the operand's canonical GEMM // matrix — A is [M][K], B is [K][N]: // A RowMajor = [M][K] storage (K-contiguous rows; the default) // A ColMajor = [K][M] storage (M-contiguous; A^T) // B RowMajor = [K][N] storage (N-contiguous; the plain a @ b operand) // B ColMajor = [N][K] storage (K-contiguous; the nn.Linear weight layout) // Empty tags: selection happens by type at compile time (see load_operand_tile). struct RowMajor {}; struct ColMajor {}; // Transpose of a layout tag: the same buffer with the rows and contract dims // swapped. B's tag is relative to the canonical [K][N] GEMM matrix, so the // stage-load (which views any operand as [rows][contract]) sees the transposed // tag — this trait makes that inversion explicit. template struct transpose_layout; template <> struct transpose_layout { using type = ColMajor; }; template <> struct transpose_layout { using type = RowMajor; }; template using transpose_layout_t = typename transpose_layout::type; // Compile-time tile configuration, mirroring KernelTraits in the attention kernels. `Fmt` selects the FP8 conversion // and the MMA PTX mnemonic; the remaining parameters shape the CTA tile and // the cp.async pipeline depth. template struct Fp8GemmTraits { static constexpr FP8Format kFormat = Fmt; static constexpr int kBlockM = BlockM; static constexpr int kBlockN = BlockN; static constexpr int kK = K; static constexpr int kStages = Stages; static constexpr bool kIsE5M2 = (Fmt == FP8Format::E5M2); static constexpr __nv_fp8_interpretation_t kNvFormat = kIsE5M2 ? __NV_E5M2 : __NV_E4M3; static constexpr float kFp8Max = kIsE5M2 ? 57344.0f : 448.0f; }; // Unified GEMM parameter POD, mirroring AttentionParams: one struct flows // through quantize / fused / pre-quantized kernels. Each kernel touches only // the fields it needs; buffers are raw pointers packed by the torch binding. // Pointer members default to null (same NSDMI rationale as AttentionParams: // bias / amax / out_scale gate optional paths via null checks, so a partially // packed struct must never hold garbage non-null pointers). Still an // aggregate, still trivially copyable. struct FP8Params { // Inputs: a/b are BF16 for the fused (quantize-in-GEMM) path, FP8 for // the pre-quantized path. Scales are quantization steps (device scalars). const void* __restrict__ a_ptr = nullptr; const void* __restrict__ b_ptr = nullptr; const void* __restrict__ bias = nullptr; const float* __restrict__ scale_a = nullptr; const float* __restrict__ scale_b = nullptr; const float* __restrict__ bias_scale = nullptr; // Output: BF16 or FP8 (E4M3). out_scale is the output quantization step // (FP8 output only). void* __restrict__ out_ptr = nullptr; const float* __restrict__ out_scale = nullptr; // Fused forward extras: bias (may be null) and amax slots (may be null). float* __restrict__ amax_a = nullptr; float* __restrict__ amax_b = nullptr; // Shapes. total is only used by the elementwise quantize kernel. `int` // covers every realistic LLM shape; the kernels promote to int64 for all // pointer arithmetic. int m, n, k; // Physical leading dimensions (column count, i.e. row stride) of A and B. // For a non-transposed operand the stride equals the contract dim; for a // transposed operand it is the operand's own column count. The binding // packs these so the kernel reads both buffers either naturally or // transposed depending on the LayoutA/LayoutB tags (see gemm.cuh). int a_ld, b_ld; int total; }; } // namespace fp8 } // namespace astrai