| -- 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 | |