module Realization.Nvidia.SM86.GEMM.HMMA.RegisterTileSM86 import Compiler.ApplicationBuilder import Realization.Nvidia.SM86.HMMAProductionSM86 import Std.Natural family RegisterTileHMMAGeometries : Type 0 constructor RegisterTileHMMAGeometriesEnd constructor RegisterTileHMMAGeometriesNext field unrestricted registerTileHMMAHead : (family HMMAProductionSM86Geometry) recursive unrestricted registerTileHMMATail end-family family RegisterTileHMMAErrorCode : Type 0 constructor RegisterTileHMMANoError constructor RegisterTileHMMAGeometryInvalid constructor RegisterTileHMMAInstructionCountMismatch constructor RegisterTileHMMAEncodedByteCountMismatch constructor RegisterTileHMMAEncodingFailed constructor RegisterTileHMMAIdentityFailed constructor RegisterTileHMMAHostFallbackForbidden constructor RegisterTileHMMAApplicationBuilderContractInvalid end-family family RegisterTileHMMATelemetry : Type 0 constructor RegisterTileHMMATelemetryValue field unrestricted registerTileHMMATelemetryM : Nat field unrestricted registerTileHMMATelemetryN : Nat field unrestricted registerTileHMMATelemetryK : Nat field unrestricted registerTileHMMATelemetryKStages : Nat field unrestricted registerTileHMMATelemetryGridX : Nat field unrestricted registerTileHMMATelemetryGridY : Nat field unrestricted registerTileHMMATelemetryGridZ : Nat field unrestricted registerTileHMMATelemetryBlockX : Nat field unrestricted registerTileHMMATelemetryInstructions : Nat field unrestricted registerTileHMMATelemetryEncodedBytes : Nat field unrestricted registerTileHMMATelemetryRegisters : Nat field unrestricted registerTileHMMATelemetrySharedBytes : Nat field unrestricted registerTileHMMATelemetryHostFallbackOperations : Nat field unrestricted registerTileHMMATelemetryApplicationBuilderContract : Nat end-family family RegisterTileHMMABuildResult : Type 0 constructor RegisterTileHMMABuildSucceeded field unrestricted registerTileHMMAEncodedBytes : Bytes field unrestricted registerTileHMMAImageSHA256 : Bytes field unrestricted registerTileHMMANativeResult : (family HMMAProductionNativeBuildResult) field unrestricted registerTileHMMABuildTelemetry : (family RegisterTileHMMATelemetry) constructor RegisterTileHMMABuildFailed field unrestricted registerTileHMMAFailureCode : (family RegisterTileHMMAErrorCode) field unrestricted registerTileHMMAFailureTelemetry : (family RegisterTileHMMATelemetry) end-family family RegisterTileHMMAGridResult : Type 0 constructor RegisterTileHMMAGridSucceeded field unrestricted registerTileHMMAGridX : Nat field unrestricted registerTileHMMAGridY : Nat field unrestricted registerTileHMMAGridZ : Nat field unrestricted registerTileHMMAGridTelemetry : (family RegisterTileHMMATelemetry) constructor RegisterTileHMMAGridFailed field unrestricted registerTileHMMAGridFailure : (family RegisterTileHMMAErrorCode) field unrestricted registerTileHMMAGridFailureTelemetry : (family RegisterTileHMMATelemetry) end-family def registerTileHMMAN8 = (byte-to-nat (byte 8)) def registerTileHMMAN16 = (byte-to-nat (byte 16)) def registerTileHMMAN21 = (byte-to-nat (byte 21)) def registerTileHMMAN32 = (byte-to-nat (byte 32)) def registerTileHMMAN40 = (byte-to-nat (byte 40)) def registerTileHMMAN64 = (byte-to-nat (byte 64)) def registerTileHMMAN69 = (byte-to-nat (byte 69)) def registerTileHMMAN128 = (naturalMultiply registerTileHMMAN16 registerTileHMMAN8) def registerTileHMMAN256 = (naturalMultiply registerTileHMMAN16 registerTileHMMAN16) def registerTileHMMAN384 = (naturalMultiply registerTileHMMAN16 (byte-to-nat (byte 24))) def registerTileHMMAN640 = (naturalMultiply registerTileHMMAN64 (byte-to-nat (byte 10))) def registerTileHMMAN1024 = (naturalMultiply registerTileHMMAN16 registerTileHMMAN64) def registerTileHMMAN3072 = (naturalMultiply (byte-to-nat (byte 48)) registerTileHMMAN64) def registerTileHMMAN4096 = (naturalMultiply registerTileHMMAN64 registerTileHMMAN64) def registerTileHMMAN6144 = (naturalMultiply (byte-to-nat (byte 96)) registerTileHMMAN64) def registerTileHMMAN8192 = (naturalMultiply registerTileHMMAN128 registerTileHMMAN64) def registerTileHMMAHostFallbackOperations : Nat = zero def registerTileHMMAErrorStableCode = (lambda unrestricted code : (family RegisterTileHMMAErrorCode) . (eliminate RegisterTileHMMAErrorCode (lambda unrestricted current : (family RegisterTileHMMAErrorCode) . Bytes) code (branch RegisterTileHMMANoError . b"HMMA-RT-000") (branch RegisterTileHMMAGeometryInvalid . b"HMMA-RT-001") (branch RegisterTileHMMAInstructionCountMismatch . b"HMMA-RT-002") (branch RegisterTileHMMAEncodedByteCountMismatch . b"HMMA-RT-003") (branch RegisterTileHMMAEncodingFailed . b"HMMA-RT-004") (branch RegisterTileHMMAIdentityFailed . b"HMMA-RT-005") (branch RegisterTileHMMAHostFallbackForbidden . b"HMMA-RT-006") (branch RegisterTileHMMAApplicationBuilderContractInvalid . b"HMMA-RT-007"))) def registerTileHMMAGeometry = (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat . (constructor HMMAProductionSM86Geometry HMMAProductionSM86GeometryValue m n k (naturalDivideUnchecked k registerTileHMMAN32))))) def registerTileHMMAGeometryEqual = (lambda unrestricted left : (family HMMAProductionSM86Geometry) . (lambda unrestricted right : (family HMMAProductionSM86Geometry) . (naturalAnd (naturalEqual (hmmaNativeGeometryM left) (hmmaNativeGeometryM right)) (naturalAnd (naturalEqual (hmmaNativeGeometryN left) (hmmaNativeGeometryN right)) (naturalAnd (naturalEqual (hmmaNativeGeometryK left) (hmmaNativeGeometryK right)) (naturalEqual (hmmaNativeGeometryStages left) (hmmaNativeGeometryStages right))))))) def registerTileHMMAInstructionCount = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (naturalAdd registerTileHMMAN69 (naturalMultiply registerTileHMMAN21 (hmmaNativeGeometryStages geometry)))) def registerTileHMMATelemetryFor = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (app (lambda unrestricted instructions : Nat . (constructor RegisterTileHMMATelemetry RegisterTileHMMATelemetryValue (hmmaNativeGeometryM geometry) (hmmaNativeGeometryN geometry) (hmmaNativeGeometryK geometry) (hmmaNativeGeometryStages geometry) (naturalDivideUnchecked (hmmaNativeGeometryM geometry) registerTileHMMAN32) (naturalDivideUnchecked (hmmaNativeGeometryN geometry) registerTileHMMAN32) (succ zero) registerTileHMMAN128 instructions (naturalMultiply instructions registerTileHMMAN16) registerTileHMMAN32 registerTileHMMAN8192 registerTileHMMAHostFallbackOperations applicationBuilderNativeOnly)) (registerTileHMMAInstructionCount geometry))) def registerTileHMMAFail = (lambda unrestricted code : (family RegisterTileHMMAErrorCode) . (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (constructor RegisterTileHMMABuildResult RegisterTileHMMABuildFailed code (registerTileHMMATelemetryFor geometry)))) def registerTileHMMAMapNativeResult = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (lambda unrestricted result : (family HMMAProductionNativeBuildResult) . (eliminate HMMAProductionNativeBuildResult (lambda unrestricted current : (family HMMAProductionNativeBuildResult) . (family RegisterTileHMMABuildResult)) result (branch HMMAProductionNativeBuildSucceeded encoded identity encodingTelemetry identityTelemetry nativeTelemetry . (constructor RegisterTileHMMABuildResult RegisterTileHMMABuildSucceeded encoded identity result (registerTileHMMATelemetryFor geometry))) (branch HMMAProductionNativeContractFailed failure nativeTelemetry . (eliminate HMMAProductionSM86FailureCode (lambda unrestricted current : (family HMMAProductionSM86FailureCode) . (family RegisterTileHMMABuildResult)) failure (branch HMMAProductionInvalidGeometry . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAGeometryInvalid) geometry)) (branch HMMAProductionInvalidGrouping . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAGeometryInvalid) geometry)) (branch HMMAProductionInstructionCountMismatch . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAInstructionCountMismatch) geometry)) (branch HMMAProductionEncodedByteCountMismatch . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAEncodedByteCountMismatch) geometry)) (branch HMMAProductionEncodingFailed . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAEncodingFailed) geometry)) (branch HMMAProductionIdentityFailed . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAIdentityFailed) geometry)) (branch HMMAProductionIdentityLengthInvalid . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAIdentityFailed) geometry)))) (branch HMMAProductionNativeEncodingFailed failure nativeTelemetry . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAEncodingFailed) geometry)) (branch HMMAProductionNativeIdentityFailed failure nativeTelemetry . (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAIdentityFailed) geometry))))) def registerTileHMMAFinalize = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (lambda unrestricted result : (family RegisterTileHMMABuildResult) . (nat-eliminate (lambda unrestricted builderValid : Nat . (family RegisterTileHMMABuildResult)) (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAApplicationBuilderContractInvalid) geometry) (lambda unrestricted builderPredecessor : Nat . (lambda unrestricted builderInduction : (family RegisterTileHMMABuildResult) . (nat-eliminate (lambda unrestricted fallbackFree : Nat . (family RegisterTileHMMABuildResult)) (registerTileHMMAFail (constructor RegisterTileHMMAErrorCode RegisterTileHMMAHostFallbackForbidden) geometry) (lambda unrestricted fallbackPredecessor : Nat . (lambda unrestricted fallbackInduction : (family RegisterTileHMMABuildResult) . result)) (naturalEqual registerTileHMMAHostFallbackOperations zero)))) applicationBuilderNativeOnly))) def emitRegisterTileHMMASM86 = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (registerTileHMMAFinalize geometry (registerTileHMMAMapNativeResult geometry (hmmaNativeBuild (constructor HMMAProductionSM86Variant HMMAProductionDenseForward) (constructor HMMAProductionSM86Grouping HMMAProductionDense) (constructor HMMAProductionSM86Orientation HMMAProductionABTransposed) geometry)))) def registerTileHMMAGrid = (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) . (nat-eliminate (lambda unrestricted valid : Nat . (family RegisterTileHMMAGridResult)) (constructor RegisterTileHMMAGridResult RegisterTileHMMAGridFailed (constructor RegisterTileHMMAErrorCode RegisterTileHMMAGeometryInvalid) (registerTileHMMATelemetryFor geometry)) (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (family RegisterTileHMMAGridResult) . (constructor RegisterTileHMMAGridResult RegisterTileHMMAGridSucceeded (naturalDivideUnchecked (hmmaNativeGeometryM geometry) registerTileHMMAN32) (naturalDivideUnchecked (hmmaNativeGeometryN geometry) registerTileHMMAN32) (succ zero) (registerTileHMMATelemetryFor geometry)))) (hmmaNativeGeometryValid geometry))) def registerTileHMMASupportedGeometries : (family RegisterTileHMMAGeometries) = (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN64 registerTileHMMAN1024) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN3072 registerTileHMMAN64) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN1024 registerTileHMMAN64) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN256 registerTileHMMAN1024) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN4096 registerTileHMMAN256) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN256 registerTileHMMAN256 registerTileHMMAN256) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN384 registerTileHMMAN640 registerTileHMMAN1024) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesNext (registerTileHMMAGeometry registerTileHMMAN384 registerTileHMMAN1024 registerTileHMMAN640) (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesEnd))))))))) def emitRegisterTileHMMAM256N256K256SM86 = (emitRegisterTileHMMASM86 (registerTileHMMAGeometry registerTileHMMAN256 registerTileHMMAN256 registerTileHMMAN256)) def registerTileHMMAM256N256K256InstructionCount = (registerTileHMMAInstructionCount (registerTileHMMAGeometry registerTileHMMAN256 registerTileHMMAN256 registerTileHMMAN256)) def registerTileHMMARegisterCount = registerTileHMMAN32 def registerTileHMMASharedBytes = registerTileHMMAN8192 def registerTileHMMAM256N256K256RegisterCount = registerTileHMMARegisterCount def registerTileHMMAM256N256K256SharedBytes = registerTileHMMASharedBytes def registerTileHMMAApplicationBuilderContract = applicationBuilderNativeOnly