Large source region · 3,184 lines
module Realization.Nvidia.SM86.HMMAProductionSM86
import Accelerator.SM86.Control
import Accelerator.SM86.Immediate
import Accelerator.SM86.Instruction
import Accelerator.SM86.InstructionEncoding
import Accelerator.SM86.Program
import Data.SHA256Digest
import Std.Natural
import Accelerator.SM86.Types
-- This is deliberately a producer only. NativeProgramPlan remains the sole
-- launch planner until this image family has independent runtime evidence.
family HMMAProductionSM86Orientation : Type 0
constructor HMMAProductionAB
constructor HMMAProductionABTransposed
constructor HMMAProductionATransposedB
end-family
family HMMAProductionSM86Grouping : Type 0
constructor HMMAProductionDense
constructor HMMAProductionGrouped
field unrestricted hmmaProductionGroupCount : Nat
end-family
family HMMAProductionSM86Variant : Type 0
constructor HMMAProductionDenseForward
constructor HMMAProductionDenseReverseInput
constructor HMMAProductionDenseReverseWeight
constructor HMMAProductionRoutedForward
constructor HMMAProductionRoutedReverseInput
constructor HMMAProductionRoutedReverseWeight
constructor HMMAProductionAttentionQKForward
constructor HMMAProductionAttentionProbabilityValueForward
constructor HMMAProductionAttentionQKReverseLeft
constructor HMMAProductionAttentionQKReverseRight
constructor HMMAProductionAttentionValueReverseLeft
constructor HMMAProductionAttentionValueReverseRight
end-family
family HMMAProductionSM86Geometry : Type 0
constructor HMMAProductionSM86GeometryValue
field unrestricted hmmaProductionM : Nat
field unrestricted hmmaProductionN : Nat
field unrestricted hmmaProductionK : Nat
field unrestricted hmmaProductionKStages : Nat
end-family
family HMMAProductionSM86ABI : Type 0
constructor HMMAProductionSM86ABIValue
field unrestricted hmmaProductionOutputOffset : Nat
field unrestricted hmmaProductionAOffset : Nat
field unrestricted hmmaProductionBOffset : Nat
end-family
-- Distances between stored rows, in elements. Transposition changes the
-- contraction, not which dimension this stride describes. A head may be a
-- column view of token-major storage without copying it into a new matrix.
family MatrixProductRowStrides : Type 0
constructor MatrixProductRowStridesValue
field unrestricted matrixProductLeftRowStride : Nat
field unrestricted matrixProductRightRowStride : Nat
field unrestricted matrixProductOutputRowStride : Nat
end-family
family HMMAProductionSM86Extents : Type 0
constructor HMMAProductionSM86ExtentsValue
field unrestricted hmmaProductionExtentM : Nat
field unrestricted hmmaProductionExtentN : Nat
field unrestricted hmmaProductionExtentK : Nat
field unrestricted hmmaProductionExtentGroups : Nat
end-family
family HMMAProductionSM86Manifest : Type 0
constructor HMMAProductionSM86ManifestValue
field unrestricted hmmaProductionManifestVariant : (family HMMAProductionSM86Variant)
field unrestricted hmmaProductionManifestOrientation : (family HMMAProductionSM86Orientation)
field unrestricted hmmaProductionManifestGrouping : (family HMMAProductionSM86Grouping)
field unrestricted hmmaProductionManifestGeometry : (family HMMAProductionSM86Geometry)
field unrestricted hmmaProductionManifestExpectedInstructions : Nat
field unrestricted hmmaProductionManifestExpectedBytes : Nat
field unrestricted hmmaProductionManifestRegisters : Nat
field unrestricted hmmaProductionManifestSharedBytes : Nat
field unrestricted hmmaProductionManifestGridX : Nat
field unrestricted hmmaProductionManifestGridY : Nat
field unrestricted hmmaProductionManifestGridZ : Nat
field unrestricted hmmaProductionManifestBlockX : Nat
field unrestricted hmmaProductionManifestABI : (family HMMAProductionSM86ABI)
field unrestricted hmmaProductionManifestExtents : (family HMMAProductionSM86Extents)
field unrestricted hmmaProductionManifestHostFallbackOperations : Nat
end-family
family HMMAProductionSM86FailureCode : Type 0
constructor HMMAProductionInvalidGeometry
constructor HMMAProductionInvalidGrouping
constructor HMMAProductionInstructionCountMismatch
constructor HMMAProductionEncodedByteCountMismatch
constructor HMMAProductionEncodingFailed
constructor HMMAProductionIdentityFailed
constructor HMMAProductionIdentityLengthInvalid
end-family
family HMMAProductionSM86Telemetry : Type 0
constructor HMMAProductionSM86TelemetryValue
field unrestricted hmmaProductionTelemetryManifest : (family HMMAProductionSM86Manifest)
field unrestricted hmmaProductionTelemetryObservedInstructions : Nat
field unrestricted hmmaProductionTelemetryObservedBytes : Nat
field unrestricted hmmaProductionTelemetryKStages : Nat
field unrestricted hmmaProductionTelemetryHostFallbackOperations : Nat
end-family
family HMMAProductionSM86BuildResult : Type 0
constructor HMMAProductionSM86BuildSucceeded
field unrestricted hmmaProductionEncodedBytes : Bytes
field unrestricted hmmaProductionImageSHA256 : Bytes
field unrestricted hmmaProductionEncodingTelemetry : (family SM86ProgramEncodingTelemetry)
field unrestricted hmmaProductionIdentityTelemetry : (family SHA256DigestTelemetry)
field unrestricted hmmaProductionBuildTelemetry : (family HMMAProductionSM86Telemetry)
constructor HMMAProductionSM86ContractFailed
field unrestricted hmmaProductionContractFailure : (family HMMAProductionSM86FailureCode)
field unrestricted hmmaProductionContractTelemetry : (family HMMAProductionSM86Telemetry)
constructor HMMAProductionSM86ImageEncodingFailed
field unrestricted hmmaProductionEncodingFailure : (family SM86ProgramEncodingResult)
field unrestricted hmmaProductionEncodingFailureTelemetry : (family HMMAProductionSM86Telemetry)
constructor HMMAProductionSM86ImageIdentityFailed
field unrestricted hmmaProductionIdentityFailure : (family SHA256HexResult)
field unrestricted hmmaProductionIdentityFailureTelemetry : (family HMMAProductionSM86Telemetry)
end-family
family HMMAProductionNativeTelemetry : Type 0
constructor HMMAProductionNativeTelemetryValue
field unrestricted hmmaProductionNativeTelemetryManifest : (family HMMAProductionSM86Manifest)
field unrestricted hmmaProductionNativeTelemetryObservedInstructions : Nat
field unrestricted hmmaProductionNativeTelemetryObservedBytes : Nat
field unrestricted hmmaProductionNativeTelemetryStages : Nat
field unrestricted hmmaProductionNativeTelemetryLastStageByteOffset : Nat
field unrestricted hmmaProductionNativeTelemetryGlobalLoads : Nat
field unrestricted hmmaProductionNativeTelemetryGlobalStores : Nat
field unrestricted hmmaProductionNativeTelemetrySharedStores : Nat
field unrestricted hmmaProductionNativeTelemetrySharedMatrixLoads : Nat
field unrestricted hmmaProductionNativeTelemetryHMMAInstructions : Nat
field unrestricted hmmaProductionNativeTelemetryBarrierInstructions : Nat
field unrestricted hmmaProductionNativeTelemetryGroupedAddressInstructions : Nat
field unrestricted hmmaProductionNativeTelemetryHostFallbackOperations : Nat
end-family
family HMMAProductionNativeBuildResult : Type 0
constructor HMMAProductionNativeBuildSucceeded
field unrestricted hmmaProductionNativeEncodedBytes : Bytes
field unrestricted hmmaProductionNativeImageSHA256 : Bytes
field unrestricted hmmaProductionNativeEncodingTelemetry : (family SM86ProgramEncodingTelemetry)
field unrestricted hmmaProductionNativeIdentityTelemetry : (family SHA256DigestTelemetry)
field unrestricted hmmaProductionNativeBuildTelemetry : (family HMMAProductionNativeTelemetry)
constructor HMMAProductionNativeContractFailed
field unrestricted hmmaProductionNativeContractFailure : (family HMMAProductionSM86FailureCode)
field unrestricted hmmaProductionNativeContractTelemetry : (family HMMAProductionNativeTelemetry)
constructor HMMAProductionNativeEncodingFailed
field unrestricted hmmaProductionNativeEncodingFailure : (family SM86ProgramEncodingResult)
field unrestricted hmmaProductionNativeEncodingFailureTelemetry : (family HMMAProductionNativeTelemetry)
constructor HMMAProductionNativeIdentityFailed
field unrestricted hmmaProductionNativeIdentityFailure : (family SHA256HexResult)
field unrestricted hmmaProductionNativeIdentityFailureTelemetry : (family HMMAProductionNativeTelemetry)
end-family
def hmmaProductionN16 =
(byte-to-nat (byte 16))
def hmmaProductionN8 =
(byte-to-nat (byte 8))
def hmmaProductionN10 =
(byte-to-nat (byte 10))
def hmmaProductionN21 =
(byte-to-nat (byte 21))
def hmmaProductionN22 =
(byte-to-nat (byte 22))
def hmmaProductionN24 =
(byte-to-nat (byte 24))
def hmmaProductionN32 =
(byte-to-nat (byte 32))
def hmmaProductionN40 =
(byte-to-nat (byte 40))
def hmmaProductionN48 =
(byte-to-nat (byte 48))
def hmmaProductionN64 =
(byte-to-nat (byte 64))
def hmmaProductionN69 =
(byte-to-nat (byte 69))
def hmmaProductionN70 =
(byte-to-nat (byte 70))
def hmmaProductionN96 =
(byte-to-nat (byte 96))
def hmmaProductionN128 =
(naturalMultiply hmmaProductionN16 hmmaProductionN8)
def hmmaProductionN256 =
(naturalMultiply hmmaProductionN16 hmmaProductionN16)
def hmmaProductionN384 =
(naturalMultiply hmmaProductionN16 hmmaProductionN24)
def hmmaProductionN512 =
(naturalMultiply hmmaProductionN32 hmmaProductionN16)
def hmmaProductionN640 =
(naturalMultiply hmmaProductionN64 hmmaProductionN10)
def hmmaProductionN1024 =
(naturalMultiply hmmaProductionN16 hmmaProductionN64)
def hmmaProductionN1536 =
(naturalMultiply hmmaProductionN16 hmmaProductionN96)
def hmmaProductionN2048 =
(naturalMultiply hmmaProductionN32 hmmaProductionN64)
def hmmaProductionN4096 =
(naturalMultiply hmmaProductionN64 hmmaProductionN64)
def hmmaProductionN6144 =
(naturalMultiply hmmaProductionN64 hmmaProductionN96)
def hmmaProductionN8192 =
(naturalMultiply hmmaProductionN128 hmmaProductionN64)
-- Natural literals above byte range are represented by arithmetic so every
-- manifest remains a typed Alpha value rather than a host-side table.
def hmmaProduction1024 =
hmmaProductionN1024
def hmmaProduction1536 =
hmmaProductionN1536
def hmmaProduction2048 =
hmmaProductionN2048
def hmmaProduction4096 =
hmmaProductionN4096
def hmmaProduction6144 =
hmmaProductionN6144
def hmmaProduction8192 =
hmmaProductionN8192
def hmmaProductionR0 =
(sm86Register (byte 0))
def hmmaProductionR1 =
(sm86Register (byte 1))
def hmmaProductionR2 =
(sm86Register (byte 2))
def hmmaProductionR3 =
(sm86Register (byte 3))
def hmmaProductionR4 =
(sm86Register (byte 4))
def hmmaProductionR6 =
(sm86Register (byte 6))
def hmmaProductionR8 =
(sm86Register (byte 8))
def hmmaProductionR9 =
(sm86Register (byte 9))
def hmmaProductionR10 =
(sm86Register (byte 10))
def hmmaProductionR11 =
(sm86Register (byte 11))
def hmmaProductionR12 =
(sm86Register (byte 12))
def hmmaProductionR13 =
(sm86Register (byte 13))
def hmmaProductionR14 =
(sm86Register (byte 14))
def hmmaProductionR15 =
(sm86Register (byte 15))
def hmmaProductionR16 =
(sm86Register (byte 16))
def hmmaProductionR18 =
(sm86Register (byte 18))
def hmmaProductionR20 =
(sm86Register (byte 20))
def hmmaProductionR22 =
(sm86Register (byte 22))
def hmmaProductionR24 =
(sm86Register (byte 24))
def hmmaProductionR25 =
(sm86Register (byte 25))
def hmmaProductionR26 =
(sm86Register (byte 26))
def hmmaProductionR27 =
(sm86Register (byte 27))
def hmmaProductionR28 =
(sm86Register (byte 28))
def hmmaProductionR29 =
(sm86Register (byte 29))
def hmmaProductionU0 =
sm86Unsigned32Zero
def hmmaProductionU4 =
(sm86Unsigned32 (byte 4) (byte 0) (byte 0) (byte 0))
def hmmaProductionU8 =
(sm86Unsigned32 (byte 8) (byte 0) (byte 0) (byte 0))
def hmmaProductionU12 =
(sm86Unsigned32 (byte 12) (byte 0) (byte 0) (byte 0))
def hmmaProductionU16 =
(sm86Unsigned32 (byte 16) (byte 0) (byte 0) (byte 0))
def hmmaProductionU20 =
(sm86Unsigned32 (byte 20) (byte 0) (byte 0) (byte 0))
def hmmaProductionU28 =
(sm86Unsigned32 (byte 28) (byte 0) (byte 0) (byte 0))
def hmmaProductionU32 =
(sm86Unsigned32 (byte 32) (byte 0) (byte 0) (byte 0))
def hmmaProductionU1024 =
(sm86Unsigned32 (byte 0) (byte 4) (byte 0) (byte 0))
def hmmaProductionOffset028 =
(sm86Unsigned32 (byte 40) (byte 0) (byte 0) (byte 0))
def hmmaProductionOffset160 =
(sm86Unsigned32 (byte 96) (byte 1) (byte 0) (byte 0))
def hmmaProductionOffset168 =
(sm86Unsigned32 (byte 104) (byte 1) (byte 0) (byte 0))
def hmmaProductionOffset170 =
(sm86Unsigned32 (byte 112) (byte 1) (byte 0) (byte 0))
def hmmaProductionSet0 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0))
def hmmaProductionSet1 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier1))
def hmmaProductionSet2 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier2))
def hmmaProductionSet3 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier3))
def hmmaProductionWait0 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0))
def hmmaProductionWait1 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier1))
def hmmaProductionWait2 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier2))
def hmmaProductionWaitFragments0 =
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86BarrierNone)
(byte 3)
(byte 0))
def hmmaProductionWaitFragments1 =
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86BarrierNone)
(byte 5)
(byte 0))
def hmmaProductionInstruction =
(lambda unrestricted body : (family SM86InstructionBody) .
(sm86ProgramSingleton (sm86Instruction body)))
-- The staged body is deliberately an Alpha structural-recursion unit. It
-- contains no host mapM/list iteration and leaves the barriers, LDSM count and
-- FP16/FP32 boundary visible to the encoder.
def hmmaProductionABTransposedStage : (family SM86Program) =
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86LoadGlobalWide
hmmaProductionR8
hmmaProductionR4
hmmaProductionU0
hmmaProductionSet0))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86LoadGlobalWide
hmmaProductionR12
hmmaProductionR6
hmmaProductionU0
hmmaProductionSet1))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR24
hmmaProductionR8
hmmaProductionU0
hmmaProductionWait0))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR24
hmmaProductionR9
hmmaProductionU4
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR24
hmmaProductionR10
hmmaProductionU8
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR24
hmmaProductionR11
hmmaProductionU12
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR27
hmmaProductionR12
(sm86Unsigned32 (byte 0) (byte 8) (byte 0) (byte 0))
hmmaProductionWait1))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR27
hmmaProductionR13
(sm86Unsigned32 (byte 4) (byte 8) (byte 0) (byte 0))
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR27
hmmaProductionR14
(sm86Unsigned32 (byte 8) (byte 8) (byte 0) (byte 0))
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86StoreShared
hmmaProductionR27
hmmaProductionR15
(sm86Unsigned32 (byte 12) (byte 8) (byte 0) (byte 0))
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor SM86InstructionBody SM86BarrierSynchronize sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaProductionR8
hmmaProductionR28
hmmaProductionU0
(constructor SM86SharedMatrixCount SM86SharedMatrix4)
(constructor SM86SharedMatrixTranspose SM86SharedMatrixNotTransposed)
hmmaProductionSet0))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaProductionR12
hmmaProductionR29
hmmaProductionU0
(constructor SM86SharedMatrixCount SM86SharedMatrix2)
(constructor
SM86SharedMatrixTranspose
SM86SharedMatrixNotTransposed)
hmmaProductionSet1))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
hmmaProductionR16
hmmaProductionR8
hmmaProductionR12
hmmaProductionR16
hmmaProductionWaitFragments0))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaProductionR8
hmmaProductionR28
hmmaProductionU32
(constructor SM86SharedMatrixCount SM86SharedMatrix4)
(constructor
SM86SharedMatrixTranspose
SM86SharedMatrixNotTransposed)
hmmaProductionSet0))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaProductionR12
hmmaProductionR29
hmmaProductionU32
(constructor SM86SharedMatrixCount SM86SharedMatrix2)
(constructor
SM86SharedMatrixTranspose
SM86SharedMatrixNotTransposed)
hmmaProductionSet1))
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
hmmaProductionR16
hmmaProductionR8
hmmaProductionR12
hmmaProductionR16
hmmaProductionWaitFragments0))))))))))))))))))
-- AB and A^T B retain their distinct layout boundaries. The former uses LDSM.T
-- for B; the latter widens selected halves and repacks in source-high/source-low
-- order before the shared transpose. These are typed as separate programs so a
-- caller cannot erase the numeric boundary into a host conversion.
def hmmaProductionABStage : (family SM86Program) =
hmmaProductionABTransposedStage
def hmmaProductionATransposedBStage : (family SM86Program) =
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaProductionR12
hmmaProductionR8
(constructor SM86HalfSelector SM86LowHalf)
hmmaProductionWait0))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaProductionR13
hmmaProductionR9
(constructor SM86HalfSelector SM86LowHalf)
hmmaProductionWait1))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaProductionR14
hmmaProductionR8
(constructor SM86HalfSelector SM86HighHalf)
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaProductionR15
hmmaProductionR9
(constructor SM86HalfSelector SM86HighHalf)
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86FloatPairToPackedHalfPair
hmmaProductionR8
hmmaProductionR13
hmmaProductionR12
sm86SafeControl))
(sm86ProgramAppend
(hmmaProductionInstruction
(constructor
SM86InstructionBody
SM86FloatPairToPackedHalfPair
hmmaProductionR9
hmmaProductionR15
hmmaProductionR14
sm86SafeControl))
hmmaProductionABTransposedStage))))))
def hmmaProductionStageFor =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) . (family SM86Program))
orientation
(branch HMMAProductionAB . hmmaProductionABStage)
(branch HMMAProductionABTransposed . hmmaProductionABTransposedStage)
(branch HMMAProductionATransposedB . hmmaProductionATransposedBStage)))
def hmmaProductionStages =
(lambda unrestricted stageCount : Nat .
(lambda unrestricted stage : (family SM86Program) .
(nat-eliminate
(lambda unrestricted current : Nat . (family SM86Program))
sm86ProgramEmpty
(lambda unrestricted predecessor : Nat .
(lambda unrestricted induction : (family SM86Program) .
(sm86ProgramAppend stage induction)))
stageCount)))
def hmmaProductionProgramFor =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate
HMMAProductionSM86Geometry
(lambda unrestricted current : (family HMMAProductionSM86Geometry) . (family SM86Program))
geometry
(branch
HMMAProductionSM86GeometryValue
m
n
k
stages
.
(hmmaProductionStages stages (hmmaProductionStageFor orientation))))))
def hmmaProductionOrientationFor =
(lambda unrestricted variant : (family HMMAProductionSM86Variant) .
(eliminate
HMMAProductionSM86Variant
(lambda unrestricted current : (family HMMAProductionSM86Variant) .
(family HMMAProductionSM86Orientation))
variant
(branch
HMMAProductionDenseForward
.
(constructor HMMAProductionSM86Orientation HMMAProductionABTransposed))
(branch
HMMAProductionDenseReverseInput
.
(constructor HMMAProductionSM86Orientation HMMAProductionAB))
(branch
HMMAProductionDenseReverseWeight
.
(constructor HMMAProductionSM86Orientation HMMAProductionATransposedB))
(branch
HMMAProductionRoutedForward
.
(constructor HMMAProductionSM86Orientation HMMAProductionABTransposed))
(branch
HMMAProductionRoutedReverseInput
.
(constructor HMMAProductionSM86Orientation HMMAProductionAB))
(branch
HMMAProductionRoutedReverseWeight
.
(constructor HMMAProductionSM86Orientation HMMAProductionATransposedB))
(branch
HMMAProductionAttentionQKForward
.
(constructor HMMAProductionSM86Orientation HMMAProductionABTransposed))
(branch
HMMAProductionAttentionProbabilityValueForward
.
(constructor HMMAProductionSM86Orientation HMMAProductionAB))
(branch
HMMAProductionAttentionQKReverseLeft
.
(constructor HMMAProductionSM86Orientation HMMAProductionAB))
(branch
HMMAProductionAttentionQKReverseRight
.
(constructor HMMAProductionSM86Orientation HMMAProductionATransposedB))
(branch
HMMAProductionAttentionValueReverseLeft
.
(constructor HMMAProductionSM86Orientation HMMAProductionABTransposed))
(branch
HMMAProductionAttentionValueReverseRight
.
(constructor HMMAProductionSM86Orientation HMMAProductionATransposedB))))
def hmmaProductionManifestFor =
(lambda unrestricted variant : (family HMMAProductionSM86Variant) .
(constructor
HMMAProductionSM86Manifest
HMMAProductionSM86ManifestValue
variant
(hmmaProductionOrientationFor variant)
(constructor HMMAProductionSM86Grouping HMMAProductionDense)
(constructor
HMMAProductionSM86Geometry
HMMAProductionSM86GeometryValue
hmmaProductionN256
hmmaProductionN256
hmmaProductionN256
hmmaProductionN8)
(naturalAdd hmmaProductionN69 (naturalMultiply hmmaProductionN21 hmmaProductionN8))
(naturalMultiply
(naturalAdd hmmaProductionN69 (naturalMultiply hmmaProductionN21 hmmaProductionN8))
hmmaProductionN16)
hmmaProductionN40
hmmaProduction8192
hmmaProductionN8
hmmaProductionN8
(succ zero)
hmmaProductionN128
(constructor
HMMAProductionSM86ABI
HMMAProductionSM86ABIValue
(naturalAdd hmmaProductionN256 hmmaProductionN96)
(naturalAdd hmmaProductionN256 (naturalAdd hmmaProductionN96 hmmaProductionN8))
(naturalAdd hmmaProductionN256 (naturalAdd hmmaProductionN96 hmmaProductionN16)))
(constructor
HMMAProductionSM86Extents
HMMAProductionSM86ExtentsValue
hmmaProductionN256
hmmaProductionN256
hmmaProductionN256
(succ zero))
zero))
def hmmaProductionTelemetryFor =
(lambda unrestricted variant : (family HMMAProductionSM86Variant) .
(lambda unrestricted instructions : Nat .
(lambda unrestricted bytes : Nat .
(constructor
HMMAProductionSM86Telemetry
HMMAProductionSM86TelemetryValue
(hmmaProductionManifestFor variant)
instructions
bytes
hmmaProductionN8
zero))))
def hmmaProductionBuild =
(lambda unrestricted variant : (family HMMAProductionSM86Variant) .
(app
(lambda unrestricted manifest : (family HMMAProductionSM86Manifest) .
(app
(lambda unrestricted program : (family SM86Program) .
(app
(lambda unrestricted instructions : Nat .
(app
(lambda unrestricted encoding : (family SM86ProgramEncodingResult) .
(eliminate
SM86ProgramEncodingResult
(lambda unrestricted current : (family SM86ProgramEncodingResult) .
(family HMMAProductionSM86BuildResult))
encoding
(branch
SM86ProgramEncodingSucceeded
bytes
encodingTelemetry
.
(eliminate
SHA256HexResult
(lambda unrestricted current : (family SHA256HexResult) .
(family HMMAProductionSM86BuildResult))
(sha256Hex bytes)
(branch
SHA256HexSucceeded
identity
identityTelemetry
.
(constructor
HMMAProductionSM86BuildResult
HMMAProductionSM86BuildSucceeded
bytes
identity
encodingTelemetry
identityTelemetry
(hmmaProductionTelemetryFor variant instructions (bytes-length bytes))))
(branch
SHA256HexFailed
error
ordinal
identityTelemetry
.
(constructor
HMMAProductionSM86BuildResult
HMMAProductionSM86ImageIdentityFailed
(sha256Hex bytes)
(hmmaProductionTelemetryFor variant instructions (bytes-length bytes))))))
(branch
SM86ProgramEncodingFailed
ordinal
failure
telemetry
.
(constructor
HMMAProductionSM86BuildResult
HMMAProductionSM86ImageEncodingFailed
encoding
(hmmaProductionTelemetryFor variant instructions zero)))))
(sm86EncodeProgram program)))
(sm86ProgramCount program)))
(hmmaProductionProgramFor
(hmmaProductionOrientationFor variant)
(constructor
HMMAProductionSM86Geometry
HMMAProductionSM86GeometryValue
hmmaProductionN256
hmmaProductionN256
hmmaProductionN256
hmmaProductionN8))))
(hmmaProductionManifestFor variant)))
-- Production register-tiled obligation family. The older generic producer
-- above remains a bootstrap reference. These six entry points are the native
-- training obligations and bind grouping, orientation, geometry, ABI, resource
-- use, exact instruction counts, encoded bytes, and image identity together.
def hmmaNativeN2 =
(byte-to-nat (byte 2))
def hmmaNativeN4 =
(byte-to-nat (byte 4))
def hmmaNativeN5 =
(byte-to-nat (byte 5))
def hmmaNativeN6 =
(byte-to-nat (byte 6))
def hmmaNativeN8 =
(byte-to-nat (byte 8))
def hmmaNativeN16 =
(byte-to-nat (byte 16))
def hmmaNativeN21 =
(byte-to-nat (byte 21))
def hmmaNativeN23 =
(byte-to-nat (byte 23))
def hmmaNativeN32 =
(byte-to-nat (byte 32))
def hmmaNativeN40 =
(byte-to-nat (byte 40))
def hmmaNativeN64 =
(byte-to-nat (byte 64))
def hmmaNativeN69 =
(byte-to-nat (byte 69))
def hmmaNativeN70 =
(byte-to-nat (byte 70))
def hmmaNativeN128 =
(naturalMultiply hmmaNativeN16 hmmaNativeN8)
def hmmaNativeN256 =
(naturalMultiply hmmaNativeN16 hmmaNativeN16)
def hmmaNativeN512 =
(naturalMultiply hmmaNativeN32 hmmaNativeN16)
def hmmaNativeN8192 =
(naturalMultiply hmmaNativeN128 hmmaNativeN64)
def hmmaNativeN65536 =
(naturalPowerOfTwo hmmaNativeN16)
def hmmaNativeU32Bytes =
(lambda unrestricted b0 : Byte .
(lambda unrestricted b1 : Byte .
(lambda unrestricted b2 : Byte .
(lambda unrestricted b3 : Byte . (sm86Unsigned32 b0 b1 b2 b3)))))
def hmmaNativeU32 =
(lambda unrestricted value : Nat .
(app
(lambda unrestricted q1 : Nat .
(app
(lambda unrestricted q2 : Nat .
(app
(lambda unrestricted q3 : Nat .
(hmmaNativeU32Bytes
(nat-to-byte (naturalModuloUnchecked value hmmaNativeN256))
(nat-to-byte (naturalModuloUnchecked q1 hmmaNativeN256))
(nat-to-byte (naturalModuloUnchecked q2 hmmaNativeN256))
(nat-to-byte (naturalModuloUnchecked q3 hmmaNativeN256))))
(naturalDivideUnchecked q2 hmmaNativeN256)))
(naturalDivideUnchecked q1 hmmaNativeN256)))
(naturalDivideUnchecked value hmmaNativeN256)))
def hmmaNativeU0 =
sm86Unsigned32Zero
def hmmaNativeU4 =
(hmmaNativeU32 hmmaNativeN4)
def hmmaNativeU32Value =
(hmmaNativeU32 hmmaNativeN32)
def hmmaNativeU512 =
(hmmaNativeU32 hmmaNativeN512)
def hmmaNativeUNeg2 =
(hmmaNativeU32Bytes (byte 254) (byte 255) (byte 255) (byte 255))
def hmmaNativeUNeg4 =
(hmmaNativeU32Bytes (byte 252) (byte 255) (byte 255) (byte 255))
def hmmaNativeUNeg8 =
(hmmaNativeU32Bytes (byte 248) (byte 255) (byte 255) (byte 255))
def hmmaNativeUNeg16 =
(hmmaNativeU32Bytes (byte 240) (byte 255) (byte 255) (byte 255))
def hmmaNativeUNeg32 =
(hmmaNativeU32Bytes (byte 224) (byte 255) (byte 255) (byte 255))
def hmmaNativeR0 =
(sm86Register (byte 0))
def hmmaNativeR1 =
(sm86Register (byte 1))
def hmmaNativeR2 =
(sm86Register (byte 2))
def hmmaNativeR3 =
(sm86Register (byte 3))
def hmmaNativeR4 =
(sm86Register (byte 4))
def hmmaNativeR5 =
(sm86Register (byte 5))
def hmmaNativeR6 =
(sm86Register (byte 6))
def hmmaNativeR7 =
(sm86Register (byte 7))
def hmmaNativeR8 =
(sm86Register (byte 8))
def hmmaNativeR9 =
(sm86Register (byte 9))
def hmmaNativeR10 =
(sm86Register (byte 10))
def hmmaNativeR11 =
(sm86Register (byte 11))
def hmmaNativeR12 =
(sm86Register (byte 12))
def hmmaNativeR13 =
(sm86Register (byte 13))
def hmmaNativeR14 =
(sm86Register (byte 14))
def hmmaNativeR15 =
(sm86Register (byte 15))
def hmmaNativeR16 =
(sm86Register (byte 16))
def hmmaNativeR17 =
(sm86Register (byte 17))
def hmmaNativeR18 =
(sm86Register (byte 18))
def hmmaNativeR19 =
(sm86Register (byte 19))
def hmmaNativeR20 =
(sm86Register (byte 20))
def hmmaNativeR21 =
(sm86Register (byte 21))
def hmmaNativeR22 =
(sm86Register (byte 22))
def hmmaNativeR23 =
(sm86Register (byte 23))
def hmmaNativeR24 =
(sm86Register (byte 24))
def hmmaNativeR25 =
(sm86Register (byte 25))
def hmmaNativeR26 =
(sm86Register (byte 26))
def hmmaNativeR27 =
(sm86Register (byte 27))
def hmmaNativeR28 =
(sm86Register (byte 28))
def hmmaNativeR29 =
(sm86Register (byte 29))
def hmmaNativeR30 =
(sm86Register (byte 30))
def hmmaNativeR31 =
(sm86Register (byte 31))
def hmmaNativeSet0 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0))
def hmmaNativeSet1 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier1))
def hmmaNativeSet2 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier2))
def hmmaNativeSet3 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier3))
def hmmaNativeWait0 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0))
def hmmaNativeWait1 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier1))
def hmmaNativeWaitFragments0 =
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86BarrierNone)
(byte 3)
(byte 0))
def hmmaNativeWaitFragments1 =
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86BarrierNone)
(byte 5)
(byte 0))
def hmmaNativeSpecialWait =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(eliminate
HMMAProductionSM86Grouping
(lambda unrestricted current : (family HMMAProductionSM86Grouping) . (family SM86Control))
grouping
(branch
HMMAProductionDense
.
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86BarrierNone)
(byte 7)
(byte 0)))
(branch
HMMAProductionGrouped
groups
.
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86BarrierNone)
(byte 15)
(byte 0)))))
def hmmaNativeCons =
(lambda unrestricted body : (family SM86InstructionBody) .
(lambda unrestricted tail : (family SM86Program) .
(constructor SM86Program SM86ProgramNext (sm86Instruction body) tail)))
def hmmaNativeAppend =
(lambda unrestricted left : (family SM86Program) .
(lambda unrestricted right : (family SM86Program) . (sm86ProgramAppend left right)))
def hmmaNativeGeometryM =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate
HMMAProductionSM86Geometry
(lambda unrestricted current : (family HMMAProductionSM86Geometry) . Nat)
geometry
(branch HMMAProductionSM86GeometryValue m n k stages . m)))
def hmmaNativeGeometryN =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate
HMMAProductionSM86Geometry
(lambda unrestricted current : (family HMMAProductionSM86Geometry) . Nat)
geometry
(branch HMMAProductionSM86GeometryValue m n k stages . n)))
def hmmaNativeGeometryK =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate
HMMAProductionSM86Geometry
(lambda unrestricted current : (family HMMAProductionSM86Geometry) . Nat)
geometry
(branch HMMAProductionSM86GeometryValue m n k stages . k)))
def hmmaNativeGeometryStages =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate
HMMAProductionSM86Geometry
(lambda unrestricted current : (family HMMAProductionSM86Geometry) . Nat)
geometry
(branch HMMAProductionSM86GeometryValue m n k stages . stages)))
def hmmaNativeGroupingCount =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(eliminate
HMMAProductionSM86Grouping
(lambda unrestricted current : (family HMMAProductionSM86Grouping) . Nat)
grouping
(branch HMMAProductionDense . (succ zero))
(branch HMMAProductionGrouped groups . groups)))
def hmmaNativeIsGrouped =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(eliminate
HMMAProductionSM86Grouping
(lambda unrestricted current : (family HMMAProductionSM86Grouping) . Nat)
grouping
(branch HMMAProductionDense . zero)
(branch HMMAProductionGrouped groups . (succ zero))))
def hmmaNativeGroupAWords =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(naturalDivideUnchecked
(naturalMultiply (hmmaNativeGeometryM geometry) (hmmaNativeGeometryK geometry))
hmmaNativeN2))
def hmmaNativeGroupBWords =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(naturalDivideUnchecked
(naturalMultiply (hmmaNativeGeometryK geometry) (hmmaNativeGeometryN geometry))
hmmaNativeN2))
def hmmaNativeGroupCWords =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(naturalMultiply (hmmaNativeGeometryM geometry) (hmmaNativeGeometryN geometry)))
-- Stored shapes determine the default row pitch and inter-group plane size.
-- AB uses [m,k] [k,n]; AB^T uses [m,k] [n,k]; A^T B uses [k,m] [k,n].
def matrixProductLeftStride = (lambda unrestricted layout : (family MatrixProductRowStrides) .
(eliminate MatrixProductRowStrides (lambda unrestricted current : (family MatrixProductRowStrides) . Nat) layout
(branch MatrixProductRowStridesValue left right output . left)))
def matrixProductRightStride = (lambda unrestricted layout : (family MatrixProductRowStrides) .
(eliminate MatrixProductRowStrides (lambda unrestricted current : (family MatrixProductRowStrides) . Nat) layout
(branch MatrixProductRowStridesValue left right output . right)))
def matrixProductOutputStride = (lambda unrestricted layout : (family MatrixProductRowStrides) .
(eliminate MatrixProductRowStrides (lambda unrestricted current : (family MatrixProductRowStrides) . Nat) layout
(branch MatrixProductRowStridesValue left right output . output)))
def matrixProductLeftStoredRows = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate HMMAProductionSM86Orientation (lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat) orientation
(branch HMMAProductionAB . (hmmaNativeGeometryM geometry))
(branch HMMAProductionABTransposed . (hmmaNativeGeometryM geometry))
(branch HMMAProductionATransposedB . (hmmaNativeGeometryK geometry)))))
def matrixProductRightStoredRows = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate HMMAProductionSM86Orientation (lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat) orientation
(branch HMMAProductionAB . (hmmaNativeGeometryK geometry))
(branch HMMAProductionABTransposed . (hmmaNativeGeometryN geometry))
(branch HMMAProductionATransposedB . (hmmaNativeGeometryK geometry)))))
def matrixProductContiguousRowStrides = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(constructor MatrixProductRowStrides MatrixProductRowStridesValue
(eliminate HMMAProductionSM86Orientation (lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat) orientation
(branch HMMAProductionAB . (hmmaNativeGeometryK geometry))
(branch HMMAProductionABTransposed . (hmmaNativeGeometryK geometry))
(branch HMMAProductionATransposedB . (hmmaNativeGeometryM geometry)))
(eliminate HMMAProductionSM86Orientation (lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat) orientation
(branch HMMAProductionAB . (hmmaNativeGeometryN geometry))
(branch HMMAProductionABTransposed . (hmmaNativeGeometryK geometry))
(branch HMMAProductionATransposedB . (hmmaNativeGeometryN geometry)))
(hmmaNativeGeometryN geometry))))
def matrixProductLeftPlaneWords = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(naturalDivideUnchecked (naturalMultiply (matrixProductLeftStoredRows orientation geometry) (matrixProductLeftStride layout)) 2))))
def matrixProductRightPlaneWords = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(naturalDivideUnchecked (naturalMultiply (matrixProductRightStoredRows orientation geometry) (matrixProductRightStride layout)) 2))))
def matrixProductOutputPlaneWords = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(naturalMultiply (hmmaNativeGeometryM geometry) (matrixProductOutputStride layout))))
def hmmaNativeGroupSpecial =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(eliminate
HMMAProductionSM86Grouping
(lambda unrestricted current : (family HMMAProductionSM86Grouping) . (family SM86Program))
grouping
(branch HMMAProductionDense . sm86ProgramEmpty)
(branch
HMMAProductionGrouped
groups
.
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
destination
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
control)
sm86ProgramEmpty))))))
def hmmaNativeGroupOffsetWith =
(lambda unrestricted groupRegister : (family SM86Register) .
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted stride : Nat .
(lambda unrestricted index : (family SM86Register) .
(eliminate
HMMAProductionSM86Grouping
(lambda unrestricted current : (family HMMAProductionSM86Grouping) .
(family SM86Program))
grouping
(branch HMMAProductionDense . sm86ProgramEmpty)
(branch
HMMAProductionGrouped
groups
.
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
index
groupRegister
(hmmaNativeU32 stride)
index
sm86SafeControl)
sm86ProgramEmpty)))))))
def hmmaNativeGroupOffset =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted stride : Nat .
(lambda unrestricted index : (family SM86Register) .
(hmmaNativeGroupOffsetWith hmmaNativeR2 grouping stride index))))
def hmmaNativeRowMajor =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted groupStride : Nat .
(lambda unrestricted pointer : (family SM86Register) .
(lambda unrestricted sharedAddress : (family SM86Register) .
(lambda unrestricted cta : (family SM86Register) .
(lambda unrestricted leading : Nat .
(lambda unrestricted constantOffset : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR20
hmmaNativeR0
(byte 2)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR21
hmmaNativeR20
hmmaNativeUNeg4
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
cta
(hmmaNativeU32 hmmaNativeN32)
hmmaNativeR20
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR23
hmmaNativeR21
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
hmmaNativeR22
(hmmaNativeU32 (naturalDivideUnchecked leading hmmaNativeN2))
hmmaNativeR23
sm86SafeControl)
(hmmaNativeAppend
(hmmaNativeGroupOffset grouping groupStride hmmaNativeR22)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddWideConstant
pointer
hmmaNativeR22
hmmaNativeR3
(byte 0)
(hmmaNativeU32 constantOffset)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR23
hmmaNativeR21
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
sharedAddress
hmmaNativeR20
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR23
sm86SafeControl)
sm86ProgramEmpty))))))))))))))))
def hmmaNativeReductionMajor =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted groupStride : Nat .
(lambda unrestricted pointer : (family SM86Register) .
(lambda unrestricted sharedAddress : (family SM86Register) .
(lambda unrestricted baseIndex : (family SM86Register) .
(lambda unrestricted cta : (family SM86Register) .
(lambda unrestricted outputElements : Nat .
(lambda unrestricted constantOffset : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR20
hmmaNativeR0
(byte 2)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR21
hmmaNativeR20
hmmaNativeUNeg4
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR23
hmmaNativeR21
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
cta
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR23
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
hmmaNativeR20
(hmmaNativeU32 (naturalDivideUnchecked outputElements hmmaNativeN2))
hmmaNativeR22
sm86SafeControl)
(hmmaNativeAppend
(hmmaNativeGroupOffset grouping groupStride hmmaNativeR22)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddWideConstant
pointer
hmmaNativeR22
hmmaNativeR3
(byte 0)
(hmmaNativeU32 constantOffset)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerAddThreeImmediate
baseIndex
hmmaNativeR22
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR23
hmmaNativeR21
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
sharedAddress
hmmaNativeR20
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR23
sm86SafeControl)
sm86ProgramEmpty))))))))))))))))))
def hmmaNativeTransposedA =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted groupStride : Nat .
(lambda unrestricted pointer : (family SM86Register) .
(lambda unrestricted sharedAddress : (family SM86Register) .
(lambda unrestricted baseIndex : (family SM86Register) .
(lambda unrestricted cta : (family SM86Register) .
(lambda unrestricted outputElements : Nat .
(lambda unrestricted constantOffset : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR20
hmmaNativeR0
(byte 4)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR21
hmmaNativeR20
hmmaNativeUNeg16
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
cta
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR21
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
hmmaNativeR20
(hmmaNativeU32 (naturalMultiply hmmaNativeN2 outputElements))
hmmaNativeR22
sm86SafeControl)
(hmmaNativeAppend
(hmmaNativeGroupOffset grouping groupStride hmmaNativeR22)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddWideConstant
pointer
hmmaNativeR22
hmmaNativeR3
(byte 0)
(hmmaNativeU32 constantOffset)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerAddThreeImmediate
baseIndex
hmmaNativeR22
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR23
hmmaNativeR20
(hmmaNativeU32 hmmaNativeN2)
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
sharedAddress
hmmaNativeR21
(hmmaNativeU32 hmmaNativeN32)
hmmaNativeR23
sm86SafeControl)
sm86ProgramEmpty)))))))))))))))))
-- For B stored K-by-N, LDSM.TRANS exchanges the within-8x8 indices,
-- not the matrix-tile addresses. Its source rows advance by 64 bytes and
-- columns by 2. The old shared address was the N-by-K address for every
-- orientation, so AB and A^T B read the wrong 8x8 tiles while AB^T passed.
def hmmaNativeBReductionRows = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(eliminate HMMAProductionSM86Orientation (lambda unrestricted o : (family HMMAProductionSM86Orientation) . Nat) orientation
(branch HMMAProductionAB . 1) (branch HMMAProductionABTransposed . 0) (branch HMMAProductionATransposedB . 1)))
def hmmaNativeBColumnLane = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(nat-eliminate (lambda unrestricted transposed : Nat . (family SM86Register)) hmmaNativeR15
(lambda unrestricted predecessor : Nat . (lambda unrestricted unused : (family SM86Register) . hmmaNativeR22))
(hmmaNativeBReductionRows orientation)))
def hmmaNativeBRowGroup = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(nat-eliminate (lambda unrestricted transposed : Nat . (family SM86Register)) hmmaNativeR22
(lambda unrestricted predecessor : Nat . (lambda unrestricted unused : (family SM86Register) . hmmaNativeR15))
(hmmaNativeBReductionRows orientation)))
def hmmaNativeBNextColumns = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(naturalSelect (hmmaNativeBReductionRows orientation) hmmaNativeN16 hmmaNativeN512))
def hmmaNativeBNextReduction = (lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(naturalSelect (hmmaNativeBReductionRows orientation) (naturalMultiply hmmaNativeN2 hmmaNativeN512) hmmaNativeN32))
def hmmaNativeFragmentAddresses =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR20
hmmaNativeR0
(byte 5)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR21
hmmaNativeR20
(byte 1)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR22
hmmaNativeR21
hmmaNativeUNeg2
hmmaNativeR20
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR23
hmmaNativeR20
hmmaNativeUNeg32
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR20
hmmaNativeR23
(byte 3)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR0
hmmaNativeR20
hmmaNativeUNeg8
hmmaNativeR23
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR14
hmmaNativeR20
(byte 1)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR15
hmmaNativeR14
hmmaNativeUNeg2
hmmaNativeR20
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR28
hmmaNativeR21
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR28
hmmaNativeR15
(hmmaNativeU32 hmmaNativeN8)
hmmaNativeR28
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR20
hmmaNativeR14
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR28
hmmaNativeR28
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR20
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR28
hmmaNativeR28
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR20
(hmmaNativeBColumnLane orientation)
(hmmaNativeU32 (naturalSelect (hmmaNativeBReductionRows orientation) hmmaNativeN8 hmmaNativeN4))
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR29
(hmmaNativeBRowGroup orientation)
(hmmaNativeU32 (naturalSelect (hmmaNativeBReductionRows orientation) hmmaNativeN8 hmmaNativeN16))
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR29
hmmaNativeR29
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR20
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerAddThreeImmediate
hmmaNativeR29
hmmaNativeR29
hmmaNativeU512
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR29
hmmaNativeR29
hmmaNativeU4
sm86ZeroRegister
sm86SafeControl)
sm86ProgramEmpty)))))))))))))))))))
def hmmaNativeZeroAccumulators =
(hmmaNativeCons
(constructor SM86InstructionBody SM86MoveImmediate hmmaNativeR16 hmmaNativeU0 sm86SafeControl)
(hmmaNativeCons
(constructor SM86InstructionBody SM86MoveImmediate hmmaNativeR17 hmmaNativeU0 sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR18
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR19
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR20
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR21
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR22
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR23
hmmaNativeU0
sm86SafeControl)
sm86ProgramEmpty))))))))
def hmmaNativeOperandA =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) .
(family SM86Program))
orientation
(branch
HMMAProductionAB
.
(hmmaNativeRowMajor
grouping
(matrixProductLeftPlaneWords orientation geometry layout)
hmmaNativeR4
hmmaNativeR24
hmmaNativeR25
(matrixProductLeftStride layout)
(byte-to-nat (byte 104))))
(branch
HMMAProductionABTransposed
.
(hmmaNativeRowMajor
grouping
(matrixProductLeftPlaneWords orientation geometry layout)
hmmaNativeR4
hmmaNativeR24
hmmaNativeR25
(matrixProductLeftStride layout)
(byte-to-nat (byte 104))))
(branch
HMMAProductionATransposedB
.
(hmmaNativeTransposedA
grouping
(matrixProductLeftPlaneWords orientation geometry layout)
hmmaNativeR4
hmmaNativeR24
hmmaNativeR30
hmmaNativeR25
(matrixProductLeftStride layout)
(byte-to-nat (byte 104)))))))))
def hmmaNativeOperandB =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) .
(family SM86Program))
orientation
(branch
HMMAProductionAB
.
(hmmaNativeReductionMajor
grouping
(matrixProductRightPlaneWords orientation geometry layout)
hmmaNativeR6
hmmaNativeR27
hmmaNativeR31
hmmaNativeR26
(matrixProductRightStride layout)
(byte-to-nat (byte 112))))
(branch
HMMAProductionABTransposed
.
(hmmaNativeRowMajor
grouping
(matrixProductRightPlaneWords orientation geometry layout)
hmmaNativeR6
hmmaNativeR27
hmmaNativeR26
(matrixProductRightStride layout)
(byte-to-nat (byte 112))))
(branch
HMMAProductionATransposedB
.
(hmmaNativeReductionMajor
grouping
(matrixProductRightPlaneWords orientation geometry layout)
hmmaNativeR6
hmmaNativeR27
hmmaNativeR31
hmmaNativeR26
(matrixProductRightStride layout)
(byte-to-nat (byte 112)))))))))
def hmmaNativePrefix =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveConstant
hmmaNativeR1
(byte 0)
(hmmaNativeU32 (byte-to-nat (byte 40)))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
hmmaNativeR0
(constructor SM86SpecialRegister SM86ThreadIdX)
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
hmmaNativeR25
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
hmmaNativeR26
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
hmmaNativeSet2)
(hmmaNativeAppend
(hmmaNativeGroupSpecial hmmaNativeR2 hmmaNativeSet3 grouping)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR3
hmmaNativeU4
(hmmaNativeSpecialWait grouping))
(hmmaNativeAppend
(hmmaNativeOperandA grouping orientation geometry layout)
(hmmaNativeAppend
(hmmaNativeOperandB grouping orientation geometry layout)
(hmmaNativeAppend (hmmaNativeFragmentAddresses orientation) hmmaNativeZeroAccumulators)))))))))))))
def hmmaNativeReductionPointer =
(lambda unrestricted pointer : (family SM86Register) .
(lambda unrestricted baseIndex : (family SM86Register) .
(lambda unrestricted stageWords : Nat .
(lambda unrestricted constantOffset : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerAddThreeImmediate
hmmaNativeR2
baseIndex
(hmmaNativeU32 stageWords)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddWideConstant
pointer
hmmaNativeR2
hmmaNativeR3
(byte 0)
(hmmaNativeU32 constantOffset)
sm86SafeControl)
sm86ProgramEmpty))))))
def hmmaNativePointerSetup =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(lambda unrestricted chunk : Nat .
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) .
(family SM86Program))
orientation
(branch
HMMAProductionAB
.
(hmmaNativeReductionPointer
hmmaNativeR6
hmmaNativeR31
(naturalMultiply chunk (naturalMultiply hmmaNativeN16 (matrixProductRightStride layout)))
(byte-to-nat (byte 112))))
(branch HMMAProductionABTransposed . sm86ProgramEmpty)
(branch
HMMAProductionATransposedB
.
(hmmaNativeAppend
(hmmaNativeReductionPointer
hmmaNativeR4
hmmaNativeR30
(naturalMultiply
chunk
(naturalMultiply hmmaNativeN16 (matrixProductLeftStride layout)))
(byte-to-nat (byte 104)))
(hmmaNativeReductionPointer
hmmaNativeR6
hmmaNativeR31
(naturalMultiply
chunk
(naturalMultiply hmmaNativeN16 (matrixProductRightStride layout)))
(byte-to-nat (byte 112))))))))))
def hmmaNativeStageA =
(lambda unrestricted byteOffset : Nat .
(lambda unrestricted bankOffset : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadGlobalWide
hmmaNativeR8
hmmaNativeR4
(hmmaNativeU32 byteOffset)
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR8
(hmmaNativeU32 bankOffset)
hmmaNativeWait0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR9
(hmmaNativeU32 (naturalAdd bankOffset hmmaNativeN4))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR10
(hmmaNativeU32 (naturalAdd bankOffset hmmaNativeN8))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR11
(hmmaNativeU32 (naturalAdd bankOffset (byte-to-nat (byte 12))))
sm86SafeControl)
sm86ProgramEmpty)))))))
def hmmaNativeStageAT =
(lambda unrestricted rowStride : Nat .
(lambda unrestricted bankOffset : Nat .
(app
(lambda unrestricted m : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadGlobal
hmmaNativeR8
hmmaNativeR4
hmmaNativeU0
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadGlobal
hmmaNativeR9
hmmaNativeR4
(hmmaNativeU32 (naturalMultiply m hmmaNativeN2))
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR12
hmmaNativeR8
(constructor SM86HalfSelector SM86LowHalf)
hmmaNativeWait0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR13
hmmaNativeR9
(constructor SM86HalfSelector SM86LowHalf)
hmmaNativeWait1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR14
hmmaNativeR8
(constructor SM86HalfSelector SM86HighHalf)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR15
hmmaNativeR9
(constructor SM86HalfSelector SM86HighHalf)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86FloatPairToPackedHalfPair
hmmaNativeR8
hmmaNativeR13
hmmaNativeR12
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86FloatPairToPackedHalfPair
hmmaNativeR9
hmmaNativeR15
hmmaNativeR14
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR8
(hmmaNativeU32 bankOffset)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR9
(hmmaNativeU32 (naturalAdd bankOffset hmmaNativeN64))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadGlobal
hmmaNativeR10
hmmaNativeR4
(hmmaNativeU32 (naturalMultiply m hmmaNativeN4))
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadGlobal
hmmaNativeR11
hmmaNativeR4
(hmmaNativeU32 (naturalMultiply m hmmaNativeN6))
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR12
hmmaNativeR10
(constructor SM86HalfSelector SM86LowHalf)
hmmaNativeWait0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR13
hmmaNativeR11
(constructor SM86HalfSelector SM86LowHalf)
hmmaNativeWait1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR14
hmmaNativeR10
(constructor SM86HalfSelector SM86HighHalf)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86HalfToFloat
hmmaNativeR15
hmmaNativeR11
(constructor SM86HalfSelector SM86HighHalf)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86FloatPairToPackedHalfPair
hmmaNativeR10
hmmaNativeR13
hmmaNativeR12
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86FloatPairToPackedHalfPair
hmmaNativeR11
hmmaNativeR15
hmmaNativeR14
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR10
(hmmaNativeU32 (naturalAdd bankOffset hmmaNativeN4))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR24
hmmaNativeR11
(hmmaNativeU32
(naturalAdd
bankOffset
(naturalAdd hmmaNativeN64 hmmaNativeN4)))
sm86SafeControl)
sm86ProgramEmpty)))))))))))))))))))))
rowStride)))
def hmmaNativeStageB =
(lambda unrestricted byteOffset : Nat .
(lambda unrestricted bankOffset : Nat .
(app
(lambda unrestricted bBase : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadGlobalWide
hmmaNativeR12
hmmaNativeR6
(hmmaNativeU32 byteOffset)
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR27
hmmaNativeR12
(hmmaNativeU32 bBase)
hmmaNativeWait1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR27
hmmaNativeR13
(hmmaNativeU32 (naturalAdd bBase hmmaNativeN4))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR27
hmmaNativeR14
(hmmaNativeU32 (naturalAdd bBase hmmaNativeN8))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreShared
hmmaNativeR27
hmmaNativeR15
(hmmaNativeU32 (naturalAdd bBase (byte-to-nat (byte 12))))
sm86SafeControl)
sm86ProgramEmpty))))))
(naturalAdd bankOffset (naturalMultiply hmmaNativeN4 hmmaNativeN512)))))
def hmmaNativeStageCommon =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted bank : Nat .
(app
(lambda unrestricted transposeB : (family SM86SharedMatrixTranspose) .
(hmmaNativeCons
(constructor SM86InstructionBody SM86BarrierSynchronize sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaNativeR8
hmmaNativeR28
(hmmaNativeU32 bank)
(constructor SM86SharedMatrixCount SM86SharedMatrix4)
(constructor SM86SharedMatrixTranspose SM86SharedMatrixNotTransposed)
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaNativeR12
hmmaNativeR29
(hmmaNativeU32 bank)
(constructor SM86SharedMatrixCount SM86SharedMatrix2)
transposeB
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaNativeR14
hmmaNativeR29
(hmmaNativeU32 (naturalAdd bank (hmmaNativeBNextColumns orientation)))
(constructor SM86SharedMatrixCount SM86SharedMatrix2)
transposeB
hmmaNativeSet2)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
hmmaNativeR16
hmmaNativeR8
hmmaNativeR12
hmmaNativeR16
hmmaNativeWaitFragments0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
hmmaNativeR20
hmmaNativeR8
hmmaNativeR14
hmmaNativeR20
hmmaNativeWaitFragments1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaNativeR8
hmmaNativeR28
(hmmaNativeU32 (naturalAdd bank hmmaNativeN32))
(constructor SM86SharedMatrixCount SM86SharedMatrix4)
(constructor SM86SharedMatrixTranspose SM86SharedMatrixNotTransposed)
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaNativeR12
hmmaNativeR29
(hmmaNativeU32 (naturalAdd bank (hmmaNativeBNextReduction orientation)))
(constructor SM86SharedMatrixCount SM86SharedMatrix2)
transposeB
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86LoadSharedMatrix
hmmaNativeR14
hmmaNativeR29
(hmmaNativeU32
(naturalAdd bank (naturalAdd (hmmaNativeBNextColumns orientation) (hmmaNativeBNextReduction orientation))))
(constructor SM86SharedMatrixCount SM86SharedMatrix2)
transposeB
hmmaNativeSet2)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
hmmaNativeR16
hmmaNativeR8
hmmaNativeR12
hmmaNativeR16
hmmaNativeWaitFragments0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
hmmaNativeR20
hmmaNativeR8
hmmaNativeR14
hmmaNativeR20
hmmaNativeWaitFragments1)
sm86ProgramEmpty))))))))))))
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) .
(family SM86SharedMatrixTranspose))
orientation
(branch
HMMAProductionAB
.
(constructor SM86SharedMatrixTranspose SM86SharedMatrixTransposed))
(branch
HMMAProductionABTransposed
.
(constructor SM86SharedMatrixTranspose SM86SharedMatrixNotTransposed))
(branch
HMMAProductionATransposedB
.
(constructor SM86SharedMatrixTranspose SM86SharedMatrixTransposed))))))
def hmmaNativeN4096 =
(naturalMultiply hmmaNativeN8 hmmaNativeN512)
def hmmaNativeStage =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(lambda unrestricted chunk : Nat .
(app
(lambda unrestricted byteOffset : Nat .
(app
(lambda unrestricted bank : Nat .
(hmmaNativeAppend
(hmmaNativePointerSetup orientation geometry layout chunk)
(hmmaNativeAppend
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) .
(family SM86Program))
orientation
(branch
HMMAProductionAB
.
(hmmaNativeAppend
(hmmaNativeStageA byteOffset bank)
(hmmaNativeStageB zero bank)))
(branch
HMMAProductionABTransposed
.
(hmmaNativeAppend
(hmmaNativeStageA byteOffset bank)
(hmmaNativeStageB byteOffset bank)))
(branch
HMMAProductionATransposedB
.
(hmmaNativeAppend
(hmmaNativeStageAT (matrixProductLeftStride layout) bank)
(hmmaNativeStageB zero bank))))
(hmmaNativeStageCommon orientation bank))))
(naturalMultiply (naturalModuloUnchecked chunk hmmaNativeN2) hmmaNativeN4096)))
(naturalMultiply chunk hmmaNativeN64))))))
def hmmaNativeStages =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(nat-eliminate
(lambda unrestricted current : Nat . (family SM86Program))
sm86ProgramEmpty
(lambda unrestricted predecessor : Nat .
(lambda unrestricted induction : (family SM86Program) .
(hmmaNativeAppend induction (hmmaNativeStage orientation geometry layout predecessor))))
(hmmaNativeGeometryStages geometry)))))
def hmmaNativeSuffix =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(app
(lambda unrestricted n : Nat .
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
hmmaNativeR0
(constructor SM86SpecialRegister SM86ThreadIdX)
hmmaNativeSet0)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
hmmaNativeR2
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
hmmaNativeSet1)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86SpecialToRegister
hmmaNativeR3
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
hmmaNativeSet2)
(hmmaNativeAppend
(hmmaNativeGroupSpecial hmmaNativeR13 hmmaNativeSet3 grouping)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86MoveImmediate
hmmaNativeR15
hmmaNativeU4
(hmmaNativeSpecialWait grouping))
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR4
hmmaNativeR0
(byte 5)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR5
hmmaNativeR4
(byte 1)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR6
hmmaNativeR5
hmmaNativeUNeg2
hmmaNativeR4
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR7
hmmaNativeR4
hmmaNativeUNeg32
hmmaNativeR0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
hmmaNativeR8
hmmaNativeR7
(byte 2)
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR9
hmmaNativeR8
hmmaNativeUNeg4
hmmaNativeR7
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR10
hmmaNativeR2
(hmmaNativeU32 hmmaNativeN32)
hmmaNativeR8
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR10
hmmaNativeR5
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR10
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR11
hmmaNativeR3
(hmmaNativeU32 hmmaNativeN32)
sm86ZeroRegister
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR11
hmmaNativeR6
(hmmaNativeU32 hmmaNativeN16)
hmmaNativeR11
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR11
hmmaNativeR9
(hmmaNativeU32 hmmaNativeN2)
hmmaNativeR11
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
hmmaNativeR12
hmmaNativeR10
(hmmaNativeU32 n)
hmmaNativeR11
sm86SafeControl)
(hmmaNativeAppend
(hmmaNativeGroupOffsetWith
hmmaNativeR13
grouping
(matrixProductOutputPlaneWords geometry layout)
hmmaNativeR12)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddWideConstant
hmmaNativeR4
hmmaNativeR12
hmmaNativeR15
(byte 0)
(hmmaNativeU32 (byte-to-nat (byte 96)))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreGlobal64
hmmaNativeR4
hmmaNativeR16
hmmaNativeU0
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreGlobal64
hmmaNativeR4
hmmaNativeR18
(hmmaNativeU32 (naturalMultiply n hmmaNativeN32))
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreGlobal64
hmmaNativeR4
hmmaNativeR20
hmmaNativeU32Value
sm86SafeControl)
(hmmaNativeCons
(constructor
SM86InstructionBody
SM86StoreGlobal64
hmmaNativeR4
hmmaNativeR22
(hmmaNativeU32
(naturalAdd (naturalMultiply n hmmaNativeN32) hmmaNativeN32))
sm86SafeControl)
(hmmaNativeCons
(constructor SM86InstructionBody SM86Exit sm86SafeControl)
sm86ProgramEmpty)))))))))))))))))))))))))
(matrixProductOutputStride layout)))))
-- where the native program reads its three pointers (output, then the two
-- operands) in its parameter block; a launch of it places them there
def hmmaNativeOutputOffset : Nat = 0x60
def hmmaNativeLeftOffset : Nat = 0x68
def hmmaNativeRightOffset : Nat = 0x70
def matrixProductProgramWithRowStrides =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(hmmaNativeAppend (hmmaNativePrefix grouping orientation geometry layout)
(hmmaNativeAppend (hmmaNativeStages orientation geometry layout)
(hmmaNativeSuffix grouping geometry layout)))))))
-- Contiguous matrices are the default view of the same implementation.
def hmmaNativeProgram =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(matrixProductProgramWithRowStrides grouping orientation geometry
(matrixProductContiguousRowStrides orientation geometry)))))
def hmmaNativeBaseInstructions =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat)
orientation
(branch HMMAProductionAB . hmmaNativeN70)
(branch HMMAProductionABTransposed . hmmaNativeN69)
(branch HMMAProductionATransposedB . hmmaNativeN70)))
def hmmaNativeStageInstructions =
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat)
orientation
(branch HMMAProductionAB . hmmaNativeN23)
(branch HMMAProductionABTransposed . hmmaNativeN21)
(branch HMMAProductionATransposedB . hmmaNativeN40)))
def hmmaNativeExpectedInstructions =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(naturalAdd
(naturalAdd
(hmmaNativeBaseInstructions orientation)
(naturalMultiply
(hmmaNativeStageInstructions orientation)
(hmmaNativeGeometryStages geometry)))
(naturalMultiply hmmaNativeN5 (hmmaNativeIsGrouped grouping))))))
def hmmaNativeGeometryValid =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(app
(lambda unrestricted m : Nat .
(app
(lambda unrestricted n : Nat .
(app
(lambda unrestricted k : Nat .
(app
(lambda unrestricted stages : Nat .
(naturalAnd
(naturalNonzero m)
(naturalAnd
(naturalNonzero n)
(naturalAnd
(naturalNonzero k)
(naturalAnd
(naturalIsZero (naturalModuloUnchecked m hmmaNativeN32))
(naturalAnd
(naturalIsZero (naturalModuloUnchecked n hmmaNativeN32))
(naturalAnd
(naturalIsZero (naturalModuloUnchecked k hmmaNativeN32))
(naturalAnd
(naturalEqual stages (naturalDivideUnchecked k hmmaNativeN32))
(naturalAnd
(naturalFitsWord32 m)
(naturalAnd
(naturalFitsWord32 n)
(naturalAnd
(naturalFitsWord32 k)
(naturalFitsWord32 (naturalMultiply n hmmaNativeN32)))))))))))))
(hmmaNativeGeometryStages geometry)))
(hmmaNativeGeometryK geometry)))
(hmmaNativeGeometryN geometry)))
(hmmaNativeGeometryM geometry)))
def hmmaNativeGroupingValid =
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(eliminate
HMMAProductionSM86Grouping
(lambda unrestricted current : (family HMMAProductionSM86Grouping) . Nat)
grouping
(branch HMMAProductionDense . (succ zero))
(branch
HMMAProductionGrouped
groups
.
(naturalAnd
(naturalNonzero groups)
(naturalAnd
(naturalLess groups hmmaNativeN65536)
(naturalAnd
(naturalAtMostTwoPower32 (naturalMultiply groups (hmmaNativeGroupAWords geometry)))
(naturalAnd
(naturalAtMostTwoPower32
(naturalMultiply groups (hmmaNativeGroupBWords geometry)))
(naturalAtMostTwoPower32
(naturalMultiply groups (hmmaNativeGroupCWords geometry)))))))))))
-- Strided views retain the vector-load/store alignment and bounded 32-bit
-- word indexing of this tile. Guard the instruction displacements as well
-- as whole-plane indexing. Buffer ownership and base-address alignment remain
-- the launch plan's obligations; these extents include stored row padding.
def matrixProductPlaneGroupsFit = (lambda unrestricted groups : Nat . (lambda unrestricted plane : Nat .
(naturalAnd (naturalNonzero plane)
(naturalLessOrEqual groups
(naturalDivideUnchecked (naturalSaturatingSubtract (naturalPowerOfTwo 32) 1)
(naturalSelect (naturalNonzero plane) plane 1))))))
def matrixProductRowStridesAdmitted = (lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted layout : (family MatrixProductRowStrides) .
(let unrestricted contiguous = (matrixProductContiguousRowStrides orientation geometry) in
(let unrestricted left = (matrixProductLeftStride layout) in
(let unrestricted right = (matrixProductRightStride layout) in
(let unrestricted output = (matrixProductOutputStride layout) in
(let unrestricted groups = (hmmaNativeGroupingCount grouping) in
(naturalAnd (hmmaNativeGeometryValid geometry)
(naturalAnd (naturalAnd (naturalNonzero groups) (naturalLess groups hmmaNativeN65536))
(naturalAnd (naturalLessOrEqual (matrixProductLeftStride contiguous) left)
(naturalAnd (naturalLessOrEqual (matrixProductRightStride contiguous) right)
(naturalAnd (naturalLessOrEqual (hmmaNativeGeometryN geometry) output)
(naturalAnd (naturalFitsWord32 left)
(naturalAnd (naturalFitsWord32 right)
(naturalAnd (naturalFitsWord32 output)
(naturalAnd (naturalIsZero (naturalModuloUnchecked left 8))
(naturalAnd (naturalIsZero (naturalModuloUnchecked right 8))
(naturalAnd (naturalIsZero (naturalModuloUnchecked output 2))
(naturalAnd (naturalLess (naturalMultiply left 6) (naturalPowerOfTwo 23))
(naturalAnd (naturalLess (naturalAdd (naturalMultiply output 32) 32) (naturalPowerOfTwo 23))
(naturalAnd (naturalLess (naturalMultiply (hmmaNativeGeometryK geometry) 2) (naturalPowerOfTwo 23))
(naturalAnd (matrixProductPlaneGroupsFit groups (matrixProductLeftPlaneWords orientation geometry layout))
(naturalAnd (matrixProductPlaneGroupsFit groups (matrixProductRightPlaneWords orientation geometry layout))
(matrixProductPlaneGroupsFit groups (matrixProductOutputPlaneWords geometry layout)))))))))))))))))))))))))))
def hmmaNativeManifest =
(lambda unrestricted variant : (family HMMAProductionSM86Variant) .
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(app
(lambda unrestricted instructions : Nat .
(constructor
HMMAProductionSM86Manifest
HMMAProductionSM86ManifestValue
variant
orientation
grouping
geometry
instructions
(naturalMultiply instructions hmmaNativeN16)
hmmaNativeN40
hmmaNativeN8192
(naturalDivideUnchecked (hmmaNativeGeometryM geometry) hmmaNativeN32)
(naturalDivideUnchecked (hmmaNativeGeometryN geometry) hmmaNativeN32)
(hmmaNativeGroupingCount grouping)
hmmaNativeN128
(constructor
HMMAProductionSM86ABI
HMMAProductionSM86ABIValue
hmmaNativeOutputOffset
hmmaNativeLeftOffset
hmmaNativeRightOffset)
(constructor
HMMAProductionSM86Extents
HMMAProductionSM86ExtentsValue
(hmmaNativeGeometryM geometry)
(hmmaNativeGeometryN geometry)
(hmmaNativeGeometryK geometry)
(hmmaNativeGroupingCount grouping))
zero))
(hmmaNativeExpectedInstructions grouping orientation geometry))))))
def hmmaNativeExpectedManifestInstructions =
(lambda unrestricted manifest : (family HMMAProductionSM86Manifest) .
(eliminate
HMMAProductionSM86Manifest
(lambda unrestricted current : (family HMMAProductionSM86Manifest) . Nat)
manifest
(branch
HMMAProductionSM86ManifestValue
variant
orientation
grouping
geometry
instructions
bytes
registers
shared
gridX
gridY
gridZ
block
abi
extents
fallback
.
instructions)))
def hmmaNativeExpectedManifestBytes =
(lambda unrestricted manifest : (family HMMAProductionSM86Manifest) .
(eliminate
HMMAProductionSM86Manifest
(lambda unrestricted current : (family HMMAProductionSM86Manifest) . Nat)
manifest
(branch
HMMAProductionSM86ManifestValue
variant
orientation
grouping
geometry
instructions
bytes
registers
shared
gridX
gridY
gridZ
block
abi
extents
fallback
.
bytes)))
def hmmaNativeTelemetry =
(lambda unrestricted manifest : (family HMMAProductionSM86Manifest) .
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted observedInstructions : Nat .
(lambda unrestricted observedBytes : Nat .
(app
(lambda unrestricted stages : Nat .
(constructor
HMMAProductionNativeTelemetry
HMMAProductionNativeTelemetryValue
manifest
observedInstructions
observedBytes
stages
(naturalMultiply (naturalPredecessor stages) hmmaNativeN64)
(naturalMultiply
stages
(eliminate
HMMAProductionSM86Orientation
(lambda unrestricted current : (family HMMAProductionSM86Orientation) . Nat)
orientation
(branch HMMAProductionAB . hmmaNativeN2)
(branch HMMAProductionABTransposed . hmmaNativeN2)
(branch HMMAProductionATransposedB . hmmaNativeN5)))
hmmaNativeN4
(naturalMultiply stages hmmaNativeN8)
(naturalMultiply stages hmmaNativeN6)
(naturalMultiply stages hmmaNativeN4)
stages
(naturalMultiply hmmaNativeN5 (hmmaNativeIsGrouped grouping))
zero))
(hmmaNativeGeometryStages geometry))))))))
def hmmaNativeFail =
(lambda unrestricted code : (family HMMAProductionSM86FailureCode) .
(lambda unrestricted telemetry : (family HMMAProductionNativeTelemetry) .
(constructor
HMMAProductionNativeBuildResult
HMMAProductionNativeContractFailed
code
telemetry)))
def hmmaNativeEncode =
(lambda unrestricted manifest : (family HMMAProductionSM86Manifest) .
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(lambda unrestricted program : (family SM86Program) .
(lambda unrestricted observedInstructions : Nat .
(app
(lambda unrestricted encoding : (family SM86ProgramEncodingResult) .
(eliminate
SM86ProgramEncodingResult
(lambda unrestricted current : (family SM86ProgramEncodingResult) .
(family HMMAProductionNativeBuildResult))
encoding
(branch
SM86ProgramEncodingSucceeded
encoded
encodingTelemetry
.
(app
(lambda unrestricted observedBytes : Nat .
(nat-eliminate
(lambda unrestricted bytesValid : Nat .
(family HMMAProductionNativeBuildResult))
(hmmaNativeFail
(constructor
HMMAProductionSM86FailureCode
HMMAProductionEncodedByteCountMismatch)
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
observedBytes))
(lambda unrestricted bytesPredecessor : Nat .
(lambda unrestricted bytesInduction : (family HMMAProductionNativeBuildResult) .
(app
(lambda unrestricted identityResult : (family SHA256HexResult) .
(eliminate
SHA256HexResult
(lambda unrestricted current : (family SHA256HexResult) .
(family HMMAProductionNativeBuildResult))
identityResult
(branch
SHA256HexSucceeded
identity
identityTelemetry
.
(nat-eliminate
(lambda unrestricted digestValid : Nat .
(family HMMAProductionNativeBuildResult))
(hmmaNativeFail
(constructor
HMMAProductionSM86FailureCode
HMMAProductionIdentityLengthInvalid)
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
observedBytes))
(lambda unrestricted digestPredecessor : Nat .
(lambda unrestricted digestInduction : (family HMMAProductionNativeBuildResult) .
(constructor
HMMAProductionNativeBuildResult
HMMAProductionNativeBuildSucceeded
encoded
identity
encodingTelemetry
identityTelemetry
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
observedBytes))))
(naturalEqual (bytes-length identity) hmmaNativeN64)))
(branch
SHA256HexFailed
error
ordinal
identityTelemetry
.
(constructor
HMMAProductionNativeBuildResult
HMMAProductionNativeIdentityFailed
identityResult
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
observedBytes)))))
(sha256Hex encoded))))
(naturalEqual observedBytes (hmmaNativeExpectedManifestBytes manifest))))
(bytes-length encoded)))
(branch
SM86ProgramEncodingFailed
ordinal
failure
encodingTelemetry
.
(constructor
HMMAProductionNativeBuildResult
HMMAProductionNativeEncodingFailed
encoding
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
zero)))))
(sm86EncodeProgram program))))))))
def hmmaNativeBuild =
(lambda unrestricted variant : (family HMMAProductionSM86Variant) .
(lambda unrestricted grouping : (family HMMAProductionSM86Grouping) .
(lambda unrestricted orientation : (family HMMAProductionSM86Orientation) .
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(app
(lambda unrestricted manifest : (family HMMAProductionSM86Manifest) .
(app
(lambda unrestricted program : (family SM86Program) .
(app
(lambda unrestricted observedInstructions : Nat .
(nat-eliminate
(lambda unrestricted geometryValid : Nat .
(family HMMAProductionNativeBuildResult))
(hmmaNativeFail
(constructor HMMAProductionSM86FailureCode HMMAProductionInvalidGeometry)
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
zero))
(lambda unrestricted geometryPredecessor : Nat .
(lambda unrestricted geometryInduction : (family HMMAProductionNativeBuildResult) .
(nat-eliminate
(lambda unrestricted groupingValid : Nat .
(family HMMAProductionNativeBuildResult))
(hmmaNativeFail
(constructor
HMMAProductionSM86FailureCode
HMMAProductionInvalidGrouping)
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
zero))
(lambda unrestricted groupingPredecessor : Nat .
(lambda unrestricted groupingInduction : (family HMMAProductionNativeBuildResult) .
(nat-eliminate
(lambda unrestricted countValid : Nat .
(family HMMAProductionNativeBuildResult))
(hmmaNativeFail
(constructor
HMMAProductionSM86FailureCode
HMMAProductionInstructionCountMismatch)
(hmmaNativeTelemetry
manifest
grouping
orientation
geometry
observedInstructions
zero))
(lambda unrestricted countPredecessor : Nat .
(lambda unrestricted countInduction : (family HMMAProductionNativeBuildResult) .
(hmmaNativeEncode
manifest
grouping
orientation
geometry
program
observedInstructions)))
(naturalEqual
observedInstructions
(hmmaNativeExpectedManifestInstructions manifest)))))
(hmmaNativeGroupingValid grouping geometry))))
(hmmaNativeGeometryValid geometry)))
(sm86ProgramCount program)))
(hmmaNativeProgram grouping orientation geometry)))
(hmmaNativeManifest variant grouping orientation geometry))))))
def hmmaProductionDenseForwardNative =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(hmmaNativeBuild
(constructor HMMAProductionSM86Variant HMMAProductionDenseForward)
(constructor HMMAProductionSM86Grouping HMMAProductionDense)
(constructor HMMAProductionSM86Orientation HMMAProductionABTransposed)
geometry))
def hmmaProductionDenseInputGradientNative =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(hmmaNativeBuild
(constructor HMMAProductionSM86Variant HMMAProductionDenseReverseInput)
(constructor HMMAProductionSM86Grouping HMMAProductionDense)
(constructor HMMAProductionSM86Orientation HMMAProductionAB)
geometry))
def hmmaProductionDenseWeightGradientNative =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(hmmaNativeBuild
(constructor HMMAProductionSM86Variant HMMAProductionDenseReverseWeight)
(constructor HMMAProductionSM86Grouping HMMAProductionDense)
(constructor HMMAProductionSM86Orientation HMMAProductionATransposedB)
geometry))
def hmmaProductionRoutedForwardNative =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(hmmaNativeBuild
(constructor HMMAProductionSM86Variant HMMAProductionRoutedForward)
(constructor HMMAProductionSM86Grouping HMMAProductionGrouped hmmaNativeN16)
(constructor HMMAProductionSM86Orientation HMMAProductionABTransposed)
geometry))
def hmmaProductionRoutedInputGradientNative =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(hmmaNativeBuild
(constructor HMMAProductionSM86Variant HMMAProductionRoutedReverseInput)
(constructor HMMAProductionSM86Grouping HMMAProductionGrouped hmmaNativeN16)
(constructor HMMAProductionSM86Orientation HMMAProductionAB)
geometry))
def hmmaProductionRoutedWeightGradientNative =
(lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
(hmmaNativeBuild
(constructor HMMAProductionSM86Variant HMMAProductionRoutedReverseWeight)
(constructor HMMAProductionSM86Grouping HMMAProductionGrouped hmmaNativeN16)
(constructor HMMAProductionSM86Orientation HMMAProductionATransposedB)
geometry))The compiler supplied declaration spans and resolved links from this source snapshot. This page does not assert that this file belongs to a checked closure.