Large source region · 2,228 lines
module Realization.Nvidia.SM86.LayerNormSM86
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 Std.Natural
import Std.Byte
family LayerNormSM86Variant : Type 0
constructor LayerNormSM86Forward64
constructor LayerNormSM86Forward256
constructor LayerNormSM86Forward512
constructor LayerNormSM86Forward1024
constructor LayerNormSM86InputBackward64
constructor LayerNormSM86InputBackward256
constructor LayerNormSM86InputBackward512
constructor LayerNormSM86InputBackward1024
constructor LayerNormSM86ParameterGradient64
constructor LayerNormSM86ParameterGradient256
constructor LayerNormSM86ParameterGradient512
constructor LayerNormSM86ParameterGradient1024
end-family
family LayerNormSM86ReductionShape : Type 0
constructor LayerNormSM86Reduction64
constructor LayerNormSM86Reduction256
constructor LayerNormSM86Reduction512
constructor LayerNormSM86Reduction1024
end-family
family LayerNormSM86FailureCode : Type 0
constructor LayerNormSM86RowsZero
constructor LayerNormSM86ParameterRowsMismatch
constructor LayerNormSM86InstructionCountMismatch
constructor LayerNormSM86EncodingFailed
constructor LayerNormSM86EncodedByteCountMismatch
constructor LayerNormSM86IdentityInvalid
end-family
family LayerNormSM86Extents : Type 0
constructor LayerNormSM86ExtentsValue
field unrestricted layerNormSM86ExtentVariant : (family LayerNormSM86Variant)
field unrestricted layerNormSM86ExtentRows : Nat
field unrestricted layerNormSM86ExtentWidth : Nat
field unrestricted layerNormSM86ExtentElements : Nat
end-family
family LayerNormSM86ScalarABI : Type 0
constructor LayerNormSM86ForwardScalarABI
field unrestricted layerNormSM86ForwardOutputPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ForwardInputPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ForwardWeightPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ForwardBiasPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ForwardStatsPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ForwardInverseWidthScalar : (family SM86Unsigned32)
field unrestricted layerNormSM86ForwardEpsilonScalar : (family SM86Unsigned32)
constructor LayerNormSM86InputBackwardScalarABI
field unrestricted layerNormSM86BackwardOutputPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86BackwardInputPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86BackwardGradientPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86BackwardWeightPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86BackwardNormalizedPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86BackwardStatsPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86BackwardInverseWidthScalar : (family SM86Unsigned32)
constructor LayerNormSM86ParameterGradientScalarABI
field unrestricted layerNormSM86ParameterWeightOutputPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ParameterGradientPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ParameterNormalizedPointer : (family SM86Unsigned32)
field unrestricted layerNormSM86ParameterBiasOutputPointer : (family SM86Unsigned32)
end-family
family LayerNormSM86Telemetry : Type 0
constructor LayerNormSM86TelemetryValue
field unrestricted layerNormSM86TelemetryVariant : (family LayerNormSM86Variant)
field unrestricted layerNormSM86TelemetryExpectedInstructions : Nat
field unrestricted layerNormSM86TelemetryActualInstructions : Nat
field unrestricted layerNormSM86TelemetryExpectedBytes : Nat
field unrestricted layerNormSM86TelemetryActualBytes : Nat
field unrestricted layerNormSM86TelemetryRegisterCount : Nat
field unrestricted layerNormSM86TelemetrySharedBytes : Nat
field unrestricted layerNormSM86TelemetryRows : Nat
field unrestricted layerNormSM86TelemetryWidth : Nat
field unrestricted layerNormSM86TelemetryElements : Nat
field unrestricted layerNormSM86TelemetryGridX : Nat
field unrestricted layerNormSM86TelemetryGridY : Nat
field unrestricted layerNormSM86TelemetryGridZ : Nat
field unrestricted layerNormSM86TelemetryBlockX : Nat
field unrestricted layerNormSM86TelemetryBlockY : Nat
field unrestricted layerNormSM86TelemetryBlockZ : Nat
field unrestricted layerNormSM86TelemetryConstantLoads : Nat
field unrestricted layerNormSM86TelemetryGlobalLoads : Nat
field unrestricted layerNormSM86TelemetryGlobalStores : Nat
field unrestricted layerNormSM86TelemetryReductionPasses : Nat
field unrestricted layerNormSM86TelemetryBarriers : Nat
field unrestricted layerNormSM86TelemetryAlgorithmPasses : Nat
field unrestricted layerNormSM86TelemetryHostFallbacks : Nat
field unrestricted layerNormSM86TelemetryEncodedFields : Nat
field unrestricted layerNormSM86TelemetryEncodedBits : Nat
field unrestricted layerNormSM86TelemetryHighestExclusiveBit : Nat
constructor LayerNormSM86TelemetryRejected
field unrestricted layerNormSM86TelemetryRejectedVariant : (family LayerNormSM86Variant)
field unrestricted layerNormSM86TelemetryFailure : (family LayerNormSM86FailureCode)
field unrestricted layerNormSM86TelemetryFailureOrdinal : Nat
end-family
family LayerNormSM86Manifest : Type 0
constructor LayerNormSM86ManifestValue
field unrestricted layerNormSM86ManifestVariant : (family LayerNormSM86Variant)
field unrestricted layerNormSM86ManifestExtents : (family LayerNormSM86Extents)
field unrestricted layerNormSM86ManifestScalarABI : (family LayerNormSM86ScalarABI)
field unrestricted layerNormSM86ManifestGridX : Nat
field unrestricted layerNormSM86ManifestGridY : Nat
field unrestricted layerNormSM86ManifestGridZ : Nat
field unrestricted layerNormSM86ManifestBlockX : Nat
field unrestricted layerNormSM86ManifestBlockY : Nat
field unrestricted layerNormSM86ManifestBlockZ : Nat
field unrestricted layerNormSM86ManifestRegisters : Nat
field unrestricted layerNormSM86ManifestSharedBytes : Nat
field unrestricted layerNormSM86ManifestExpectedInstructions : Nat
field unrestricted layerNormSM86ManifestExpectedEncodedBytes : Nat
field unrestricted layerNormSM86ManifestProgram : (family SM86Program)
field unrestricted layerNormSM86ManifestTelemetry : (family LayerNormSM86Telemetry)
end-family
family LayerNormSM86PlanResult : Type 0
constructor LayerNormSM86PlanReady
field unrestricted layerNormSM86ReadyManifest : (family LayerNormSM86Manifest)
constructor LayerNormSM86PlanFailed
field unrestricted layerNormSM86PlanFailure : (family LayerNormSM86FailureCode)
field unrestricted layerNormSM86PlanFailureTelemetry : (family LayerNormSM86Telemetry)
end-family
family LayerNormSM86ImageResult : Type 0
constructor LayerNormSM86ImageReady
field unrestricted layerNormSM86ImageVariant : (family LayerNormSM86Variant)
field unrestricted layerNormSM86ImageBytes : Bytes
field unrestricted layerNormSM86ImageSHA256 : Bytes
field unrestricted layerNormSM86ImageManifest : (family LayerNormSM86Manifest)
field unrestricted layerNormSM86ImageTelemetry : (family LayerNormSM86Telemetry)
constructor LayerNormSM86ImageFailed
field unrestricted layerNormSM86ImageFailure : (family LayerNormSM86FailureCode)
field unrestricted layerNormSM86ImageFailureInstruction : Nat
field unrestricted layerNormSM86ImageFailureCode : Bytes
field unrestricted layerNormSM86ImageFailureTelemetry : (family LayerNormSM86Telemetry)
end-family
def layerNormSM86NaturalOne =
(succ zero)
def layerNormSM86NaturalTwo =
(byte-to-nat (byte 2))
def layerNormSM86NaturalThree =
(byte-to-nat (byte 3))
def layerNormSM86NaturalFive =
(byte-to-nat (byte 5))
def layerNormSM86NaturalSixteen =
(byte-to-nat (byte 16))
def layerNormSM86NaturalTwentyFour =
(byte-to-nat (byte 24))
def layerNormSM86NaturalThirtyTwo =
(byte-to-nat (byte 32))
def layerNormSM86NaturalSixtyFour =
(byte-to-nat (byte 64))
def layerNormSM86NaturalNinetyThree =
(byte-to-nat (byte 93))
def layerNormSM86NaturalNinetySeven =
(byte-to-nat (byte 97))
def layerNormSM86NaturalOneHundredOne =
(byte-to-nat (byte 101))
def layerNormSM86NaturalOneHundredTwentyOne =
(byte-to-nat (byte 121))
def layerNormSM86NaturalOneHundredTwentyEight =
(byte-to-nat (byte 128))
def layerNormSM86NaturalTwoHundredFiftySix =
byteNaturalTwoHundredFiftySix
-- coppelius D24: width 512 (16 warps).
def layerNormSM86NaturalFiveHundredTwelve =
(naturalAdd byteNaturalTwoHundredFiftySix byteNaturalTwoHundredFiftySix)
def layerNormSM86NaturalOneThousandTwentyFour =
(naturalMultiply (byte-to-nat (byte 4)) byteNaturalTwoHundredFiftySix)
def layerNormSM86NaturalEight =
(byte-to-nat (byte 8))
def layerNormSM86NaturalSeventyThree =
(byte-to-nat (byte 73))
def layerNormSM86NaturalSixThousandOneHundredFortyFour =
(naturalMultiply (byte-to-nat (byte 24)) byteNaturalTwoHundredFiftySix)
def layerNormSM86Register =
(lambda unrestricted value : Byte . (sm86Register value))
def layerNormSM86Predicate0 =
(constructor SM86Predicate SM86Predicate0)
def layerNormSM86Unsigned32 =
(lambda unrestricted byte0 : Byte .
(lambda unrestricted byte1 : Byte .
(lambda unrestricted byte2 : Byte .
(lambda unrestricted byte3 : Byte . (sm86Unsigned32 byte0 byte1 byte2 byte3)))))
def layerNormSM86Zero =
(layerNormSM86Unsigned32 (byte 0) (byte 0) (byte 0) (byte 0))
def layerNormSM86One =
(layerNormSM86Unsigned32 (byte 1) (byte 0) (byte 0) (byte 0))
def layerNormSM86Two =
(layerNormSM86Unsigned32 (byte 2) (byte 0) (byte 0) (byte 0))
def layerNormSM86Four =
(layerNormSM86Unsigned32 (byte 4) (byte 0) (byte 0) (byte 0))
def layerNormSM86ThirtyOne =
(layerNormSM86Unsigned32 (byte 31) (byte 0) (byte 0) (byte 0))
def layerNormSM86ConstantBase =
(layerNormSM86Unsigned32 (byte 40) (byte 0) (byte 0) (byte 0))
def layerNormSM86OutputPointer =
(layerNormSM86Unsigned32 (byte 96) (byte 1) (byte 0) (byte 0))
def layerNormSM86InputPointer =
(layerNormSM86Unsigned32 (byte 104) (byte 1) (byte 0) (byte 0))
def layerNormSM86GradientPointer =
(layerNormSM86Unsigned32 (byte 112) (byte 1) (byte 0) (byte 0))
def layerNormSM86WeightPointer =
(layerNormSM86Unsigned32 (byte 120) (byte 1) (byte 0) (byte 0))
def layerNormSM86NormalizedPointer =
(layerNormSM86Unsigned32 (byte 128) (byte 1) (byte 0) (byte 0))
def layerNormSM86BackwardStatsPointer =
(layerNormSM86Unsigned32 (byte 136) (byte 1) (byte 0) (byte 0))
def layerNormSM86InverseWidthScalar =
(layerNormSM86Unsigned32 (byte 144) (byte 1) (byte 0) (byte 0))
def layerNormSM86EpsilonScalar =
(layerNormSM86Unsigned32 (byte 148) (byte 1) (byte 0) (byte 0))
def layerNormSM86Term1024 =
(layerNormSM86Unsigned32 (byte 0) (byte 4) (byte 0) (byte 0))
def layerNormSM86Term2048 =
(layerNormSM86Unsigned32 (byte 0) (byte 8) (byte 0) (byte 0))
def layerNormSM86Term3072 =
(layerNormSM86Unsigned32 (byte 0) (byte 12) (byte 0) (byte 0))
def layerNormSM86Term4096 =
(layerNormSM86Unsigned32 (byte 0) (byte 16) (byte 0) (byte 0))
def layerNormSM86Term5120 =
(layerNormSM86Unsigned32 (byte 0) (byte 20) (byte 0) (byte 0))
def layerNormSM86SetBarrier0 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0))
def layerNormSM86SetBarrier1 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier1))
def layerNormSM86SetBarrier2 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier2))
def layerNormSM86SetBarrier3 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier3))
def layerNormSM86SetBarrier4 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier4))
def layerNormSM86SetBarrier5 =
(sm86SetBarrierControl (constructor SM86Barrier SM86Barrier5))
def layerNormSM86WaitBarrier0 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0))
def layerNormSM86WaitBarrier1 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier1))
def layerNormSM86WaitBarrier3 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier3))
def layerNormSM86WaitBarrier4 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier4))
def layerNormSM86WaitBarrier5 =
(sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier5))
def layerNormSM86WaitSharedSetShuffle4 : (family SM86Control) =
(constructor
SM86Control
SM86ControlValue
(byte 1)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86Barrier4)
(constructor SM86Barrier SM86BarrierNone)
(byte 4)
(byte 0))
def layerNormSM86SetSharedAddressRead : (family SM86Control) =
(constructor
SM86Control
SM86ControlValue
(byte 7)
(constructor SM86YieldMode SM86Continue)
(constructor SM86Barrier SM86BarrierNone)
(constructor SM86Barrier SM86Barrier4)
(byte 0)
(byte 0))
def layerNormSM86MoveConstant =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted bank : Byte .
(lambda unrestricted offset : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86MoveConstant destination bank offset control)))))
def layerNormSM86SpecialToRegister =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted special : (family SM86SpecialRegister) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86SpecialToRegister destination special control))))
def layerNormSM86MoveImmediate =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted value : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86MoveImmediate destination value control))))
def layerNormSM86IMADImmediate =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted value : (family SM86Unsigned32) .
(lambda unrestricted addend : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddImmediate
destination
left
value
addend
control))))))
def layerNormSM86IMADWideConstant =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted right : (family SM86Register) .
(lambda unrestricted bank : Byte .
(lambda unrestricted offset : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86IntegerMultiplyAddWideConstant
destination
left
right
bank
offset
control)))))))
def layerNormSM86IADD3Immediate =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted value : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86IntegerAddThreeImmediate destination left value control)))))
def layerNormSM86ShiftRight =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted source : (family SM86Register) .
(lambda unrestricted amount : Byte .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86ShiftRightImmediate
destination
source
amount
control)))))
def layerNormSM86Logic3 =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted right : (family SM86Register) .
(lambda unrestricted truth : Byte .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86LogicThreeInputTruthTable destination left right truth control))))))
def layerNormSM86FloatAdd =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted right : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86FloatAdd destination left right control)))))
def layerNormSM86FloatMultiply =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted right : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86FloatMultiply destination left right control)))))
def layerNormSM86FloatFMA =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted left : (family SM86Register) .
(lambda unrestricted right : (family SM86Register) .
(lambda unrestricted addend : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86FloatFusedMultiplyAdd
destination
left
right
addend
control))))))
def layerNormSM86FloatNegate =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted source : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86FloatNegate destination source control))))
def layerNormSM86MultiFunction =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted source : (family SM86Register) .
(lambda unrestricted operation : (family SM86MultiFunction) .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86MultiFunctionUnitApproximation
destination
source
operation
control)))))
def layerNormSM86LoadGlobal =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted address : (family SM86Register) .
(lambda unrestricted offset : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86LoadGlobal destination address offset control)))))
def layerNormSM86Shuffle =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted source : (family SM86Register) .
(lambda unrestricted lane : Byte .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86WarpShuffle
destination
source
lane
layerNormSM86ThirtyOne
(constructor SM86ShuffleMode SM86ShuffleButterfly)
control)))))
def layerNormSM86LoadShared =
(lambda unrestricted destination : (family SM86Register) .
(lambda unrestricted address : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86LoadShared
destination
address
layerNormSM86Zero
control))))
def layerNormSM86StoreShared =
(lambda unrestricted address : (family SM86Register) .
(lambda unrestricted value : (family SM86Register) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86StoreShared address value layerNormSM86Zero control))))
def layerNormSM86Barrier =
(constructor SM86InstructionBody SM86BarrierSynchronize sm86SafeControl)
def layerNormSM86StoreGlobal =
(lambda unrestricted address : (family SM86Register) .
(lambda unrestricted value : (family SM86Register) .
(lambda unrestricted offset : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor SM86InstructionBody SM86StoreGlobal address value offset control)))))
def layerNormSM86PredicateGreater =
(lambda unrestricted source : (family SM86Register) .
(lambda unrestricted immediate : (family SM86Unsigned32) .
(lambda unrestricted control : (family SM86Control) .
(constructor
SM86InstructionBody
SM86PredicateGreaterThanImmediate
layerNormSM86Predicate0
source
immediate
control))))
def layerNormSM86Exit =
(constructor SM86InstructionBody SM86Exit sm86SafeControl)
def layerNormSM86Next =
(lambda unrestricted body : (family SM86InstructionBody) .
(lambda unrestricted tail : (family SM86Program) .
(constructor SM86Program SM86ProgramNext (sm86Instruction body) tail)))
def layerNormSM86NextWhenNotPredicate0 =
(lambda unrestricted body : (family SM86InstructionBody) .
(lambda unrestricted tail : (family SM86Program) .
(constructor
SM86Program
SM86ProgramNext
(sm86NegatedPredicatedInstruction layerNormSM86Predicate0 body)
tail)))
def layerNormSM86End =
(constructor SM86Program SM86ProgramEnd)
def layerNormSM86WidthImmediate =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . (family SM86Unsigned32))
variant
(branch
LayerNormSM86Forward64
.
(layerNormSM86Unsigned32 (byte 64) (byte 0) (byte 0) (byte 0)))
(branch
LayerNormSM86Forward256
.
(layerNormSM86Unsigned32 (byte 0) (byte 1) (byte 0) (byte 0)))
(branch
LayerNormSM86Forward512
.
(layerNormSM86Unsigned32 (byte 0) (byte 2) (byte 0) (byte 0)))
(branch LayerNormSM86Forward1024 . layerNormSM86Term1024)
(branch
LayerNormSM86InputBackward64
.
(layerNormSM86Unsigned32 (byte 64) (byte 0) (byte 0) (byte 0)))
(branch
LayerNormSM86InputBackward256
.
(layerNormSM86Unsigned32 (byte 0) (byte 1) (byte 0) (byte 0)))
(branch
LayerNormSM86InputBackward512
.
(layerNormSM86Unsigned32 (byte 0) (byte 2) (byte 0) (byte 0)))
(branch LayerNormSM86InputBackward1024 . layerNormSM86Term1024)
(branch
LayerNormSM86ParameterGradient64
.
(layerNormSM86Unsigned32 (byte 64) (byte 0) (byte 0) (byte 0)))
(branch
LayerNormSM86ParameterGradient256
.
(layerNormSM86Unsigned32 (byte 0) (byte 1) (byte 0) (byte 0)))
(branch
LayerNormSM86ParameterGradient512
.
(layerNormSM86Unsigned32 (byte 0) (byte 2) (byte 0) (byte 0)))
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86Term1024)))
def layerNormSM86Width =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalSixtyFour)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalTwoHundredFiftySix)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalFiveHundredTwelve)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalOneThousandTwentyFour)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalSixtyFour)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalTwoHundredFiftySix)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalFiveHundredTwelve)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalOneThousandTwentyFour)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalSixtyFour)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalTwoHundredFiftySix)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalFiveHundredTwelve)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalOneThousandTwentyFour)))
-- the parameter gradient's block: a block per column, a thread per 256th
-- row (the row terms' immediates are term x 256: layerNormSM86ParameterTermImmediate)
def layerNormSM86ParameterBlockThreads : Nat = layerNormSM86NaturalTwoHundredFiftySix
-- The threads of a variant's block: a thread per column for the forward and
-- the input backward, and 256 for the parameter gradient (a block per
-- column, a thread per 256th row; layerNormSM86ParameterTerms).
def layerNormSM86BlockX =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalSixtyFour)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalTwoHundredFiftySix)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalFiveHundredTwelve)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalOneThousandTwentyFour)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalSixtyFour)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalTwoHundredFiftySix)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalFiveHundredTwelve)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalOneThousandTwentyFour)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86ParameterBlockThreads)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86ParameterBlockThreads)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86ParameterBlockThreads)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86ParameterBlockThreads)))
-- The block reduction's shape is the block's: a warp's partial for each of
-- its warps, and every other lane of the second stage reading the identity.
-- (The parameter gradient's was the 1024-thread shape under its 256-thread
-- block: lanes 8..31 of the second stage read shared memory no warp had
-- written -- what an earlier block on the SM left there -- and every norm's
-- gain and shift gradient carried it, differently from run to run.)
def layerNormSM86ShapeWhen =
(lambda unrestricted flag : Nat .
(lambda unrestricted chosen : (family LayerNormSM86ReductionShape) .
(lambda unrestricted otherwise : (family LayerNormSM86ReductionShape) .
(nat-eliminate (lambda unrestricted current : Nat . (family LayerNormSM86ReductionShape)) otherwise
(lambda unrestricted predecessor : Nat . (lambda unrestricted ignored : (family LayerNormSM86ReductionShape) . chosen))
flag))))
def layerNormSM86ReductionShapeFor =
(lambda unrestricted threads : Nat .
(layerNormSM86ShapeWhen (naturalEqual threads layerNormSM86NaturalSixtyFour)
(constructor LayerNormSM86ReductionShape LayerNormSM86Reduction64)
(layerNormSM86ShapeWhen (naturalEqual threads layerNormSM86NaturalTwoHundredFiftySix)
(constructor LayerNormSM86ReductionShape LayerNormSM86Reduction256)
(layerNormSM86ShapeWhen (naturalEqual threads layerNormSM86NaturalFiveHundredTwelve)
(constructor LayerNormSM86ReductionShape LayerNormSM86Reduction512)
(constructor LayerNormSM86ReductionShape LayerNormSM86Reduction1024)))))
def layerNormSM86ReductionShape =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(layerNormSM86ReductionShapeFor (layerNormSM86BlockX variant)))
-- the warps whose partials a reduction shape's second stage reads
def layerNormSM86ReductionWarps =
(lambda unrestricted shape : (family LayerNormSM86ReductionShape) .
(eliminate LayerNormSM86ReductionShape (lambda unrestricted current : (family LayerNormSM86ReductionShape) . Nat) shape
(branch LayerNormSM86Reduction64 . 2)
(branch LayerNormSM86Reduction256 . 8)
(branch LayerNormSM86Reduction512 . 16)
(branch LayerNormSM86Reduction1024 . 32)))
def layerNormSM86Butterfly =
(lambda unrestricted lane : Byte .
(lambda unrestricted producer : (family SM86Control) .
(lambda unrestricted consumer : (family SM86Control) .
(lambda unrestricted accumulator : (family SM86Register) .
(lambda unrestricted scratch : (family SM86Register) .
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86Shuffle scratch accumulator lane producer)
(layerNormSM86Next
(layerNormSM86FloatAdd accumulator accumulator scratch consumer)
tail))))))))
def layerNormSM86Butterflies =
(lambda unrestricted firstProducer : (family SM86Control) .
(lambda unrestricted secondProducer : (family SM86Control) .
(lambda unrestricted firstConsumer : (family SM86Control) .
(lambda unrestricted secondConsumer : (family SM86Control) .
(lambda unrestricted accumulator : (family SM86Register) .
(lambda unrestricted scratch : (family SM86Register) .
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Butterfly
(byte 16)
firstProducer
firstConsumer
accumulator
scratch
(layerNormSM86Butterfly
(byte 8)
secondProducer
secondConsumer
accumulator
scratch
(layerNormSM86Butterfly
(byte 4)
firstProducer
firstConsumer
accumulator
scratch
(layerNormSM86Butterfly
(byte 2)
secondProducer
secondConsumer
accumulator
scratch
(layerNormSM86Butterfly
(byte 1)
firstProducer
firstConsumer
accumulator
scratch
tail))))))))))))
def layerNormSM86ReductionPartial =
(lambda unrestricted shape : (family LayerNormSM86ReductionShape) .
(lambda unrestricted accumulator : (family SM86Register) .
(lambda unrestricted scratch : (family SM86Register) .
(lambda unrestricted tail : (family SM86Program) .
(eliminate
LayerNormSM86ReductionShape
(lambda unrestricted current : (family LayerNormSM86ReductionShape) .
(family SM86Program))
shape
(branch
LayerNormSM86Reduction64
.
(layerNormSM86Next
(layerNormSM86MoveImmediate accumulator layerNormSM86Zero sm86SafeControl)
(layerNormSM86Next
(layerNormSM86PredicateGreater scratch layerNormSM86One sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86LoadShared accumulator scratch layerNormSM86SetBarrier2)
tail))))
(branch
LayerNormSM86Reduction256
.
(layerNormSM86Next
(layerNormSM86MoveImmediate accumulator layerNormSM86Zero sm86SafeControl)
(layerNormSM86Next
(layerNormSM86PredicateGreater
scratch
(layerNormSM86Unsigned32 (byte 7) (byte 0) (byte 0) (byte 0))
sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86LoadShared accumulator scratch layerNormSM86SetBarrier2)
tail))))
(branch
LayerNormSM86Reduction512
.
(layerNormSM86Next
(layerNormSM86MoveImmediate accumulator layerNormSM86Zero sm86SafeControl)
(layerNormSM86Next
(layerNormSM86PredicateGreater
scratch
(layerNormSM86Unsigned32 (byte 15) (byte 0) (byte 0) (byte 0))
sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86LoadShared accumulator scratch layerNormSM86SetBarrier2)
tail))))
(branch
LayerNormSM86Reduction1024
.
(layerNormSM86Next
(layerNormSM86LoadShared accumulator scratch layerNormSM86SetBarrier2)
tail)))))))
def layerNormSM86Reduction =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted thread : (family SM86Register) .
(lambda unrestricted accumulator : (family SM86Register) .
(lambda unrestricted scratch : (family SM86Register) .
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
layerNormSM86Barrier
(layerNormSM86Butterflies
layerNormSM86SetBarrier4
layerNormSM86SetBarrier5
layerNormSM86WaitBarrier4
layerNormSM86WaitBarrier5
accumulator
scratch
(layerNormSM86Next
(layerNormSM86MoveImmediate scratch layerNormSM86ThirtyOne sm86SafeControl)
(layerNormSM86Next
(layerNormSM86Logic3 scratch thread scratch (byte 192) sm86SafeControl)
(layerNormSM86Next
(layerNormSM86PredicateGreater scratch layerNormSM86Zero sm86SafeControl)
(layerNormSM86Next
(layerNormSM86ShiftRight scratch thread (byte 5) sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86StoreShared
scratch
accumulator
layerNormSM86SetSharedAddressRead)
(layerNormSM86Next
layerNormSM86Barrier
(layerNormSM86Next
(layerNormSM86MoveImmediate
scratch
layerNormSM86ThirtyOne
layerNormSM86WaitBarrier4)
(layerNormSM86Next
(layerNormSM86Logic3
scratch
thread
scratch
(byte 192)
sm86SafeControl)
(layerNormSM86ReductionPartial
(layerNormSM86ReductionShape variant)
accumulator
scratch
(layerNormSM86Butterflies
layerNormSM86WaitSharedSetShuffle4
layerNormSM86SetBarrier5
layerNormSM86WaitBarrier4
layerNormSM86WaitBarrier5
accumulator
scratch
tail)))))))))))))))))
def layerNormSM86ForwardPrefixBase =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86MoveConstant
(layerNormSM86Register (byte 1))
(byte 0)
layerNormSM86ConstantBase
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86SpecialToRegister
(layerNormSM86Register (byte 0))
(constructor SM86SpecialRegister SM86ThreadIdX)
layerNormSM86SetBarrier0)
(layerNormSM86Next
(layerNormSM86SpecialToRegister
(layerNormSM86Register (byte 1))
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
layerNormSM86SetBarrier0)
(layerNormSM86Next
(layerNormSM86MoveImmediate
(layerNormSM86Register (byte 5))
layerNormSM86Four
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADImmediate
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 1))
(layerNormSM86WidthImmediate variant)
(layerNormSM86Register (byte 0))
layerNormSM86WaitBarrier0)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86InputPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 2))
layerNormSM86Zero
layerNormSM86SetBarrier1)
tail)))))))))
def layerNormSM86ForwardAffineLoads =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 20))
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86GradientPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 20))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 22))
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86WeightPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 18))
(layerNormSM86Register (byte 22))
layerNormSM86Zero
layerNormSM86SetBarrier1)
tail)))))
def layerNormSM86ForwardMeanInput =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86IADD3Immediate
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 4))
layerNormSM86Zero
layerNormSM86WaitBarrier1)
tail))
def layerNormSM86ForwardMiddle =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86MoveConstant
(layerNormSM86Register (byte 12))
(byte 0)
layerNormSM86InverseWidthScalar
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 12))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatNegate
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 8))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 9))
sm86SafeControl)
(layerNormSM86Next
layerNormSM86Barrier
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
sm86SafeControl)
tail)))))))
def layerNormSM86ForwardNormalize =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86MoveConstant
(layerNormSM86Register (byte 13))
(byte 0)
layerNormSM86EpsilonScalar
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 12))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 13))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86MultiFunction
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 9))
(constructor SM86MultiFunction SM86ReciprocalSquareRoot)
layerNormSM86SetBarrier3)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 9))
layerNormSM86WaitBarrier3)
tail))))))
def layerNormSM86ForwardStats =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86PredicateGreater
(layerNormSM86Register (byte 0))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADImmediate
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 1))
layerNormSM86Two
sm86ZeroRegister
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86NormalizedPointer
sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 8))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 9))
layerNormSM86Four
sm86SafeControl)
tail))))))
def layerNormSM86ForwardSuffix =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 17))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 18))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 10))
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86OutputPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 10))
(layerNormSM86Register (byte 4))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next layerNormSM86Exit tail))))))
def layerNormSM86ForwardProgram =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(layerNormSM86ForwardPrefixBase
variant
(layerNormSM86ForwardAffineLoads
(layerNormSM86ForwardMeanInput
(layerNormSM86Reduction
variant
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 7))
(layerNormSM86ForwardMiddle
(layerNormSM86Reduction
variant
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 7))
(layerNormSM86ForwardNormalize
(layerNormSM86ForwardStats (layerNormSM86ForwardSuffix layerNormSM86End))))))))))
def layerNormSM86BackwardPrefix =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86MoveConstant
(layerNormSM86Register (byte 1))
(byte 0)
layerNormSM86ConstantBase
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86SpecialToRegister
(layerNormSM86Register (byte 0))
(constructor SM86SpecialRegister SM86ThreadIdX)
layerNormSM86SetBarrier0)
(layerNormSM86Next
(layerNormSM86SpecialToRegister
(layerNormSM86Register (byte 1))
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
layerNormSM86SetBarrier0)
(layerNormSM86Next
(layerNormSM86MoveImmediate
(layerNormSM86Register (byte 5))
layerNormSM86Four
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADImmediate
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 1))
(layerNormSM86WidthImmediate variant)
(layerNormSM86Register (byte 0))
layerNormSM86WaitBarrier0)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86InputPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 22))
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86GradientPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 24))
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86WeightPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 2))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 22))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 18))
(layerNormSM86Register (byte 24))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86MoveConstant
(layerNormSM86Register (byte 12))
(byte 0)
layerNormSM86InverseWidthScalar
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADImmediate
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 1))
layerNormSM86Two
sm86ZeroRegister
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 28))
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86BackwardStatsPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 15))
(layerNormSM86Register (byte 28))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 19))
(layerNormSM86Register (byte 28))
layerNormSM86Four
layerNormSM86SetBarrier1)
tail))))))))))))))))))
def layerNormSM86BackwardPrepare =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86FloatNegate
(layerNormSM86Register (byte 14))
(layerNormSM86Register (byte 15))
layerNormSM86WaitBarrier1)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 14))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 19))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 26))
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86NormalizedPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 26))
(layerNormSM86Register (byte 4))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 18))
sm86SafeControl)
(layerNormSM86Next
layerNormSM86Barrier
(layerNormSM86Next
(layerNormSM86IADD3Immediate
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 17))
layerNormSM86Zero
sm86SafeControl)
tail)))))))))
def layerNormSM86BackwardMiddle =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 20))
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 12))
sm86SafeControl)
(layerNormSM86Next
layerNormSM86Barrier
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 4))
sm86SafeControl)
tail))))
def layerNormSM86BackwardSuffix =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 21))
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 12))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatNegate
(layerNormSM86Register (byte 14))
(layerNormSM86Register (byte 20))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 14))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 14))
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 21))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatNegate
(layerNormSM86Register (byte 14))
(layerNormSM86Register (byte 14))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 14))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86FloatMultiply
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 17))
(layerNormSM86Register (byte 19))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 10))
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 5))
(byte 0)
layerNormSM86OutputPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 10))
(layerNormSM86Register (byte 17))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next layerNormSM86Exit tail)))))))))))
def layerNormSM86BackwardProgram =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(layerNormSM86BackwardPrefix
variant
(layerNormSM86BackwardPrepare
(layerNormSM86Reduction
variant
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 9))
(layerNormSM86BackwardMiddle
(layerNormSM86Reduction
variant
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 9))
(layerNormSM86BackwardSuffix layerNormSM86End)))))))
def layerNormSM86ParameterPrefix =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86MoveConstant
(layerNormSM86Register (byte 1))
(byte 0)
layerNormSM86ConstantBase
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86SpecialToRegister
(layerNormSM86Register (byte 0))
(constructor SM86SpecialRegister SM86ThreadIdX)
layerNormSM86SetBarrier0)
(layerNormSM86Next
(layerNormSM86SpecialToRegister
(layerNormSM86Register (byte 1))
(constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
layerNormSM86SetBarrier0)
(layerNormSM86Next
(layerNormSM86MoveImmediate
(layerNormSM86Register (byte 14))
layerNormSM86Four
layerNormSM86WaitBarrier0)
(layerNormSM86Next
(layerNormSM86MoveImmediate
(layerNormSM86Register (byte 10))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86MoveImmediate
(layerNormSM86Register (byte 11))
layerNormSM86Zero
sm86SafeControl)
tail)))))))
def layerNormSM86ParameterRowTerm =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted term : (family SM86Unsigned32) .
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86IADD3Immediate
(layerNormSM86Register (byte 13))
(layerNormSM86Register (byte 0))
term
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADImmediate
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 13))
(layerNormSM86WidthImmediate variant)
(layerNormSM86Register (byte 1))
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 4))
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 14))
(byte 0)
layerNormSM86InputPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 6))
(layerNormSM86Register (byte 2))
(layerNormSM86Register (byte 14))
(byte 0)
layerNormSM86GradientPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 4))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86LoadGlobal
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 6))
layerNormSM86Zero
layerNormSM86SetBarrier1)
(layerNormSM86Next
(layerNormSM86FloatFMA
(layerNormSM86Register (byte 10))
(layerNormSM86Register (byte 8))
(layerNormSM86Register (byte 9))
(layerNormSM86Register (byte 10))
layerNormSM86WaitBarrier1)
(layerNormSM86Next
(layerNormSM86FloatAdd
(layerNormSM86Register (byte 11))
(layerNormSM86Register (byte 11))
(layerNormSM86Register (byte 8))
sm86SafeControl)
tail)))))))))))
def layerNormSM86ParameterSuffix =
(lambda unrestricted tail : (family SM86Program) .
(layerNormSM86Next
(layerNormSM86PredicateGreater
(layerNormSM86Register (byte 0))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 1))
(layerNormSM86Register (byte 14))
(byte 0)
layerNormSM86OutputPointer
sm86SafeControl)
(layerNormSM86Next
(layerNormSM86IMADWideConstant
(layerNormSM86Register (byte 18))
(layerNormSM86Register (byte 1))
(layerNormSM86Register (byte 14))
(byte 0)
layerNormSM86WeightPointer
sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 16))
(layerNormSM86Register (byte 10))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86NextWhenNotPredicate0
(layerNormSM86StoreGlobal
(layerNormSM86Register (byte 18))
(layerNormSM86Register (byte 11))
layerNormSM86Zero
sm86SafeControl)
(layerNormSM86Next layerNormSM86Exit tail)))))))
-- the parameter-gradient kernel sums dy (bias) and dy·x̂ (weight) over the rows:
-- a block of 256 threads per column, thread t covering rows t, t+256, t+512, …
-- (one row term per 256 rows), then a block reduction. Rows must be a nonzero
-- multiple of 256 up to 6144 (coppelius: 1024 → 4 terms; the D-ladder: 6144 → 24).
-- D26: the previous fixed six terms were spaced 1024 apart under a 256-thread
-- block, so only a quarter of the rows were summed; the spacing now follows
-- the block size.
def layerNormSM86ParameterTerms =
(lambda unrestricted rows : Nat .
(naturalDivideUnchecked rows layerNormSM86ParameterBlockThreads))
def layerNormSM86ParameterTermImmediate =
(lambda unrestricted term : Nat .
(layerNormSM86Unsigned32 (byte 0) (nat-to-byte term) (byte 0) (byte 0)))
def layerNormSM86ParameterRowTerms =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted terms : Nat .
(nat-eliminate
(lambda unrestricted current : Nat .
(pi unrestricted tail : (family SM86Program) . (family SM86Program)))
(lambda unrestricted tail : (family SM86Program) . tail)
(lambda unrestricted term : Nat .
(lambda unrestricted inner : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) .
(lambda unrestricted tail : (family SM86Program) .
(inner
(layerNormSM86ParameterRowTerm
variant
(layerNormSM86ParameterTermImmediate term)
tail)))))
terms)))
def layerNormSM86ParameterProgram =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(layerNormSM86ParameterPrefix
(layerNormSM86ParameterRowTerms
variant
(layerNormSM86ParameterTerms rows)
(layerNormSM86Reduction
variant
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 10))
(layerNormSM86Register (byte 12))
(layerNormSM86Next
layerNormSM86Barrier
(layerNormSM86Reduction
variant
(layerNormSM86Register (byte 0))
(layerNormSM86Register (byte 11))
(layerNormSM86Register (byte 12))
(layerNormSM86ParameterSuffix layerNormSM86End))))))))
def layerNormSM86Program =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . (family SM86Program))
variant
(branch LayerNormSM86Forward64 . (layerNormSM86ForwardProgram variant))
(branch LayerNormSM86Forward256 . (layerNormSM86ForwardProgram variant))
(branch LayerNormSM86Forward512 . (layerNormSM86ForwardProgram variant))
(branch LayerNormSM86Forward1024 . (layerNormSM86ForwardProgram variant))
(branch LayerNormSM86InputBackward64 . (layerNormSM86BackwardProgram variant))
(branch LayerNormSM86InputBackward256 . (layerNormSM86BackwardProgram variant))
(branch LayerNormSM86InputBackward512 . (layerNormSM86BackwardProgram variant))
(branch LayerNormSM86InputBackward1024 . (layerNormSM86BackwardProgram variant))
(branch LayerNormSM86ParameterGradient64 . (layerNormSM86ParameterProgram variant rows))
(branch LayerNormSM86ParameterGradient256 . (layerNormSM86ParameterProgram variant rows))
(branch LayerNormSM86ParameterGradient512 . (layerNormSM86ParameterProgram variant rows))
(branch LayerNormSM86ParameterGradient1024 . (layerNormSM86ParameterProgram variant rows)))))
def layerNormSM86ExpectedInstructions =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalNinetySeven)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalNinetySeven)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalNinetySeven)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalNinetyThree)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalOneHundredOne)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalOneHundredOne)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalOneHundredOne)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalNinetySeven)
(branch
LayerNormSM86ParameterGradient64
.
(naturalAdd
layerNormSM86NaturalSeventyThree
(naturalMultiply layerNormSM86NaturalEight (layerNormSM86ParameterTerms rows))))
(branch
LayerNormSM86ParameterGradient256
.
(naturalAdd
layerNormSM86NaturalSeventyThree
(naturalMultiply layerNormSM86NaturalEight (layerNormSM86ParameterTerms rows))))
(branch
LayerNormSM86ParameterGradient512
.
(naturalAdd
layerNormSM86NaturalSeventyThree
(naturalMultiply layerNormSM86NaturalEight (layerNormSM86ParameterTerms rows))))
(branch
LayerNormSM86ParameterGradient1024
.
(naturalAdd
layerNormSM86NaturalSeventyThree
(naturalMultiply layerNormSM86NaturalEight (layerNormSM86ParameterTerms rows)))))))
def layerNormSM86ExpectedEncodedBytes =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(naturalMultiply (layerNormSM86ExpectedInstructions variant rows) layerNormSM86NaturalSixteen)))
def layerNormSM86RegisterCount =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalThirtyTwo)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalTwentyFour)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalTwentyFour)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalTwentyFour)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalTwentyFour)))
def layerNormSM86SharedBytes =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
layerNormSM86NaturalOneHundredTwentyEight)
def layerNormSM86GridX =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . rows)
(branch LayerNormSM86Forward256 . rows)
(branch LayerNormSM86Forward512 . rows)
(branch LayerNormSM86Forward1024 . rows)
(branch LayerNormSM86InputBackward64 . rows)
(branch LayerNormSM86InputBackward256 . rows)
(branch LayerNormSM86InputBackward512 . rows)
(branch LayerNormSM86InputBackward1024 . rows)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalSixtyFour)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalTwoHundredFiftySix)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalFiveHundredTwelve)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalOneThousandTwentyFour))))
def layerNormSM86RowsValid =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalOne)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalOne)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalOne)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalOne)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalOne)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalOne)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalOne)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalOne)
(branch
LayerNormSM86ParameterGradient64
.
(naturalAnd
(naturalIsZero (naturalModuloUnchecked rows layerNormSM86NaturalTwoHundredFiftySix))
(naturalLessOrEqual rows layerNormSM86NaturalSixThousandOneHundredFortyFour)))
(branch
LayerNormSM86ParameterGradient256
.
(naturalAnd
(naturalIsZero (naturalModuloUnchecked rows layerNormSM86NaturalTwoHundredFiftySix))
(naturalLessOrEqual rows layerNormSM86NaturalSixThousandOneHundredFortyFour)))
(branch
LayerNormSM86ParameterGradient512
.
(naturalAnd
(naturalIsZero (naturalModuloUnchecked rows layerNormSM86NaturalTwoHundredFiftySix))
(naturalLessOrEqual rows layerNormSM86NaturalSixThousandOneHundredFortyFour)))
(branch
LayerNormSM86ParameterGradient1024
.
(naturalAnd
(naturalIsZero (naturalModuloUnchecked rows layerNormSM86NaturalTwoHundredFiftySix))
(naturalLessOrEqual rows layerNormSM86NaturalSixThousandOneHundredFortyFour))))))
def layerNormSM86ScalarABI =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) .
(family LayerNormSM86ScalarABI))
variant
(branch
LayerNormSM86Forward64
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ForwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86InverseWidthScalar
layerNormSM86EpsilonScalar))
(branch
LayerNormSM86Forward256
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ForwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86InverseWidthScalar
layerNormSM86EpsilonScalar))
(branch
LayerNormSM86Forward512
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ForwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86InverseWidthScalar
layerNormSM86EpsilonScalar))
(branch
LayerNormSM86Forward1024
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ForwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86InverseWidthScalar
layerNormSM86EpsilonScalar))
(branch
LayerNormSM86InputBackward64
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86InputBackwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86BackwardStatsPointer
layerNormSM86InverseWidthScalar))
(branch
LayerNormSM86InputBackward256
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86InputBackwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86BackwardStatsPointer
layerNormSM86InverseWidthScalar))
(branch
LayerNormSM86InputBackward512
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86InputBackwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86BackwardStatsPointer
layerNormSM86InverseWidthScalar))
(branch
LayerNormSM86InputBackward1024
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86InputBackwardScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer
layerNormSM86NormalizedPointer
layerNormSM86BackwardStatsPointer
layerNormSM86InverseWidthScalar))
(branch
LayerNormSM86ParameterGradient64
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ParameterGradientScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer))
(branch
LayerNormSM86ParameterGradient256
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ParameterGradientScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer))
(branch
LayerNormSM86ParameterGradient512
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ParameterGradientScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer))
(branch
LayerNormSM86ParameterGradient1024
.
(constructor
LayerNormSM86ScalarABI
LayerNormSM86ParameterGradientScalarABI
layerNormSM86OutputPointer
layerNormSM86InputPointer
layerNormSM86GradientPointer
layerNormSM86WeightPointer))))
def layerNormSM86ConstantLoads =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalThree)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalTwo)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalTwo)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalTwo)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalTwo)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalOne)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalOne)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalOne)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalOne)))
def layerNormSM86GlobalLoads =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalThree)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalFive)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalFive)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalFive)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalFive)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalTwentyFour)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalTwentyFour)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalTwentyFour)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalTwentyFour)))
def layerNormSM86GlobalStores =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalThree)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalThree)
(branch LayerNormSM86InputBackward64 . layerNormSM86NaturalTwo)
(branch LayerNormSM86InputBackward256 . layerNormSM86NaturalTwo)
(branch LayerNormSM86InputBackward512 . layerNormSM86NaturalTwo)
(branch LayerNormSM86InputBackward1024 . layerNormSM86NaturalTwo)
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalTwo)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalTwo)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalTwo)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalTwo)))
def layerNormSM86Barriers =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(eliminate
LayerNormSM86Variant
(lambda unrestricted current : (family LayerNormSM86Variant) . Nat)
variant
(branch LayerNormSM86Forward64 . layerNormSM86NaturalFive)
(branch LayerNormSM86Forward256 . layerNormSM86NaturalFive)
(branch LayerNormSM86Forward512 . layerNormSM86NaturalFive)
(branch LayerNormSM86Forward1024 . layerNormSM86NaturalFive)
(branch LayerNormSM86InputBackward64 . (byte-to-nat (byte 6)))
(branch LayerNormSM86InputBackward256 . (byte-to-nat (byte 6)))
(branch LayerNormSM86InputBackward512 . (byte-to-nat (byte 6)))
(branch LayerNormSM86InputBackward1024 . (byte-to-nat (byte 6)))
(branch LayerNormSM86ParameterGradient64 . layerNormSM86NaturalFive)
(branch LayerNormSM86ParameterGradient256 . layerNormSM86NaturalFive)
(branch LayerNormSM86ParameterGradient512 . layerNormSM86NaturalFive)
(branch LayerNormSM86ParameterGradient1024 . layerNormSM86NaturalFive)))
def layerNormSM86Telemetry =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(lambda unrestricted actualInstructions : Nat .
(lambda unrestricted actualBytes : Nat .
(lambda unrestricted encodedFields : Nat .
(lambda unrestricted encodedBits : Nat .
(lambda unrestricted highestExclusiveBit : Nat .
(constructor
LayerNormSM86Telemetry
LayerNormSM86TelemetryValue
variant
(layerNormSM86ExpectedInstructions variant rows)
actualInstructions
(layerNormSM86ExpectedEncodedBytes variant rows)
actualBytes
(layerNormSM86RegisterCount variant)
(layerNormSM86SharedBytes variant)
rows
(layerNormSM86Width variant)
(naturalMultiply rows (layerNormSM86Width variant))
(layerNormSM86GridX variant rows)
layerNormSM86NaturalOne
layerNormSM86NaturalOne
(layerNormSM86BlockX variant)
layerNormSM86NaturalOne
layerNormSM86NaturalOne
(layerNormSM86ConstantLoads variant)
(layerNormSM86GlobalLoads variant)
(layerNormSM86GlobalStores variant)
layerNormSM86NaturalTwo
(layerNormSM86Barriers variant)
layerNormSM86NaturalTwo
zero
encodedFields
encodedBits
highestExclusiveBit))))))))
def layerNormSM86RejectedTelemetry =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted code : (family LayerNormSM86FailureCode) .
(lambda unrestricted ordinal : Nat .
(constructor LayerNormSM86Telemetry LayerNormSM86TelemetryRejected variant code ordinal))))
def layerNormSM86ExtentsFor =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(constructor
LayerNormSM86Extents
LayerNormSM86ExtentsValue
variant
rows
(layerNormSM86Width variant)
(naturalMultiply rows (layerNormSM86Width variant)))))
def layerNormSM86ManifestFor =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(lambda unrestricted program : (family SM86Program) .
(lambda unrestricted telemetry : (family LayerNormSM86Telemetry) .
(constructor
LayerNormSM86Manifest
LayerNormSM86ManifestValue
variant
(layerNormSM86ExtentsFor variant rows)
(layerNormSM86ScalarABI variant)
(layerNormSM86GridX variant rows)
layerNormSM86NaturalOne
layerNormSM86NaturalOne
(layerNormSM86BlockX variant)
layerNormSM86NaturalOne
layerNormSM86NaturalOne
(layerNormSM86RegisterCount variant)
(layerNormSM86SharedBytes variant)
(layerNormSM86ExpectedInstructions variant rows)
(layerNormSM86ExpectedEncodedBytes variant rows)
program
telemetry)))))
def layerNormSM86FailureCodeBytes =
(lambda unrestricted code : (family LayerNormSM86FailureCode) .
(eliminate
LayerNormSM86FailureCode
(lambda unrestricted current : (family LayerNormSM86FailureCode) . Bytes)
code
(branch LayerNormSM86RowsZero . b"ALPHA-SM86-LN-001")
(branch LayerNormSM86ParameterRowsMismatch . b"ALPHA-SM86-LN-002")
(branch LayerNormSM86InstructionCountMismatch . b"ALPHA-SM86-LN-003")
(branch LayerNormSM86EncodingFailed . b"ALPHA-SM86-LN-004")
(branch LayerNormSM86EncodedByteCountMismatch . b"ALPHA-SM86-LN-005")
(branch LayerNormSM86IdentityInvalid . b"ALPHA-SM86-LN-006")))
def layerNormSM86PlanCounted =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(app
(lambda unrestricted program : (family SM86Program) .
(app
(lambda unrestricted actual : Nat .
(nat-eliminate
(lambda unrestricted valid : Nat . (family LayerNormSM86PlanResult))
(constructor
LayerNormSM86PlanResult
LayerNormSM86PlanFailed
(constructor LayerNormSM86FailureCode LayerNormSM86InstructionCountMismatch)
(layerNormSM86RejectedTelemetry
variant
(constructor LayerNormSM86FailureCode LayerNormSM86InstructionCountMismatch)
actual))
(lambda unrestricted predecessor : Nat .
(lambda unrestricted induction : (family LayerNormSM86PlanResult) .
(app
(lambda unrestricted telemetry : (family LayerNormSM86Telemetry) .
(constructor
LayerNormSM86PlanResult
LayerNormSM86PlanReady
(layerNormSM86ManifestFor variant rows program telemetry)))
(layerNormSM86Telemetry variant rows actual zero zero zero zero))))
(naturalEqual actual (layerNormSM86ExpectedInstructions variant rows))))
(sm86ProgramCount program)))
(layerNormSM86Program variant rows))))
def layerNormSM86Plan =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(nat-eliminate
(lambda unrestricted nonzero : Nat . (family LayerNormSM86PlanResult))
(constructor
LayerNormSM86PlanResult
LayerNormSM86PlanFailed
(constructor LayerNormSM86FailureCode LayerNormSM86RowsZero)
(layerNormSM86RejectedTelemetry
variant
(constructor LayerNormSM86FailureCode LayerNormSM86RowsZero)
rows))
(lambda unrestricted nonzeroPredecessor : Nat .
(lambda unrestricted nonzeroInduction : (family LayerNormSM86PlanResult) .
(nat-eliminate
(lambda unrestricted rowsValid : Nat . (family LayerNormSM86PlanResult))
(constructor
LayerNormSM86PlanResult
LayerNormSM86PlanFailed
(constructor LayerNormSM86FailureCode LayerNormSM86ParameterRowsMismatch)
(layerNormSM86RejectedTelemetry
variant
(constructor LayerNormSM86FailureCode LayerNormSM86ParameterRowsMismatch)
rows))
(lambda unrestricted rowsPredecessor : Nat .
(lambda unrestricted rowsInduction : (family LayerNormSM86PlanResult) .
(layerNormSM86PlanCounted variant rows)))
(layerNormSM86RowsValid variant rows))))
(naturalNonzero rows))))
def layerNormSM86ImageIdentity =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(lambda unrestricted program : (family SM86Program) .
(lambda unrestricted image : Bytes .
(lambda unrestricted encodingTelemetry : (family SM86ProgramEncodingTelemetry) .
(eliminate
SM86ProgramEncodingTelemetry
(lambda unrestricted current : (family SM86ProgramEncodingTelemetry) .
(family LayerNormSM86ImageResult))
encodingTelemetry
(branch
SM86ProgramEncodingTelemetryValue
instructions
bytes
fields
bits
highest
.
(app
(lambda unrestricted telemetry : (family LayerNormSM86Telemetry) .
(nat-eliminate
(lambda unrestricted instructionValid : Nat .
(family LayerNormSM86ImageResult))
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageFailed
(constructor LayerNormSM86FailureCode LayerNormSM86InstructionCountMismatch)
instructions
(layerNormSM86FailureCodeBytes
(constructor
LayerNormSM86FailureCode
LayerNormSM86InstructionCountMismatch))
telemetry)
(lambda unrestricted instructionPredecessor : Nat .
(lambda unrestricted instructionInduction : (family LayerNormSM86ImageResult) .
(nat-eliminate
(lambda unrestricted encodedBytesValid : Nat .
(family LayerNormSM86ImageResult))
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageFailed
(constructor
LayerNormSM86FailureCode
LayerNormSM86EncodedByteCountMismatch)
bytes
(layerNormSM86FailureCodeBytes
(constructor
LayerNormSM86FailureCode
LayerNormSM86EncodedByteCountMismatch))
telemetry)
(lambda unrestricted encodedBytesPredecessor : Nat .
(lambda unrestricted encodedBytesInduction : (family LayerNormSM86ImageResult) .
(nat-eliminate
(lambda unrestricted imageBytesValid : Nat .
(family LayerNormSM86ImageResult))
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageFailed
(constructor
LayerNormSM86FailureCode
LayerNormSM86EncodedByteCountMismatch)
(bytes-length image)
(layerNormSM86FailureCodeBytes
(constructor
LayerNormSM86FailureCode
LayerNormSM86EncodedByteCountMismatch))
telemetry)
(lambda unrestricted imageBytesPredecessor : Nat .
(lambda unrestricted imageBytesInduction : (family LayerNormSM86ImageResult) .
(app
(lambda unrestricted identity : Bytes .
(nat-eliminate
(lambda unrestricted identityValid : Nat .
(family LayerNormSM86ImageResult))
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageFailed
(constructor
LayerNormSM86FailureCode
LayerNormSM86IdentityInvalid)
zero
(layerNormSM86FailureCodeBytes
(constructor
LayerNormSM86FailureCode
LayerNormSM86IdentityInvalid))
telemetry)
(lambda unrestricted identityPredecessor : Nat .
(lambda unrestricted identityInduction : (family LayerNormSM86ImageResult) .
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageReady
variant
image
identity
(layerNormSM86ManifestFor variant rows program telemetry)
telemetry)))
(naturalEqual
(bytes-length identity)
layerNormSM86NaturalSixtyFour)))
(sha256HexBytesOrEmpty (sha256Hex image)))))
(naturalEqual
(bytes-length image)
(layerNormSM86ExpectedEncodedBytes variant rows)))))
(naturalEqual bytes (layerNormSM86ExpectedEncodedBytes variant rows)))))
(naturalEqual instructions (layerNormSM86ExpectedInstructions variant rows))))
(layerNormSM86Telemetry variant rows instructions bytes fields bits highest)))))))))
def layerNormSM86Image =
(lambda unrestricted variant : (family LayerNormSM86Variant) .
(lambda unrestricted rows : Nat .
(eliminate
LayerNormSM86PlanResult
(lambda unrestricted current : (family LayerNormSM86PlanResult) .
(family LayerNormSM86ImageResult))
(layerNormSM86Plan variant rows)
(branch
LayerNormSM86PlanReady
manifest
.
(eliminate
LayerNormSM86Manifest
(lambda unrestricted current : (family LayerNormSM86Manifest) .
(family LayerNormSM86ImageResult))
manifest
(branch
LayerNormSM86ManifestValue
plannedVariant
extents
scalarABI
gridX
gridY
gridZ
blockX
blockY
blockZ
registers
shared
expectedInstructions
expectedBytes
program
telemetry
.
(eliminate
SM86ProgramEncodingResult
(lambda unrestricted current : (family SM86ProgramEncodingResult) .
(family LayerNormSM86ImageResult))
(sm86EncodeProgram program)
(branch
SM86ProgramEncodingSucceeded
image
encodingTelemetry
.
(layerNormSM86ImageIdentity plannedVariant rows program image encodingTelemetry))
(branch
SM86ProgramEncodingFailed
index
failure
encodingTelemetry
.
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageFailed
(constructor LayerNormSM86FailureCode LayerNormSM86EncodingFailed)
index
(sm86InstructionEncodingStableCode failure)
telemetry))))))
(branch
LayerNormSM86PlanFailed
code
telemetry
.
(constructor
LayerNormSM86ImageResult
LayerNormSM86ImageFailed
code
zero
(layerNormSM86FailureCodeBytes code)
telemetry)))))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.