module Realization.Nvidia.SM86.SoftmaxSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.InstructionEncoding import Accelerator.SM86.NumericSemantics import Accelerator.SM86.Program import Accelerator.SM86.Types import Data.SHA256Digest import Realization.Nvidia.SM86.ReductionSM86 import Std.Natural family SoftmaxSM86Variant : Type 0 constructor SoftmaxSM86Forward constructor SoftmaxSM86CausalForward constructor SoftmaxSM86CausalForward1024 constructor SoftmaxSM86CausalBackward1024 constructor SoftmaxSM86Backward end-family family SoftmaxSM86InputContract : Type 0 constructor SoftmaxSM86FP32ScoresInput constructor SoftmaxSM86FP32ProbabilitiesAndAdjointsInput end-family family SoftmaxSM86OutputContract : Type 0 constructor SoftmaxSM86FP32ProbabilitiesOutput constructor SoftmaxSM86FP32ScoreAdjointsOutput end-family family SoftmaxSM86MaskContract : Type 0 constructor SoftmaxSM86Unmasked constructor SoftmaxSM86CausalLowEightBitRowMask constructor SoftmaxSM86CausalFullRowMask end-family family SoftmaxSM86ABI : Type 0 constructor SoftmaxSM86ABIValue field unrestricted softmaxSM86ConstantBaseOffset : Nat field unrestricted softmaxSM86OutputPointerOffset : Nat field unrestricted softmaxSM86PrimaryInputPointerOffset : Nat field unrestricted softmaxSM86SecondaryInputPointerOffset : Nat field unrestricted softmaxSM86Log2EOffset : Nat end-family family SoftmaxSM86Extents : Type 0 constructor SoftmaxSM86ExtentsValue field unrestricted softmaxSM86RowWidth : Nat field unrestricted softmaxSM86PrimaryInputElementsPerRow : Nat field unrestricted softmaxSM86SecondaryInputElementsPerRow : Nat field unrestricted softmaxSM86OutputElementsPerRow : Nat field unrestricted softmaxSM86ElementsPerThread : Nat field unrestricted softmaxSM86BlocksPerRow : Nat field unrestricted softmaxSM86ThreadsPerBlock : Nat field unrestricted softmaxSM86InputContract : (family SoftmaxSM86InputContract) field unrestricted softmaxSM86OutputContract : (family SoftmaxSM86OutputContract) field unrestricted softmaxSM86MaskContract : (family SoftmaxSM86MaskContract) end-family family SoftmaxSM86Manifest : Type 0 constructor SoftmaxSM86ManifestValue field unrestricted softmaxSM86ManifestVariant : (family SoftmaxSM86Variant) field unrestricted softmaxSM86ManifestExpectedInstructions : Nat field unrestricted softmaxSM86ManifestExpectedBytes : Nat field unrestricted softmaxSM86ManifestRegisters : Nat field unrestricted softmaxSM86ManifestSharedBytes : Nat field unrestricted softmaxSM86ManifestBlocksPerRow : Nat field unrestricted softmaxSM86ManifestThreadsPerBlock : Nat field unrestricted softmaxSM86ManifestABI : (family SoftmaxSM86ABI) field unrestricted softmaxSM86ManifestExtents : (family SoftmaxSM86Extents) field unrestricted softmaxSM86ManifestPrefixInstructions : Nat field unrestricted softmaxSM86ManifestMaximumReductionInstructions : Nat field unrestricted softmaxSM86ManifestTransformInstructions : Nat field unrestricted softmaxSM86ManifestSumReductionInstructions : Nat field unrestricted softmaxSM86ManifestSuffixInstructions : Nat field unrestricted softmaxSM86ManifestGlobalLoads : Nat field unrestricted softmaxSM86ManifestGlobalStores : Nat field unrestricted softmaxSM86ManifestHostFallbackOperations : Nat field unrestricted softmaxSM86ManifestExpectedSHA256 : Bytes end-family family SoftmaxSM86FailureCode : Type 0 constructor SoftmaxSM86InstructionCountMismatch constructor SoftmaxSM86EncodedByteCountMismatch constructor SoftmaxSM86RegisterCountMismatch constructor SoftmaxSM86SharedByteCountMismatch constructor SoftmaxSM86BlockCountMismatch constructor SoftmaxSM86ThreadCountMismatch constructor SoftmaxSM86SectionCountMismatch constructor SoftmaxSM86GlobalLoadCountMismatch constructor SoftmaxSM86GlobalStoreCountMismatch constructor SoftmaxSM86HostFallbackDetected constructor SoftmaxSM86EncodingFailed constructor SoftmaxSM86IdentityFailed constructor SoftmaxSM86IdentityLengthInvalid constructor SoftmaxSM86IdentityMismatch end-family family SoftmaxSM86Telemetry : Type 0 constructor SoftmaxSM86TelemetryValue field unrestricted softmaxSM86TelemetryManifest : (family SoftmaxSM86Manifest) field unrestricted softmaxSM86TelemetryObservedInstructions : Nat field unrestricted softmaxSM86TelemetryObservedEncodedBytes : Nat field unrestricted softmaxSM86TelemetryEncodingFields : Nat field unrestricted softmaxSM86TelemetryEncodingBits : Nat field unrestricted softmaxSM86TelemetryHighestEncodedBit : Nat field unrestricted softmaxSM86TelemetryIdentityInputBytes : Nat field unrestricted softmaxSM86TelemetryIdentityOutputBytes : Nat field unrestricted softmaxSM86TelemetryGlobalLoads : Nat field unrestricted softmaxSM86TelemetryGlobalStores : Nat field unrestricted softmaxSM86TelemetryHostFallbackOperations : Nat end-family family SoftmaxSM86ResourceCheck : Type 0 constructor SoftmaxSM86ResourceCheckValue field unrestricted softmaxSM86ResourceObserved : Nat field unrestricted softmaxSM86ResourceExpected : Nat field unrestricted softmaxSM86ResourceFailure : (family SoftmaxSM86FailureCode) end-family family SoftmaxSM86ResourceChecks : Type 0 constructor SoftmaxSM86ResourceChecksEnd constructor SoftmaxSM86ResourceChecksNext field unrestricted softmaxSM86ResourceCheckHead : (family SoftmaxSM86ResourceCheck) recursive unrestricted softmaxSM86ResourceCheckTail end-family family SoftmaxSM86ResourceGateResult : Type 0 constructor SoftmaxSM86ResourcesExact constructor SoftmaxSM86ResourcesRejected field unrestricted softmaxSM86ResourceFailureCode : (family SoftmaxSM86FailureCode) end-family family SoftmaxSM86BuildResult : Type 0 constructor SoftmaxSM86BuildSucceeded field unrestricted softmaxSM86EncodedBytes : Bytes field unrestricted softmaxSM86ImageSHA256 : Bytes field unrestricted softmaxSM86ProgramEncodingTelemetry : (family SM86ProgramEncodingTelemetry) field unrestricted softmaxSM86IdentityTelemetry : (family SHA256DigestTelemetry) field unrestricted softmaxSM86BuildTelemetry : (family SoftmaxSM86Telemetry) constructor SoftmaxSM86ContractFailed field unrestricted softmaxSM86ContractFailure : (family SoftmaxSM86FailureCode) field unrestricted softmaxSM86ContractFailureTelemetry : (family SoftmaxSM86Telemetry) constructor SoftmaxSM86ImageEncodingFailed field unrestricted softmaxSM86EncodingFailure : (family SoftmaxSM86FailureCode) field unrestricted softmaxSM86FailedEncoding : (family SM86ProgramEncodingResult) field unrestricted softmaxSM86EncodingFailureTelemetry : (family SoftmaxSM86Telemetry) constructor SoftmaxSM86ImageIdentityFailed field unrestricted softmaxSM86IdentityFailure : (family SoftmaxSM86FailureCode) field unrestricted softmaxSM86FailedIdentity : (family SHA256HexResult) field unrestricted softmaxSM86IdentityFailureTelemetry : (family SoftmaxSM86Telemetry) end-family def softmaxSM86FailureCodeBytes = (lambda unrestricted code : (family SoftmaxSM86FailureCode) . (eliminate SoftmaxSM86FailureCode (lambda unrestricted current : (family SoftmaxSM86FailureCode) . Bytes) code (branch SoftmaxSM86InstructionCountMismatch . b"ALPHA-SM86-SOFTMAX-001") (branch SoftmaxSM86EncodedByteCountMismatch . b"ALPHA-SM86-SOFTMAX-002") (branch SoftmaxSM86RegisterCountMismatch . b"ALPHA-SM86-SOFTMAX-003") (branch SoftmaxSM86SharedByteCountMismatch . b"ALPHA-SM86-SOFTMAX-004") (branch SoftmaxSM86BlockCountMismatch . b"ALPHA-SM86-SOFTMAX-005") (branch SoftmaxSM86ThreadCountMismatch . b"ALPHA-SM86-SOFTMAX-006") (branch SoftmaxSM86SectionCountMismatch . b"ALPHA-SM86-SOFTMAX-007") (branch SoftmaxSM86GlobalLoadCountMismatch . b"ALPHA-SM86-SOFTMAX-008") (branch SoftmaxSM86GlobalStoreCountMismatch . b"ALPHA-SM86-SOFTMAX-009") (branch SoftmaxSM86HostFallbackDetected . b"ALPHA-SM86-SOFTMAX-010") (branch SoftmaxSM86EncodingFailed . b"ALPHA-SM86-SOFTMAX-011") (branch SoftmaxSM86IdentityFailed . b"ALPHA-SM86-SOFTMAX-012") (branch SoftmaxSM86IdentityLengthInvalid . b"ALPHA-SM86-SOFTMAX-013") (branch SoftmaxSM86IdentityMismatch . b"ALPHA-SM86-SOFTMAX-014"))) def softmaxSM86N1 = (succ zero) def softmaxSM86N2 = (byte-to-nat (byte 2)) def softmaxSM86N5 = (byte-to-nat (byte 5)) def softmaxSM86N6 = (byte-to-nat (byte 6)) def softmaxSM86N7 = (byte-to-nat (byte 7)) def softmaxSM86N8 = (byte-to-nat (byte 8)) def softmaxSM86N12 = (byte-to-nat (byte 12)) def softmaxSM86N15 = (byte-to-nat (byte 15)) def softmaxSM86N16 = (byte-to-nat (byte 16)) def softmaxSM86N24 = (byte-to-nat (byte 24)) def softmaxSM86N34 = (byte-to-nat (byte 34)) def softmaxSM86N40 = (byte-to-nat (byte 40)) def softmaxSM86N53 = (byte-to-nat (byte 53)) def softmaxSM86N64 = (byte-to-nat (byte 64)) def softmaxSM86N87 = (byte-to-nat (byte 87)) def softmaxSM86N94 = (byte-to-nat (byte 94)) def softmaxSM86N96 = (byte-to-nat (byte 96)) def softmaxSM86N104 = (byte-to-nat (byte 104)) def softmaxSM86N112 = (byte-to-nat (byte 112)) def softmaxSM86N128 = (byte-to-nat (byte 128)) def softmaxSM86N144 = (byte-to-nat (byte 144)) def softmaxSM86N256 = (succ (byte-to-nat (byte 255))) def softmaxSM86N352 = (naturalAdd softmaxSM86N256 softmaxSM86N96) def softmaxSM86N360 = (naturalAdd softmaxSM86N256 softmaxSM86N104) def softmaxSM86N368 = (naturalAdd softmaxSM86N256 softmaxSM86N112) def softmaxSM86N400 = (naturalAdd softmaxSM86N256 softmaxSM86N144) def softmaxSM86N848 = (naturalMultiply softmaxSM86N53 softmaxSM86N16) def softmaxSM86N72 = (byte-to-nat (byte 72)) def softmaxSM86N1152 = (naturalMultiply softmaxSM86N72 softmaxSM86N16) def softmaxSM86N127 = (naturalAdd softmaxSM86N96 (byte-to-nat (byte 31))) def softmaxSM86N2032 = (naturalMultiply softmaxSM86N127 softmaxSM86N16) def softmaxSM86N30 = (byte-to-nat (byte 30)) def softmaxSM86N17 = (byte-to-nat (byte 17)) def softmaxSM86N11 = (byte-to-nat (byte 11)) def softmaxSM86U4Natural = (byte-to-nat (byte 4)) def softmaxSM86N1024 = (naturalMultiply softmaxSM86N256 (byte-to-nat (byte 4))) def softmaxSM86N1392 = (naturalMultiply softmaxSM86N87 softmaxSM86N16) def softmaxSM86N1504 = (naturalMultiply softmaxSM86N94 softmaxSM86N16) def softmaxSM86R0 = (sm86Register (byte 0)) def softmaxSM86R1 = (sm86Register (byte 1)) def softmaxSM86R2 = (sm86Register (byte 2)) def softmaxSM86R4 = (sm86Register (byte 4)) def softmaxSM86R5 = (sm86Register (byte 5)) def softmaxSM86R6 = (sm86Register (byte 6)) def softmaxSM86R7 = (sm86Register (byte 7)) def softmaxSM86R8 = (sm86Register (byte 8)) def softmaxSM86R9 = (sm86Register (byte 9)) def softmaxSM86R10 = (sm86Register (byte 10)) def softmaxSM86R11 = (sm86Register (byte 11)) def softmaxSM86R16 = (sm86Register (byte 16)) def softmaxSM86R12 = (sm86Register (byte 12)) def softmaxSM86R13 = (sm86Register (byte 13)) def softmaxSM86R14 = (sm86Register (byte 14)) def softmaxSM86R17 = (sm86Register (byte 17)) def softmaxSM86R18 = (sm86Register (byte 18)) def softmaxSM86R19 = (sm86Register (byte 19)) def softmaxSM86R20 = (sm86Register (byte 20)) def softmaxSM86R21 = (sm86Register (byte 21)) def softmaxSM86R22 = (sm86Register (byte 22)) def softmaxSM86P0 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate0) def softmaxSM86U0 = sm86Unsigned32Zero def softmaxSM86U4 = (sm86Unsigned32 (byte 4) (byte 0) (byte 0) (byte 0)) def softmaxSM86U40 = (sm86Unsigned32 (byte 40) (byte 0) (byte 0) (byte 0)) def softmaxSM86U255 = (sm86Unsigned32 (byte 255) (byte 0) (byte 0) (byte 0)) def softmaxSM86U256 = (sm86Unsigned32 (byte 0) (byte 1) (byte 0) (byte 0)) def softmaxSM86U352 = (sm86Unsigned32 (byte 96) (byte 1) (byte 0) (byte 0)) def softmaxSM86U360 = (sm86Unsigned32 (byte 104) (byte 1) (byte 0) (byte 0)) def softmaxSM86U368 = (sm86Unsigned32 (byte 112) (byte 1) (byte 0) (byte 0)) def softmaxSM86U400 = (sm86Unsigned32 (byte 144) (byte 1) (byte 0) (byte 0)) def softmaxSM86U8 = (sm86Unsigned32 (byte 8) (byte 0) (byte 0) (byte 0)) def softmaxSM86U12 = (sm86Unsigned32 (byte 12) (byte 0) (byte 0) (byte 0)) def softmaxSM86U1024 = (sm86Unsigned32 (byte 0) (byte 4) (byte 0) (byte 0)) def softmaxSM86UNegativeTwo = (sm86Unsigned32 (byte 254) (byte 255) (byte 255) (byte 255)) def softmaxSM86UNegativeThree = (sm86Unsigned32 (byte 253) (byte 255) (byte 255) (byte 255)) def softmaxSM86UNegativeOne = (sm86Unsigned32 (byte 255) (byte 255) (byte 255) (byte 255)) def softmaxSM86UNegativeInfinity = (sm86Unsigned32 (byte 0) (byte 0) (byte 128) (byte 255)) def softmaxSM86Set0 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0)) def softmaxSM86Set1 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier1)) def softmaxSM86Set2 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier2)) def softmaxSM86Set3 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier3)) def softmaxSM86Wait0 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0)) def softmaxSM86Wait1 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier1)) def softmaxSM86Wait2 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier2)) def softmaxSM86Wait3 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier3)) def softmaxSM86ThreadIdX : (family SM86SpecialRegister) = (constructor SM86SpecialRegister SM86ThreadIdX) def softmaxSM86BlockIdX : (family SM86SpecialRegister) = (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) def softmaxSM86Exp2 : (family SM86MultiFunction) = (constructor SM86MultiFunction SM86ExponentialBase2) def softmaxSM86Reciprocal : (family SM86MultiFunction) = (constructor SM86MultiFunction SM86Reciprocal) def softmaxSM86Next = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86Instruction body) tail))) def softmaxSM86PredicatedNext = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction predicate body) tail)))) def softmaxSM86ForwardPrefix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R1 (byte 0) softmaxSM86U40 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R0 softmaxSM86ThreadIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R1 softmaxSM86BlockIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R5 softmaxSM86U4 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R16 softmaxSM86R1 softmaxSM86U256 softmaxSM86R0 softmaxSM86Wait0) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R2 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U360 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R4 softmaxSM86R2 softmaxSM86U0 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R6 softmaxSM86R4 softmaxSM86U0 softmaxSM86Wait1) (constructor SM86Program SM86ProgramEnd))))))))) def softmaxSM86CausalPrefix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R1 (byte 0) softmaxSM86U40 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R0 softmaxSM86ThreadIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R1 softmaxSM86BlockIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R5 softmaxSM86U4 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R16 softmaxSM86R1 softmaxSM86U256 softmaxSM86R0 softmaxSM86Wait0) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R2 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U360 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R4 softmaxSM86R2 softmaxSM86U0 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R4 softmaxSM86R4 softmaxSM86U0 softmaxSM86Wait1) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R10 softmaxSM86U255 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86LogicThreeInputTruthTable softmaxSM86R10 softmaxSM86R1 softmaxSM86R10 (byte 192) sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R11 softmaxSM86R0 softmaxSM86UNegativeOne softmaxSM86R10 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86ShiftRightImmediate softmaxSM86R11 softmaxSM86R11 (byte 31) sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate softmaxSM86P0 softmaxSM86R11 softmaxSM86U0 sm86SafeControl) (softmaxSM86PredicatedNext softmaxSM86P0 (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R4 softmaxSM86UNegativeInfinity sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R6 softmaxSM86R4 softmaxSM86U0 sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))))))))))))))) def softmaxSM86ExponentProgram : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86FloatNegate softmaxSM86R8 softmaxSM86R6 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd softmaxSM86R4 softmaxSM86R4 softmaxSM86R8 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R9 (byte 0) softmaxSM86U400 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply softmaxSM86R4 softmaxSM86R4 softmaxSM86R9 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation softmaxSM86R4 softmaxSM86R4 softmaxSM86Exp2 softmaxSM86Set3) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R6 softmaxSM86R4 softmaxSM86U0 softmaxSM86Wait3) (constructor SM86Program SM86ProgramEnd))))))) def softmaxSM86ForwardSuffix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation softmaxSM86R6 softmaxSM86R6 softmaxSM86Reciprocal softmaxSM86Set3) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply softmaxSM86R4 softmaxSM86R4 softmaxSM86R6 softmaxSM86Wait3) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R10 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U352 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86StoreGlobal softmaxSM86R10 softmaxSM86R4 softmaxSM86U0 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))))) def softmaxSM86BackwardPrefix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R1 (byte 0) softmaxSM86U40 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R0 softmaxSM86ThreadIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R1 softmaxSM86BlockIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R8 softmaxSM86U4 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R16 softmaxSM86R1 softmaxSM86U256 softmaxSM86R0 softmaxSM86Wait0) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R2 softmaxSM86R16 softmaxSM86R8 (byte 0) softmaxSM86U360 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R4 softmaxSM86R2 softmaxSM86U0 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R4 softmaxSM86R4 softmaxSM86U0 softmaxSM86Wait1) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R2 softmaxSM86R16 softmaxSM86R8 (byte 0) softmaxSM86U368 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R5 softmaxSM86R2 softmaxSM86U0 softmaxSM86Set2) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R5 softmaxSM86R5 softmaxSM86U0 softmaxSM86Wait2) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply softmaxSM86R6 softmaxSM86R4 softmaxSM86R5 sm86SafeControl) (constructor SM86Program SM86ProgramEnd))))))))))))) def softmaxSM86BackwardSuffix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86FloatNegate softmaxSM86R7 softmaxSM86R6 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd softmaxSM86R5 softmaxSM86R5 softmaxSM86R7 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply softmaxSM86R4 softmaxSM86R4 softmaxSM86R5 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R8 softmaxSM86U4 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R10 softmaxSM86R16 softmaxSM86R8 (byte 0) softmaxSM86U352 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86StoreGlobal softmaxSM86R10 softmaxSM86R4 softmaxSM86U0 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))))))) def softmaxSM86MaximumReduction : (family SM86Program) = (reductionSM86ParameterizedIdentityReduction reductionSM86MaximumCombine softmaxSM86R0 softmaxSM86R6 softmaxSM86R7 softmaxSM86R8 softmaxSM86UNegativeInfinity) def softmaxSM86SumReduction : (family SM86Program) = (reductionSM86ParameterizedIdentityReduction reductionSM86SumCombine softmaxSM86R0 softmaxSM86R6 softmaxSM86R7 softmaxSM86R8 softmaxSM86U0) -- coppelius D24: causal softmax over 1024-wide rows (attention scores at -- context 1024), 256 threads x FOUR consecutive elements each, one block per -- row, the row index taken in full from CTAID.X (the 256-wide variant masks -- from the low eight bits). Registers: R0 tid, R1 row, R5 = 4, R17 = tid*4 -- (first column), R16 = row*1024 + R17 (element index), R2:R3 input address, -- R4/R12/R13/R14 the four scores, R11 = row - column0, R9 scratch, R6 -- accumulator, R7/R8 reduction scratch, R10:R11 output address. def softmaxSM86MaskOne1024 = (lambda unrestricted delta : (family SM86Unsigned32) . (lambda unrestricted element : (family SM86Register) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerAddThreeImmediate softmaxSM86R9 softmaxSM86R11 delta sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86ShiftRightImmediate softmaxSM86R9 softmaxSM86R9 (byte 31) sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate softmaxSM86P0 softmaxSM86R9 softmaxSM86U0 sm86SafeControl) (softmaxSM86PredicatedNext softmaxSM86P0 (constructor SM86InstructionBody SM86MoveImmediate element softmaxSM86UNegativeInfinity sm86SafeControl) tail))))))) def softmaxSM86Causal1024Prefix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R1 (byte 0) softmaxSM86U40 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R0 softmaxSM86ThreadIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R1 softmaxSM86BlockIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R5 softmaxSM86U4 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R17 softmaxSM86R0 softmaxSM86U4 sm86ZeroRegister softmaxSM86Wait0) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R16 softmaxSM86R1 softmaxSM86U1024 softmaxSM86R17 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R2 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U360 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R4 softmaxSM86R2 softmaxSM86U0 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R12 softmaxSM86R2 softmaxSM86U4 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R13 softmaxSM86R2 softmaxSM86U8 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal softmaxSM86R14 softmaxSM86R2 softmaxSM86U12 softmaxSM86Set1) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R11 softmaxSM86R17 softmaxSM86UNegativeOne softmaxSM86R1 softmaxSM86Wait1) (softmaxSM86MaskOne1024 softmaxSM86U0 softmaxSM86R4 (softmaxSM86MaskOne1024 softmaxSM86UNegativeOne softmaxSM86R12 (softmaxSM86MaskOne1024 softmaxSM86UNegativeTwo softmaxSM86R13 (softmaxSM86MaskOne1024 softmaxSM86UNegativeThree softmaxSM86R14 (constructor SM86Program SM86ProgramEnd))))))))))))))))) def softmaxSM86Extremum1024 = (lambda unrestricted left : (family SM86Register) . (lambda unrestricted right : (family SM86Register) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMinimumOrMaximum softmaxSM86R6 left right (constructor SM86FloatExtremum SM86FloatMaximum) sm86SafeControl) tail)))) def softmaxSM86LocalMaximum1024 : (family SM86Program) = (softmaxSM86Extremum1024 softmaxSM86R4 softmaxSM86R12 (softmaxSM86Extremum1024 softmaxSM86R6 softmaxSM86R13 (softmaxSM86Extremum1024 softmaxSM86R6 softmaxSM86R14 (constructor SM86Program SM86ProgramEnd)))) def softmaxSM86ExponentOne1024 = (lambda unrestricted element : (family SM86Register) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd element element softmaxSM86R8 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply element element softmaxSM86R9 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation element element softmaxSM86Exp2 softmaxSM86Set3) tail))))) def softmaxSM86Exponent1024Program : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86FloatNegate softmaxSM86R8 softmaxSM86R6 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R9 (byte 0) softmaxSM86U400 sm86SafeControl) (softmaxSM86ExponentOne1024 softmaxSM86R4 (softmaxSM86ExponentOne1024 softmaxSM86R12 (softmaxSM86ExponentOne1024 softmaxSM86R13 (softmaxSM86ExponentOne1024 softmaxSM86R14 (constructor SM86Program SM86ProgramEnd))))))) def softmaxSM86LocalSum1024 : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd softmaxSM86R6 softmaxSM86R4 softmaxSM86R12 softmaxSM86Wait3) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd softmaxSM86R6 softmaxSM86R6 softmaxSM86R13 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd softmaxSM86R6 softmaxSM86R6 softmaxSM86R14 sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))) def softmaxSM86Scale1024 = (lambda unrestricted element : (family SM86Register) . (lambda unrestricted control : (family SM86Control) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply element element softmaxSM86R6 control) tail)))) def softmaxSM86Store1024 = (lambda unrestricted element : (family SM86Register) . (lambda unrestricted offset : (family SM86Unsigned32) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86StoreGlobal softmaxSM86R10 element offset sm86SafeControl) tail)))) def softmaxSM86Suffix1024 : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation softmaxSM86R6 softmaxSM86R6 softmaxSM86Reciprocal softmaxSM86Set3) (softmaxSM86Scale1024 softmaxSM86R4 softmaxSM86Wait3 (softmaxSM86Scale1024 softmaxSM86R12 sm86SafeControl (softmaxSM86Scale1024 softmaxSM86R13 sm86SafeControl (softmaxSM86Scale1024 softmaxSM86R14 sm86SafeControl (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R10 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U352 sm86SafeControl) (softmaxSM86Store1024 softmaxSM86R4 softmaxSM86U0 (softmaxSM86Store1024 softmaxSM86R12 softmaxSM86U4 (softmaxSM86Store1024 softmaxSM86R13 softmaxSM86U8 (softmaxSM86Store1024 softmaxSM86R14 softmaxSM86U12 (softmaxSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))))))))))) -- coppelius S5 rung 4d: the softmax BACKWARD over 1024-wide rows, the adjoint -- of SoftmaxSM86CausalForward1024 (256 threads x four elements, one block per -- row): dS = P * (dP - sum(P * dP)). Probabilities P at the primary input -- (0x168), adjoints dP at the secondary input (0x170), score adjoints to the -- output (0x160). Masked entries carry P = 0 and so dS = 0 with no mask logic. def softmaxSM86Backward1024Load = (lambda unrestricted element : (family SM86Register) . (lambda unrestricted address : (family SM86Register) . (lambda unrestricted offset : (family SM86Unsigned32) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86LoadGlobal element address offset softmaxSM86Set1) tail))))) def softmaxSM86Backward1024Prefix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86MoveConstant softmaxSM86R1 (byte 0) softmaxSM86U40 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R0 softmaxSM86ThreadIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86SpecialToRegister softmaxSM86R1 softmaxSM86BlockIdX softmaxSM86Set0) (softmaxSM86Next (constructor SM86InstructionBody SM86MoveImmediate softmaxSM86R5 softmaxSM86U4 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R17 softmaxSM86R0 softmaxSM86U4 sm86ZeroRegister softmaxSM86Wait0) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate softmaxSM86R16 softmaxSM86R1 softmaxSM86U1024 softmaxSM86R17 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R2 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U360 sm86SafeControl) (softmaxSM86Backward1024Load softmaxSM86R4 softmaxSM86R2 softmaxSM86U0 (softmaxSM86Backward1024Load softmaxSM86R12 softmaxSM86R2 softmaxSM86U4 (softmaxSM86Backward1024Load softmaxSM86R13 softmaxSM86R2 softmaxSM86U8 (softmaxSM86Backward1024Load softmaxSM86R14 softmaxSM86R2 softmaxSM86U12 (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R10 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U368 sm86SafeControl) (softmaxSM86Backward1024Load softmaxSM86R18 softmaxSM86R10 softmaxSM86U0 (softmaxSM86Backward1024Load softmaxSM86R19 softmaxSM86R10 softmaxSM86U4 (softmaxSM86Backward1024Load softmaxSM86R20 softmaxSM86R10 softmaxSM86U8 (softmaxSM86Backward1024Load softmaxSM86R21 softmaxSM86R10 softmaxSM86U12 (constructor SM86Program SM86ProgramEnd))))))))))))))))) -- the products P * dP summed into R6 (seven instructions) def softmaxSM86Backward1024Product = (lambda unrestricted probability : (family SM86Register) . (lambda unrestricted adjoint : (family SM86Register) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply softmaxSM86R22 probability adjoint sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd softmaxSM86R6 softmaxSM86R6 softmaxSM86R22 sm86SafeControl) tail))))) def softmaxSM86Backward1024Transform : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply softmaxSM86R6 softmaxSM86R4 softmaxSM86R18 softmaxSM86Wait1) (softmaxSM86Backward1024Product softmaxSM86R12 softmaxSM86R19 (softmaxSM86Backward1024Product softmaxSM86R13 softmaxSM86R20 (softmaxSM86Backward1024Product softmaxSM86R14 softmaxSM86R21 (constructor SM86Program SM86ProgramEnd))))) -- dS = P * (dP - sum): fifteen instructions def softmaxSM86Backward1024Adjoint = (lambda unrestricted probability : (family SM86Register) . (lambda unrestricted adjoint : (family SM86Register) . (lambda unrestricted tail : (family SM86Program) . (softmaxSM86Next (constructor SM86InstructionBody SM86FloatAdd adjoint adjoint softmaxSM86R7 sm86SafeControl) (softmaxSM86Next (constructor SM86InstructionBody SM86FloatMultiply probability probability adjoint sm86SafeControl) tail))))) def softmaxSM86Backward1024Suffix : (family SM86Program) = (softmaxSM86Next (constructor SM86InstructionBody SM86FloatNegate softmaxSM86R7 softmaxSM86R6 sm86SafeControl) (softmaxSM86Backward1024Adjoint softmaxSM86R4 softmaxSM86R18 (softmaxSM86Backward1024Adjoint softmaxSM86R12 softmaxSM86R19 (softmaxSM86Backward1024Adjoint softmaxSM86R13 softmaxSM86R20 (softmaxSM86Backward1024Adjoint softmaxSM86R14 softmaxSM86R21 (softmaxSM86Next (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant softmaxSM86R10 softmaxSM86R16 softmaxSM86R5 (byte 0) softmaxSM86U352 sm86SafeControl) (softmaxSM86Store1024 softmaxSM86R4 softmaxSM86U0 (softmaxSM86Store1024 softmaxSM86R12 softmaxSM86U4 (softmaxSM86Store1024 softmaxSM86R13 softmaxSM86U8 (softmaxSM86Store1024 softmaxSM86R14 softmaxSM86U12 (softmaxSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))))))))))) def softmaxSM86CausalBackward1024Program : (family SM86Program) = (sm86ProgramAppend softmaxSM86Backward1024Prefix (sm86ProgramAppend softmaxSM86Backward1024Transform (sm86ProgramAppend softmaxSM86SumReduction softmaxSM86Backward1024Suffix))) def softmaxSM86CausalForward1024Program : (family SM86Program) = (sm86ProgramAppend softmaxSM86Causal1024Prefix (sm86ProgramAppend softmaxSM86LocalMaximum1024 (sm86ProgramAppend softmaxSM86MaximumReduction (sm86ProgramAppend softmaxSM86Exponent1024Program (sm86ProgramAppend softmaxSM86LocalSum1024 (sm86ProgramAppend softmaxSM86SumReduction softmaxSM86Suffix1024)))))) def softmaxSM86ForwardProgram : (family SM86Program) = (sm86ProgramAppend softmaxSM86ForwardPrefix (sm86ProgramAppend softmaxSM86MaximumReduction (sm86ProgramAppend softmaxSM86ExponentProgram (sm86ProgramAppend softmaxSM86SumReduction softmaxSM86ForwardSuffix)))) def softmaxSM86CausalForwardProgram : (family SM86Program) = (sm86ProgramAppend softmaxSM86CausalPrefix (sm86ProgramAppend softmaxSM86MaximumReduction (sm86ProgramAppend softmaxSM86ExponentProgram (sm86ProgramAppend softmaxSM86SumReduction softmaxSM86ForwardSuffix)))) def softmaxSM86BackwardProgram : (family SM86Program) = (sm86ProgramAppend softmaxSM86BackwardPrefix (sm86ProgramAppend softmaxSM86SumReduction softmaxSM86BackwardSuffix)) def softmaxSM86ProgramFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . (family SM86Program)) variant (branch SoftmaxSM86Forward . softmaxSM86ForwardProgram) (branch SoftmaxSM86CausalForward . softmaxSM86CausalForwardProgram) (branch SoftmaxSM86CausalForward1024 . softmaxSM86CausalForward1024Program) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86CausalBackward1024Program) (branch SoftmaxSM86Backward . softmaxSM86BackwardProgram))) def softmaxSM86ExpectedInstructionsFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N87) (branch SoftmaxSM86CausalForward . softmaxSM86N94) (branch SoftmaxSM86CausalForward1024 . softmaxSM86N127) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86N72) (branch SoftmaxSM86Backward . softmaxSM86N53))) def softmaxSM86ExpectedBytesFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N1392) (branch SoftmaxSM86CausalForward . softmaxSM86N1504) (branch SoftmaxSM86CausalForward1024 . softmaxSM86N2032) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86N1152) (branch SoftmaxSM86Backward . softmaxSM86N848))) def softmaxSM86ExpectedPrefixFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N8) (branch SoftmaxSM86CausalForward . softmaxSM86N15) (branch SoftmaxSM86CausalForward1024 . softmaxSM86N30) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86N16) (branch SoftmaxSM86Backward . softmaxSM86N12))) def softmaxSM86ExpectedMaximumReductionFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N34) (branch SoftmaxSM86CausalForward . softmaxSM86N34) (branch SoftmaxSM86CausalForward1024 . softmaxSM86N34) (branch SoftmaxSM86CausalBackward1024 . zero) (branch SoftmaxSM86Backward . zero))) def softmaxSM86ExpectedTransformFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N6) (branch SoftmaxSM86CausalForward . softmaxSM86N6) (branch SoftmaxSM86CausalForward1024 . softmaxSM86N17) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86N7) (branch SoftmaxSM86Backward . zero))) def softmaxSM86ExpectedSumReductionFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . softmaxSM86N34) def softmaxSM86ExpectedSuffixFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N5) (branch SoftmaxSM86CausalForward . softmaxSM86N5) (branch SoftmaxSM86CausalForward1024 . softmaxSM86N11) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86N15) (branch SoftmaxSM86Backward . softmaxSM86N7))) def softmaxSM86ExpectedLoadsFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Nat) variant (branch SoftmaxSM86Forward . softmaxSM86N1) (branch SoftmaxSM86CausalForward . softmaxSM86N1) (branch SoftmaxSM86CausalForward1024 . softmaxSM86U4Natural) (branch SoftmaxSM86CausalBackward1024 . softmaxSM86N8) (branch SoftmaxSM86Backward . softmaxSM86N2))) def softmaxSM86ExpectedSHA256For = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . Bytes) variant (branch SoftmaxSM86Forward . b"fc735688a7f0ef2ad998b8a5277c9ab9dd76abda85f43b0507895d9cd0554a54") (branch SoftmaxSM86CausalForward . b"116bfb4e6e0f75509a8177052d266ec16cd6415229bd52205ef483049d2f24c1") (branch SoftmaxSM86CausalForward1024 . b"3d5e834c9cc584ec239d3ddf71d13a4dc4784b53914b4093132fdaaa4a5cf231") (branch SoftmaxSM86CausalBackward1024 . b"7acd75b64174ef9001bf0ddef140557e8b768109414b7e9d4c99f5d8d65d602d") (branch SoftmaxSM86Backward . b"f56e51ab99ad911774d47e6c9859be3bc4532736be97a07e0a6ad84bca4c40b3"))) def softmaxSM86ABIValue : (family SoftmaxSM86ABI) = (constructor SoftmaxSM86ABI SoftmaxSM86ABIValue softmaxSM86N40 softmaxSM86N352 softmaxSM86N360 softmaxSM86N368 softmaxSM86N400) def softmaxSM86ExtentsFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (eliminate SoftmaxSM86Variant (lambda unrestricted current : (family SoftmaxSM86Variant) . (family SoftmaxSM86Extents)) variant (branch SoftmaxSM86Forward . (constructor SoftmaxSM86Extents SoftmaxSM86ExtentsValue softmaxSM86N256 softmaxSM86N256 zero softmaxSM86N256 softmaxSM86N1 softmaxSM86N1 softmaxSM86N256 (constructor SoftmaxSM86InputContract SoftmaxSM86FP32ScoresInput) (constructor SoftmaxSM86OutputContract SoftmaxSM86FP32ProbabilitiesOutput) (constructor SoftmaxSM86MaskContract SoftmaxSM86Unmasked))) (branch SoftmaxSM86CausalForward . (constructor SoftmaxSM86Extents SoftmaxSM86ExtentsValue softmaxSM86N256 softmaxSM86N256 zero softmaxSM86N256 softmaxSM86N1 softmaxSM86N1 softmaxSM86N256 (constructor SoftmaxSM86InputContract SoftmaxSM86FP32ScoresInput) (constructor SoftmaxSM86OutputContract SoftmaxSM86FP32ProbabilitiesOutput) (constructor SoftmaxSM86MaskContract SoftmaxSM86CausalLowEightBitRowMask))) (branch SoftmaxSM86CausalForward1024 . (constructor SoftmaxSM86Extents SoftmaxSM86ExtentsValue softmaxSM86N1024 softmaxSM86N1024 zero softmaxSM86N1024 softmaxSM86U4Natural softmaxSM86N1 softmaxSM86N256 (constructor SoftmaxSM86InputContract SoftmaxSM86FP32ScoresInput) (constructor SoftmaxSM86OutputContract SoftmaxSM86FP32ProbabilitiesOutput) (constructor SoftmaxSM86MaskContract SoftmaxSM86CausalFullRowMask))) (branch SoftmaxSM86CausalBackward1024 . (constructor SoftmaxSM86Extents SoftmaxSM86ExtentsValue softmaxSM86N1024 softmaxSM86N1024 softmaxSM86N1024 softmaxSM86N1024 softmaxSM86U4Natural softmaxSM86N1 softmaxSM86N256 (constructor SoftmaxSM86InputContract SoftmaxSM86FP32ProbabilitiesAndAdjointsInput) (constructor SoftmaxSM86OutputContract SoftmaxSM86FP32ScoreAdjointsOutput) (constructor SoftmaxSM86MaskContract SoftmaxSM86Unmasked))) (branch SoftmaxSM86Backward . (constructor SoftmaxSM86Extents SoftmaxSM86ExtentsValue softmaxSM86N256 softmaxSM86N256 softmaxSM86N256 softmaxSM86N256 softmaxSM86N1 softmaxSM86N1 softmaxSM86N256 (constructor SoftmaxSM86InputContract SoftmaxSM86FP32ProbabilitiesAndAdjointsInput) (constructor SoftmaxSM86OutputContract SoftmaxSM86FP32ScoreAdjointsOutput) (constructor SoftmaxSM86MaskContract SoftmaxSM86Unmasked))))) def softmaxSM86ManifestFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (constructor SoftmaxSM86Manifest SoftmaxSM86ManifestValue variant (softmaxSM86ExpectedInstructionsFor variant) (softmaxSM86ExpectedBytesFor variant) softmaxSM86N24 softmaxSM86N128 softmaxSM86N1 softmaxSM86N256 softmaxSM86ABIValue (softmaxSM86ExtentsFor variant) (softmaxSM86ExpectedPrefixFor variant) (softmaxSM86ExpectedMaximumReductionFor variant) (softmaxSM86ExpectedTransformFor variant) (softmaxSM86ExpectedSumReductionFor variant) (softmaxSM86ExpectedSuffixFor variant) (softmaxSM86ExpectedLoadsFor variant) softmaxSM86N1 zero (softmaxSM86ExpectedSHA256For variant))) def softmaxSM86ResourceCheckCons = (lambda unrestricted observed : Nat . (lambda unrestricted expected : Nat . (lambda unrestricted failure : (family SoftmaxSM86FailureCode) . (lambda unrestricted tail : (family SoftmaxSM86ResourceChecks) . (constructor SoftmaxSM86ResourceChecks SoftmaxSM86ResourceChecksNext (constructor SoftmaxSM86ResourceCheck SoftmaxSM86ResourceCheckValue observed expected failure) tail))))) def softmaxSM86ResourceChecksFor = (lambda unrestricted manifest : (family SoftmaxSM86Manifest) . (eliminate SoftmaxSM86Manifest (lambda unrestricted current : (family SoftmaxSM86Manifest) . (family SoftmaxSM86ResourceChecks)) manifest (branch SoftmaxSM86ManifestValue variant instructions encodedBytes registers sharedBytes blocksPerRow threadsPerBlock abi extents prefix maximumReduction transform sumReduction suffix loads stores hostFallback expectedIdentity . (softmaxSM86ResourceCheckCons instructions (softmaxSM86ExpectedInstructionsFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86InstructionCountMismatch) (softmaxSM86ResourceCheckCons encodedBytes (softmaxSM86ExpectedBytesFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86EncodedByteCountMismatch) (softmaxSM86ResourceCheckCons registers softmaxSM86N24 (constructor SoftmaxSM86FailureCode SoftmaxSM86RegisterCountMismatch) (softmaxSM86ResourceCheckCons sharedBytes softmaxSM86N128 (constructor SoftmaxSM86FailureCode SoftmaxSM86SharedByteCountMismatch) (softmaxSM86ResourceCheckCons blocksPerRow softmaxSM86N1 (constructor SoftmaxSM86FailureCode SoftmaxSM86BlockCountMismatch) (softmaxSM86ResourceCheckCons threadsPerBlock softmaxSM86N256 (constructor SoftmaxSM86FailureCode SoftmaxSM86ThreadCountMismatch) (softmaxSM86ResourceCheckCons prefix (softmaxSM86ExpectedPrefixFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86SectionCountMismatch) (softmaxSM86ResourceCheckCons maximumReduction (softmaxSM86ExpectedMaximumReductionFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86SectionCountMismatch) (softmaxSM86ResourceCheckCons transform (softmaxSM86ExpectedTransformFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86SectionCountMismatch) (softmaxSM86ResourceCheckCons sumReduction (softmaxSM86ExpectedSumReductionFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86SectionCountMismatch) (softmaxSM86ResourceCheckCons suffix (softmaxSM86ExpectedSuffixFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86SectionCountMismatch) (softmaxSM86ResourceCheckCons loads (softmaxSM86ExpectedLoadsFor variant) (constructor SoftmaxSM86FailureCode SoftmaxSM86GlobalLoadCountMismatch) (softmaxSM86ResourceCheckCons stores softmaxSM86N1 (constructor SoftmaxSM86FailureCode SoftmaxSM86GlobalStoreCountMismatch) (softmaxSM86ResourceCheckCons hostFallback zero (constructor SoftmaxSM86FailureCode SoftmaxSM86HostFallbackDetected) (constructor SoftmaxSM86ResourceChecks SoftmaxSM86ResourceChecksEnd)))))))))))))))))) def softmaxSM86RunResourceChecks = (lambda unrestricted checks : (family SoftmaxSM86ResourceChecks) . (eliminate SoftmaxSM86ResourceChecks (lambda unrestricted current : (family SoftmaxSM86ResourceChecks) . (family SoftmaxSM86ResourceGateResult)) checks (branch SoftmaxSM86ResourceChecksEnd . (constructor SoftmaxSM86ResourceGateResult SoftmaxSM86ResourcesExact)) (branch SoftmaxSM86ResourceChecksNext check tail induction . (eliminate SoftmaxSM86ResourceCheck (lambda unrestricted current : (family SoftmaxSM86ResourceCheck) . (family SoftmaxSM86ResourceGateResult)) check (branch SoftmaxSM86ResourceCheckValue observed expected failure . (nat-eliminate (lambda unrestricted exact : Nat . (family SoftmaxSM86ResourceGateResult)) (constructor SoftmaxSM86ResourceGateResult SoftmaxSM86ResourcesRejected failure) (lambda unrestricted predecessor : Nat . (lambda unrestricted checkInduction : (family SoftmaxSM86ResourceGateResult) . induction)) (naturalEqual observed expected))))))) def softmaxSM86ResourceGate = (lambda unrestricted manifest : (family SoftmaxSM86Manifest) . (softmaxSM86RunResourceChecks (softmaxSM86ResourceChecksFor manifest))) def softmaxSM86TelemetryFor = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (lambda unrestricted observedInstructions : Nat . (lambda unrestricted observedEncodedBytes : Nat . (lambda unrestricted encodingFields : Nat . (lambda unrestricted encodingBits : Nat . (lambda unrestricted highestEncodedBit : Nat . (lambda unrestricted identityInputBytes : Nat . (lambda unrestricted identityOutputBytes : Nat . (constructor SoftmaxSM86Telemetry SoftmaxSM86TelemetryValue (softmaxSM86ManifestFor variant) observedInstructions observedEncodedBytes encodingFields encodingBits highestEncodedBit identityInputBytes identityOutputBytes (softmaxSM86ExpectedLoadsFor variant) softmaxSM86N1 zero))))))))) def softmaxSM86EmptyTelemetry = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (lambda unrestricted observedInstructions : Nat . (softmaxSM86TelemetryFor variant observedInstructions zero zero zero zero zero zero))) def softmaxSM86BuildValidated = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (lambda unrestricted program : (family SM86Program) . (app (lambda unrestricted observedInstructions : Nat . (nat-eliminate (lambda unrestricted instructionCountExact : Nat . (family SoftmaxSM86BuildResult)) (constructor SoftmaxSM86BuildResult SoftmaxSM86ContractFailed (constructor SoftmaxSM86FailureCode SoftmaxSM86InstructionCountMismatch) (softmaxSM86EmptyTelemetry variant observedInstructions)) (lambda unrestricted instructionPredecessor : Nat . (lambda unrestricted instructionInduction : (family SoftmaxSM86BuildResult) . (app (lambda unrestricted encoding : (family SM86ProgramEncodingResult) . (eliminate SM86ProgramEncodingResult (lambda unrestricted current : (family SM86ProgramEncodingResult) . (family SoftmaxSM86BuildResult)) encoding (branch SM86ProgramEncodingSucceeded bytes encodingTelemetry . (eliminate SM86ProgramEncodingTelemetry (lambda unrestricted current : (family SM86ProgramEncodingTelemetry) . (family SoftmaxSM86BuildResult)) encodingTelemetry (branch SM86ProgramEncodingTelemetryValue encodedInstructions encodedBytes encodingFields encodingBits highestEncodedBit . (app (lambda unrestricted telemetry : (family SoftmaxSM86Telemetry) . (nat-eliminate (lambda unrestricted byteCountExact : Nat . (family SoftmaxSM86BuildResult)) (constructor SoftmaxSM86BuildResult SoftmaxSM86ContractFailed (constructor SoftmaxSM86FailureCode SoftmaxSM86EncodedByteCountMismatch) telemetry) (lambda unrestricted bytePredecessor : Nat . (lambda unrestricted byteInduction : (family SoftmaxSM86BuildResult) . (app (lambda unrestricted identityResult : (family SHA256HexResult) . (eliminate SHA256HexResult (lambda unrestricted current : (family SHA256HexResult) . (family SoftmaxSM86BuildResult)) identityResult (branch SHA256HexSucceeded identity identityTelemetry . (app (lambda unrestricted identityTelemetryValue : (family SoftmaxSM86Telemetry) . (nat-eliminate (lambda unrestricted identityLengthExact : Nat . (family SoftmaxSM86BuildResult)) (constructor SoftmaxSM86BuildResult SoftmaxSM86ImageIdentityFailed (constructor SoftmaxSM86FailureCode SoftmaxSM86IdentityLengthInvalid) identityResult identityTelemetryValue) (lambda unrestricted identityLengthPredecessor : Nat . (lambda unrestricted identityLengthInduction : (family SoftmaxSM86BuildResult) . (nat-eliminate (lambda unrestricted identityExact : Nat . (family SoftmaxSM86BuildResult)) (constructor SoftmaxSM86BuildResult SoftmaxSM86ImageIdentityFailed (constructor SoftmaxSM86FailureCode SoftmaxSM86IdentityMismatch) identityResult identityTelemetryValue) (lambda unrestricted identityPredecessor : Nat . (lambda unrestricted identityInduction : (family SoftmaxSM86BuildResult) . (constructor SoftmaxSM86BuildResult SoftmaxSM86BuildSucceeded bytes identity encodingTelemetry identityTelemetry identityTelemetryValue))) (bytes-equal identity (softmaxSM86ExpectedSHA256For variant))))) (naturalEqual (bytes-length identity) softmaxSM86N64))) (softmaxSM86TelemetryFor variant observedInstructions (bytes-length bytes) encodingFields encodingBits highestEncodedBit (bytes-length bytes) (bytes-length identity)))) (branch SHA256HexFailed error ordinal identityTelemetry . (constructor SoftmaxSM86BuildResult SoftmaxSM86ImageIdentityFailed (constructor SoftmaxSM86FailureCode SoftmaxSM86IdentityFailed) identityResult (softmaxSM86TelemetryFor variant observedInstructions (bytes-length bytes) encodingFields encodingBits highestEncodedBit (bytes-length bytes) zero))))) (sha256Hex bytes)))) (naturalEqual (bytes-length bytes) (softmaxSM86ExpectedBytesFor variant)))) (softmaxSM86TelemetryFor variant observedInstructions (bytes-length bytes) encodingFields encodingBits highestEncodedBit zero zero))))) (branch SM86ProgramEncodingFailed instructionIndex failure encodingTelemetry . (constructor SoftmaxSM86BuildResult SoftmaxSM86ImageEncodingFailed (constructor SoftmaxSM86FailureCode SoftmaxSM86EncodingFailed) encoding (softmaxSM86EmptyTelemetry variant observedInstructions))))) (sm86EncodeProgram program)))) (naturalEqual observedInstructions (softmaxSM86ExpectedInstructionsFor variant)))) (sm86ProgramCount program)))) def softmaxSM86Build = (lambda unrestricted variant : (family SoftmaxSM86Variant) . (app (lambda unrestricted program : (family SM86Program) . (eliminate SoftmaxSM86ResourceGateResult (lambda unrestricted current : (family SoftmaxSM86ResourceGateResult) . (family SoftmaxSM86BuildResult)) (softmaxSM86ResourceGate (softmaxSM86ManifestFor variant)) (branch SoftmaxSM86ResourcesExact . (softmaxSM86BuildValidated variant program)) (branch SoftmaxSM86ResourcesRejected code . (constructor SoftmaxSM86BuildResult SoftmaxSM86ContractFailed code (softmaxSM86EmptyTelemetry variant (sm86ProgramCount program)))))) (softmaxSM86ProgramFor variant))) def softmaxSM86ForwardBuild : (family SoftmaxSM86BuildResult) = (softmaxSM86Build (constructor SoftmaxSM86Variant SoftmaxSM86Forward)) def softmaxSM86CausalForwardBuild : (family SoftmaxSM86BuildResult) = (softmaxSM86Build (constructor SoftmaxSM86Variant SoftmaxSM86CausalForward)) def softmaxSM86BackwardBuild : (family SoftmaxSM86BuildResult) = (softmaxSM86Build (constructor SoftmaxSM86Variant SoftmaxSM86Backward))