Source/Systems

Coppelius.Build.Graph

systems/coppelius/src/Coppelius/Build/Graph.alpha

3,461 lines536 declarations124.9 KiBSHA-256 6b9f923bd170

Complete file · line 1678

Graph.alpha

Definition view

Large source region · 3,461 lines

module Coppelius.Build.Graph

import Coppelius.ArenaPlan
import Coppelius.Learner
import Coppelius.Model
import Coppelius.TrainingRun
import Coppelius.Build.DeviceImages
import Component.RotaryFrequency
import Float32Exact
import Float32Model
import FloatLiteralSpec
import Learning.Checked.MomentCorrection
import Hardware.Nvidia.SM86.Command.QMD
import Hardware.Nvidia.SM86.Command.AffineLaunchSchedule
import Hardware.Nvidia.SM86.Command.WholeProgramPlan
import Platform.Linux.Nvidia.PlanHostRequest
import Realization.Nvidia.SM86.HMMAProductionSM86
import Realization.Nvidia.SM86.TiledProductSM86
import Realization.Nvidia.SM86.GatedGELUSM86
import Runtime.ArenaCertificate
import Runtime.DeviceArenaCertificate
import Model.Config
import Model.Parameter
import Model.Word32
import Model.Word64
import Std.Foundation
import Std.List
import Std.Natural
import Std.Word

-- Coppelius's launch graph, derived: every launch of a training run
-- (Coppelius.TrainingRun: initialize or load, predict, the update the run
-- issues once per step, the final forward, save) computed from the model's
-- shape
-- (Coppelius.Model), the learner's hyperparameters (Coppelius.Learner),
-- the device images (Coppelius.Build.DeviceImages) and the staging window's
-- layout (Coppelius.TrainingRun) -- no address, count or constant is
-- written down that one of those determines.
--
-- The graph is the one the legacy generator emitted
-- (systems/coppelius/reference/oracles/legacy-generated/tools/gen_train.py,
-- `unified10`), from which the plan was once imported as a recording; it
-- was re-derived here launch for launch, and the realized tables were
-- byte-identical to the recording's (the migration, 2026-09-24).  Changes
-- since: `loss.correct.zero` is the target-logit gather, ten loss-record
-- copies close the updates (docs/observability PRD item 1); and the
-- workspace is a named layout per phase, certified on the device side, which
-- moved the tensors the recording aliased while live -- the cross-entropy
-- backward wrote the logits gradient over the loss rows it was reading, and
-- two backward phases reused a buffer they still held -- and the output's
-- and two feed-forward phases' planes with them.
-- a 32-bit slot of a parameter block, by its word index (byte offset / 4)
family CgSlot : Type 0
constructor CgSlotValue
field unrestricted cgSlotIndex : Nat
field unrestricted cgSlotValue : Nat

end-family

-- a range of the launch order: its first launch and how many
family CgRange : Type 0
constructor CgRangeValue
field unrestricted cgRangeFirst : Nat
field unrestricted cgRangeCount : Nat

end-family

-- ---- the shape ----
def cgSeq : Nat =
  (stdU32ToNatural modelCoppeliusBlock)

def cgWidth : Nat =
  (stdU32ToNatural modelCoppeliusWidth)

def cgHeads : Nat =
  (stdU32ToNatural modelCoppeliusHeads)

def cgHeadWidth : Nat =
  (naturalDivideUnchecked cgWidth cgHeads)

def cgFfn : Nat =
  (stdU32ToNatural modelCoppeliusFfnWidth)

def cgVocab : Nat =
  (stdU32ToNatural modelCoppeliusVocab)

def cgLayers : Nat =
  (stdU32ToNatural modelCoppeliusLayers)

-- ---- arithmetic ----
def cgAlign =
  (lambda unrestricted x : Nat .
    (lambda unrestricted a : Nat .
      (naturalMultiply (naturalDivideUnchecked (naturalAdd x (naturalSaturatingSubtract a 1)) a) a)))

def cgMinimum =
  (lambda unrestricted a : Nat .
    (lambda unrestricted b : Nat . (naturalSelect (naturalLessOrEqual a b) a b)))

def cgCeilDivide =
  (lambda unrestricted a : Nat .
    (lambda unrestricted b : Nat .
      (naturalDivideUnchecked (naturalAdd a (naturalSaturatingSubtract b 1)) b)))

-- A bump allocator: the cursor after `count` items (each aligned to
-- `alignment` before it is taken), and the offset of item `index`.
def cgBumpEnd =
  (lambda unrestricted alignment : Nat .
    (lambda unrestricted start : Nat .
      (lambda unrestricted size : (pi unrestricted index : Nat . Nat) .
        (lambda unrestricted count : Nat .
          (app
            (nat-eliminate
              (lambda unrestricted current : Nat .
                (pi unrestricted index : Nat . (pi unrestricted cursor : Nat . Nat)))
              (lambda unrestricted index : Nat . (lambda unrestricted cursor : Nat . cursor))
              (lambda unrestricted p : Nat .
                (lambda unrestricted induction : (pi unrestricted index : Nat . (pi unrestricted cursor : Nat . Nat)) .
                  (lambda unrestricted index : Nat .
                    (lambda unrestricted cursor : Nat .
                      (induction (succ index) (naturalAdd (cgAlign cursor alignment) (size index)))))))
              count)
            0
            start)))))

def cgBumpAt =
  (lambda unrestricted alignment : Nat .
    (lambda unrestricted start : Nat .
      (lambda unrestricted size : (pi unrestricted index : Nat . Nat) .
        (lambda unrestricted index : Nat .
          (cgAlign (cgBumpEnd alignment start size index) alignment)))))

-- ---- the parameters (bytes from the video arena's base) ----
-- entry 0 the embedding (vocabulary x width); entry 1 + 9 l + k layer l's
-- k-th item (ln1 gamma, ln1 beta, qkv, o, ln2 gamma, ln2 beta, up, gate,
-- down); entries 145, 146 the final gamma and beta
-- the offsets of items 0 .. count - 1, computed in one pass (the plan
-- looks them up for every launch)
def cgBumpTable =
  (lambda unrestricted alignment : Nat .
    (lambda unrestricted start : Nat .
      (lambda unrestricted size : (pi unrestricted index : Nat . Nat) .
        (lambda unrestricted count : Nat .
          (app
            (nat-eliminate
              (lambda unrestricted current : Nat .
                (pi unrestricted index : Nat .
                  (pi unrestricted cursor : Nat . (family StdList Nat))))
              (lambda unrestricted index : Nat .
                (lambda unrestricted cursor : Nat . (constructor StdList StdListEmpty Nat)))
              (lambda unrestricted p : Nat .
                (lambda unrestricted induction : (pi unrestricted index : Nat . (pi unrestricted cursor : Nat . (family StdList Nat))) .
                  (lambda unrestricted index : Nat .
                    (lambda unrestricted cursor : Nat .
                      (let unrestricted offset =
                        (cgAlign cursor alignment)
                        in
                        (constructor
                          StdList
                          StdListCons
                          Nat
                          offset
                          (induction (succ index) (naturalAdd offset (size index)))))))))
              count)
            0
            start)))))

def cgIndex =
  (lambda unrestricted values : (family StdList Nat) .
    (lambda unrestricted p : Nat .
      (eliminate
        StdOption
        (lambda unrestricted current : (family StdOption Nat) . Nat)
        (stdListIndex Nat values p)
        (branch StdNone . 0)
        (branch StdSome value . value))))

-- a layer's items, in their order in the parameters
def cgItemNorm1Gain : Nat =
  zero

def cgItemNorm1Shift : Nat =
  (succ cgItemNorm1Gain)

def cgItemQkv : Nat =
  (succ cgItemNorm1Shift)

def cgItemOutput : Nat =
  (succ cgItemQkv)

def cgItemNorm2Gain : Nat =
  (succ cgItemOutput)

def cgItemNorm2Shift : Nat =
  (succ cgItemNorm2Gain)

def cgItemUp : Nat =
  (succ cgItemNorm2Shift)

def cgItemGate : Nat =
  (succ cgItemUp)

def cgItemDown : Nat =
  (succ cgItemGate)

def cgLayerItems : Nat =
  (succ cgItemDown)

-- the qkv projection's three parts
def cgQkvParts : Nat =
  3

def cgLayerItemElements =
  (lambda unrestricted k : Nat .
    (naturalSelect
      (naturalEqual k cgItemQkv)
      (naturalMultiply (naturalMultiply cgQkvParts cgWidth) cgWidth)
      (naturalSelect
        (naturalEqual k cgItemOutput)
        (naturalMultiply cgWidth cgWidth)
        (naturalSelect
          (naturalOr
            (naturalEqual k cgItemUp)
            (naturalOr (naturalEqual k cgItemGate) (naturalEqual k cgItemDown)))
          (naturalMultiply cgFfn cgWidth)
          cgWidth))))

-- the entries: the embedding, each layer's items, the final gain and shift
def cgEmbeddingEntry : Nat =
  zero

def cgFirstLayerEntry : Nat =
  (succ cgEmbeddingEntry)

def cgFinalGamma : Nat =
  (naturalAdd cgFirstLayerEntry (naturalMultiply cgLayerItems cgLayers))

def cgFinalBeta : Nat =
  (succ cgFinalGamma)

def cgEntries : Nat =
  (succ cgFinalBeta)

def cgEntryElements =
  (lambda unrestricted j : Nat .
    (naturalSelect
      (naturalEqual j cgEmbeddingEntry)
      (naturalMultiply cgVocab cgWidth)
      (naturalSelect
        (naturalLess j cgFinalGamma)
        (cgLayerItemElements
          (naturalModuloUnchecked (naturalSaturatingSubtract j cgFirstLayerEntry) cgLayerItems))
        cgWidth)))

def cgEntryBytes =
  (lambda unrestricted j : Nat . (naturalMultiply 4 (cgEntryElements j)))

def cgParameterAlignment : Nat =
  0x100

def cgEntryOffsets : (family StdList Nat) =
  (cgBumpTable cgParameterAlignment 0 cgEntryBytes cgEntries)

def cgEntryOffset =
  (lambda unrestricted j : Nat . (cgIndex cgEntryOffsets j))

def cgParameterBytes : Nat =
  (cgBumpEnd cgParameterAlignment 0 cgEntryBytes cgEntries)

def cgLayerEntry =
  (lambda unrestricted l : Nat .
    (lambda unrestricted k : Nat .
      (naturalAdd cgFirstLayerEntry (naturalAdd (naturalMultiply cgLayerItems l) k))))

-- the four banks: parameters, gradients, first and second moments
def cgBankAlignment : Nat =
  0x1000

def cgBankP : Nat =
  0

def cgBankG : Nat =
  (cgAlign (naturalAdd cgBankP cgParameterBytes) cgBankAlignment)

def cgBankM : Nat =
  (cgAlign (naturalAdd cgBankG cgParameterBytes) cgBankAlignment)

def cgBankV : Nat =
  (cgAlign (naturalAdd cgBankM cgParameterBytes) cgBankAlignment)

-- the matrices (entries whose element count is not the width), in order:
-- matrix 0 the embedding, 1 + 5 l + t layer l's qkv, o, up, gate, down
-- a layer's matrices, in the order of their half copies
def cgMatrixQkv : Nat =
  zero

def cgMatrixOutput : Nat =
  (succ cgMatrixQkv)

def cgMatrixUp : Nat =
  (succ cgMatrixOutput)

def cgMatrixGate : Nat =
  (succ cgMatrixUp)

def cgMatrixDown : Nat =
  (succ cgMatrixGate)

def cgLayerMatrices : Nat =
  (succ cgMatrixDown)

def cgEmbeddingMatrix : Nat =
  zero

def cgFirstLayerMatrix : Nat =
  (succ cgEmbeddingMatrix)

def cgMatrices : Nat =
  (naturalAdd cgFirstLayerMatrix (naturalMultiply cgLayerMatrices cgLayers))

def cgMatrixItem =
  (lambda unrestricted t : Nat .
    (naturalSelect
      (naturalEqual t cgMatrixQkv)
      cgItemQkv
      (naturalSelect
        (naturalEqual t cgMatrixOutput)
        cgItemOutput
        (naturalSelect
          (naturalEqual t cgMatrixUp)
          cgItemUp
          (naturalSelect (naturalEqual t cgMatrixGate) cgItemGate cgItemDown)))))

def cgMatrixEntry =
  (lambda unrestricted index : Nat .
    (naturalSelect
      (naturalEqual index cgEmbeddingMatrix)
      cgEmbeddingEntry
      (cgLayerEntry
        (naturalDivideUnchecked
          (naturalSaturatingSubtract index cgFirstLayerMatrix)
          cgLayerMatrices)
        (cgMatrixItem
          (naturalModuloUnchecked
            (naturalSaturatingSubtract index cgFirstLayerMatrix)
            cgLayerMatrices)))))

def cgHalfStart : Nat =
  (cgAlign (naturalAdd cgBankV cgParameterBytes) cgBankAlignment)

-- the half bank parallel to the parameter bank (element i's half at byte
-- 2 i), so AdamW writes each parameter's half beside its update
-- (Realization.Nvidia.SM86.AdamWHalfSM86): matrix index's copy at half its
-- entry's offset
def cgHalfOffset =
  (lambda unrestricted index : Nat .
    (naturalAdd cgHalfStart (naturalDivideUnchecked (cgEntryOffset (cgMatrixEntry index)) 2)))

def cgHalfEnd : Nat =
  (naturalAdd cgHalfStart (naturalDivideUnchecked cgParameterBytes 2))

-- the layer boundaries: the residual stream entering layer l (l = 0 ..
-- layers), each a sequence x width binary32 plane
def cgPlaneBytes : Nat =
  (naturalMultiply 4 (naturalMultiply cgSeq cgWidth))

def cgBoundaryStart : Nat =
  (cgAlign cgHalfEnd cgBankAlignment)

def cgBoundaryOffset =
  (lambda unrestricted l : Nat .
    (cgBumpAt cgBankAlignment cgBoundaryStart (lambda unrestricted i : Nat . cgPlaneBytes) l))

def cgBoundaryEnd : Nat =
  (cgBumpEnd
    cgBankAlignment
    cgBoundaryStart
    (lambda unrestricted i : Nat . cgPlaneBytes)
    (succ cgLayers))

-- the rotary tables (Coppelius.Model's positions): a row per position and
-- a column per head-dimension pair, the cosines then the sines; computed at
-- the start of every invocation (cgRotaryTables) and read by the attention
-- layouts, so they are placed outside the workspace the phases reuse
def cgRotaryColumns : Nat =
  (naturalDivideUnchecked cgHeadWidth 2)

def cgRotaryRowBytes : Nat =
  (naturalMultiply 4 cgRotaryColumns)

def cgRotaryTableBytes : Nat =
  (naturalMultiply cgSeq cgRotaryRowBytes)

def cgRotaryStart : Nat =
  (cgAlign cgBoundaryEnd cgBankAlignment)

def cgWorkBase : Nat =
  (cgAlign (naturalAdd cgRotaryStart (naturalMultiply 2 cgRotaryTableBytes)) cgBankAlignment)

-- device addresses: the arena, and the staging window
def cgArena =
  (lambda unrestricted offset : Nat . (naturalAdd coppeliusSM86CompatVideoBaseNatural offset))

def cgStaging =
  (lambda unrestricted offset : Nat . (naturalAdd coppeliusStagingBase offset))

-- ---- the workspace: a layer's tensors (the forward's, recomputed in the
-- backward), each placed after the previous, page-aligned ----
def cgNext =
  (lambda unrestricted address : Nat .
    (lambda unrestricted bytes : Nat . (cgAlign (naturalAdd address bytes) cgBankAlignment)))

def cgHalfPlaneBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgSeq cgWidth))

-- each row's mean and inverse deviation
def cgStatisticsBytes : Nat =
  (naturalMultiply 8 cgSeq)

def cgQkvBytes : Nat =
  (naturalMultiply 4 (naturalMultiply cgSeq (naturalMultiply cgQkvParts cgWidth)))

def cgHeadsHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgHeads (naturalMultiply cgSeq cgHeadWidth)))

def cgScoresBytes : Nat =
  (naturalMultiply 4 (naturalMultiply cgHeads (naturalMultiply cgSeq cgSeq)))

def cgScoresHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgHeads (naturalMultiply cgSeq cgSeq)))

def cgContextBytes : Nat =
  (naturalMultiply 4 (naturalMultiply cgHeads (naturalMultiply cgSeq cgHeadWidth)))

def cgFfnBytes : Nat =
  (naturalMultiply 4 (naturalMultiply cgSeq cgFfn))

def cgFfnHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgSeq cgFfn))

def cgNorm1 : Nat =
  (cgArena (cgAlign cgWorkBase cgBankAlignment))

def cgNorm1Statistics : Nat =
  (cgNext cgNorm1 cgPlaneBytes)

def cgNorm1Half : Nat =
  (cgNext cgNorm1Statistics cgStatisticsBytes)

def cgQkv : Nat =
  (cgNext cgNorm1Half cgHalfPlaneBytes)

def cgQueries : Nat =
  (cgNext cgQkv cgQkvBytes)

def cgKeys : Nat =
  (cgNext cgQueries cgHeadsHalfBytes)

def cgValuesT : Nat =
  (cgNext cgKeys cgHeadsHalfBytes)

def cgScores : Nat =
  (cgNext cgValuesT cgHeadsHalfBytes)

-- the streaming attention's: each row's base-2
-- log-sum-exp and the backward's row terms ([head][T] binary32), and V in
-- its heads' layout for the backward ([head][T][64] half)
def cgLogSumExpBytes : Nat =
  (naturalMultiply 4 (naturalMultiply cgHeads cgSeq))

def cgLogSumExp : Nat =
  (cgNext cgScores cgScoresBytes)

def cgRowTerms : Nat =
  (cgNext cgLogSumExp cgLogSumExpBytes)

def cgValuesHeads : Nat =
  (cgNext cgRowTerms cgLogSumExpBytes)

-- a binary32 plane the norms' backward phases use (the second norm's input
-- gradient, the first norm's gradient)
def cgNormScratch : Nat =
  (cgNext cgValuesHeads cgHeadsHalfBytes)

def cgMergedHalf : Nat =
  (cgNext cgNormScratch cgPlaneBytes)

def cgProjection : Nat =
  (cgNext cgMergedHalf cgHalfPlaneBytes)

def cgNorm2 : Nat =
  (cgNext cgProjection cgPlaneBytes)

def cgNorm2Statistics : Nat =
  (cgNext cgNorm2 cgPlaneBytes)

def cgNorm2Half : Nat =
  (cgNext cgNorm2Statistics cgStatisticsBytes)

def cgUp : Nat =
  (cgNext cgNorm2Half cgHalfPlaneBytes)

def cgGate : Nat =
  (cgNext cgUp cgFfnBytes)

def cgGelu : Nat =
  (cgNext cgGate cgFfnBytes)

def cgGated : Nat =
  (cgNext cgGelu cgFfnBytes)

def cgGatedHalf : Nat =
  (cgNext cgGated cgFfnBytes)

def cgDown : Nat =
  (cgNext cgGatedHalf cgFfnHalfBytes)

def cgRowBytes : Nat =
  (naturalMultiply 4 cgVocab)

-- ---- the output's and the backward's planes ----
-- Each phase's tensors, placed in the workspace where the phase's inputs
-- are not: the output from the workspace's start (the final norm where a
-- layer's first norm is), the backward's scratch in the scores item (dead
-- once the recompute has its probabilities).  Which item a plane reuses is
-- the placement; that no plane is overwritten while a later phase needs it
-- is what Proof.CoppeliusArenaCertificate decides.
def cgLogitsBytes : Nat =
  (naturalMultiply cgSeq cgRowBytes)

def cgLogitsHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgSeq cgVocab))

def cgRowsBytes : Nat =
  (naturalMultiply 4 cgSeq)

def cgFinalNorm : Nat =
  cgNorm1

def cgFinalStatistics : Nat =
  (cgNext cgFinalNorm cgPlaneBytes)

def cgFinalHalf : Nat =
  (cgNext cgFinalStatistics cgStatisticsBytes)

def cgLogits : Nat =
  (cgNext cgFinalHalf cgHalfPlaneBytes)

-- the rows' cross-entropy, and `correct` (each row's target logit)
def cgLoss : Nat =
  (cgNext cgLogits cgLogitsBytes)

def cgCorrect : Nat =
  (cgNext cgLoss cgRowsBytes)

def cgLogitsGradient : Nat =
  (cgNext cgCorrect cgRowsBytes)

-- the output's gradient: its halves where the logits were
def cgLogitsGradientHalf : Nat =
  cgLogits

def cgLogitsGradientT : Nat =
  (cgNext cgLogitsGradientHalf cgLogitsHalfBytes)

def cgFinalHalfT : Nat =
  (cgNext cgLogitsGradient cgLogitsBytes)

def cgFinalGradient : Nat =
  (cgNext cgFinalHalfT cgHalfPlaneBytes)

-- the tied embedding's binary16 copy transposed (width x vocab): the final
-- norm's gradient is dlogits E, and the product takes its right operand as
-- B^T -- the vocab x width copy itself would be read as E^T
def cgEmbeddingHalfTBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgVocab cgWidth))

def cgEmbeddingHalfT : Nat =
  (cgNext cgFinalGradient cgPlaneBytes)

-- the workspace ends past a layer's tensors and past the output's and the
-- output gradient's (placed from its start, below), whichever reach further
def cgWorkEnd : Nat =
  (naturalSaturatingSubtract
    (naturalSelect
      (naturalLess (naturalAdd cgDown cgPlaneBytes) (naturalAdd cgEmbeddingHalfT cgEmbeddingHalfTBytes))
      (naturalAdd cgEmbeddingHalfT cgEmbeddingHalfTBytes)
      (naturalAdd cgDown cgPlaneBytes))
    (cgArena zero))

def cgDy0 : Nat =
  (cgAlign cgWorkEnd cgBankAlignment)

def cgDy1 : Nat =
  (naturalAdd cgDy0 cgPlaneBytes)

def cgXHat : Nat =
  (naturalAdd cgDy1 cgPlaneBytes)

def cgReport : Nat =
  (naturalAdd cgXHat cgPlaneBytes)

-- ---- the saved activations: what each layer's backward reads of its
-- forward, kept per layer from the forward, so the backward does not
-- recompute the layer (16 layers of 29 MiB) ----
def cgSavedStart : Nat =
  (cgAlign (naturalAdd cgReport (naturalAdd coppeliusResultLossBlockBytes cgRowBytes)) cgBankAlignment)

-- a layer's planes, from its block's start
def cgSvNorm1Statistics : Nat = 0
def cgSvNorm1Half : Nat = (cgNext cgSvNorm1Statistics cgStatisticsBytes)
def cgSvQueries : Nat = (cgNext cgSvNorm1Half cgHalfPlaneBytes)
def cgSvKeys : Nat = (cgNext cgSvQueries cgHeadsHalfBytes)
def cgSvValuesT : Nat = (cgNext cgSvKeys cgHeadsHalfBytes)
def cgSvLogSumExp : Nat = (cgNext cgSvValuesT cgHeadsHalfBytes)
def cgSvMergedHalf : Nat = (cgNext cgSvLogSumExp cgLogSumExpBytes)
def cgSvProjection : Nat = (cgNext cgSvMergedHalf cgHalfPlaneBytes)
def cgSvNorm2Statistics : Nat = (cgNext cgSvProjection cgPlaneBytes)
def cgSvNorm2Half : Nat = (cgNext cgSvNorm2Statistics cgStatisticsBytes)
def cgSvUp : Nat = (cgNext cgSvNorm2Half cgHalfPlaneBytes)
def cgSvGate : Nat = (cgNext cgSvUp cgFfnBytes)
def cgSvGelu : Nat = (cgNext cgSvGate cgFfnBytes)
def cgSvGatedHalf : Nat = (cgNext cgSvGelu cgFfnBytes)
def cgSavedLayerBytes : Nat = (cgNext cgSvGatedHalf cgFfnHalfBytes)
def cgSavedBytes : Nat = (naturalMultiply cgLayers cgSavedLayerBytes)

-- layer l's plane at `offset` of its block
def cgSaved =
  (lambda unrestricted l : Nat .
    (lambda unrestricted offset : Nat .
      (cgArena (naturalAdd cgSavedStart (naturalAdd (naturalMultiply l cgSavedLayerBytes) offset)))))
def cgNorm1StatisticsOf = (lambda unrestricted l : Nat . (cgSaved l cgSvNorm1Statistics))
def cgNorm1HalfOf = (lambda unrestricted l : Nat . (cgSaved l cgSvNorm1Half))
def cgQueriesOf = (lambda unrestricted l : Nat . (cgSaved l cgSvQueries))
def cgKeysOf = (lambda unrestricted l : Nat . (cgSaved l cgSvKeys))
def cgValuesTOf = (lambda unrestricted l : Nat . (cgSaved l cgSvValuesT))
def cgLogSumExpOf = (lambda unrestricted l : Nat . (cgSaved l cgSvLogSumExp))
def cgMergedHalfOf = (lambda unrestricted l : Nat . (cgSaved l cgSvMergedHalf))
def cgProjectionOf = (lambda unrestricted l : Nat . (cgSaved l cgSvProjection))
def cgNorm2StatisticsOf = (lambda unrestricted l : Nat . (cgSaved l cgSvNorm2Statistics))
def cgNorm2HalfOf = (lambda unrestricted l : Nat . (cgSaved l cgSvNorm2Half))
def cgUpOf = (lambda unrestricted l : Nat . (cgSaved l cgSvUp))
def cgGateOf = (lambda unrestricted l : Nat . (cgSaved l cgSvGate))
def cgGeluOf = (lambda unrestricted l : Nat . (cgSaved l cgSvGelu))
def cgGatedHalfOf = (lambda unrestricted l : Nat . (cgSaved l cgSvGatedHalf))

-- the arena holds the saved activations
def cgSavedEnd : Nat = (naturalAdd cgSavedStart cgSavedBytes)
def cgArenaHoldsTheSaved : (equal Nat (naturalLessOrEqual cgSavedEnd coppeliusSM86CompatVideoBytesNatural) 1) =
  (refl Nat 1)

-- device addresses: the arena, and the staging window
def cgParameter =
  (lambda unrestricted bank : Nat .
    (lambda unrestricted j : Nat . (cgArena (naturalAdd bank (cgEntryOffset j)))))

def cgHalf =
  (lambda unrestricted index : Nat . (cgArena (cgHalfOffset index)))

def cgBoundary =
  (lambda unrestricted l : Nat . (cgArena (cgBoundaryOffset l)))

def cgIds : Nat =
  (cgStaging 0)

def cgTargets : Nat =
  (cgStaging coppeliusInputTargetsOffset)

def cgCosine : Nat =
  (cgArena cgRotaryStart)

def cgSine : Nat =
  (cgArena (naturalAdd cgRotaryStart cgRotaryTableBytes))

def cgDiagnostic : Nat =
  (cgStaging coppeliusResultLossBlockOffset)

def cgInference : Nat =
  (cgStaging coppeliusResultRecordOffset)

def cgScratch : Nat =
  cgScores

def cgFfnWeightHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgFfn cgWidth))

def cgOutputWeightHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgWidth cgWidth))

def cgQkvHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply cgSeq (naturalMultiply cgQkvParts cgWidth)))

def cgQkvWeightHalfBytes : Nat =
  (naturalMultiply 2 (naturalMultiply (naturalMultiply cgQkvParts cgWidth) cgWidth))

-- the feed-forward output projection's backward
def cgDownGatedT : Nat =
  cgScratch

def cgDownWeightT : Nat =
  (cgNext cgDownGatedT cgFfnHalfBytes)

def cgDownDyT : Nat =
  (cgNext cgDownWeightT cgFfnWeightHalfBytes)

def cgDownDyHalf : Nat =
  (cgNext cgDownDyT cgHalfPlaneBytes)

def cgGatedGradient : Nat =
  cgQkv

-- the gated GELU's, the up and gate projections', the second norm's
def cgUpGradient : Nat =
  cgScratch

def cgGeluGradient : Nat =
  (cgNext cgUpGradient cgFfnBytes)

def cgGateGradient : Nat =
  (cgNext cgGeluGradient cgFfnBytes)

def cgFfnInputT : Nat =
  (cgNext cgGateGradient cgFfnBytes)

-- the gate's gradient in half (Realization.Nvidia.SM86.GatedGELUSM86's
-- backward writes it beside the up's): where its binary32 was
def cgGateGradientHalf : Nat =
  cgGateGradient

def cgFfnGradientHalf : Nat =
  (cgNext cgFfnInputT cgHalfPlaneBytes)

def cgFfnGradientT : Nat =
  (cgNext cgFfnGradientHalf cgFfnHalfBytes)

def cgNorm2UpGradient : Nat =
  (cgNext cgFfnGradientT cgFfnHalfBytes)

-- the up (then the gate) projection's binary16 weights transposed (width x
-- ffn): the second norm's gradient is dup Wup + dgate Wgate, and the product
-- takes its right operand as B^T -- the ffn x width copy itself would be
-- read as W^T
def cgFfnWeightT : Nat =
  (cgNext cgNorm2UpGradient cgPlaneBytes)

def cgNorm2Gradient : Nat =
  cgDown

def cgNorm2InputGradient : Nat =
  cgNormScratch

-- the attention output projection's
def cgOutputDyHalf : Nat =
  cgNorm2Half

def cgOutputDyT : Nat =
  cgScratch

def cgOutputMergedT : Nat =
  (cgNext cgOutputDyT cgHalfPlaneBytes)

def cgOutputWeightT : Nat =
  (cgNext cgOutputMergedT cgHalfPlaneBytes)

def cgMergedGradient : Nat =
  cgQkv

-- the heads' shared layouts
def cgContextGradientHeads : Nat =
  cgGatedHalf

def cgQueriesT : Nat =
  cgUp

def cgKeysT : Nat =
  (cgNext cgQueriesT cgHeadsHalfBytes)

def cgContextGradientT : Nat =
  (cgNext cgKeysT cgHeadsHalfBytes)

-- the heads' gradients
def cgQueriesGradient : Nat =
  cgGate

def cgKeysGradient : Nat =
  (cgNext cgQueriesGradient cgContextBytes)

def cgValuesGradient : Nat =
  (cgNext cgKeysGradient cgContextBytes)

-- the qkv projection's and the first norm's
def cgQkvGradientHalf : Nat =
  cgGated

def cgQkvGradientT : Nat =
  cgScratch

def cgQkvInputT : Nat =
  (cgNext cgQkvGradientT cgQkvHalfBytes)

def cgQkvWeightT : Nat =
  (cgNext cgQkvInputT cgHalfPlaneBytes)

def cgNorm1Gradient : Nat =
  cgNormScratch

def cgNorm1InputGradient : Nat =
  cgDown

-- the diagnostic block's windows
def cgWindow =
  (lambda unrestricted index : Nat .
    (naturalAdd cgDiagnostic (naturalMultiply index coppeliusResultWindowBytes)))

def cgWindowWords : Nat =
  (naturalDivideUnchecked coppeliusResultWindowBytes 4)

-- ---- binary32 constants ----
-- Each constant is the binary32 nearest its exact value, computed from the
-- quantities that determine it (the shape, the learner's hyperparameters,
-- the realization's choices, pi and ln 2) by Float32Exact and the
-- functions below, and evaluated by the compiler (`compile-time`): a changed
-- hyperparameter changes the word with nothing else to edit.
def cgExactAdamStepSize =
  (lambda unrestricted step : Nat .
    (momentCorrectedStepSizeWord
      coppeliusLearningRateNumerator coppeliusLearningRateDenominator
      coppeliusBeta1Numerator coppeliusBeta1Denominator
      coppeliusBeta2Numerator coppeliusBeta2Denominator step))

def cgExactAdamEpsilon =
  (lambda unrestricted step : Nat .
    (momentScaledEpsilonWord
      (naturalMultiply coppeliusEpsilonNumerator coppeliusLossScale)
      coppeliusEpsilonDenominator
      coppeliusBeta2Numerator coppeliusBeta2Denominator step))

-- the realization's choices: LayerNorm's epsilon, GELU's tanh approximation
def cgLayerNormEpsilonNumerator : Nat =
  1

def cgLayerNormEpsilonDenominator : Nat =
  100000

def cgGeluCubicNumerator : Nat =
  44715

def cgGeluCubicDenominator : Nat =
  1000000

def cgZero : Nat =
  0

def cgOne : Nat =
  (compile-time (float32ExactRational 1 1))

def cgHalf32 : Nat =
  (compile-time (float32ExactRational 1 2))

def cgInverseWidth : Nat =
  (compile-time (float32ExactRational 1 cgWidth))

def cgInverseSeq : Nat =
  (compile-time (float32ExactRational 1 cgSeq))

-- the logits' gradient: the mean over the rows, times the loss scale
def cgLossGradientScale : Nat =
  (compile-time (float32ExactRational coppeliusLossScale cgSeq))

def cgScoreScale : Nat =
  (compile-time (float32ExactInverseRoot cgHeadWidth))

def cgLayerNormEpsilon : Nat =
  (compile-time (float32ExactRational cgLayerNormEpsilonNumerator cgLayerNormEpsilonDenominator))

-- sqrt(2 / pi), pi in [P, P + 1] / 10^40
def cgSqrtTwoOverPi : Nat =
  (compile-time (float32ExactBetween
        (float32ExactRootLow (naturalMultiply 2 float32Places) (succ float32PiDigits))
        (naturalMultiply (succ float32PiDigits) float32RootScale)
        (succ (float32ExactRootLow (naturalMultiply 2 float32Places) float32PiDigits))
        (naturalMultiply float32PiDigits float32RootScale)))

def cgGeluCubic : Nat =
  (compile-time (float32ExactRational cgGeluCubicNumerator cgGeluCubicDenominator))

def cgGeluCubicSlope : Nat =
  (compile-time (float32ExactRational (naturalMultiply 3 cgGeluCubicNumerator) cgGeluCubicDenominator))

def cgBeta1 : Nat =
  (compile-time (float32ExactRational coppeliusBeta1Numerator coppeliusBeta1Denominator))

def cgOneMinusBeta1 : Nat =
  (compile-time (float32ExactRational
        (naturalSaturatingSubtract coppeliusBeta1Denominator coppeliusBeta1Numerator)
        coppeliusBeta1Denominator))

def cgBeta2 : Nat =
  (compile-time (float32ExactRational coppeliusBeta2Numerator coppeliusBeta2Denominator))

def cgOneMinusBeta2 : Nat =
  (compile-time (float32ExactRational
        (naturalSaturatingSubtract coppeliusBeta2Denominator coppeliusBeta2Numerator)
        coppeliusBeta2Denominator))

-- 1 - lr x weight decay
def cgDecay : Nat =
  (compile-time (float32ExactRational
        (naturalSaturatingSubtract
          (naturalMultiply coppeliusLearningRateDenominator coppeliusWeightDecayDenominator)
          (naturalMultiply coppeliusLearningRateNumerator coppeliusWeightDecayNumerator))
        (naturalMultiply coppeliusLearningRateDenominator coppeliusWeightDecayDenominator)))

-- the update pieces' step sizes and epsilons, from the first step on
def cgExactTable =
  (lambda unrestricted exact : (pi unrestricted step : Nat . Nat) .
    (nat-eliminate
      (lambda unrestricted current : Nat . (family StdList Nat))
      (constructor StdList StdListEmpty Nat)
      (lambda unrestricted p : Nat .
        (lambda unrestricted rest : (family StdList Nat) .
          (constructor
            StdList
            StdListCons
            Nat
            (exact
              (naturalAdd
                coppeliusFirstStep
                (naturalSaturatingSubtract
                  (naturalSaturatingSubtract coppeliusUpdatePieces 1)
                  p)))
            rest)))
      coppeliusUpdatePieces))

def cgAdamStepSizes : (family StdList Nat) =
  (compile-time (cgExactTable cgExactAdamStepSize))

def cgAdamEpsilons : (family StdList Nat) =
  (compile-time (cgExactTable cgExactAdamEpsilon))

def cgTableAt =
  (lambda unrestricted values : (family StdList Nat) .
    (lambda unrestricted p : Nat .
      (eliminate
        StdOption
        (lambda unrestricted current : (family StdOption Nat) . Nat)
        (stdListIndex Nat values p)
        (branch StdNone . float32NotRounded)
        (branch StdSome value . value))))

def cgAdamStepSize =
  (lambda unrestricted step : Nat .
    (cgTableAt cgAdamStepSizes (naturalSaturatingSubtract step coppeliusFirstStep)))

def cgAdamEpsilon =
  (lambda unrestricted step : Nat .
    (cgTableAt cgAdamEpsilons (naturalSaturatingSubtract step coppeliusFirstStep)))

-- ---- parameter blocks ----
def cgNoSlots : (family StdList (family CgSlot)) =
  (constructor StdList StdListEmpty (family CgSlot))

-- ---- where the images read their arguments ----
-- A launch's parameter block against constant bank 0: pointer argument i
-- in the i-th 8-byte slot from the SM86 parameter base; the elementwise,
-- normalization and attention images read their binary32 scalars after six
-- pointer slots, AdamW after its four; the HMMA products read their
-- pointers where their emitter puts them (hmmaNative*Offset).
def cgArgument =
  (lambda unrestricted i : Nat . (naturalAdd qmdKernelParameterBase (naturalMultiply 8 i)))

def cgScalarsAfter =
  (lambda unrestricted pointers : Nat .
    (lambda unrestricted j : Nat . (naturalAdd (cgArgument pointers) (naturalMultiply 4 j))))

def cgElementwisePointers : Nat =
  6

def cgAdamWPointers : Nat =
  4

def cgScalar =
  (cgScalarsAfter cgElementwisePointers)

def cgAdamWScalar =
  (cgScalarsAfter cgAdamWPointers)

def cgSlot =
  (lambda unrestricted offset : Nat .
    (lambda unrestricted value : Nat .
      (lambda unrestricted rest : (family StdList (family CgSlot)) .
        (constructor
          StdList
          StdListCons
          (family CgSlot)
          (constructor CgSlot CgSlotValue (naturalDivideUnchecked offset 4) value)
          rest))))

def cgWordRange : Nat =
  4294967296

def cgPointer =
  (lambda unrestricted offset : Nat .
    (lambda unrestricted address : Nat .
      (lambda unrestricted rest : (family StdList (family CgSlot)) .
        (cgSlot
          offset
          (naturalModuloUnchecked address cgWordRange)
          (cgSlot (naturalAdd offset 4) (naturalDivideUnchecked address cgWordRange) rest)))))

-- the sum of the slots at a word index (there is at most one)
def cgSlotAt =
  (lambda unrestricted slots : (family StdList (family CgSlot)) .
    (lambda unrestricted index : Nat .
      (eliminate
        StdList
        (lambda unrestricted current : (family StdList (family CgSlot)) . Nat)
        slots
        (branch StdListEmpty . 0)
        (branch
          StdListCons
          head
          tail
          induction
          .
          (eliminate
            CgSlot
            (lambda unrestricted current : (family CgSlot) . Nat)
            head
            (branch
              CgSlotValue
              at
              value
              .
              (naturalAdd (naturalSelect (naturalEqual at index) value 0) induction)))))))

-- whether a slot sits at a word index
def cgSlotPresent =
  (lambda unrestricted slots : (family StdList (family CgSlot)) .
    (lambda unrestricted index : Nat .
      (eliminate
        StdList
        (lambda unrestricted current : (family StdList (family CgSlot)) . Nat)
        slots
        (branch StdListEmpty . 0)
        (branch
          StdListCons
          head
          tail
          induction
          .
          (eliminate
            CgSlot
            (lambda unrestricted current : (family CgSlot) . Nat)
            head
            (branch CgSlotValue at value . (naturalOr (naturalEqual at index) induction)))))))

-- the block's 64-bit words, one per word a slot names (the even slot of a
-- pair speaks for both halves), those that are zero left out
def cgBlock =
  (lambda unrestricted slots : (family StdList (family CgSlot)) .
    (constructor
      NvidiaParameterBlock
      NvidiaParameterBlockValue
      (eliminate
        StdList
        (lambda unrestricted current : (family StdList (family CgSlot)) .
          (family NvidiaParameterPatches))
        slots
        (branch StdListEmpty . (constructor NvidiaParameterPatches NvidiaParameterPatchesEnd))
        (branch
          StdListCons
          head
          tail
          induction
          .
          (eliminate
            CgSlot
            (lambda unrestricted current : (family CgSlot) . (family NvidiaParameterPatches))
            head
            (branch
              CgSlotValue
              at
              value
              .
              (let unrestricted even =
                (naturalMultiply 2 (naturalDivideUnchecked at 2))
                in
                (let unrestricted speaks =
                  (naturalOr (naturalEqual at even) (naturalIsZero (cgSlotPresent slots even)))
                  in
                  (let unrestricted word =
                    (naturalAdd
                      (cgSlotAt slots even)
                      (naturalMultiply cgWordRange (cgSlotAt slots (succ even))))
                    in
                    (nat-eliminate
                      (lambda unrestricted current : Nat . (family NvidiaParameterPatches))
                      induction
                      (lambda unrestricted q : Nat .
                        (lambda unrestricted ignored : (family NvidiaParameterPatches) .
                          (constructor
                            NvidiaParameterPatches
                            NvidiaParameterPatchesNext
                            (constructor
                              NvidiaParameterPatch
                              NvidiaParameterPatchValue
                              (naturalDivideUnchecked at 2)
                              (modelWord64FromNaturalTruncated word))
                            induction)))
                      (naturalAnd speaks (naturalNonzero word))))))))))))

def cgParameterStride : Nat =
  512

def cgImage =
  (lambda unrestricted identity : Bytes . (coppeliusDeviceImageAddress identity))

-- ---- phases ----
-- Every launch names its phase in its source identity (the backend does
-- not read it; the device arena certificate does): a launch is made
-- unphased, and `cgPhase` claims the unphased launches of a schedule, so
-- the innermost phase wins.
def cgUnphased : Bytes =
  b"coppelius"

def cgPhase =
  (lambda unrestricted phase : Bytes .
    (lambda unrestricted schedule : (family NvidiaLaunchSchedule) .
      (eliminate
        NvidiaLaunchSchedule
        (lambda unrestricted current : (family NvidiaLaunchSchedule) .
          (family NvidiaLaunchSchedule))
        schedule
        (branch
          NvidiaLaunchScheduleEmpty
          .
          (constructor NvidiaLaunchSchedule NvidiaLaunchScheduleEmpty))
        (branch
          NvidiaLaunchScheduleOne
          template
          .
          (constructor
            NvidiaLaunchSchedule
            NvidiaLaunchScheduleOne
            (eliminate
              NvidiaLaunchTemplate
              (lambda unrestricted current : (family NvidiaLaunchTemplate) .
                (family NvidiaLaunchTemplate))
              template
              (branch
                NvidiaLaunchTemplateValue
                identity
                kernel
                block
                .
                (constructor
                  NvidiaLaunchTemplate
                  NvidiaLaunchTemplateValue
                  (nat-eliminate
                    (lambda unrestricted unphased : Nat . Bytes)
                    identity
                    (lambda unrestricted q : Nat . (lambda unrestricted ignored : Bytes . phase))
                    (bytes-equal identity cgUnphased))
                  kernel
                  block)))))
        (branch
          NvidiaLaunchScheduleAppend
          left
          right
          il
          ir
          .
          (constructor NvidiaLaunchSchedule NvidiaLaunchScheduleAppend il ir))
        (branch
          NvidiaLaunchScheduleRepeat
          count
          iteration
          body
          ib
          .
          (constructor NvidiaLaunchSchedule NvidiaLaunchScheduleRepeat count iteration ib)))))

def cgPhaseInitialize : Bytes =
  b"coppelius/initialize"

def cgPhaseRecast : Bytes =
  b"coppelius/recast"

def cgPhaseRotary : Bytes =
  b"coppelius/rotary"

def cgPhaseCheckpoint : Bytes =
  b"coppelius/checkpoint"

def cgPhaseReport : Bytes =
  b"coppelius/report"

def cgPhaseGradientZero : Bytes =
  b"coppelius/gradient-zero"

def cgPhaseEmbed : Bytes =
  b"coppelius/embed"

def cgPhaseLayer : Bytes =
  b"coppelius/layer"

def cgPhaseOutput : Bytes =
  b"coppelius/output"

def cgPhaseOutputGradient : Bytes =
  b"coppelius/output-gradient"

def cgPhaseFfnDown : Bytes =
  b"coppelius/ffn-down"

def cgPhaseFfnGate : Bytes =
  b"coppelius/ffn-gate"

def cgPhaseAttentionOutput : Bytes =
  b"coppelius/attention-output"

def cgPhaseHeadLayout : Bytes =
  b"coppelius/head-layout"

def cgPhaseAttentionHead : Bytes =
  b"coppelius/attention-head"

def cgPhaseQkv : Bytes =
  b"coppelius/qkv"

def cgPhaseScatter : Bytes =
  b"coppelius/scatter"

def cgPhaseAdamW : Bytes =
  b"coppelius/adamw"

def cgPhaseDiagnostics : Bytes =
  b"coppelius/diagnostics"

def cgPhaseLossRecord : Bytes =
  b"coppelius/loss-record"

def cgLaunch =
  (lambda unrestricted image : Nat .
    (lambda unrestricted gx : Nat .
      (lambda unrestricted gy : Nat .
        (lambda unrestricted gz : Nat .
          (lambda unrestricted slots : (family StdList (family CgSlot)) .
            (constructor
              NvidiaLaunchSchedule
              NvidiaLaunchScheduleOne
              (constructor
                NvidiaLaunchTemplate
                NvidiaLaunchTemplateValue
                cgUnphased
                (constructor
                  NvidiaLaunchKernel
                  NvidiaLaunchKernelValue
                  (modelWord64FromNaturalTruncated image)
                  gx
                  gy
                  gz
                  1
                  1
                  cgParameterStride)
                (cgBlock slots))))))))

def cgThen =
  (lambda unrestricted a : (family NvidiaLaunchSchedule) .
    (lambda unrestricted b : (family NvidiaLaunchSchedule) .
      (constructor NvidiaLaunchSchedule NvidiaLaunchScheduleAppend a b)))

def cgNothing : (family NvidiaLaunchSchedule) =
  (constructor NvidiaLaunchSchedule NvidiaLaunchScheduleEmpty)

def cgIf =
  (lambda unrestricted condition : Nat .
    (lambda unrestricted whenTrue : (family NvidiaLaunchSchedule) .
      (lambda unrestricted whenFalse : (family NvidiaLaunchSchedule) .
        (nat-eliminate
          (lambda unrestricted current : Nat . (family NvidiaLaunchSchedule))
          whenFalse
          (lambda unrestricted p : Nat .
            (lambda unrestricted induction : (family NvidiaLaunchSchedule) . whenTrue))
          condition))))

-- body 0, body 1, ..., body (count - 1)
def cgFor =
  (lambda unrestricted count : Nat .
    (lambda unrestricted body : (pi unrestricted index : Nat . (family NvidiaLaunchSchedule)) .
      (nat-eliminate
        (lambda unrestricted current : Nat . (family NvidiaLaunchSchedule))
        cgNothing
        (lambda unrestricted p : Nat .
          (lambda unrestricted induction : (family NvidiaLaunchSchedule) .
            (cgThen
              (body (naturalSaturatingSubtract (naturalSaturatingSubtract count 1) p))
              induction)))
        count)))

-- A range of units in pieces of at most `piece`; affine launch repetition is shared.
def cgPieces =
  (lambda unrestricted total : Nat .
    (lambda unrestricted piece : Nat .
      (lambda unrestricted body : (pi unrestricted offset : Nat . (pi unrestricted units : Nat . (family NvidiaLaunchSchedule))) .
        (let unrestricted full =
          (naturalDivideUnchecked total piece)
          in
          (let unrestricted rest =
            (naturalModuloUnchecked total piece)
            in
            (cgThen
              (nvidiaAffineLaunchRepeat full (lambda unrestricted i : Nat . (body (naturalMultiply i piece) piece)))
              (cgIf (naturalNonzero rest) (body (naturalMultiply full piece) rest) cgNothing)))))))

def cgThreads : Nat =
  256

def cgImageFill : Nat =
  (cgImage b"fill")

def cgImageCopy : Nat =
  (cgImage b"copy")

def cgImageCastDown : Nat =
  (cgImage b"cast-down")

def cgImageAdd : Nat =
  (cgImage b"add")

def cgImageScale : Nat =
  (cgImage b"scale")

def cgImageArgmax : Nat =
  (cgImage b"argmax")

def cgImageGather : Nat =
  (cgImage b"gather")

def cgImageLayerNorm : Nat =
  (cgImage b"layernorm-forward")

def cgImageGatedGeluForward : Nat =
  (cgImage b"gated-gelu-forward")

def cgImageGatedGeluBackward : Nat =
  (cgImage b"gated-gelu-backward")

def cgImageQKRope : Nat =
  (cgImage b"qk-rope")

def cgImageValueTranspose : Nat =
  (cgImage b"value-transpose")

def cgImageHeadToToken : Nat =
  (cgImage b"head-to-token")

def cgImageAttentionForward : Nat =
  (cgImage b"attention-forward")

def cgImageAttentionQuery : Nat =
  (cgImage b"attention-query")

def cgImageAttentionKey : Nat =
  (cgImage b"attention-key")

def cgImageTokenToHead : Nat =
  (cgImage b"token-to-head")

def cgImageInverseRope : Nat =
  (cgImage b"inverse-rope-qkv")

def cgImageSoftmax : Nat =
  (cgImage b"softmax-forward")

def cgImageSoftmaxBackward : Nat =
  (cgImage b"softmax-backward")

def cgImageLayerNormBackward : Nat =
  (cgImage b"layernorm-backward")

def cgImageLayerNormParameters : Nat =
  (cgImage b"layernorm-parameters")

def cgImageCrossEntropy : Nat =
  (cgImage b"cross-entropy")

def cgImageCrossEntropyBackward : Nat =
  (cgImage b"cross-entropy-backward")

def cgImageScatter : Nat =
  (cgImage b"embedding-scatter")

def cgImageAdamW : Nat =
  (cgImage b"adamw")

def cgImageTargetGather : Nat =
  (cgImage b"loss-target-gather")

def cgDecimal =
  (lambda unrestricted n : Nat . (naturalDecimalBytesWithin 8 n))

def cgImageTranspose =
  (lambda unrestricted groups : Nat .
    (lambda unrestricted rows : Nat .
      (lambda unrestricted columns : Nat .
        (cgImage
          (bytes-append
            (nat-eliminate
              (lambda unrestricted current : Nat . Bytes)
              (bytes-append b"transpose-g" (bytes-append (cgDecimal groups) b"-r"))
              (lambda unrestricted q : Nat . (lambda unrestricted ignored : Bytes . b"transpose-r"))
              (naturalEqual groups 1))
            (bytes-append (cgDecimal rows) (bytes-append b"-c" (cgDecimal columns))))))))

def cgUnaryWith =
  (lambda unrestricted image : Nat .
    (lambda unrestricted output : Nat .
      (lambda unrestricted input : Nat .
        (lambda unrestricted elements : Nat .
          (lambda unrestricted extra : (family StdList (family CgSlot)) .
            (cgLaunch
              image
              (naturalDivideUnchecked elements cgThreads)
              1
              1
              (cgSlot
                qmdBlockDimensionXOffset
                cgThreads
                (cgPointer (cgArgument 0) output (cgPointer (cgArgument 1) input extra)))))))))

def cgUnary =
  (lambda unrestricted image : Nat .
    (lambda unrestricted output : Nat .
      (lambda unrestricted input : Nat .
        (lambda unrestricted elements : Nat . (cgUnaryWith image output input elements cgNoSlots)))))

def cgBinary =
  (lambda unrestricted image : Nat .
    (lambda unrestricted output : Nat .
      (lambda unrestricted left : Nat .
        (lambda unrestricted right : Nat .
          (lambda unrestricted elements : Nat .
            (cgLaunch
              image
              (naturalDivideUnchecked elements cgThreads)
              1
              1
              (cgSlot
                0
                cgThreads
                (cgPointer
                  (cgArgument 0)
                  output
                  (cgPointer (cgArgument 1) left (cgPointer (cgArgument 2) right cgNoSlots))))))))))

-- binary32 -> binary16: each thread converts one pair of elements
-- (Realization.Nvidia.SM86.Cast's geometry: elements / 2 threads). A launch
-- of one thread per element writes a second output's worth past the plane.
def cgCastDown =
  (lambda unrestricted output : Nat .
    (lambda unrestricted input : Nat .
      (lambda unrestricted elements : Nat .
        (cgUnary cgImageCastDown output input (naturalDivideUnchecked elements 2)))))

def cgCopy =
  (cgUnary cgImageCopy)

-- C (m x n) = A B, the tiled products (Coppelius.Build.DeviceImages): a block
-- per tile of C, each operand read where it is stored
def cgTiledLaunch =
  (lambda unrestricted residual : Nat .
    (lambda unrestricted a : (family TiledProductLayout) .
      (lambda unrestricted b : (family TiledProductLayout) .
        (lambda unrestricted m : Nat .
          (lambda unrestricted n : Nat .
            (lambda unrestricted k : Nat .
              (lambda unrestricted slots : (family StdList (family CgSlot)) .
                (let unrestricted tile = (coppeliusTiledTile residual a b m n k) in
                (cgLaunch
                  (cgImage (coppeliusTiledName residual a b m n k))
                  (naturalDivideUnchecked n (tiledProductSM86BlockN tile))
                  (naturalDivideUnchecked m (tiledProductSM86BlockM tile))
                  1
                  slots)))))))))

def cgTiled =
  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
    (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
      (lambda unrestricted c : Nat . (lambda unrestricted left : Nat . (lambda unrestricted right : Nat .
        (cgTiledLaunch 0 a b m n k
          (cgPointer (cgArgument 0) c (cgPointer (cgArgument 1) left (cgPointer (cgArgument 2) right cgNoSlots))))))))))))

-- C = A B + R, the residual added in the product's epilogue (C and R
-- distinct planes)
def cgTiledAdding =
  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
    (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
      (lambda unrestricted c : Nat . (lambda unrestricted left : Nat . (lambda unrestricted right : Nat . (lambda unrestricted r : Nat .
        (cgTiledLaunch 1 a b m n k
          (cgPointer (cgArgument 0) c (cgPointer (cgArgument 1) left (cgPointer (cgArgument 2) right
            (cgPointer (cgArgument 3) r cgNoSlots))))))))))))))

-- C = X W^T (the forward): X (m x k) and W (n x k) as they are stored
def cgProductForward = (cgTiled coppeliusTiledRows coppeliusTiledRows)
-- C = dY W (a data gradient): dY (m x k), W (k x n)
def cgProductData = (cgTiled coppeliusTiledRows coppeliusTiledColumns)
-- C = dY^T X (a weight gradient): dY (k x m), X (k x n)
def cgProductWeight = (cgTiled coppeliusTiledColumns coppeliusTiledColumns)
-- the forward's and a data gradient's, plus a residual R
def cgProductForwardAdding = (cgTiledAdding coppeliusTiledRows coppeliusTiledRows)
def cgProductDataAdding = (cgTiledAdding coppeliusTiledRows coppeliusTiledColumns)

def cgTransposeTiled =
  (lambda unrestricted groups : Nat .
    (lambda unrestricted tile : Nat .
      (lambda unrestricted rows : Nat .
        (lambda unrestricted columns : Nat .
          (lambda unrestricted output : Nat .
            (lambda unrestricted input : Nat .
              (cgLaunch
                (cgImageTranspose groups rows columns)
                (naturalDivideUnchecked rows 2)
                (naturalDivideUnchecked columns tile)
                groups
                (cgPointer (cgArgument 0) output (cgPointer (cgArgument 1) input cgNoSlots)))))))))

def cgHeadTranspose =
  (cgTransposeTiled cgHeads coppeliusTransposeHeadTile)

def cgFillWith =
  (lambda unrestricted output : Nat .
    (lambda unrestricted elements : Nat .
      (lambda unrestricted value : Nat .
        (cgLaunch
          cgImageFill
          (naturalDivideUnchecked elements cgThreads)
          1
          1
          (cgSlot
            qmdBlockDimensionXOffset
            cgThreads
            (cgPointer (cgArgument 0) output (cgSlot (cgScalar 0) value cgNoSlots)))))))

def cgLayerNormScalars =
  (lambda unrestricted rest : (family StdList (family CgSlot)) .
    (cgSlot (cgScalar 0) cgInverseWidth (cgSlot (cgScalar 1) cgLayerNormEpsilon rest)))

def cgGeluScalars : (family StdList (family CgSlot)) =
  (cgSlot
    (cgScalar 0)
    cgSqrtTwoOverPi
    (cgSlot
      (cgScalar 1)
      cgGeluCubic
      (cgSlot
        (cgScalar 2)
        cgHalf32
        (cgSlot (cgScalar 3) cgOne (cgSlot (cgScalar 4) cgGeluCubicSlope cgNoSlots)))))

-- the largest grid the channel accepted for one launch of 256-thread blocks
-- the gated activation's forward and backward in one launch each
-- (Realization.Nvidia.SM86.GatedGELUSM86): bit for bit the GELU, the
-- product and the cast (the backward's two products, GELU' and two casts)
-- they replace
def cgGatedGeluForward =
  (lambda unrestricted gelu : Nat . (lambda unrestricted gated : Nat .
    (lambda unrestricted up : Nat . (lambda unrestricted gate : Nat . (lambda unrestricted elements : Nat .
      (cgLaunch cgImageGatedGeluForward (gatedGELUSM86Blocks elements) 1 1
        (cgPointer (cgArgument 0) gelu (cgPointer (cgArgument 1) gated
          (cgPointer (cgArgument 2) up (cgPointer (cgArgument 3) gate cgGeluScalars))))))))))

def cgGatedGeluBackward =
  (lambda unrestricted upGradient : Nat . (lambda unrestricted gateGradient : Nat .
    (lambda unrestricted gatedGradient : Nat . (lambda unrestricted gelu : Nat .
      (lambda unrestricted up : Nat . (lambda unrestricted gate : Nat . (lambda unrestricted elements : Nat .
        (cgLaunch cgImageGatedGeluBackward (gatedGELUSM86Blocks elements) 1 1
          (cgPointer (cgArgument 0) upGradient (cgPointer (cgArgument 1) gateGradient
            (cgPointer (cgArgument 2) gatedGradient (cgPointer (cgArgument 3) gelu
              (cgPointer (cgArgument 4) up (cgPointer (cgArgument 5) gate cgGeluScalars))))))))))))))

def cgMaximumBlocks : Nat =
  1024

def cgMaximumElements : Nat =
  (naturalMultiply cgMaximumBlocks cgThreads)

def cgParameterElements : Nat =
  (naturalDivideUnchecked cgParameterBytes 4)

-- every matrix's binary16 copy from its binary32 parameters
def cgRecastMatrix =
  (lambda unrestricted index : Nat .
    (cgCastDown
      (cgHalf index)
      (cgParameter cgBankP (cgMatrixEntry index))
      (cgEntryElements (cgMatrixEntry index))))

-- the embedding's, then each layer's (one repeat: the layers are alike)
def cgRecast : (family NvidiaLaunchSchedule) =
  (cgPhase
    cgPhaseRecast
    (cgThen
      (cgRecastMatrix cgEmbeddingMatrix)
      (nvidiaAffineLaunchRepeat
        cgLayers
        (lambda unrestricted l : Nat .
          (cgFor
            cgLayerMatrices
            (lambda unrestricted j : Nat .
              (cgRecastMatrix
                (naturalAdd cgFirstLayerMatrix (naturalAdd (naturalMultiply l cgLayerMatrices) j)))))))))

-- ---- initialization ----
-- the parameters drawn from the seeded random-normal generator
-- (cgInitialization); the other banks zero; every normalization gain one and bias zero; the
-- matrices' half copies
-- the banks a fresh run starts at zero: gradients and both moments
def cgZeroedBanks : (family StdList Nat) =
  (constructor
    StdList
    StdListCons
    Nat
    cgBankG
    (constructor
      StdList
      StdListCons
      Nat
      cgBankM
      (constructor StdList StdListCons Nat cgBankV (constructor StdList StdListEmpty Nat))))

-- an entry's initial value: the layer norms' gains one, their shifts zero
-- (the matrices come from the initialization input)
def cgEntryInitialization =
  (lambda unrestricted j : Nat .
    (cgIf
      (naturalEqual (cgEntryElements j) cgWidth)
      (cgFillWith
        (cgParameter cgBankP j)
        cgWidth
        (naturalSelect
          (naturalOr
            (naturalEqual j cgFinalGamma)
            (naturalAnd
              (naturalLessOrEqual cgFirstLayerEntry j)
              (naturalOr
                (naturalEqual
                  (naturalModuloUnchecked
                    (naturalSaturatingSubtract j cgFirstLayerEntry)
                    cgLayerItems)
                  cgItemNorm1Gain)
                (naturalEqual
                  (naturalModuloUnchecked
                    (naturalSaturatingSubtract j cgFirstLayerEntry)
                    cgLayerItems)
                  cgItemNorm2Gain))))
          cgOne
          cgZero))
      cgNothing))

-- a fresh model: every parameter drawn from the random-normal image (a
-- stream of three hashes per element, so the counter advances three per
-- element; each draw is twelve uniform bytes centred, standard deviation
-- sqrt(65535)), scaled to the model's deviation; then the norms' gains one
-- and their shifts zero
def cgImageRandomNormal : Nat =
  (cgImage b"random-normal")

def cgRandomStreams : Nat =
  3

-- the element the streams start at: past the high word of the plan's
-- highest address window.  A launch's parameter block is patched in
-- eight-byte words, and the device certificate reads every word inside a
-- window as an address; a lower counter would pair with the seed below it
-- into one
def cgRandomStreamOrigin : Nat =
  (succ
    (naturalDivideUnchecked
      (naturalAdd coppeliusSM86CompatVideoBaseNatural coppeliusSM86CompatVideoBytesNatural)
      cgWordRange))

def cgInitializationScale : Nat =
  (compile-time
    (float32ExactInverseRoot
      (naturalMultiply
        (naturalMultiply coppeliusInitializationDeviationDenominator coppeliusInitializationDeviationDenominator)
        65535)))

def cgInitialization : (family NvidiaLaunchSchedule) =
  (cgPhase
    cgPhaseInitialize
    (cgThen
      (cgPieces
        cgParameterElements
        cgMaximumElements
        (lambda unrestricted offset : Nat .
          (lambda unrestricted elements : Nat .
            (cgLaunch
              cgImageRandomNormal
              (naturalDivideUnchecked elements cgThreads)
              1
              1
              (cgSlot
                qmdBlockDimensionXOffset
                cgThreads
                (cgPointer
                  (cgArgument 0)
                  (cgArena (naturalAdd cgBankP (naturalMultiply 4 offset)))
                  (cgSlot
                    (cgScalar 0)
                    coppeliusRunSeed
                    (cgSlot
                      (cgScalar 1)
                      (naturalMultiply cgRandomStreams (naturalAdd cgRandomStreamOrigin offset))
                      (cgSlot (cgScalar 2) cgInitializationScale cgNoSlots)))))))))
      (cgThen
        (nvidiaAffineLaunchRepeat
          (stdListLength Nat cgZeroedBanks)
          (lambda unrestricted bank : Nat .
            (cgPieces
              cgParameterElements
              cgMaximumElements
              (lambda unrestricted offset : Nat .
                (lambda unrestricted elements : Nat .
                  (cgFillWith
                    (cgArena (naturalAdd (cgIndex cgZeroedBanks bank) (naturalMultiply 4 offset)))
                    elements
                    cgZero))))))
        (cgThen
          (cgThen
            (cgEntryInitialization cgEmbeddingEntry)
            (cgThen
              (nvidiaAffineLaunchRepeat
                cgLayers
                (lambda unrestricted l : Nat .
                  (cgFor
                    cgLayerItems
                    (lambda unrestricted k : Nat . (cgEntryInitialization (cgLayerEntry l k))))))
              (cgThen (cgEntryInitialization cgFinalGamma) (cgEntryInitialization cgFinalBeta))))
          cgRecast))))

-- ---- the rotary tables ----
-- theta_i = base^(-i / columns) for column i, known to 96 bits (the
-- largest r with r^columns base^i <= 2^(96 columns)), in two words: the
-- binary32 nearest it and the binary32 nearest what that misses
def cgRotaryResolution : Nat =
  rotaryFrequencyResolution

def cgRotaryThetaRoot =
  (lambda unrestricted i : Nat .
    (rotaryFrequencyRoot coppeliusRotaryBase cgRotaryColumns i))

def cgRotaryThetaHigh =
  (lambda unrestricted i : Nat .
    (rotaryFrequencyHigh coppeliusRotaryBase cgRotaryColumns i))

def cgRotaryThetaLow =
  (lambda unrestricted i : Nat .
    (rotaryFrequencyLow coppeliusRotaryBase cgRotaryColumns i))

-- every column's two words, in column order
def cgRotaryThetas : (family StdList Nat) =
  (compile-time (rotaryFrequencyWords coppeliusRotaryBase cgRotaryColumns))

-- every word rounded: an interval both of whose ends did not round to one
-- binary32 would leave float32NotRounded (a NaN) in the table
def cgRotaryThetasRounded :
  (equal Nat (stdListFold Nat Nat (lambda unrestricted w : Nat . (lambda unrestricted n : Nat . (naturalAdd n (naturalSelect (naturalEqual w float32NotRounded) 1 0)))) 0 cgRotaryThetas) 0) =
  (refl Nat 0)

-- the columns whose frequency is rational, held to the definition: column 0
-- is base^0 = 1, exact in one word; column columns/2 is base^(-1/2) = 1/100
-- (base 10000), the binary32 nearest a hundredth
def cgRotaryThetaColumnZero :
  (equal Nat (naturalAdd (cgIndex cgRotaryThetas 0) (cgIndex cgRotaryThetas 1)) (float32ExactRational 1 1)) =
  (refl Nat (float32ExactRational 1 1))

def cgRotaryThetaHalfColumns :
  (equal Nat (cgIndex cgRotaryThetas cgRotaryColumns) (float32ExactRational 1 100)) =
  (refl Nat (float32ExactRational 1 100))

-- 1.5 x 2^23: added to a binary32 below 2^22, it leaves the nearest integer
def cgRotaryRoundingMagic : Nat =
  rotaryFrequencyRoundingMagic

def cgImageRotaryTable : Nat =
  (cgImage b"rotary-table")

-- a launch per column, a thread per position
def cgRotaryTables : (family NvidiaLaunchSchedule) =
  (cgPhase
    cgPhaseRotary
    (cgFor
      cgRotaryColumns
      (lambda unrestricted i : Nat .
        (cgLaunch
          cgImageRotaryTable
          (naturalDivideUnchecked cgSeq cgThreads)
          1
          1
          (cgSlot
            qmdBlockDimensionXOffset
            cgThreads
            (cgPointer
              (cgArgument 0)
              (naturalAdd cgCosine (naturalMultiply 4 i))
              (cgPointer
                (cgArgument 1)
                (naturalAdd cgSine (naturalMultiply 4 i))
                (cgSlot
                  (cgScalar 0)
                  (cgIndex cgRotaryThetas (naturalMultiply 2 i))
                  (cgSlot
                    (cgScalar 1)
                    (cgIndex cgRotaryThetas (succ (naturalMultiply 2 i)))
                    (cgSlot
                      (cgScalar 2)
                      float32InverseTwoPi
                      (cgSlot
                        (cgScalar 3)
                        float32TwoPiHigh
                        (cgSlot
                          (cgScalar 4)
                          float32TwoPiLow
                          (cgSlot
                            (cgScalar 5)
                            cgRotaryRoundingMagic
                            (cgSlot (cgScalar 6) cgRotaryRowBytes cgNoSlots))))))))))))))

-- ---- the checkpoint through the staging window ----
-- the logical checkpoint is P, then M, then V (each the parameter bytes);
-- chunk c (the staging window's size) is copied in pieces of at most a
-- mebibyte, cut at the banks' ends
def cgCheckpointBanks : (family StdList Nat) =
  (constructor
    StdList
    StdListCons
    Nat
    cgBankP
    (constructor
      StdList
      StdListCons
      Nat
      cgBankM
      (constructor StdList StdListCons Nat cgBankV (constructor StdList StdListEmpty Nat))))

def cgCheckpointBytes : Nat =
  (naturalMultiply (stdListLength Nat cgCheckpointBanks) cgParameterBytes)

def cgChunkBytes : Nat =
  coppeliusSM86CompatCheckpointStagingBytesNatural

def cgPieceBytes : Nat =
  (naturalMultiply 4 cgMaximumElements)

def cgBankOf =
  (lambda unrestricted logical : Nat . (naturalDivideUnchecked logical cgParameterBytes))

def cgBankBase =
  (lambda unrestricted bank : Nat . (cgIndex cgCheckpointBanks bank))

def cgTransferPieces : Nat =
  8

def cgChunkTransfer =
  (lambda unrestricted load : Nat .
    (lambda unrestricted chunk : Nat .
      (cgPhase
        cgPhaseCheckpoint
        (let unrestricted start =
          (naturalMultiply chunk cgChunkBytes)
          in
          (let unrestricted end =
            (cgMinimum (naturalAdd start cgChunkBytes) cgCheckpointBytes)
            in
            (app
              (nat-eliminate
                (lambda unrestricted current : Nat .
                  (pi unrestricted cursor : Nat . (family NvidiaLaunchSchedule)))
                (lambda unrestricted cursor : Nat . cgNothing)
                (lambda unrestricted p : Nat .
                  (lambda unrestricted induction : (pi unrestricted cursor : Nat . (family NvidiaLaunchSchedule)) .
                    (lambda unrestricted cursor : Nat .
                      (cgIf
                        (naturalLess cursor end)
                        (let unrestricted bank =
                          (cgBankOf cursor)
                          in
                          (let unrestricted bankEnd =
                            (naturalMultiply (succ bank) cgParameterBytes)
                            in
                            (let unrestricted bytes =
                              (cgMinimum
                                cgPieceBytes
                                (cgMinimum
                                  (naturalSaturatingSubtract bankEnd cursor)
                                  (naturalSaturatingSubtract end cursor)))
                              in
                              (let unrestricted video =
                                (cgArena
                                  (naturalAdd
                                    (cgBankBase bank)
                                    (naturalSaturatingSubtract
                                      cursor
                                      (naturalMultiply bank cgParameterBytes))))
                                in
                                (let unrestricted host =
                                  (cgStaging (naturalSaturatingSubtract cursor start))
                                  in
                                  (cgThen
                                    (cgLaunch
                                      cgImageCopy
                                      (naturalDivideUnchecked
                                        (naturalDivideUnchecked bytes 4)
                                        cgThreads)
                                      1
                                      1
                                      (cgPointer
                                        (cgArgument 0)
                                        (naturalSelect load video host)
                                        (cgPointer
                                        (cgArgument 1)
                                        (naturalSelect load host video)
                                        cgNoSlots)))
                                    (induction (naturalAdd cursor bytes))))))))
                        cgNothing))))
                cgTransferPieces)
              start))))))

def cgCheckpointTransfer =
  (lambda unrestricted load : Nat .
    (cgFor (cgCeilDivide cgCheckpointBytes cgChunkBytes) (cgChunkTransfer load)))

-- the final diagnostics and logits kept in the arena while checkpoint
-- chunks pass through the staging window, and put back
def cgReportTransfer =
  (lambda unrestricted restore : Nat .
    (cgPhase
      cgPhaseReport
      (cgThen
        (cgLaunch
          cgImageCopy
          (naturalDivideUnchecked
            (naturalDivideUnchecked coppeliusResultLossBlockBytes 4)
            cgThreads)
          1
          1
          (cgPointer
            (cgArgument 0)
            (naturalSelect restore cgDiagnostic (cgArena cgReport))
            (cgPointer
              (cgArgument 1)
              (naturalSelect restore (cgArena cgReport) cgDiagnostic)
              cgNoSlots)))
        (cgLaunch
          cgImageCopy
          (naturalDivideUnchecked (naturalDivideUnchecked cgRowBytes 4) cgThreads)
          1
          1
          (cgPointer
            (cgArgument 0)
            (naturalSelect
              restore
              cgInference
              (cgArena (naturalAdd cgReport coppeliusResultLossBlockBytes)))
            (cgPointer
              (cgArgument 1)
              (naturalSelect
                restore
                (cgArena (naturalAdd cgReport coppeliusResultLossBlockBytes))
                cgInference)
              cgNoSlots))))))

-- ---- the forward pass ----
def cgLayerParameter =
  (lambda unrestricted l : Nat .
    (lambda unrestricted k : Nat . (cgParameter cgBankP (cgLayerEntry l k))))

def cgLayerGradient =
  (lambda unrestricted l : Nat .
    (lambda unrestricted k : Nat . (cgParameter cgBankG (cgLayerEntry l k))))

-- the half copy of layer l's matrix t (qkv 0, o 1, up 2, gate 3, down 4)
def cgLayerHalf =
  (lambda unrestricted l : Nat .
    (lambda unrestricted t : Nat .
      (cgHalf (naturalAdd cgFirstLayerMatrix (naturalAdd (naturalMultiply cgLayerMatrices l) t)))))

def cgLayerNorm =
  (lambda unrestricted output : Nat .
    (lambda unrestricted input : Nat .
      (lambda unrestricted gamma : Nat .
        (lambda unrestricted beta : Nat .
          (lambda unrestricted statistics : Nat .
            (cgLaunch
              cgImageLayerNorm
              cgSeq
              1
              1
              (cgPointer
                (cgArgument 0)
                output
                (cgPointer
                  (cgArgument 1)
                  input
                  (cgPointer
                    (cgArgument 2)
                    gamma
                    (cgPointer
                      (cgArgument 3)
                      beta
                      (cgPointer (cgArgument 4) statistics (cgLayerNormScalars cgNoSlots))))))))))))

def cgStreamingScale : Nat =
  (naturalSaturatingSubtract float32Log2E (naturalMultiply 3 8388608))
def cgScoreScaleIsAPowerOfTwo : (equal Nat cgScoreScale 1040187392) = (refl Nat 1040187392)

-- layer l's heads' attention as one streaming launch (Realization.Nvidia.
-- SM86.StreamingAttentionSM86): a block per 64 queries
-- of a head, the output written into the layer's merged half, each row's
-- base-2 log-sum-exp into its log-sum-exp (the backward's)
def cgStreamingAttention = (lambda unrestricted l : Nat .
  (cgLaunch
    cgImageAttentionForward
    (naturalDivideUnchecked cgSeq 64)
    cgHeads
    1
    (cgPointer
      (cgArgument 0)
      (cgQueriesOf l)
      (cgPointer
        (cgArgument 1)
        (cgKeysOf l)
        (cgPointer
          (cgArgument 2)
          (cgValuesTOf l)
          (cgPointer
            (cgArgument 3)
            (cgMergedHalfOf l)
            (cgPointer (cgArgument 4) (cgLogSumExpOf l) (cgSlot (cgArgument 5) cgStreamingScale cgNoSlots))))))))

-- layer l's heads' attention backward as two streaming launches
-- (Realization.Nvidia.SM86.StreamingAttentionSM86): the query launch forms
-- each row's terms D from dO and the forward's output O (the layer's merged
-- half) and dQ; the key launch dK and dV^T.  P is recomputed from the
-- layer's log-sum-exp.

def cgStreamingAttentionBackward = (lambda unrestricted l : Nat .
  (cgThen
    (cgLaunch
      cgImageAttentionQuery
      (naturalDivideUnchecked cgSeq 64)
      cgHeads
      1
      (cgPointer (cgArgument 0) (cgQueriesOf l)
      (cgPointer (cgArgument 1) (cgKeysOf l)
      (cgPointer (cgArgument 2) cgValuesHeads
      (cgPointer (cgArgument 3) cgKeysT
      (cgPointer (cgArgument 4) cgContextGradientHeads
      (cgPointer (cgArgument 5) (cgMergedHalfOf l)
      (cgPointer (cgArgument 6) (cgLogSumExpOf l)
      (cgPointer (cgArgument 7) cgQueriesGradient
      (cgPointer (cgArgument 8) cgRowTerms
      (cgSlot (cgArgument 9) cgStreamingScale
      (cgSlot (cgArgument 10) cgScoreScale cgNoSlots))))))))))))
    (cgLaunch
      cgImageAttentionKey
      (naturalDivideUnchecked cgSeq 32)
      cgHeads
      1
      (cgPointer (cgArgument 0) (cgQueriesOf l)
      (cgPointer (cgArgument 1) (cgKeysOf l)
      (cgPointer (cgArgument 2) cgValuesHeads
      (cgPointer (cgArgument 3) cgContextGradientHeads
      (cgPointer (cgArgument 4) cgQueriesT
      (cgPointer (cgArgument 5) cgContextGradientT
      (cgPointer (cgArgument 6) (cgLogSumExpOf l)
      (cgPointer (cgArgument 7) cgRowTerms
      (cgPointer (cgArgument 8) cgKeysGradient
      (cgPointer (cgArgument 9) cgValuesGradient
      (cgSlot (cgArgument 10) cgStreamingScale
      (cgSlot (cgArgument 11) cgScoreScale cgNoSlots)))))))))))))))

-- A layer's forward with `attention` the heads' attention into the layer's
-- merged half (cgStreamingAttention); what the backward reads it writes into
-- the layer's saved planes (cgSaved)

def cgLayerForwardWith =
  (lambda unrestricted attention : (pi unrestricted layer : Nat . (family NvidiaLaunchSchedule)) .
  (lambda unrestricted l : Nat .
    (lambda unrestricted x : Nat .
      (lambda unrestricted out : Nat .
        (cgPhase
          cgPhaseLayer
          (cgThen
            (cgLayerNorm
              cgNorm1
              x
              (cgLayerParameter l cgItemNorm1Gain)
              (cgLayerParameter l cgItemNorm1Shift)
              (cgNorm1StatisticsOf l))
            (cgThen
              (cgCastDown (cgNorm1HalfOf l) cgNorm1 (naturalMultiply cgSeq cgWidth))
              (cgThen
                (cgProductForward
                  cgSeq
                  (naturalMultiply cgQkvParts cgWidth)
                  cgWidth
                  cgQkv
                  (cgNorm1HalfOf l)
                  (cgLayerHalf l cgMatrixQkv))
                (cgThen
                  (cgLaunch
                    cgImageQKRope
                    cgSeq
                    1
                    1
                    (cgPointer
                      (cgArgument 0)
                      (cgQueriesOf l)
                      (cgPointer
                        (cgArgument 1)
                        (cgKeysOf l)
                        (cgPointer
                          (cgArgument 2)
                          cgQkv
                          (cgPointer
                            (cgArgument 3)
                            cgCosine
                            (cgPointer (cgArgument 4) cgSine cgNoSlots))))))
                  (cgThen
                    (cgLaunch
                      cgImageValueTranspose
                      (naturalDivideUnchecked cgSeq 2)
                      1
                      1
                      (cgPointer
                        (cgArgument 0)
                        (cgValuesTOf l)
                        (cgPointer (cgArgument 2) cgQkv cgNoSlots)))
                    (cgThen
                      (attention l)
                      (cgThen (cgProductForwardAdding cgSeq cgWidth cgWidth (cgProjectionOf l) (cgMergedHalfOf l) (cgLayerHalf l cgMatrixOutput) x) (cgThen (cgLayerNorm
                                cgNorm2
                                (cgProjectionOf l)
                                (cgLayerParameter l cgItemNorm2Gain)
                                (cgLayerParameter l cgItemNorm2Shift)
                                (cgNorm2StatisticsOf l)) (cgThen (cgCastDown (cgNorm2HalfOf l) cgNorm2 (naturalMultiply cgSeq cgWidth)) (cgThen (cgProductForward
                                    cgSeq
                                    cgFfn
                                    cgWidth
                                    (cgUpOf l)
                                    (cgNorm2HalfOf l)
                                    (cgLayerHalf l cgMatrixUp)) (cgThen (cgProductForward
                                      cgSeq
                                      cgFfn
                                      cgWidth
                                      (cgGateOf l)
                                      (cgNorm2HalfOf l)
                                      (cgLayerHalf l cgMatrixGate)) (cgThen (cgGatedGeluForward (cgGeluOf l) (cgGatedHalfOf l) (cgUpOf l) (cgGateOf l) (naturalMultiply cgSeq cgFfn)) (cgProductForwardAdding cgSeq cgWidth cgFfn out (cgGatedHalfOf l) (cgLayerHalf l cgMatrixDown) (cgProjectionOf l)))))))))))))))))))

def cgLayerForward = (cgLayerForwardWith cgStreamingAttention)

def cgForward : (family NvidiaLaunchSchedule) =
  (cgThen
    (cgPhase
      cgPhaseEmbed
      (cgLaunch
        cgImageGather
        1
        cgSeq
        1
        (cgPointer
          (cgArgument 0)
          (cgBoundary 0)
          (cgPointer
            (cgArgument 1)
            cgIds
            (cgPointer (cgArgument 2) (cgParameter cgBankP cgEmbeddingEntry) cgNoSlots)))))
    (cgThen
      (nvidiaAffineLaunchRepeat
        cgLayers
        (lambda unrestricted l : Nat . (cgLayerForward l (cgBoundary l) (cgBoundary (succ l)))))
      (cgPhase
        cgPhaseOutput
        (cgThen
          (cgLayerNorm
            cgFinalNorm
            (cgBoundary cgLayers)
            (cgParameter cgBankP cgFinalGamma)
            (cgParameter cgBankP cgFinalBeta)
            cgFinalStatistics)
          (cgThen
            (cgCastDown cgFinalHalf cgFinalNorm (naturalMultiply cgSeq cgWidth))
            (cgThen
              (cgProductForward cgSeq cgVocab cgWidth cgLogits cgFinalHalf (cgHalf cgEmbeddingMatrix))
              (cgCopy
                cgInference
                (naturalAdd
                  cgLogits
                  (naturalMultiply (naturalSaturatingSubtract cgSeq 1) cgRowBytes))
                cgVocab)))))))

-- the prediction rows' logits and their greedy tokens
def cgPredictionRows : (family StdList Nat) =
  (constructor
    StdList
    StdListCons
    Nat
    coppeliusPredictionRow0
    (constructor
      StdList
      StdListCons
      Nat
      coppeliusPredictionRow1
      (constructor
        StdList
        StdListCons
        Nat
        coppeliusPredictionRow2
        (constructor StdList StdListEmpty Nat))))

def cgPredictionRow =
  (lambda unrestricted index : Nat . (cgIndex cgPredictionRows index))

def cgPrediction : (family NvidiaLaunchSchedule) =
  (cgThen
    cgForward
    (cgPhase
      cgPhaseOutput
      (cgFor
        coppeliusPredictionRows
        (lambda unrestricted index : Nat .
          (let unrestricted rowLogits =
            (naturalAdd cgInference (naturalMultiply index cgRowBytes))
            in
            (cgThen
              (cgCopy
                rowLogits
                (naturalAdd cgLogits (naturalMultiply (cgPredictionRow index) cgRowBytes))
                cgVocab)
              (cgLaunch
                cgImageArgmax
                1
                1
                1
                (cgPointer
                  (cgArgument 0)
                  (cgStaging
                    (naturalAdd
                      coppeliusResultPredictionsOffset
                      (naturalMultiply index coppeliusResultPredictionBytes)))
                  (cgPointer (cgArgument 1) rowLogits cgNoSlots)))))))))

-- ---- the backward pass ----
def cgLayerNormBackward =
  (lambda unrestricted gamma : Nat .
    (lambda unrestricted gammaGradient : Nat .
      (lambda unrestricted betaGradient : Nat .
        (lambda unrestricted x : Nat .
          (lambda unrestricted dy : Nat .
            (lambda unrestricted statistics : Nat .
              (lambda unrestricted dx : Nat .
                (cgThen
                  (cgLaunch
                    cgImageLayerNormBackward
                    cgSeq
                    1
                    1
                    (cgPointer
                      (cgArgument 0)
                      dx
                      (cgPointer
                        (cgArgument 1)
                        x
                        (cgPointer
                          (cgArgument 2)
                          dy
                          (cgPointer
                            (cgArgument 3)
                            gamma
                            (cgPointer
                              (cgArgument 4)
                              (cgArena cgXHat)
                              (cgPointer (cgArgument 5) statistics (cgLayerNormScalars cgNoSlots))))))))
                  (cgLaunch
                    cgImageLayerNormParameters
                    cgWidth
                    1
                    1
                    (cgPointer
                      (cgArgument 0)
                      gammaGradient
                      (cgPointer
                        (cgArgument 1)
                        dy
                        (cgPointer
                          (cgArgument 2)
                          (cgArena cgXHat)
                          (cgPointer
                            (cgArgument 3)
                            betaGradient
                            (cgSlot (cgScalar 0) cgInverseWidth cgNoSlots))))))))))))))

-- ---- the loss ----
-- `correct` holds each row's target logit (the gather), so the rows hold
-- each row's cross-entropy; window 0 of the diagnostic block keeps them;
-- the backward reads loss + correct, the log-sum-exp
def cgLossForward : (family NvidiaLaunchSchedule) =
  (cgPhase
    cgPhaseOutput
    (cgThen
      (cgLaunch
        cgImageTargetGather
        (naturalDivideUnchecked cgSeq cgThreads)
        1
        1
        (cgPointer
          (cgArgument 0)
          cgCorrect
          (cgPointer (cgArgument 1) cgLogits (cgPointer (cgArgument 2) cgTargets cgNoSlots))))
      (cgThen
        (cgLaunch
          cgImageCrossEntropy
          cgSeq
          1
          1
          (cgPointer
            (cgArgument 0)
            cgLoss
            (cgPointer
              (cgArgument 1)
              cgLogits
              (cgPointer
                (cgArgument 2)
                cgCorrect
                (cgSlot
                  (cgScalar 0)
                  float32Log2E
                  (cgSlot (cgScalar 1) float32Ln2 (cgSlot (cgScalar 2) cgInverseSeq cgNoSlots)))))))
        (cgCopy cgDiagnostic cgLoss cgWindowWords))))

-- layer l's backward: `current` holds dy entering the layer from above and
-- `other` the other residual gradient plane; the layer is recomputed first,
-- then in order: the feed-forward output projection; the gated GELU, the up
-- and gate projections and the second norm; the attention output
-- projection; the heads' shared layouts, then the heads' attention
-- backward (cgStreamingAttentionBackward); the qkv projection and the first
-- norm
def cgLayerBackward =
  (lambda unrestricted l : Nat .
    (lambda unrestricted current : Nat .
      (lambda unrestricted other : Nat .
        (let unrestricted x =
          (cgBoundary l)
          in
          (cgThen
              (cgPhase
                cgPhaseFfnDown
                (cgThen (cgCastDown cgDownDyHalf current (naturalMultiply cgSeq cgWidth)) (cgThen (cgProductData cgSeq cgFfn cgWidth cgGatedGradient cgDownDyHalf (cgLayerHalf l cgMatrixDown)) (cgProductWeight cgWidth cgFfn cgSeq (cgLayerGradient l cgItemDown) cgDownDyHalf (cgGatedHalfOf l)))))
              (cgThen
                (cgPhase
                  cgPhaseFfnGate
                  (cgThen (cgGatedGeluBackward cgFfnGradientHalf cgGateGradientHalf cgGatedGradient (cgGeluOf l) (cgUpOf l) (cgGateOf l) (naturalMultiply cgSeq cgFfn)) (cgThen (cgProductWeight cgFfn cgWidth cgSeq (cgLayerGradient l cgItemUp) cgFfnGradientHalf (cgNorm2HalfOf l)) (cgThen (cgProductData cgSeq cgWidth cgFfn cgNorm2UpGradient cgFfnGradientHalf (cgLayerHalf l cgMatrixUp)) (cgThen (cgProductWeight cgFfn cgWidth cgSeq (cgLayerGradient l cgItemGate) cgGateGradientHalf (cgNorm2HalfOf l)) (cgThen (cgProductDataAdding cgSeq cgWidth cgFfn cgNorm2Gradient cgGateGradientHalf (cgLayerHalf l cgMatrixGate) cgNorm2UpGradient) (cgThen (cgLayerNormBackward
                                        (cgLayerParameter l cgItemNorm2Gain)
                                        (cgLayerGradient l cgItemNorm2Gain)
                                        (cgLayerGradient l cgItemNorm2Shift)
                                        (cgProjectionOf l)
                                        cgNorm2Gradient
                                        (cgNorm2StatisticsOf l)
                                        cgNorm2InputGradient) (cgBinary
                                        cgImageAdd
                                        other
                                        current
                                        cgNorm2InputGradient
                                        (naturalMultiply cgSeq cgWidth)))))))))
                (cgThen
                  (cgPhase
                    cgPhaseAttentionOutput
                    (cgThen (cgCastDown cgOutputDyHalf other (naturalMultiply cgSeq cgWidth)) (cgThen (cgProductWeight cgWidth cgWidth cgSeq (cgLayerGradient l cgItemOutput) cgOutputDyHalf (cgMergedHalfOf l)) (cgProductData cgSeq cgWidth cgWidth cgMergedGradient cgOutputDyHalf (cgLayerHalf l cgMatrixOutput)))))
                  (cgThen
                    (cgPhase
                      cgPhaseHeadLayout
                      (cgThen
                        (cgLaunch
                          cgImageTokenToHead
                          cgSeq
                          1
                          1
                          (cgPointer
                            (cgArgument 0)
                            cgContextGradientHeads
                            (cgPointer (cgArgument 2) cgMergedGradient cgNoSlots)))
                        (cgThen
                          (cgHeadTranspose cgHeadWidth cgSeq cgValuesHeads (cgValuesTOf l))
                          (cgThen
                            (cgHeadTranspose cgSeq cgHeadWidth cgQueriesT (cgQueriesOf l))
                            (cgThen
                              (cgHeadTranspose cgSeq cgHeadWidth cgKeysT (cgKeysOf l))
                              (cgHeadTranspose
                                cgSeq
                                cgHeadWidth
                                cgContextGradientT
                                cgContextGradientHeads))))))
                    (cgThen
                      (cgPhase cgPhaseAttentionHead (cgStreamingAttentionBackward l))
                      (cgPhase
                        cgPhaseQkv
                        (cgThen (cgLaunch
                            cgImageInverseRope
                            cgSeq
                            1
                            1
                            (cgPointer
                              (cgArgument 0)
                              cgQkvGradientHalf
                              (cgPointer
                                (cgArgument 1)
                                cgQueriesGradient
                                (cgPointer
                                  (cgArgument 2)
                                  cgKeysGradient
                                  (cgPointer
                                    (cgArgument 3)
                                    cgCosine
                                    (cgPointer
                                      (cgArgument 4)
                                      cgSine
                                      (cgPointer (cgArgument 5) cgValuesGradient cgNoSlots))))))) (cgThen (cgProductWeight (naturalMultiply cgQkvParts cgWidth) cgWidth cgSeq (cgLayerGradient l cgItemQkv) cgQkvGradientHalf (cgNorm1HalfOf l)) (cgThen (cgProductData cgSeq cgWidth (naturalMultiply cgQkvParts cgWidth) cgNorm1Gradient cgQkvGradientHalf (cgLayerHalf l cgMatrixQkv)) (cgThen (cgLayerNormBackward
                                        (cgLayerParameter l cgItemNorm1Gain)
                                        (cgLayerGradient l cgItemNorm1Gain)
                                        (cgLayerGradient l cgItemNorm1Shift)
                                        x
                                        cgNorm1Gradient
                                        (cgNorm1StatisticsOf l)
                                        cgNorm1InputGradient) (cgBinary
                                        cgImageAdd
                                        current
                                        other
                                        cgNorm1InputGradient
                                        (naturalMultiply cgSeq cgWidth))))))))))))))))

-- the residual gradient: `current` (dy0) carries it down the layers and
-- `other` (dy1) holds the attention branch's share; each layer reads and
-- rewrites both, in place
def cgBackward : (family NvidiaLaunchSchedule) =
  (cgThen
    cgLossForward
    (cgThen
      (cgPhase
        cgPhaseOutput
        (cgLaunch
          cgImageCrossEntropyBackward
          cgSeq
          (naturalDivideUnchecked cgVocab cgThreads)
          1
          (cgPointer
            (cgArgument 0)
            cgLogitsGradient
            (cgPointer
              (cgArgument 1)
              cgLogits
              (cgPointer
                (cgArgument 2)
                cgLoss
                (cgPointer
                  (cgArgument 3)
                  cgCorrect
                  (cgPointer
                    (cgArgument 4)
                    cgTargets
                    (cgSlot (cgScalar 0) float32Log2E (cgSlot (cgScalar 2) cgLossGradientScale cgNoSlots)))))))))
      (cgThen
        (cgPhase
          cgPhaseOutputGradient
          (cgThen (cgCastDown cgLogitsGradientHalf cgLogitsGradient (naturalMultiply cgSeq cgVocab)) (cgThen (cgProductWeight cgVocab cgWidth cgSeq (cgParameter cgBankG cgEmbeddingEntry) cgLogitsGradientHalf cgFinalHalf) (cgThen (cgProductData cgSeq cgWidth cgVocab cgFinalGradient cgLogitsGradientHalf (cgHalf cgEmbeddingMatrix)) (cgLayerNormBackward
                      (cgParameter cgBankP cgFinalGamma)
                      (cgParameter cgBankG cgFinalGamma)
                      (cgParameter cgBankG cgFinalBeta)
                      (cgBoundary cgLayers)
                      cgFinalGradient
                      cgFinalStatistics
                      (cgArena cgDy0))))))
        (cgThen
          (nvidiaAffineLaunchRepeat
            cgLayers
            (lambda unrestricted fromTop : Nat .
              (cgLayerBackward
                (naturalSaturatingSubtract (naturalSaturatingSubtract cgLayers 1) fromTop)
                (cgArena cgDy0)
                (cgArena cgDy1))))
          (cgPhase
            cgPhaseScatter
            (cgLaunch
              cgImageScatter
              cgSeq
              1
              1
              (cgPointer
                (cgArgument 0)
                (cgParameter cgBankG cgEmbeddingEntry)
                (cgPointer
                  (cgArgument 1)
                  (cgArena cgDy0)
                  (cgPointer (cgArgument 2) cgIds cgNoSlots)))))))))

-- ---- AdamW ----
def cgAdamScalars =
  (lambda unrestricted step : Nat .
    (cgSlot
      (cgAdamWScalar 0)
      cgOne
      (cgSlot
        (cgAdamWScalar 1)
        cgBeta1
        (cgSlot
          (cgAdamWScalar 2)
          cgOneMinusBeta1
          (cgSlot
            (cgAdamWScalar 3)
            cgBeta2
            (cgSlot
              (cgAdamWScalar 4)
              cgOneMinusBeta2
              (cgSlot
                (cgAdamWScalar 5)
                cgDecay
                (cgSlot
                  (cgAdamWScalar 6)
                  (cgAdamStepSize step)
                  (cgSlot
                    (cgAdamWScalar 7)
                    (cgAdamEpsilon step)
                    (cgSlot (cgAdamWScalar 8) cgOne cgNoSlots))))))))))

-- one launch over the whole bank (Coppelius.Build.DeviceImages.
-- coppeliusAdamWImage: a grid stride over element pairs, each pair's half
-- copies written with its update, so no cast follows)
def cgAdamW =
  (lambda unrestricted step : Nat .
    (cgPhase
      cgPhaseAdamW
      (cgLaunch
        cgImageAdamW
        coppeliusAdamWBlocks
        1
        1
        (cgPointer
          (cgArgument 0)
          (cgArena cgBankP)
          (cgPointer
            (cgArgument 1)
            (cgArena cgBankG)
            (cgPointer
              (cgArgument 2)
              (cgArena cgBankM)
              (cgPointer
                (cgArgument 3)
                (cgArena cgBankV)
                (cgPointer (cgArgument 9) (cgArena cgHalfStart) (cgAdamScalars step)))))))))

-- the launch covers the bank: its image walks exactly the bank's elements
def cgAdamWCoversTheBank : (equal Nat coppeliusAdamWElements cgParameterElements) =
  (refl Nat cgParameterElements)

-- ---- the invocation ----
-- one update's shared part: gradients zeroed, forward, backward, and the
-- embedding gradient's first scalars to window 1
-- (the gradients are zero: the initialization clears them, and every
-- AdamW clears each after its update)
def cgUpdate : (family NvidiaLaunchSchedule) =
  (cgThen
    cgForward
    (cgThen
      cgBackward
      (cgPhase
        cgPhaseDiagnostics
        (cgCopy (cgWindow 1) (cgParameter cgBankG cgEmbeddingEntry) cgWindowWords))))

-- after the final update: the forward on the result, the final
-- diagnostics, the report kept aside
def cgFinal : (family NvidiaLaunchSchedule) =
  (cgThen
    cgForward
    (cgPhase
      cgPhaseDiagnostics
      (cgThen
        (cgCopy (cgWindow 1) (cgParameter cgBankG cgEmbeddingEntry) cgWindowWords)
        (cgThen
          (cgCopy (cgWindow 2) (cgParameter cgBankP cgEmbeddingEntry) cgWindowWords)
          (cgThen
            (cgCopy (cgWindow 3) (cgParameter cgBankM cgEmbeddingEntry) cgWindowWords)
            (cgReportTransfer 0))))))

-- the loss record: the step's rows (window 0) to its slot, which the host
-- appends to the result after every step
def cgLossRecord : (family NvidiaLaunchSchedule) =
  (cgPhase
    cgPhaseLossRecord
    (nvidiaAffineLaunchRepeat
      coppeliusUpdatePieces
      (lambda unrestricted k : Nat .
        (cgCopy
          (cgStaging
            (naturalAdd
              coppeliusResultLossRecordOffset
              (naturalMultiply k coppeliusResultWindowBytes)))
          (cgWindow 0)
          cgWindowWords))))

-- In launch order: initialize; the default checkpoint saved and the
-- request's loaded (chunk by chunk); the half copies; the prediction; one
-- update's shared part; each update's AdamW (which writes the half
-- copies); the final forward; the checkpoint saved; the report put back;
-- the loss record.
-- coppeliusSubmissionSchedule (below) submits them.
def cgBeforeAdamW : (family NvidiaLaunchSchedule) =
  (cgThen
    cgInitialization
    (cgThen
      (cgCheckpointTransfer 0)
      (cgThen
        (cgCheckpointTransfer 1)
        (cgThen cgRecast (cgThen cgRotaryTables (cgThen cgPrediction cgUpdate))))))

def coppeliusLaunchSchedule : (family NvidiaLaunchSchedule) =
  (cgThen
    cgBeforeAdamW
    (cgThen
      (cgFor
        coppeliusUpdatePieces
        (lambda unrestricted k : Nat . (cgAdamW (naturalAdd coppeliusFirstStep k))))
      (cgThen
        cgFinal
        (cgThen (cgCheckpointTransfer 0) (cgThen (cgReportTransfer 1) cgLossRecord)))))

-- The launches of update k's AdamW (from 0; each update's pieces in
-- order): after the launches before the first, the earlier updates'.  The
-- host writes each update's step size and epsilon into these launches'
-- parameter blocks (Platform.Linux.Nvidia.PlanHostAdamW); the values baked
-- here are the first invocation's.
def coppeliusAdamWPieces : Nat =
  (nvidiaLaunchScheduleCount (cgAdamW coppeliusFirstStep))

def coppeliusAdamWFirstLaunch : Nat =
  (nvidiaLaunchScheduleCount cgBeforeAdamW)

def coppeliusAdamWLaunch =
  (lambda unrestricted k : Nat .
    (lambda unrestricted piece : Nat .
      (naturalAdd coppeliusAdamWFirstLaunch
        (naturalAdd (naturalMultiply k coppeliusAdamWPieces) piece))))

-- ---- counting ----
def coppeliusLaunchCount : Nat =
  (nvidiaLaunchScheduleCount coppeliusLaunchSchedule)

-- ---- submissions ----
-- Each submission is a list of ranges of the launch order (adjacent ranges
-- merged) and its semaphore slot, one per submission.  The phases' launch
-- counts place them: initialize; for every checkpoint chunk but the last,
-- its default save then its load; the last save; the last load with the
-- half copies; the prediction; the update (its shared part, its AdamW, the
-- half copies, its loss-record copy); the final forward; every save chunk;
-- the report put back.  The host issues the update once per step and the
-- save chunks at every checkpoint (Coppelius.Build.NativeHost).
def cgChunks : Nat =
  (cgCeilDivide cgCheckpointBytes cgChunkBytes)

def cgChunkLaunches =
  (lambda unrestricted chunk : Nat . (nvidiaLaunchScheduleCount (cgChunkTransfer 0 chunk)))

-- launches before each chunk of a transfer, in one pass (the plan looks
-- them up for every chunk's submission)
def cgChunkStarts : (family StdList Nat) =
  (cgBumpTable 1 0 cgChunkLaunches (succ cgChunks))

def cgChunkStart =
  (lambda unrestricted chunk : Nat . (cgIndex cgChunkStarts chunk))

def cgTransferLaunches : Nat =
  (cgChunkStart cgChunks)

def cgInitializationLaunches : Nat =
  (nvidiaLaunchScheduleCount cgInitialization)

def cgRecastLaunches : Nat =
  (nvidiaLaunchScheduleCount cgRecast)

def cgDefaultSaveStart : Nat =
  cgInitializationLaunches

def cgLoadStart : Nat =
  (naturalAdd cgDefaultSaveStart cgTransferLaunches)

def cgRecastStart : Nat =
  (naturalAdd cgLoadStart cgTransferLaunches)

def cgRotaryTablesStart : Nat =
  (naturalAdd cgRecastStart cgRecastLaunches)

def cgPredictionStart : Nat =
  (naturalAdd cgRotaryTablesStart (nvidiaLaunchScheduleCount cgRotaryTables))

def cgUpdateStart : Nat =
  (naturalAdd cgPredictionStart (nvidiaLaunchScheduleCount cgPrediction))

def cgUpdateLaunches : Nat =
  (nvidiaLaunchScheduleCount cgUpdate)

def cgAdamStart : Nat =
  (naturalAdd cgUpdateStart cgUpdateLaunches)

def cgAdamLaunches : Nat =
  (nvidiaLaunchScheduleCount (cgAdamW coppeliusFirstStep))

def cgFinalStart : Nat =
  (naturalAdd cgAdamStart (naturalMultiply coppeliusUpdatePieces cgAdamLaunches))

def cgSaveStart : Nat =
  (naturalAdd cgFinalStart (nvidiaLaunchScheduleCount cgFinal))

def cgRestoreStart : Nat =
  (naturalAdd cgSaveStart cgTransferLaunches)

def cgLossRecordStart : Nat =
  (naturalAdd cgRestoreStart (nvidiaLaunchScheduleCount (cgReportTransfer 1)))

-- ranges in reverse order of addition; a range that starts where the
-- last one ends extends it
def cgPush =
  (lambda unrestricted ranges : (family StdList (family CgRange)) .
    (lambda unrestricted first : Nat .
      (lambda unrestricted count : Nat .
        (eliminate
          StdList
          (lambda unrestricted current : (family StdList (family CgRange)) .
            (family StdList (family CgRange)))
          ranges
          (branch
            StdListEmpty
            .
            (constructor
              StdList
              StdListCons
              (family CgRange)
              (constructor CgRange CgRangeValue first count)
              ranges))
          (branch
            StdListCons
            head
            tail
            induction
            .
            (eliminate
              CgRange
              (lambda unrestricted current : (family CgRange) . (family StdList (family CgRange)))
              head
              (branch
                CgRangeValue
                lastFirst
                lastCount
                .
                (nat-eliminate
                  (lambda unrestricted current : Nat . (family StdList (family CgRange)))
                  (constructor
                    StdList
                    StdListCons
                    (family CgRange)
                    (constructor CgRange CgRangeValue first count)
                    ranges)
                  (lambda unrestricted q : Nat .
                    (lambda unrestricted ignored : (family StdList (family CgRange)) .
                      (constructor
                        StdList
                        StdListCons
                        (family CgRange)
                        (constructor CgRange CgRangeValue lastFirst (naturalAdd lastCount count))
                        tail)))
                  (naturalEqual (naturalAdd lastFirst lastCount) first)))))))))

def cgNoRanges : (family StdList (family CgRange)) =
  (constructor StdList StdListEmpty (family CgRange))

def cgReferences =
  (lambda unrestricted ranges : (family StdList (family CgRange)) .
    (eliminate
      StdList
      (lambda unrestricted current : (family StdList (family CgRange)) .
        (family NvidiaLaunchReferences))
      ranges
      (branch StdListEmpty . (constructor NvidiaLaunchReferences NvidiaLaunchReferencesEmpty))
      (branch
        StdListCons
        head
        tail
        induction
        .
        (eliminate
          CgRange
          (lambda unrestricted current : (family CgRange) . (family NvidiaLaunchReferences))
          head
          (branch
            CgRangeValue
            first
            count
            .
            (constructor
              NvidiaLaunchReferences
              NvidiaLaunchReferencesAppend
              induction
              (constructor NvidiaLaunchReferences NvidiaLaunchReferencesRange first count)))))))

-- the submissions, each with its ranges; the semaphore slot follows the
-- submission's place
def cgSubmission =
  (lambda unrestricted index : Nat .
    (lambda unrestricted ranges : (family StdList (family CgRange)) .
      (constructor
        NvidiaSubmissionSchedule
        NvidiaSubmissionScheduleOne
        (constructor
          NvidiaSubmissionBatch
          NvidiaSubmissionBatchValue
          b"coppelius/submission"
          (naturalMultiply index prSemaphoreStride)
          (cgReferences ranges)))))

def cgRange =
  (lambda unrestricted first : Nat .
    (lambda unrestricted count : Nat . (cgPush cgNoRanges first count)))

def cgSubmissionThen =
  (lambda unrestricted a : (family NvidiaSubmissionSchedule) .
    (lambda unrestricted b : (family NvidiaSubmissionSchedule) .
      (constructor NvidiaSubmissionSchedule NvidiaSubmissionScheduleAppend a b)))

def cgSubmissionFor =
  (lambda unrestricted count : Nat .
    (lambda unrestricted body : (pi unrestricted index : Nat . (family NvidiaSubmissionSchedule)) .
      (nat-eliminate
        (lambda unrestricted current : Nat . (family NvidiaSubmissionSchedule))
        (constructor NvidiaSubmissionSchedule NvidiaSubmissionScheduleEmpty)
        (lambda unrestricted p : Nat .
          (lambda unrestricted induction : (family NvidiaSubmissionSchedule) .
            (cgSubmissionThen
              (body (naturalSaturatingSubtract (naturalSaturatingSubtract count 1) p))
              induction)))
        count)))

def cgChunkRange =
  (lambda unrestricted transferStart : Nat .
    (lambda unrestricted chunk : Nat .
      (cgRange (naturalAdd transferStart (cgChunkStart chunk)) (cgChunkLaunches chunk))))

def cgLastChunk : Nat =
  (naturalSaturatingSubtract cgChunks 1)

-- the submissions up to the prediction: initialize; for every checkpoint
-- chunk but the last, its default save then its load; the last save; the
-- last load with the half copies; the rotary tables and the prediction
def cgPrefixSubmissions : (family NvidiaSubmissionSchedule) =
  (cgSubmissionThen
    (cgSubmission 0 (cgRange 0 cgInitializationLaunches))
    (cgSubmissionThen
      (cgSubmissionFor
        cgLastChunk
        (lambda unrestricted chunk : Nat .
          (cgSubmissionThen
            (cgSubmission
              (naturalAdd 1 (naturalMultiply 2 chunk))
              (cgChunkRange cgDefaultSaveStart chunk))
            (cgSubmission (naturalAdd 2 (naturalMultiply 2 chunk)) (cgChunkRange cgLoadStart chunk)))))
      (cgSubmissionThen
        (cgSubmission
          (naturalAdd 1 (naturalMultiply 2 cgLastChunk))
          (cgChunkRange cgDefaultSaveStart cgLastChunk))
        (cgSubmissionThen
          (cgSubmission
            (naturalAdd 2 (naturalMultiply 2 cgLastChunk))
            (cgPush (cgChunkRange cgLoadStart cgLastChunk) cgRecastStart cgRecastLaunches))
          (cgSubmission
            (naturalAdd 3 (naturalMultiply 2 cgLastChunk))
            (cgRange cgRotaryTablesStart (naturalAdd (nvidiaLaunchScheduleCount cgRotaryTables) (nvidiaLaunchScheduleCount cgPrediction))))))))

-- the update: its shared part, its AdamW (the half copies with it), its
-- loss copy
def cgUpdateSubmissions : (family NvidiaSubmissionSchedule) =
  (cgSubmissionFor
    coppeliusUpdatePieces
    (lambda unrestricted k : Nat .
      (cgSubmission
        (naturalAdd 4 (naturalAdd (naturalMultiply 2 cgLastChunk) k))
        (cgPush
          (cgPush
            (cgRange cgUpdateStart cgUpdateLaunches)
            (naturalAdd cgAdamStart (naturalMultiply k cgAdamLaunches))
            cgAdamLaunches)
          (naturalAdd cgLossRecordStart k)
          1))))

def cgFinalSubmission : (family NvidiaSubmissionSchedule) =
  (cgSubmission
    (naturalAdd 4 (naturalAdd (naturalMultiply 2 cgLastChunk) coppeliusUpdatePieces))
    (cgRange cgFinalStart (nvidiaLaunchScheduleCount cgFinal)))

-- every save chunk
def cgGatherSubmissions : (family NvidiaSubmissionSchedule) =
  (cgSubmissionFor
    cgChunks
    (lambda unrestricted chunk : Nat .
      (cgSubmission
        (naturalAdd 5 (naturalAdd (naturalMultiply 2 cgLastChunk) (naturalAdd coppeliusUpdatePieces chunk)))
        (cgChunkRange cgSaveStart chunk))))

-- the report put back
def cgRestoreSubmission : (family NvidiaSubmissionSchedule) =
  (cgSubmission
    (naturalAdd 5 (naturalAdd (naturalMultiply 2 cgLastChunk) (naturalAdd coppeliusUpdatePieces cgChunks)))
    (cgRange cgRestoreStart (nvidiaLaunchScheduleCount (cgReportTransfer 1))))

def coppeliusSubmissionSchedule : (family NvidiaSubmissionSchedule) =
  (cgSubmissionThen cgPrefixSubmissions
    (cgSubmissionThen cgUpdateSubmissions
      (cgSubmissionThen cgFinalSubmission
        (cgSubmissionThen cgGatherSubmissions cgRestoreSubmission))))

-- The order the run issues the plan's submissions in (Coppelius.Build.
-- NativeHost), at its smallest: the prefix; two steps; a checkpoint; two
-- steps; the final forward; the last checkpoint; the report put back.
-- Every transition a longer run makes is one of these (a step after a step,
-- a step after a checkpoint, the final forward after a step), so the device
-- certificate over this order (coppeliusRunDeviceHazard) covers any run.
def coppeliusRunSubmissionSample : (family NvidiaSubmissionSchedule) =
  (cgSubmissionThen cgPrefixSubmissions
    (cgSubmissionThen cgUpdateSubmissions
      (cgSubmissionThen cgUpdateSubmissions
        (cgSubmissionThen cgGatherSubmissions
          (cgSubmissionThen cgUpdateSubmissions
            (cgSubmissionThen cgUpdateSubmissions
              (cgSubmissionThen cgFinalSubmission
                (cgSubmissionThen cgGatherSubmissions cgRestoreSubmission))))))))

def cgSubmissionCount =
  (lambda unrestricted schedule : (family NvidiaSubmissionSchedule) .
    (eliminate
      NvidiaSubmissionSchedule
      (lambda unrestricted current : (family NvidiaSubmissionSchedule) . Nat)
      schedule
      (branch NvidiaSubmissionScheduleEmpty . 0)
      (branch NvidiaSubmissionScheduleOne batch . 1)
      (branch
        NvidiaSubmissionScheduleAppend
        left
        right
        leftCount
        rightCount
        .
        (naturalAdd leftCount rightCount))
      (branch
        NvidiaSubmissionScheduleRepeat
        count
        referenceStride
        semaphoreStride
        body
        bodyCount
        .
        (naturalMultiply count bodyCount))
      (branch NvidiaSubmissionScheduleProfiled offset body bodyCount . bodyCount)))

def cgReferenceCount =
  (lambda unrestricted references : (family NvidiaLaunchReferences) .
    (eliminate
      NvidiaLaunchReferences
      (lambda unrestricted current : (family NvidiaLaunchReferences) . Nat)
      references
      (branch NvidiaLaunchReferencesEmpty . 0)
      (branch NvidiaLaunchReferencesRange first count . count)
      (branch
        NvidiaLaunchReferencesAppend
        left
        right
        leftCount
        rightCount
        .
        (naturalAdd leftCount rightCount))))

def cgSubmissionReferences =
  (lambda unrestricted schedule : (family NvidiaSubmissionSchedule) .
    (eliminate
      NvidiaSubmissionSchedule
      (lambda unrestricted current : (family NvidiaSubmissionSchedule) . Nat)
      schedule
      (branch NvidiaSubmissionScheduleEmpty . 0)
      (branch
        NvidiaSubmissionScheduleOne
        batch
        .
        (eliminate
          NvidiaSubmissionBatch
          (lambda unrestricted current : (family NvidiaSubmissionBatch) . Nat)
          batch
          (branch
            NvidiaSubmissionBatchValue
            identity
            semaphore
            references
            .
            (cgReferenceCount references))))
      (branch
        NvidiaSubmissionScheduleAppend
        left
        right
        leftCount
        rightCount
        .
        (naturalAdd leftCount rightCount))
      (branch
        NvidiaSubmissionScheduleRepeat
        count
        referenceStride
        semaphoreStride
        body
        bodyCount
        .
        (naturalMultiply count bodyCount))
      (branch NvidiaSubmissionScheduleProfiled offset body bodyCount . bodyCount)))

def coppeliusSubmissionCount : Nat =
  (cgSubmissionCount coppeliusSubmissionSchedule)

def coppeliusSubmissionReferenceCount : Nat =
  (cgSubmissionReferences coppeliusSubmissionSchedule)

-- ---- the device side of the arena ----
-- Every tensor the plan places, named: the globals, live for the whole
-- program, and each phase's planes (fresh, or carried from an earlier
-- phase).  Runtime.DeviceArenaCertificate decides, on the launch and
-- submission schedules above, that every pointer a launch passes lies in a
-- plane of its phase or in a global, that the planes live together never
-- share a byte, and that no carried plane is overwritten before its phase
-- runs; the host is admitted only then (Coppelius.Build.NativeHost).
-- Staging planes are declared fresh: the host writes them between
-- submissions, in the order the host side of the certificate decides.
def cgPlane =
  (lambda unrestricted name : Bytes .
    (lambda unrestricted address : Nat .
      (lambda unrestricted bytes : Nat .
        (constructor ArenaResident ArenaResidentValue name address bytes cgParameterAlignment))))

def cgNoPlanes : (family ArenaResidents) =
  (constructor ArenaResidents ArenaResidentsEnd)

def cgAnd =
  (lambda unrestricted plane : (family ArenaResident) .
    (lambda unrestricted rest : (family ArenaResidents) .
      (constructor ArenaResidents ArenaResidentsNext plane rest)))

def cgResidentNorm1 : (family ArenaResident) =
  (cgPlane b"norm1" cgNorm1 cgPlaneBytes)

def cgResidentNorm1Statistics : (family ArenaResident) =
  (cgPlane b"norm1-statistics" cgNorm1Statistics cgStatisticsBytes)

def cgResidentNorm1Half : (family ArenaResident) =
  (cgPlane b"norm1-half" cgNorm1Half cgHalfPlaneBytes)

def cgResidentQkv : (family ArenaResident) =
  (cgPlane b"qkv" cgQkv cgQkvBytes)

def cgResidentQueries : (family ArenaResident) =
  (cgPlane b"queries" cgQueries cgHeadsHalfBytes)

def cgResidentKeys : (family ArenaResident) =
  (cgPlane b"keys" cgKeys cgHeadsHalfBytes)

def cgResidentValuesT : (family ArenaResident) =
  (cgPlane b"values-t" cgValuesT cgHeadsHalfBytes)

def cgResidentLogSumExp : (family ArenaResident) =
  (cgPlane b"log-sum-exp" cgLogSumExp cgLogSumExpBytes)

def cgResidentRowTerms : (family ArenaResident) =
  (cgPlane b"row-terms" cgRowTerms cgLogSumExpBytes)

def cgResidentMergedHalf : (family ArenaResident) =
  (cgPlane b"merged-half" cgMergedHalf cgHalfPlaneBytes)

def cgResidentProjection : (family ArenaResident) =
  (cgPlane b"projection" cgProjection cgPlaneBytes)

def cgResidentNorm2 : (family ArenaResident) =
  (cgPlane b"norm2" cgNorm2 cgPlaneBytes)

def cgResidentNorm2Statistics : (family ArenaResident) =
  (cgPlane b"norm2-statistics" cgNorm2Statistics cgStatisticsBytes)

def cgResidentNorm2Half : (family ArenaResident) =
  (cgPlane b"norm2-half" cgNorm2Half cgHalfPlaneBytes)

def cgResidentUp : (family ArenaResident) =
  (cgPlane b"up" cgUp cgFfnBytes)

def cgResidentGate : (family ArenaResident) =
  (cgPlane b"gate" cgGate cgFfnBytes)

def cgResidentGelu : (family ArenaResident) =
  (cgPlane b"gelu" cgGelu cgFfnBytes)

def cgResidentGated : (family ArenaResident) =
  (cgPlane b"gated" cgGated cgFfnBytes)

def cgResidentGatedHalf : (family ArenaResident) =
  (cgPlane b"gated-half" cgGatedHalf cgFfnHalfBytes)

def cgResidentDown : (family ArenaResident) =
  (cgPlane b"down" cgDown cgPlaneBytes)

def cgResidentFinalNorm : (family ArenaResident) =
  (cgPlane b"final-norm" cgFinalNorm cgPlaneBytes)

def cgResidentFinalStatistics : (family ArenaResident) =
  (cgPlane b"final-statistics" cgFinalStatistics cgStatisticsBytes)

def cgResidentFinalHalf : (family ArenaResident) =
  (cgPlane b"final-half" cgFinalHalf cgHalfPlaneBytes)

def cgResidentLogits : (family ArenaResident) =
  (cgPlane b"logits" cgLogits cgLogitsBytes)

def cgResidentLoss : (family ArenaResident) =
  (cgPlane b"loss" cgLoss cgRowsBytes)

def cgResidentCorrect : (family ArenaResident) =
  (cgPlane b"correct" cgCorrect cgRowsBytes)

def cgResidentLogitsGradient : (family ArenaResident) =
  (cgPlane b"logits-gradient" cgLogitsGradient cgLogitsBytes)

def cgResidentLogitsGradientHalf : (family ArenaResident) =
  (cgPlane b"logits-gradient-half" cgLogitsGradientHalf cgLogitsHalfBytes)

def cgResidentLogitsGradientT : (family ArenaResident) =
  (cgPlane b"logits-gradient-t" cgLogitsGradientT cgLogitsHalfBytes)

def cgResidentFinalHalfT : (family ArenaResident) =
  (cgPlane b"final-half-t" cgFinalHalfT cgHalfPlaneBytes)

def cgResidentFinalGradient : (family ArenaResident) =
  (cgPlane b"final-gradient" cgFinalGradient cgPlaneBytes)

def cgResidentEmbeddingHalfT : (family ArenaResident) =
  (cgPlane b"embedding-half-t" cgEmbeddingHalfT cgEmbeddingHalfTBytes)

def cgResidentDownGatedT : (family ArenaResident) =
  (cgPlane b"down-gated-t" cgDownGatedT cgFfnHalfBytes)

def cgResidentDownWeightT : (family ArenaResident) =
  (cgPlane b"down-weight-t" cgDownWeightT cgFfnWeightHalfBytes)

def cgResidentDownDyT : (family ArenaResident) =
  (cgPlane b"down-dy-t" cgDownDyT cgHalfPlaneBytes)

def cgResidentDownDyHalf : (family ArenaResident) =
  (cgPlane b"down-dy-half" cgDownDyHalf cgHalfPlaneBytes)

def cgResidentGatedGradient : (family ArenaResident) =
  (cgPlane b"gated-gradient" cgGatedGradient cgFfnBytes)

def cgResidentUpGradient : (family ArenaResident) =
  (cgPlane b"up-gradient" cgUpGradient cgFfnBytes)

def cgResidentGeluGradient : (family ArenaResident) =
  (cgPlane b"gelu-gradient" cgGeluGradient cgFfnBytes)

def cgResidentGateGradient : (family ArenaResident) =
  (cgPlane b"gate-gradient" cgGateGradient cgFfnBytes)

def cgResidentFfnInputT : (family ArenaResident) =
  (cgPlane b"ffn-input-t" cgFfnInputT cgHalfPlaneBytes)

def cgResidentFfnGradientHalf : (family ArenaResident) =
  (cgPlane b"ffn-gradient-half" cgFfnGradientHalf cgFfnHalfBytes)

def cgResidentFfnGradientT : (family ArenaResident) =
  (cgPlane b"ffn-gradient-t" cgFfnGradientT cgFfnHalfBytes)

def cgResidentNorm2UpGradient : (family ArenaResident) =
  (cgPlane b"norm2-up-gradient" cgNorm2UpGradient cgPlaneBytes)

def cgResidentFfnWeightT : (family ArenaResident) =
  (cgPlane b"ffn-weight-t" cgFfnWeightT cgFfnWeightHalfBytes)

def cgResidentNorm2Gradient : (family ArenaResident) =
  (cgPlane b"norm2-gradient" cgNorm2Gradient cgPlaneBytes)

def cgResidentNorm2InputGradient : (family ArenaResident) =
  (cgPlane b"norm2-input-gradient" cgNorm2InputGradient cgPlaneBytes)

def cgResidentOutputDyHalf : (family ArenaResident) =
  (cgPlane b"output-dy-half" cgOutputDyHalf cgHalfPlaneBytes)

def cgResidentOutputDyT : (family ArenaResident) =
  (cgPlane b"output-dy-t" cgOutputDyT cgHalfPlaneBytes)

def cgResidentOutputMergedT : (family ArenaResident) =
  (cgPlane b"output-merged-t" cgOutputMergedT cgHalfPlaneBytes)

def cgResidentOutputWeightT : (family ArenaResident) =
  (cgPlane b"output-weight-t" cgOutputWeightT cgOutputWeightHalfBytes)

def cgResidentMergedGradient : (family ArenaResident) =
  (cgPlane b"merged-gradient" cgMergedGradient cgPlaneBytes)

def cgResidentContextGradientHeads : (family ArenaResident) =
  (cgPlane b"context-gradient-heads" cgContextGradientHeads cgHeadsHalfBytes)

def cgResidentValuesHeads : (family ArenaResident) =
  (cgPlane b"values-heads" cgValuesHeads cgHeadsHalfBytes)

def cgResidentQueriesT : (family ArenaResident) =
  (cgPlane b"queries-t" cgQueriesT cgHeadsHalfBytes)

def cgResidentKeysT : (family ArenaResident) =
  (cgPlane b"keys-t" cgKeysT cgHeadsHalfBytes)

def cgResidentContextGradientT : (family ArenaResident) =
  (cgPlane b"context-gradient-t" cgContextGradientT cgHeadsHalfBytes)

def cgResidentQueriesGradient : (family ArenaResident) =
  (cgPlane b"queries-gradient" cgQueriesGradient cgContextBytes)

def cgResidentKeysGradient : (family ArenaResident) =
  (cgPlane b"keys-gradient" cgKeysGradient cgContextBytes)

def cgResidentValuesGradient : (family ArenaResident) =
  (cgPlane b"values-gradient" cgValuesGradient cgContextBytes)

def cgResidentQkvGradientHalf : (family ArenaResident) =
  (cgPlane b"qkv-gradient-half" cgQkvGradientHalf cgQkvHalfBytes)

def cgResidentQkvGradientT : (family ArenaResident) =
  (cgPlane b"qkv-gradient-t" cgQkvGradientT cgQkvHalfBytes)

def cgResidentQkvInputT : (family ArenaResident) =
  (cgPlane b"qkv-input-t" cgQkvInputT cgHalfPlaneBytes)

def cgResidentQkvWeightT : (family ArenaResident) =
  (cgPlane b"qkv-weight-t" cgQkvWeightT cgQkvWeightHalfBytes)

def cgResidentNorm1Gradient : (family ArenaResident) =
  (cgPlane b"norm1-gradient" cgNorm1Gradient cgPlaneBytes)

def cgResidentNorm1InputGradient : (family ArenaResident) =
  (cgPlane b"norm1-input-gradient" cgNorm1InputGradient cgPlaneBytes)

def cgResidentInputIds : (family ArenaResident) =
  (cgPlane b"input-ids" cgIds cgRowsBytes)

def cgResidentInputTargets : (family ArenaResident) =
  (cgPlane b"input-targets" cgTargets cgRowsBytes)

def cgResidentInputCosine : (family ArenaResident) =
  (cgPlane b"rotary-cosine" cgCosine cgRotaryTableBytes)

def cgResidentInputSine : (family ArenaResident) =
  (cgPlane b"rotary-sine" cgSine cgRotaryTableBytes)

def cgResidentResultDiagnostics : (family ArenaResident) =
  (cgPlane b"result-diagnostics" cgDiagnostic coppeliusResultLossBlockBytes)

def cgResidentResultLossRecord : (family ArenaResident) =
  (cgPlane
    b"result-loss-record"
    (cgStaging coppeliusResultLossRecordOffset)
    (naturalMultiply coppeliusUpdatePieces coppeliusResultWindowBytes))

def cgResidentResultPredictions : (family ArenaResident) =
  (cgPlane
    b"result-predictions"
    (cgStaging coppeliusResultPredictionsOffset)
    coppeliusResultPredictionsBytes)

def cgResidentResultRecord : (family ArenaResident) =
  (cgPlane
    b"result-record"
    cgInference
    (naturalAdd coppeliusResultRecordForwardBytes coppeliusResultRecordFinalBytes))

def cgResidentCheckpointWindow : (family ArenaResident) =
  (cgPlane b"checkpoint-window" (cgStaging zero) cgChunkBytes)

def cgResidentParameters : (family ArenaResident) =
  (cgPlane b"parameters" (cgArena cgBankP) cgParameterBytes)

def cgResidentGradients : (family ArenaResident) =
  (cgPlane b"gradients" (cgArena cgBankG) cgParameterBytes)

def cgResidentFirstMoments : (family ArenaResident) =
  (cgPlane b"first-moments" (cgArena cgBankM) cgParameterBytes)

def cgResidentSecondMoments : (family ArenaResident) =
  (cgPlane b"second-moments" (cgArena cgBankV) cgParameterBytes)

def cgResidentHalfCopies : (family ArenaResident) =
  (cgPlane b"half-copies" (cgArena cgHalfStart) (naturalSaturatingSubtract cgHalfEnd cgHalfStart))

def cgResidentBoundaries : (family ArenaResident) =
  (cgPlane
    b"boundaries"
    (cgArena cgBoundaryStart)
    (naturalSaturatingSubtract cgBoundaryEnd cgBoundaryStart))

def cgResidentResidualGradient : (family ArenaResident) =
  (cgPlane b"residual-gradient" (cgArena cgDy0) cgPlaneBytes)

def cgResidentAttentionResidualGradient : (family ArenaResident) =
  (cgPlane b"attention-residual-gradient" (cgArena cgDy1) cgPlaneBytes)

def cgResidentNormalizedInput : (family ArenaResident) =
  (cgPlane b"normalized-input" (cgArena cgXHat) cgPlaneBytes)

def cgResidentReport : (family ArenaResident) =
  (cgPlane b"report" (cgArena cgReport) (naturalAdd coppeliusResultLossBlockBytes cgRowBytes))

-- every layer's saved activations (its forward's tensors its backward reads)
def cgResidentSaved : (family ArenaResident) =
  (cgPlane b"saved-activations" (cgArena cgSavedStart) cgSavedBytes)

-- the address ranges whose words are pointers: the arena and the staging
-- window
def coppeliusDeviceWindows : (family ArenaResidents) =
  (cgAnd
    (cgPlane b"arena" (cgArena zero) coppeliusSM86CompatVideoBytesNatural)
    (cgAnd (cgPlane b"staging" (cgStaging zero) cgChunkBytes) cgNoPlanes))

def coppeliusDeviceGlobals : (family ArenaResidents) =
  (cgAnd
    cgResidentParameters
    (cgAnd
      cgResidentGradients
      (cgAnd
        cgResidentFirstMoments
        (cgAnd
          cgResidentSecondMoments
          (cgAnd
            cgResidentHalfCopies
            (cgAnd
              cgResidentBoundaries
              (cgAnd
                cgResidentResidualGradient
                (cgAnd
                  cgResidentAttentionResidualGradient
                  (cgAnd cgResidentNormalizedInput (cgAnd cgResidentReport (cgAnd cgResidentSaved cgNoPlanes)))))))))))

def coppeliusDevicePhases : (family DeviceArenaPhases) =
  (constructor
    DeviceArenaPhases
    DeviceArenaPhasesNext
    (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseInitialize cgNoPlanes cgNoPlanes)
    (constructor
      DeviceArenaPhases
      DeviceArenaPhasesNext
      (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseRecast cgNoPlanes cgNoPlanes)
      (constructor
        DeviceArenaPhases
        DeviceArenaPhasesNext
        (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseCheckpoint (cgAnd cgResidentCheckpointWindow cgNoPlanes) cgNoPlanes)
        (constructor
          DeviceArenaPhases
          DeviceArenaPhasesNext
          (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseReport (cgAnd cgResidentResultDiagnostics (cgAnd cgResidentResultRecord cgNoPlanes)) cgNoPlanes)
          (constructor
            DeviceArenaPhases
            DeviceArenaPhasesNext
            (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseGradientZero cgNoPlanes cgNoPlanes)
            (constructor
              DeviceArenaPhases
              DeviceArenaPhasesNext
              (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseEmbed (cgAnd cgResidentInputIds cgNoPlanes) cgNoPlanes)
              (constructor
                DeviceArenaPhases
                DeviceArenaPhasesNext
                (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseLayer (cgAnd cgResidentNorm1 (cgAnd cgResidentQkv (cgAnd cgResidentNorm2 (cgAnd cgResidentGated (cgAnd cgResidentDown (cgAnd cgResidentInputCosine (cgAnd cgResidentInputSine cgNoPlanes))))))) cgNoPlanes)
                (constructor
                  DeviceArenaPhases
                  DeviceArenaPhasesNext
                  (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseOutput (cgAnd cgResidentFinalNorm (cgAnd cgResidentFinalStatistics (cgAnd cgResidentFinalHalf (cgAnd cgResidentLogits (cgAnd cgResidentLoss (cgAnd cgResidentCorrect (cgAnd cgResidentLogitsGradient (cgAnd cgResidentResultRecord (cgAnd cgResidentResultPredictions (cgAnd cgResidentInputTargets (cgAnd cgResidentResultDiagnostics cgNoPlanes))))))))))) cgNoPlanes)
                  (constructor
                    DeviceArenaPhases
                    DeviceArenaPhasesNext
                    (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseOutputGradient (cgAnd cgResidentLogitsGradientHalf (cgAnd cgResidentLogitsGradientT (cgAnd cgResidentFinalHalfT (cgAnd cgResidentFinalGradient (cgAnd cgResidentEmbeddingHalfT cgNoPlanes))))) (cgAnd cgResidentLogitsGradient (cgAnd cgResidentFinalHalf (cgAnd cgResidentFinalStatistics cgNoPlanes))))
                    (constructor
                      DeviceArenaPhases
                      DeviceArenaPhasesNext
                      (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseFfnDown (cgAnd cgResidentDownGatedT (cgAnd cgResidentDownWeightT (cgAnd cgResidentDownDyT (cgAnd cgResidentDownDyHalf (cgAnd cgResidentGatedGradient cgNoPlanes))))) cgNoPlanes)
                      (constructor
                        DeviceArenaPhases
                        DeviceArenaPhasesNext
                        (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseFfnGate (cgAnd cgResidentUpGradient (cgAnd cgResidentGeluGradient (cgAnd cgResidentGateGradient (cgAnd cgResidentFfnInputT (cgAnd cgResidentFfnGradientHalf (cgAnd cgResidentFfnGradientT (cgAnd cgResidentNorm2UpGradient (cgAnd cgResidentFfnWeightT (cgAnd cgResidentNorm2Gradient (cgAnd cgResidentNorm2InputGradient cgNoPlanes)))))))))) (cgAnd cgResidentGatedGradient cgNoPlanes))
                        (constructor
                          DeviceArenaPhases
                          DeviceArenaPhasesNext
                          (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseAttentionOutput (cgAnd cgResidentOutputDyHalf (cgAnd cgResidentOutputDyT (cgAnd cgResidentOutputMergedT (cgAnd cgResidentOutputWeightT (cgAnd cgResidentMergedGradient cgNoPlanes))))) cgNoPlanes)
                          (constructor
                            DeviceArenaPhases
                            DeviceArenaPhasesNext
                            (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseHeadLayout (cgAnd cgResidentContextGradientHeads (cgAnd cgResidentValuesHeads (cgAnd cgResidentQueriesT (cgAnd cgResidentKeysT (cgAnd cgResidentContextGradientT cgNoPlanes))))) (cgAnd cgResidentMergedGradient cgNoPlanes))
                            (constructor
                              DeviceArenaPhases
                              DeviceArenaPhasesNext
                              (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseAttentionHead (cgAnd cgResidentRowTerms (cgAnd cgResidentQueriesGradient (cgAnd cgResidentKeysGradient (cgAnd cgResidentValuesGradient cgNoPlanes)))) (cgAnd cgResidentContextGradientHeads (cgAnd cgResidentValuesHeads (cgAnd cgResidentQueriesT (cgAnd cgResidentKeysT (cgAnd cgResidentContextGradientT cgNoPlanes))))))
                              (constructor
                                DeviceArenaPhases
                                DeviceArenaPhasesNext
                                (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseQkv (cgAnd cgResidentQkvGradientHalf (cgAnd cgResidentQkvGradientT (cgAnd cgResidentQkvInputT (cgAnd cgResidentQkvWeightT (cgAnd cgResidentNorm1Gradient (cgAnd cgResidentNorm1InputGradient (cgAnd cgResidentInputCosine (cgAnd cgResidentInputSine cgNoPlanes)))))))) (cgAnd cgResidentQueriesGradient (cgAnd cgResidentKeysGradient (cgAnd cgResidentValuesGradient cgNoPlanes))))
                                (constructor
                                  DeviceArenaPhases
                                  DeviceArenaPhasesNext
                                  (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseScatter (cgAnd cgResidentInputIds cgNoPlanes) cgNoPlanes)
                                  (constructor
                                    DeviceArenaPhases
                                    DeviceArenaPhasesNext
                                    (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseAdamW cgNoPlanes cgNoPlanes)
                                    (constructor
                                      DeviceArenaPhases
                                      DeviceArenaPhasesNext
                                      (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseDiagnostics (cgAnd cgResidentResultDiagnostics cgNoPlanes) cgNoPlanes)
                                      (constructor
                                        DeviceArenaPhases
                                        DeviceArenaPhasesNext
                                        (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseLossRecord (cgAnd cgResidentResultDiagnostics (cgAnd cgResidentResultLossRecord cgNoPlanes)) cgNoPlanes)
                                        (constructor DeviceArenaPhases DeviceArenaPhasesNext (constructor DeviceArenaPhase DeviceArenaPhaseValue cgPhaseRotary (cgAnd cgResidentInputCosine (cgAnd cgResidentInputSine cgNoPlanes)) cgNoPlanes) (constructor DeviceArenaPhases DeviceArenaPhasesEnd)))))))))))))))))))))

-- 0 when the device side is certified (Runtime.DeviceArenaCertificate); a
-- function, so a build computes it only where it is applied (an admission
-- that has already failed does not)

def coppeliusDeviceHazardOf =
  (lambda unrestricted ignored : Nat .
    (deviceArenaHazard
      coppeliusDeviceWindows
      coppeliusDeviceGlobals
      coppeliusDevicePhases
      coppeliusLaunchSchedule
      coppeliusSubmissionSchedule))

def coppeliusDeviceHazard : Nat =
  (coppeliusDeviceHazardOf zero)

-- the same certificate over the run's order (coppeliusRunSubmissionSample)
def coppeliusRunDeviceHazardOf =
  (lambda unrestricted ignored : Nat .
    (deviceArenaHazard
      coppeliusDeviceWindows
      coppeliusDeviceGlobals
      coppeliusDevicePhases
      coppeliusLaunchSchedule
      coppeliusRunSubmissionSample))

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.