File size: 1,521 Bytes
ef6eb55 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 | -- PAX WMMA — mma.sync.aligned.m16n8k8 FP16→FP32 semantics
-- Ahmad Ali Parr · PAX Architecture · sm_86
-- Proof obligation PO3: SIMT divergence reconvergence
namespace PAX.WMMA
/-- Abstract matrix tile: m×n×k WMMA fragment -/
structure WMMAFragment (m n k : ℕ) (α β : Type*) where
aFrag : Fin m → Fin k → α -- A matrix (FP16)
bFrag : Fin k → Fin n → α -- B matrix (FP16)
cFrag : Fin m → Fin n → β -- accumulator (FP32)
/-- Functional GEMM spec: C += A × B -/
def gemmSpec [Add β] [Mul α] [HMul α α β] [Zero β]
{m n k : ℕ} (frag : WMMAFragment m n k α β) : Fin m → Fin n → β :=
fun i j =>
frag.cFrag i j +
Finset.univ.sum (fun (l : Fin k) => frag.aFrag i l * frag.bFrag l j)
/-- mma.sync abstract model — 32-thread warp computes 16×8 tile -/
structure MMASyncResult (m n : ℕ) (β : Type*) where
result : Fin m → Fin n → β
/-- PO3: mma.sync result equals functional spec -/
axiom mma_sync_correct [Add β] [Mul Float Float] [HMul Float Float β] [Zero β]
{m n k : ℕ} (frag : WMMAFragment m n k Float β) :
∀ i j, (mmaSync frag).result i j = gemmSpec frag i j
/-- warp_gemm: issue mma.sync, accumulate 8 tiles per warp -/
def warpGEMM [Add β] [Mul Float Float] [HMul Float Float β] [Zero β]
{tiles : ℕ} (frags : Fin tiles → WMMAFragment 16 8 8 Float β) :
Fin 16 → Fin 8 → β :=
fun i j =>
Finset.univ.sum (fun t => (mmaSync (frags t)).result i j)
end PAX.WMMA
|