module Realization.Nvidia.SM86.Cast.SM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.InstructionEncoding import Accelerator.SM86.Program import Accelerator.SM86.Types import Compiler.ApplicationBuilder import Data.SHA256Digest import Std.Natural import Std.Byte family CastDirection : Type 0 constructor Float32ToFloat16 constructor Float16ToFloat32 end-family family CastGeometry : Type 0 constructor CastGeometryValue field unrestricted castElementCount : Nat field unrestricted castBlockThreads : Nat end-family family CastSM86ErrorCode : Type 0 constructor CastSM86ElementCountZero constructor CastSM86ElementCountOdd constructor CastSM86BlockThreadsZero constructor CastSM86BlockThreadsTooLarge constructor CastSM86PairCountNotDivisible constructor CastSM86InstructionCountMismatch constructor CastSM86EncodingFailed constructor CastSM86ByteCountMismatch constructor CastSM86IdentityInvalid constructor CastSM86HostFallbackObserved end-family family CastSM86Telemetry : Type 0 constructor CastSM86TelemetryValue field unrestricted castTelemetryDirection : (family CastDirection) field unrestricted castTelemetryElements : Nat field unrestricted castTelemetryPairs : Nat field unrestricted castTelemetryBlockThreads : Nat field unrestricted castTelemetryGridX : Nat field unrestricted castTelemetryExpectedInstructions : Nat field unrestricted castTelemetryActualInstructions : Nat field unrestricted castTelemetryExpectedBytes : Nat field unrestricted castTelemetryActualBytes : Nat field unrestricted castTelemetryRegisters : Nat field unrestricted castTelemetryGlobalLoads : Nat field unrestricted castTelemetryGlobalStores : Nat field unrestricted castTelemetryConversions : Nat field unrestricted castTelemetryHostFallbackCalls : Nat field unrestricted castTelemetryApplicationBuilderContract : Nat end-family family CastSM86BuildResult : Type 0 constructor CastSM86BuildReady field unrestricted castReadyDirection : (family CastDirection) field unrestricted castReadyGeometry : (family CastGeometry) field unrestricted castReadyGridX : Nat field unrestricted castReadyProgram : (family SM86Program) field unrestricted castReadyImage : Bytes field unrestricted castReadySHA256 : Bytes field unrestricted castReadyTelemetry : (family CastSM86Telemetry) constructor CastSM86BuildRejected field unrestricted castRejectedCode : (family CastSM86ErrorCode) field unrestricted castRejectedOrdinal : Nat field unrestricted castRejectedDetail : Bytes end-family family CastGeometryValidation : Type 0 constructor CastGeometryAccepted field unrestricted castAcceptedPairs : Nat field unrestricted castAcceptedGridX : Nat constructor CastGeometryRejected field unrestricted castGeometryRejectedCode : (family CastSM86ErrorCode) end-family def castSM86N1 = (succ zero) def castSM86N2 = (byte-to-nat (byte 2)) def castSM86N13 = (byte-to-nat (byte 13)) def castSM86N14 = (byte-to-nat (byte 14)) def castSM86N16 = (byte-to-nat (byte 16)) def castSM86N24 = (byte-to-nat (byte 24)) def castSM86N1024 = (naturalMultiply byteNaturalTwoHundredFiftySix (byte-to-nat (byte 4))) -- The instruction stream reads the pair stride from the driver's CB0 block-X -- word and reads output/input pointers from the first two launch arguments. -- WholeProgramPlan owns that driver's word offset; callers patch it with the -- admitted block size before launching the image. def castSM86OutputArgument : Nat = 0 def castSM86InputArgument : Nat = 1 def castSM86ArgumentCount : Nat = 2 def castSM86PairsPerThread : Nat = 1 def castSM86BlockThreads : Nat = byteNaturalTwoHundredFiftySix def castSM86Word32 = (lambda unrestricted b0 : Byte . (lambda unrestricted b1 : Byte . (lambda unrestricted b2 : Byte . (lambda unrestricted b3 : Byte . (sm86Unsigned32 b0 b1 b2 b3))))) def castSM86Zero32 = (castSM86Word32 (byte 0) (byte 0) (byte 0) (byte 0)) def castSM86Two32 = (castSM86Word32 (byte 2) (byte 0) (byte 0) (byte 0)) def castSM86Four32 = (castSM86Word32 (byte 4) (byte 0) (byte 0) (byte 0)) def castSM86ConstantBase = (castSM86Word32 (byte 40) (byte 0) (byte 0) (byte 0)) -- Constant-bank word zero is the CTA stride in element pairs. A 256-thread -- launch therefore binds 256 here; leaving it zero repeats the first tile. def castSM86OutputPointer = (castSM86Word32 (byte 96) (byte 1) (byte 0) (byte 0)) def castSM86InputPointer = (castSM86Word32 (byte 104) (byte 1) (byte 0) (byte 0)) def castSM86Register = (lambda unrestricted index : Byte . (sm86Register index)) def castSM86Instruction = (lambda unrestricted body : (family SM86InstructionBody) . (constructor SM86Instruction SM86InstructionValue (constructor SM86InstructionGuard SM86InstructionAlways) body)) def castSM86Next = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (castSM86Instruction body) tail))) def castSM86End = (constructor SM86Program SM86ProgramEnd) def castSM86Set0 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0)) def castSM86Wait0 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0)) def castSM86Prefix = (lambda unrestricted tail : (family SM86Program) . (castSM86Next (constructor SM86InstructionBody SM86MoveConstant (castSM86Register (byte 1)) (byte 0) castSM86ConstantBase sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86SpecialToRegister (castSM86Register (byte 0)) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) castSM86Set0) (castSM86Next (constructor SM86InstructionBody SM86MoveImmediate (castSM86Register (byte 3)) castSM86Four32 castSM86Wait0) (castSM86Next (constructor SM86InstructionBody SM86SpecialToRegister (castSM86Register (byte 2)) (constructor SM86SpecialRegister SM86ThreadIdX) castSM86Set0) (castSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddConstant (castSM86Register (byte 0)) (castSM86Register (byte 0)) (byte 0) castSM86Zero32 (castSM86Register (byte 2)) castSM86Wait0) (castSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (castSM86Register (byte 4)) (castSM86Register (byte 0)) castSM86Two32 sm86ZeroRegister sm86SafeControl) tail))))))) def castSM86F32ToF16Tail : (family SM86Program) = (castSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (castSM86Register (byte 6)) (castSM86Register (byte 4)) (castSM86Register (byte 3)) (byte 0) castSM86InputPointer sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86LoadGlobal (castSM86Register (byte 8)) (castSM86Register (byte 6)) castSM86Zero32 castSM86Set0) (castSM86Next (constructor SM86InstructionBody SM86LoadGlobal (castSM86Register (byte 9)) (castSM86Register (byte 6)) castSM86Four32 castSM86Set0) (castSM86Next (constructor SM86InstructionBody SM86FloatPairToPackedHalfPair (castSM86Register (byte 10)) (castSM86Register (byte 9)) (castSM86Register (byte 8)) castSM86Wait0) (castSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (castSM86Register (byte 12)) (castSM86Register (byte 0)) (castSM86Register (byte 3)) (byte 0) castSM86OutputPointer sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86StoreGlobal (castSM86Register (byte 12)) (castSM86Register (byte 10)) castSM86Zero32 sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) castSM86End))))))) def castSM86F16ToF32Tail : (family SM86Program) = (castSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (castSM86Register (byte 6)) (castSM86Register (byte 0)) (castSM86Register (byte 3)) (byte 0) castSM86InputPointer sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86LoadGlobal (castSM86Register (byte 10)) (castSM86Register (byte 6)) castSM86Zero32 castSM86Set0) (castSM86Next (constructor SM86InstructionBody SM86HalfToFloat (castSM86Register (byte 14)) (castSM86Register (byte 10)) (constructor SM86HalfSelector SM86LowHalf) castSM86Wait0) (castSM86Next (constructor SM86InstructionBody SM86HalfToFloat (castSM86Register (byte 15)) (castSM86Register (byte 10)) (constructor SM86HalfSelector SM86HighHalf) sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (castSM86Register (byte 12)) (castSM86Register (byte 4)) (castSM86Register (byte 3)) (byte 0) castSM86OutputPointer sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86StoreGlobal (castSM86Register (byte 12)) (castSM86Register (byte 14)) castSM86Zero32 sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86StoreGlobal (castSM86Register (byte 12)) (castSM86Register (byte 15)) castSM86Four32 sm86SafeControl) (castSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) castSM86End)))))))) def emitCastSM86 = (lambda unrestricted direction : (family CastDirection) . (eliminate CastDirection (lambda unrestricted current : (family CastDirection) . (family SM86Program)) direction (branch Float32ToFloat16 . (castSM86Prefix castSM86F32ToF16Tail)) (branch Float16ToFloat32 . (castSM86Prefix castSM86F16ToF32Tail)))) def castInstructionCount = (lambda unrestricted direction : (family CastDirection) . (eliminate CastDirection (lambda unrestricted current : (family CastDirection) . Nat) direction (branch Float32ToFloat16 . castSM86N13) (branch Float16ToFloat32 . castSM86N14))) def castRegisterCount : Nat = castSM86N24 def castSM86HostFallbackCalls : Nat = zero def castSM86ApplicationBuilderContract : Nat = Compiler.ApplicationBuilder/applicationBuilderNativeOnly def castSM86ErrorStableCode = (lambda unrestricted code : (family CastSM86ErrorCode) . (eliminate CastSM86ErrorCode (lambda unrestricted current : (family CastSM86ErrorCode) . Bytes) code (branch CastSM86ElementCountZero . b"CAST-001") (branch CastSM86ElementCountOdd . b"CAST-002") (branch CastSM86BlockThreadsZero . b"CAST-003") (branch CastSM86BlockThreadsTooLarge . b"CAST-004") (branch CastSM86PairCountNotDivisible . b"CAST-005") (branch CastSM86InstructionCountMismatch . b"CAST-006") (branch CastSM86EncodingFailed . b"CAST-007") (branch CastSM86ByteCountMismatch . b"CAST-008") (branch CastSM86IdentityInvalid . b"CAST-009") (branch CastSM86HostFallbackObserved . b"CAST-010"))) def validateCastGeometry = (lambda unrestricted geometry : (family CastGeometry) . (eliminate CastGeometry (lambda unrestricted current : (family CastGeometry) . (family CastGeometryValidation)) geometry (branch CastGeometryValue elements blockThreads . (nat-eliminate (lambda unrestricted elementsNonzero : Nat . (family CastGeometryValidation)) (constructor CastGeometryValidation CastGeometryRejected (constructor CastSM86ErrorCode CastSM86ElementCountZero)) (lambda unrestricted elementPredecessor : Nat . (lambda unrestricted elementInduction : (family CastGeometryValidation) . (nat-eliminate (lambda unrestricted evenElements : Nat . (family CastGeometryValidation)) (constructor CastGeometryValidation CastGeometryRejected (constructor CastSM86ErrorCode CastSM86ElementCountOdd)) (lambda unrestricted evenPredecessor : Nat . (lambda unrestricted evenInduction : (family CastGeometryValidation) . (nat-eliminate (lambda unrestricted blockNonzero : Nat . (family CastGeometryValidation)) (constructor CastGeometryValidation CastGeometryRejected (constructor CastSM86ErrorCode CastSM86BlockThreadsZero)) (lambda unrestricted blockPredecessor : Nat . (lambda unrestricted blockInduction : (family CastGeometryValidation) . (nat-eliminate (lambda unrestricted blockFits : Nat . (family CastGeometryValidation)) (constructor CastGeometryValidation CastGeometryRejected (constructor CastSM86ErrorCode CastSM86BlockThreadsTooLarge)) (lambda unrestricted fitsPredecessor : Nat . (lambda unrestricted fitsInduction : (family CastGeometryValidation) . (app (lambda unrestricted pairs : Nat . (nat-eliminate (lambda unrestricted divisible : Nat . (family CastGeometryValidation)) (constructor CastGeometryValidation CastGeometryRejected (constructor CastSM86ErrorCode CastSM86PairCountNotDivisible)) (lambda unrestricted divisiblePredecessor : Nat . (lambda unrestricted divisibleInduction : (family CastGeometryValidation) . (constructor CastGeometryValidation CastGeometryAccepted pairs (naturalDivideUnchecked pairs blockThreads)))) (naturalIsZero (naturalModuloUnchecked pairs blockThreads)))) (naturalDivideUnchecked elements castSM86N2)))) (naturalLessOrEqual blockThreads castSM86N1024)))) (naturalNonzero blockThreads)))) (naturalIsZero (naturalModuloUnchecked elements castSM86N2))))) (naturalNonzero elements))))) def castGridX = (lambda unrestricted geometry : (family CastGeometry) . (validateCastGeometry geometry)) -- Only an accepted geometry supplies a launch grid. A rejected geometry -- evaluates to zero, so callers must prove a nonzero grid before emission. def castSM86AcceptedGridX = (lambda unrestricted geometry : (family CastGeometry) . (eliminate CastGeometryValidation (lambda unrestricted current : (family CastGeometryValidation) . Nat) (validateCastGeometry geometry) (branch CastGeometryAccepted pairs gridX . gridX) (branch CastGeometryRejected code . 0))) def castSM86TelemetryFor = (lambda unrestricted direction : (family CastDirection) . (lambda unrestricted geometry : (family CastGeometry) . (lambda unrestricted pairs : Nat . (lambda unrestricted gridX : Nat . (lambda unrestricted actualInstructions : Nat . (lambda unrestricted actualBytes : Nat . (eliminate CastGeometry (lambda unrestricted current : (family CastGeometry) . (family CastSM86Telemetry)) geometry (branch CastGeometryValue elements blockThreads . (constructor CastSM86Telemetry CastSM86TelemetryValue direction elements pairs blockThreads gridX (castInstructionCount direction) actualInstructions (naturalMultiply (castInstructionCount direction) castSM86N16) actualBytes castRegisterCount (eliminate CastDirection (lambda unrestricted current : (family CastDirection) . Nat) direction (branch Float32ToFloat16 . castSM86N2) (branch Float16ToFloat32 . castSM86N1)) (eliminate CastDirection (lambda unrestricted current : (family CastDirection) . Nat) direction (branch Float32ToFloat16 . castSM86N1) (branch Float16ToFloat32 . castSM86N2)) (eliminate CastDirection (lambda unrestricted current : (family CastDirection) . Nat) direction (branch Float32ToFloat16 . castSM86N1) (branch Float16ToFloat32 . castSM86N2)) castSM86HostFallbackCalls castSM86ApplicationBuilderContract))))))))) def castImageSHA256 = (lambda unrestricted direction : (family CastDirection) . (lambda unrestricted geometry : (family CastGeometry) . (eliminate CastGeometryValidation (lambda unrestricted current : (family CastGeometryValidation) . (family CastSM86BuildResult)) (validateCastGeometry geometry) (branch CastGeometryAccepted pairs gridX . (app (lambda unrestricted program : (family SM86Program) . (app (lambda unrestricted actualInstructions : Nat . (nat-eliminate (lambda unrestricted countMatches : Nat . (family CastSM86BuildResult)) (constructor CastSM86BuildResult CastSM86BuildRejected (constructor CastSM86ErrorCode CastSM86InstructionCountMismatch) actualInstructions (castSM86ErrorStableCode (constructor CastSM86ErrorCode CastSM86InstructionCountMismatch))) (lambda unrestricted countPredecessor : Nat . (lambda unrestricted countInduction : (family CastSM86BuildResult) . (eliminate SM86ProgramEncodingResult (lambda unrestricted current : (family SM86ProgramEncodingResult) . (family CastSM86BuildResult)) (sm86EncodeProgram program) (branch SM86ProgramEncodingSucceeded image encodingTelemetry . (app (lambda unrestricted actualBytes : Nat . (nat-eliminate (lambda unrestricted bytesMatch : Nat . (family CastSM86BuildResult)) (constructor CastSM86BuildResult CastSM86BuildRejected (constructor CastSM86ErrorCode CastSM86ByteCountMismatch) actualBytes (castSM86ErrorStableCode (constructor CastSM86ErrorCode CastSM86ByteCountMismatch))) (lambda unrestricted bytePredecessor : Nat . (lambda unrestricted byteInduction : (family CastSM86BuildResult) . (app (lambda unrestricted identity : Bytes . (nat-eliminate (lambda unrestricted identityValid : Nat . (family CastSM86BuildResult)) (constructor CastSM86BuildResult CastSM86BuildRejected (constructor CastSM86ErrorCode CastSM86IdentityInvalid) zero (castSM86ErrorStableCode (constructor CastSM86ErrorCode CastSM86IdentityInvalid))) (lambda unrestricted identityPredecessor : Nat . (lambda unrestricted identityInduction : (family CastSM86BuildResult) . (constructor CastSM86BuildResult CastSM86BuildReady direction geometry gridX program image identity (castSM86TelemetryFor direction geometry pairs gridX actualInstructions actualBytes)))) (naturalEqual (bytes-length identity) (byte-to-nat (byte 64))))) (sha256HexBytesOrEmpty (sha256Hex image))))) (naturalEqual actualBytes (naturalMultiply (castInstructionCount direction) castSM86N16)))) (bytes-length image))) (branch SM86ProgramEncodingFailed ordinal failure encodingTelemetry . (constructor CastSM86BuildResult CastSM86BuildRejected (constructor CastSM86ErrorCode CastSM86EncodingFailed) ordinal (sm86InstructionEncodingStableCode failure)))))) (naturalEqual actualInstructions (castInstructionCount direction)))) (sm86ProgramCount program))) (emitCastSM86 direction))) (branch CastGeometryRejected code . (constructor CastSM86BuildResult CastSM86BuildRejected code zero (castSM86ErrorStableCode code))))))