Source/Packages

Realization.Nvidia.SM86.LayerNormSM86

packages/realizations/cooperative/nvidia-sm86/src/Realization/Nvidia/SM86/LayerNormSM86.alpha

2,228 lines251 declarations94.1 KiBSHA-256 676a74b45b30

Complete file · line 153

LayerNormSM86.alpha

Definition view

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.