module Realization.Nvidia.SM86.Elementwise.GeneralSM86 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 family Elementwise256 : Type 0 constructor Add256 constructor Multiply256 constructor Scale256 end-family family Elementwise256ErrorCode : Type 0 constructor Elementwise256InstructionCountMismatch constructor Elementwise256EncodingFailed constructor Elementwise256ByteCountMismatch constructor Elementwise256IdentityInvalid constructor Elementwise256HostFallbackObserved end-family family Elementwise256Telemetry : Type 0 constructor Elementwise256TelemetryValue field unrestricted elementwise256TelemetryOperation : (family Elementwise256) field unrestricted elementwise256TelemetryExpectedInstructions : Nat field unrestricted elementwise256TelemetryActualInstructions : Nat field unrestricted elementwise256TelemetryExpectedBytes : Nat field unrestricted elementwise256TelemetryActualBytes : Nat field unrestricted elementwise256TelemetryRegisters : Nat field unrestricted elementwise256TelemetryGlobalLoads : Nat field unrestricted elementwise256TelemetryGlobalStores : Nat field unrestricted elementwise256TelemetryFloatAdds : Nat field unrestricted elementwise256TelemetryFloatMultiplies : Nat field unrestricted elementwise256TelemetryHostFallbackCalls : Nat field unrestricted elementwise256TelemetryApplicationBuilderContract : Nat end-family family Elementwise256BuildResult : Type 0 constructor Elementwise256BuildReady field unrestricted elementwise256ReadyOperation : (family Elementwise256) field unrestricted elementwise256ReadyProgram : (family SM86Program) field unrestricted elementwise256ReadyImage : Bytes field unrestricted elementwise256ReadySHA256 : Bytes field unrestricted elementwise256ReadyTelemetry : (family Elementwise256Telemetry) constructor Elementwise256BuildRejected field unrestricted elementwise256RejectedCode : (family Elementwise256ErrorCode) field unrestricted elementwise256RejectedOrdinal : Nat field unrestricted elementwise256RejectedDetail : Bytes field unrestricted elementwise256RejectedTelemetry : (family Elementwise256Telemetry) end-family def elementwise256N1 = (succ zero) def elementwise256N2 = (byte-to-nat (byte 2)) def elementwise256N12 = (byte-to-nat (byte 12)) def elementwise256N13 = (byte-to-nat (byte 13)) def elementwise256N16 = (byte-to-nat (byte 16)) def elementwise256N24 = (byte-to-nat (byte 24)) def elementwise256Word32 = (lambda unrestricted b0 : Byte . (lambda unrestricted b1 : Byte . (lambda unrestricted b2 : Byte . (lambda unrestricted b3 : Byte . (sm86Unsigned32 b0 b1 b2 b3))))) def elementwise256Zero32 = (elementwise256Word32 (byte 0) (byte 0) (byte 0) (byte 0)) def elementwise256Four32 = (elementwise256Word32 (byte 4) (byte 0) (byte 0) (byte 0)) def elementwise256ConstantBase = (elementwise256Word32 (byte 40) (byte 0) (byte 0) (byte 0)) def elementwise256Output0 = (elementwise256Word32 (byte 96) (byte 1) (byte 0) (byte 0)) def elementwise256Input0 = (elementwise256Word32 (byte 104) (byte 1) (byte 0) (byte 0)) def elementwise256Input1 = (elementwise256Word32 (byte 112) (byte 1) (byte 0) (byte 0)) def elementwise256Scalar0 = (elementwise256Word32 (byte 144) (byte 1) (byte 0) (byte 0)) def elementwise256Register = (lambda unrestricted index : Byte . (sm86Register index)) def elementwise256Instruction = (lambda unrestricted body : (family SM86InstructionBody) . (constructor SM86Instruction SM86InstructionValue (constructor SM86InstructionGuard SM86InstructionAlways) body)) def elementwise256Next = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (elementwise256Instruction body) tail))) def elementwise256End = (constructor SM86Program SM86ProgramEnd) def elementwise256Set0 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0)) def elementwise256Set1 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier1)) def elementwise256Wait0 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0)) def elementwise256WaitBoth = (constructor SM86Control SM86ControlValue (byte 7) (constructor SM86YieldMode SM86Continue) (constructor SM86Barrier SM86BarrierNone) (constructor SM86Barrier SM86BarrierNone) (byte 3) (byte 0)) -- The prologue's schedule is a choice of controls. Existing kernels retain -- their qualified controls; generated expression kernels use the fixed- -- latency model's conservative schedule and can later propose a shorter one. def elementwise256PrefixWith = (lambda unrestricted fixed : (family SM86Control) . (lambda unrestricted set : (family SM86Control) . (lambda unrestricted wait : (family SM86Control) . (lambda unrestricted tail : (family SM86Program) . (elementwise256Next (constructor SM86InstructionBody SM86MoveConstant (elementwise256Register (byte 1)) (byte 0) elementwise256ConstantBase fixed) (elementwise256Next (constructor SM86InstructionBody SM86SpecialToRegister (elementwise256Register (byte 0)) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) set) (elementwise256Next (constructor SM86InstructionBody SM86MoveImmediate (elementwise256Register (byte 3)) elementwise256Four32 wait) (elementwise256Next (constructor SM86InstructionBody SM86SpecialToRegister (elementwise256Register (byte 2)) (constructor SM86SpecialRegister SM86ThreadIdX) set) (elementwise256Next (constructor SM86InstructionBody SM86IntegerMultiplyAddConstant (elementwise256Register (byte 0)) (elementwise256Register (byte 0)) (byte 0) elementwise256Zero32 (elementwise256Register (byte 2)) wait) tail))))))))) def elementwise256Prefix = (elementwise256PrefixWith sm86SafeControl elementwise256Set0 elementwise256Wait0) def elementwise256BinaryTail = (lambda unrestricted operation : (family Elementwise256) . (elementwise256Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (elementwise256Register (byte 6)) (elementwise256Register (byte 0)) (elementwise256Register (byte 3)) (byte 0) elementwise256Input0 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86LoadGlobal (elementwise256Register (byte 8)) (elementwise256Register (byte 6)) elementwise256Zero32 elementwise256Set0) (elementwise256Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (elementwise256Register (byte 10)) (elementwise256Register (byte 0)) (elementwise256Register (byte 3)) (byte 0) elementwise256Input1 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86LoadGlobal (elementwise256Register (byte 12)) (elementwise256Register (byte 10)) elementwise256Zero32 elementwise256Set1) (elementwise256Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (elementwise256Register (byte 14)) (elementwise256Register (byte 0)) (elementwise256Register (byte 3)) (byte 0) elementwise256Output0 sm86SafeControl) (elementwise256Next (eliminate Elementwise256 (lambda unrestricted current : (family Elementwise256) . (family SM86InstructionBody)) operation (branch Add256 . (constructor SM86InstructionBody SM86FloatAdd (elementwise256Register (byte 16)) (elementwise256Register (byte 8)) (elementwise256Register (byte 12)) elementwise256WaitBoth)) (branch Multiply256 . (constructor SM86InstructionBody SM86FloatMultiply (elementwise256Register (byte 16)) (elementwise256Register (byte 8)) (elementwise256Register (byte 12)) elementwise256WaitBoth)) (branch Scale256 . (constructor SM86InstructionBody SM86FloatMultiply (elementwise256Register (byte 16)) (elementwise256Register (byte 8)) (elementwise256Register (byte 12)) elementwise256WaitBoth))) (elementwise256Next (constructor SM86InstructionBody SM86StoreGlobal (elementwise256Register (byte 14)) (elementwise256Register (byte 16)) elementwise256Zero32 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) elementwise256End))))))))) def elementwise256ScaleTail : (family SM86Program) = (elementwise256Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (elementwise256Register (byte 6)) (elementwise256Register (byte 0)) (elementwise256Register (byte 3)) (byte 0) elementwise256Input0 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86LoadGlobal (elementwise256Register (byte 8)) (elementwise256Register (byte 6)) elementwise256Zero32 elementwise256Set0) (elementwise256Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (elementwise256Register (byte 10)) (elementwise256Register (byte 0)) (elementwise256Register (byte 3)) (byte 0) elementwise256Output0 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86MoveConstant (elementwise256Register (byte 12)) (byte 0) elementwise256Scalar0 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86FloatMultiply (elementwise256Register (byte 14)) (elementwise256Register (byte 8)) (elementwise256Register (byte 12)) elementwise256Wait0) (elementwise256Next (constructor SM86InstructionBody SM86StoreGlobal (elementwise256Register (byte 10)) (elementwise256Register (byte 14)) elementwise256Zero32 sm86SafeControl) (elementwise256Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) elementwise256End))))))) def emitElementwise256SM86 = (lambda unrestricted operation : (family Elementwise256) . (eliminate Elementwise256 (lambda unrestricted current : (family Elementwise256) . (family SM86Program)) operation (branch Add256 . (elementwise256Prefix (elementwise256BinaryTail operation))) (branch Multiply256 . (elementwise256Prefix (elementwise256BinaryTail operation))) (branch Scale256 . (elementwise256Prefix elementwise256ScaleTail)))) def elementwise256InstructionCount = (lambda unrestricted operation : (family Elementwise256) . (eliminate Elementwise256 (lambda unrestricted current : (family Elementwise256) . Nat) operation (branch Add256 . elementwise256N13) (branch Multiply256 . elementwise256N13) (branch Scale256 . elementwise256N12))) def elementwise256RegisterCount : Nat = elementwise256N24 def elementwise256HostFallbackCalls : Nat = zero def elementwise256ApplicationBuilderContract : Nat = Compiler.ApplicationBuilder/applicationBuilderNativeOnly def elementwise256ErrorStableCode = (lambda unrestricted code : (family Elementwise256ErrorCode) . (eliminate Elementwise256ErrorCode (lambda unrestricted current : (family Elementwise256ErrorCode) . Bytes) code (branch Elementwise256InstructionCountMismatch . b"ELT-001") (branch Elementwise256EncodingFailed . b"ELT-002") (branch Elementwise256ByteCountMismatch . b"ELT-003") (branch Elementwise256IdentityInvalid . b"ELT-004") (branch Elementwise256HostFallbackObserved . b"ELT-005"))) def elementwise256TelemetryFor = (lambda unrestricted operation : (family Elementwise256) . (lambda unrestricted actualInstructions : Nat . (lambda unrestricted actualBytes : Nat . (constructor Elementwise256Telemetry Elementwise256TelemetryValue operation (elementwise256InstructionCount operation) actualInstructions (naturalMultiply (elementwise256InstructionCount operation) elementwise256N16) actualBytes elementwise256RegisterCount (eliminate Elementwise256 (lambda unrestricted current : (family Elementwise256) . Nat) operation (branch Add256 . elementwise256N2) (branch Multiply256 . elementwise256N2) (branch Scale256 . elementwise256N1)) elementwise256N1 (eliminate Elementwise256 (lambda unrestricted current : (family Elementwise256) . Nat) operation (branch Add256 . elementwise256N1) (branch Multiply256 . zero) (branch Scale256 . zero)) (eliminate Elementwise256 (lambda unrestricted current : (family Elementwise256) . Nat) operation (branch Add256 . zero) (branch Multiply256 . elementwise256N1) (branch Scale256 . elementwise256N1)) elementwise256HostFallbackCalls elementwise256ApplicationBuilderContract)))) def elementwise256Reject = (lambda unrestricted operation : (family Elementwise256) . (lambda unrestricted code : (family Elementwise256ErrorCode) . (lambda unrestricted ordinal : Nat . (lambda unrestricted actualInstructions : Nat . (lambda unrestricted actualBytes : Nat . (constructor Elementwise256BuildResult Elementwise256BuildRejected code ordinal (elementwise256ErrorStableCode code) (elementwise256TelemetryFor operation actualInstructions actualBytes))))))) def elementwise256ImageSHA256 = (lambda unrestricted operation : (family Elementwise256) . (app (lambda unrestricted program : (family SM86Program) . (app (lambda unrestricted actualInstructions : Nat . (nat-eliminate (lambda unrestricted countMatches : Nat . (family Elementwise256BuildResult)) (elementwise256Reject operation (constructor Elementwise256ErrorCode Elementwise256InstructionCountMismatch) actualInstructions actualInstructions zero) (lambda unrestricted countPredecessor : Nat . (lambda unrestricted countInduction : (family Elementwise256BuildResult) . (eliminate SM86ProgramEncodingResult (lambda unrestricted current : (family SM86ProgramEncodingResult) . (family Elementwise256BuildResult)) (sm86EncodeProgram program) (branch SM86ProgramEncodingSucceeded image encodingTelemetry . (app (lambda unrestricted actualBytes : Nat . (nat-eliminate (lambda unrestricted byteMatches : Nat . (family Elementwise256BuildResult)) (elementwise256Reject operation (constructor Elementwise256ErrorCode Elementwise256ByteCountMismatch) actualBytes actualInstructions actualBytes) (lambda unrestricted bytePredecessor : Nat . (lambda unrestricted byteInduction : (family Elementwise256BuildResult) . (app (lambda unrestricted identity : Bytes . (nat-eliminate (lambda unrestricted valid : Nat . (family Elementwise256BuildResult)) (elementwise256Reject operation (constructor Elementwise256ErrorCode Elementwise256IdentityInvalid) zero actualInstructions actualBytes) (lambda unrestricted validPredecessor : Nat . (lambda unrestricted validInduction : (family Elementwise256BuildResult) . (constructor Elementwise256BuildResult Elementwise256BuildReady operation program image identity (elementwise256TelemetryFor operation actualInstructions actualBytes)))) (naturalEqual (bytes-length identity) (byte-to-nat (byte 64))))) (sha256HexBytesOrEmpty (sha256Hex image))))) (naturalEqual actualBytes (naturalMultiply (elementwise256InstructionCount operation) elementwise256N16)))) (bytes-length image))) (branch SM86ProgramEncodingFailed ordinal failure encodingTelemetry . (elementwise256Reject operation (constructor Elementwise256ErrorCode Elementwise256EncodingFailed) ordinal actualInstructions zero))))) (naturalEqual actualInstructions (elementwise256InstructionCount operation)))) (sm86ProgramCount program))) (emitElementwise256SM86 operation)))