Source/Packages

Realization.Nvidia.SM86.HMMAProductionSM86

packages/realizations/cooperative/nvidia-sm86/src/Realization/Nvidia/SM86/HMMAProductionSM86.alpha

3,184 lines350 declarations131.1 KiBSHA-256 ca61a44ca043

Complete file · line 133

HMMAProductionSM86.alpha

Definition view

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.