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)))))