Large source region · 2,114 lines
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))The compiler supplied declaration spans and resolved links from this source snapshot. This page does not assert that this file belongs to a checked closure.