Source/Packages

Accelerator.SM121.Lowering

packages/hardware/architectures/nvidia-sm121/src/Accelerator/SM121/Lowering.alpha

2,621 lines365 declarations134.1 KiBSHA-256 b7b3bbc05e9c

Complete file · line 1748

Lowering.alpha

Definition view

Large source region · 2,621 lines

module Accelerator.SM121.Lowering

import Accelerator.SM86.Instruction
import Accelerator.SM86.InstructionEncoding
import Accelerator.SM86.Types
import Hardware.Nvidia.SM86.Command.WholeProgramPlan
import Std.Natural

-- The sm_121 realization of an sm_86 program: what a Blackwell GB10 runs for
-- it.  Each instruction keeps its SM86 word (Accelerator.SM86's encoder)
-- wherever the two architectures mean the same thing; four rules rewrite the
-- rest, and a scheduler repairs the program's timing for sm_121's fixed
-- latencies:
--
--   * sm_121 has no constant-bank ALU operand: MOV Rd, c[b][o] becomes
--     LDC Rd, c[b][o], and IMAD / IMAD.WIDE with a constant operand become an
--     LDC (LDC.64) into a register followed by the register form.  A constant
--     whose destination is also a source is loaded into a register (pair)
--     the program never names below its declared count, and the lowering
--     refuses when there is none.  The LDC is variable-latency and signals
--     the program's free scoreboard.
--   * sm_86's RED is sm_121's REDG, with the descriptor and modifier bits
--     ptxas sets.
--   * Shared memory is based at 0x400 (the first kilobyte is reserved), so
--     every LDS, STS and LDSM offset moves up by it.
--   * The global-memory descriptor UR4:UR5 is zeroed in a two-instruction
--     prologue before the first access.
--
-- The scheduler walks the lowered instructions in order with the cycle each
-- is issued at.  A register written by a fixed-latency instruction may be
-- read (or written again) only after the read-after-write latency of the
-- writer's class to the reader's (NVIDIA's sm100 table plus Mesa NAK's
-- sm_120 padding); a guard predicate after the predicate latency.  The
-- instruction before is made to stall longer, up to fifteen cycles, and NOPs
-- carry the rest.  Variable-latency results keep the program's own
-- scoreboards, except the LDCs above and a variable-latency instruction's
-- source registers that the SM86 program does not guard with a read
-- barrier: both go on the free scoreboard, and the next instruction touching
-- a loaded register or overwriting such a source waits on it.  (Unguarded,
-- the block reductions' identity store -- STS of -inf, then the register
-- reused for the warp index -- stored 0 on sm_121 now and then: a causal
-- softmax row whose one score was far below zero summed to zero and came out
-- NaN.)  A program with a branch is refused: its schedule would need the
-- latencies across the branch.

family SM121LowerClass : Type 0
constructor SM121LowerAlu
constructor SM121LowerDualAlu
constructor SM121LowerFma
constructor SM121LowerWide
constructor SM121LowerHalfToFloat
constructor SM121LowerTensor
constructor SM121LowerVariable
constructor SM121LowerMemory
constructor SM121LowerNoResult
end-family

family SM121LowerRegisters : Type 0
constructor SM121LowerRegistersEnd
constructor SM121LowerRegistersNext
field unrestricted sm121LowerRegistersHead : Nat
recursive unrestricted sm121LowerRegistersTail
end-family

-- The registers a program names, as a mask of four 64-bit words.
family SM121LowerMask : Type 0
constructor SM121LowerMaskValue
field unrestricted sm121LowerMask0 : Nat
field unrestricted sm121LowerMask1 : Nat
field unrestricted sm121LowerMask2 : Nat
field unrestricted sm121LowerMask3 : Nat
end-family

family SM121LowerOp : Type 0
constructor SM121LowerOpValue
field unrestricted sm121LowerOpLow : Nat
field unrestricted sm121LowerOpHigh : Nat
field unrestricted sm121LowerOpClass : (family SM121LowerClass)
field unrestricted sm121LowerOpReads : (family SM121LowerRegisters)
field unrestricted sm121LowerOpWrites : (family SM121LowerRegisters)
field unrestricted sm121LowerOpPredicateReads : (family SM121LowerRegisters)
field unrestricted sm121LowerOpPredicateWrites : (family SM121LowerRegisters)
field unrestricted sm121LowerOpConstantLoad : Nat
-- the SM86 instruction's index + 1 on the first op it lowers to (0 on the
-- others and on inserted ops): where a branch to it lands
field unrestricted sm121LowerOpLabel : Nat
-- 1 on a branch and on the first op of a branch's target: a join, where the
-- schedule is drained (every result ready, every scoreboard a register rides
-- on waited) so that every path arrives with no register in flight
field unrestricted sm121LowerOpJoin : Nat
-- a branch's target's SM86 index + 1 (0 on every other op)
field unrestricted sm121LowerOpTarget : Nat
end-family

-- An instruction's two 64-bit words.
family SM121LowerWords : Type 0
constructor SM121LowerWordsValue
field unrestricted sm121LowerWordsLow : Nat
field unrestricted sm121LowerWordsHigh : Nat
end-family

-- Where each labelled op was placed: its label and its position in the
-- lowered program.
family SM121LowerLabels : Type 0
constructor SM121LowerLabelsEnd
constructor SM121LowerLabelsNext
field unrestricted sm121LowerLabelsLabel : Nat
field unrestricted sm121LowerLabelsPosition : Nat
recursive unrestricted sm121LowerLabelsTail
end-family

family SM121LowerOps : Type 0
constructor SM121LowerOpsEnd
constructor SM121LowerOpsNext
field unrestricted sm121LowerOpsHead : (family SM121LowerOp)
recursive unrestricted sm121LowerOpsTail
end-family

family SM121LowerCounted : Type 0
constructor SM121LowerCountedValue
field unrestricted sm121LowerCountedCount : Nat
field unrestricted sm121LowerCountedLabels : (family SM121LowerLabels)
end-family

-- A placed program with its branches relocated (and its length), or the
-- refusal of a branch whose target was not placed.
family SM121LowerRelocation : Type 0
constructor SM121LowerRelocated
field unrestricted sm121LowerRelocatedCount : Nat
field unrestricted sm121LowerRelocatedOps : (family SM121LowerOps)
constructor SM121LowerRelocationRefused
end-family


-- What the lowering refuses, by name.
family SM121LowerRefusal : Type 0
-- Accelerator.SM86's encoder refused an instruction
constructor SM121LowerRefusedEncoding
constructor SM121LowerRefusedBranchTarget
constructor SM121LowerRefusedRegisterBeyondCount
constructor SM121LowerRefusedNoFreeRegister
constructor SM121LowerRefusedNoFreePair
constructor SM121LowerRefusedSharedOffset
constructor SM121LowerRefusedNoStall
-- the static scoreboard check found a hazard in the schedule
constructor SM121LowerRefusedHazard
end-family

-- How a program's plain LDS and STS address shared memory.  Accelerator.SM86's
-- encoder writes them as [Ra.X4+off] (nvdisasm reads width byte 0x48 so on
-- sm_86 and sm_121 alike): the address register counts 32-bit words.  A
-- program whose addresses are bytes -- the typed block models' convention --
-- is lowered with the scale bit cleared.  LDSM has its own address form and
-- is never rescaled.
family SM121SharedAddressing : Type 0
constructor SM121SharedAddressScaled
constructor SM121SharedAddressBytes
end-family

-- How an SM86 instruction is lowered.
family SM121LowerRule : Type 0
-- its SM86 word
constructor SM121LowerRuleSame
constructor SM121LowerRuleConstantMove
constructor SM121LowerRuleConstantMultiplyAdd
constructor SM121LowerRuleConstantMultiplyAddWide
constructor SM121LowerRuleReduction
constructor SM121LowerRuleShared
constructor SM121LowerRuleSharedMatrix
-- LDGSTS: its shared offset rebased, its descriptor bits sm_121's
constructor SM121LowerRuleAsyncCopy
constructor SM121LowerRuleBranch
end-family

-- An SM86 instruction: its rule, latency class and the registers and
-- predicates it reads and writes (its guard aside).
family SM121LowerShape : Type 0
constructor SM121LowerShapeValue
field unrestricted sm121LowerShapeRule : (family SM121LowerRule)
field unrestricted sm121LowerShapeClass : (family SM121LowerClass)
field unrestricted sm121LowerShapeReads : (family SM121LowerRegisters)
field unrestricted sm121LowerShapeWrites : (family SM121LowerRegisters)
field unrestricted sm121LowerShapePredicateWrites : (family SM121LowerRegisters)
end-family

family SM121LowerDecoded : Type 0
constructor SM121LowerDecodedValue
field unrestricted sm121LowerDecodedLow : Nat
field unrestricted sm121LowerDecodedHigh : Nat
field unrestricted sm121LowerDecodedGuardReads : (family SM121LowerRegisters)
field unrestricted sm121LowerDecodedShape : (family SM121LowerShape)
end-family

family SM121LowerDecodedList : Type 0
constructor SM121LowerDecodedEnd
constructor SM121LowerDecodedNext
field unrestricted sm121LowerDecodedHead : (family SM121LowerDecoded)
recursive unrestricted sm121LowerDecodedTail
constructor SM121LowerDecodedRefused
field unrestricted sm121LowerDecodedRefusal : (family SM121LowerRefusal)
end-family

family SM121LowerContext : Type 0
constructor SM121LowerContextValue
-- the scoreboard the program names nowhere (5 when it names all six)
field unrestricted sm121LowerContextBarrier : Nat
-- the least register the program never names, below its count (the count
-- when there is none)
field unrestricted sm121LowerContextFree : Nat
-- the least even register that and the next are never named
field unrestricted sm121LowerContextFreePair : Nat
field unrestricted sm121LowerContextCount : Nat
field unrestricted sm121LowerContextAddressing : (family SM121SharedAddressing)
-- the wait mask naming every scoreboard a register rides on (one an
-- instruction reading or writing registers names) and the free one: a join
-- waits on all of them.  A scoreboard only register-free instructions name
-- -- LDGDEPBAR's, cp.async's count of groups in flight -- holds no register
-- the schedule tracks, and a join leaves it running: the copies stay in
-- flight across a loop's back edge, as the program's own waits (DEPBAR)
-- intend
field unrestricted sm121LowerContextWaitAll : Nat
end-family

family SM121LowerExpansion : Type 0
constructor SM121LowerExpanded
field unrestricted sm121LowerExpandedOps : (family SM121LowerOps)
constructor SM121LowerExpansionRefused
field unrestricted sm121LowerExpansionRefusal : (family SM121LowerRefusal)
end-family

-- Results still in flight: a register (or predicate), the cycle its writer
-- issued at and the writer's class.  An entry is dropped once every reader
-- would find it ready.
family SM121LowerPending : Type 0
constructor SM121LowerPendingEnd
constructor SM121LowerPendingNext
field unrestricted sm121LowerPendingRegister : Nat
field unrestricted sm121LowerPendingTime : Nat
field unrestricted sm121LowerPendingClass : (family SM121LowerClass)
recursive unrestricted sm121LowerPendingTail
end-family

family SM121LowerSchedule : Type 0
constructor SM121LowerScheduleValue
-- the cycle the next instruction may issue at
field unrestricted sm121LowerScheduleTime : Nat
field unrestricted sm121LowerScheduleRegisters : (family SM121LowerPending)
field unrestricted sm121LowerSchedulePredicates : (family SM121LowerPending)
-- registers an LDC on the free scoreboard is still writing
field unrestricted sm121LowerScheduleLoads : (family SM121LowerRegisters)
-- registers an unguarded variable-latency instruction may still be reading
field unrestricted sm121LowerScheduleGuardedReads : (family SM121LowerRegisters)
-- the instructions placed so far, the last first
field unrestricted sm121LowerSchedulePlaced : (family SM121LowerOps)
constructor SM121LowerScheduleRefused
field unrestricted sm121LowerScheduleRefusal : (family SM121LowerRefusal)
end-family

family SM121LowerResult : Type 0
constructor SM121LowerLowered
field unrestricted sm121LowerLoweredBytes : Bytes
constructor SM121LowerRefused
field unrestricted sm121LowerRefusedBecause : (family SM121LowerRefusal)
end-family

-- Registers on a scoreboard: a variable-latency result not yet waited for,
-- or a source an instruction may still be reading (the barrier it signals,
-- or none: sm121LowerUnguarded).
family SM121LowerBoards : Type 0
constructor SM121LowerBoardsEnd
constructor SM121LowerBoardsNext
field unrestricted sm121LowerBoardsRegister : Nat
field unrestricted sm121LowerBoardsBarrier : Nat
recursive unrestricted sm121LowerBoardsTail
end-family

-- The static scoreboard check's walk of a lowered program.
family SM121LowerTrace : Type 0
constructor SM121LowerTraceValue
-- the cycle the instruction issues at, from the stalls placed before it
field unrestricted sm121LowerTraceTime : Nat
field unrestricted sm121LowerTraceFixed : (family SM121LowerPending)
field unrestricted sm121LowerTracePredicates : (family SM121LowerPending)
field unrestricted sm121LowerTraceWritten : (family SM121LowerBoards)
field unrestricted sm121LowerTraceReading : (family SM121LowerBoards)
field unrestricted sm121LowerTraceHazards : Nat
-- the position of the op being checked
field unrestricted sm121LowerTracePosition : Nat
end-family

-- The joins of a lowered program read off its words: each branch and each
-- position one lands on (`bad` counts the displacements that land outside
-- the program or between instructions).
family SM121LowerJoins : Type 0
constructor SM121LowerJoinsValue
field unrestricted sm121LowerJoinsCount : Nat
field unrestricted sm121LowerJoinsPositions : (family SM121LowerRegisters)
field unrestricted sm121LowerJoinsBad : Nat
end-family

-- ---------------------------------------------------------------------------
-- The instruction word's fields

-- A field of a 64-bit word, given the place value of its lowest bit and the
-- number of values it holds.
def sm121LowerField =
  (lambda unrestricted word : Nat .
    (lambda unrestricted place : Nat .
      (lambda unrestricted span : Nat .
        (nat-modulo (nat-divide word place) span))))

def sm121LowerWithField =
  (lambda unrestricted word : Nat .
    (lambda unrestricted place : Nat .
      (lambda unrestricted span : Nat .
        (lambda unrestricted value : Nat .
          (nat-add
            (nat-subtract word (nat-multiply (sm121LowerField word place span) place))
            (nat-multiply (nat-modulo value span) place))))))

def sm121LowerPlace =
  (lambda unrestricted bit : Nat . (naturalPowerOfTwo bit))

-- the low word
def sm121LowerOpcodeSpan : Nat = (compile-time (sm121LowerPlace 12))
def sm121LowerGuardPlace : Nat = sm121LowerOpcodeSpan
def sm121LowerGuardSpan : Nat = (compile-time (sm121LowerPlace 4))
def sm121LowerDestinationPlace : Nat = (compile-time (sm121LowerPlace 16))
def sm121LowerSourcePlace : Nat = (compile-time (sm121LowerPlace 24))
def sm121LowerRegisterSpan : Nat = (compile-time (sm121LowerPlace 8))
def sm121LowerHalfWordSpan : Nat = (compile-time (sm121LowerPlace 32))
def sm121LowerSecondSourcePlace : Nat = sm121LowerHalfWordSpan
def sm121LowerConstantOffsetPlace : Nat = (compile-time (sm121LowerPlace 38))
def sm121LowerConstantOffsetSpan : Nat = (compile-time (sm121LowerPlace 16))
def sm121LowerConstantBankPlace : Nat = (compile-time (sm121LowerPlace 54))
def sm121LowerConstantBankSpan : Nat = (compile-time (sm121LowerPlace 5))
def sm121LowerSharedOffsetPlace : Nat = (compile-time (sm121LowerPlace 40))
def sm121LowerSharedOffsetSpan : Nat = (compile-time (sm121LowerPlace 24))
-- the shared offset is signed: a non-negative one stays below 2^23
def sm121LowerSharedOffsetLimit : Nat = (compile-time (sm121LowerPlace 23))
-- LDS/STS's address scale (.X4), bit 78: nvdisasm, width byte 0x48 against 0x08
def sm121LowerSharedScalePlace : Nat = (compile-time (sm121LowerPlace 14))

-- the high word
def sm121LowerThirdSourcePlace : Nat = 1
def sm121LowerConstantWidthPlace : Nat = (compile-time (sm121LowerPlace 9))
def sm121LowerStallPlace : Nat = (compile-time (sm121LowerPlace 41))
def sm121LowerStallSpan : Nat = (compile-time (sm121LowerPlace 4))
def sm121LowerYieldPlace : Nat = (compile-time (sm121LowerPlace 45))
def sm121LowerYieldSpan : Nat = 2
def sm121LowerWriteBarrierPlace : Nat = (compile-time (sm121LowerPlace 46))
def sm121LowerBarrierSpan : Nat = (compile-time (sm121LowerPlace 3))
def sm121LowerReadBarrierPlace : Nat = (compile-time (sm121LowerPlace 49))
def sm121LowerWaitBit : Nat = 52
def sm121LowerWaitPlace : Nat = (compile-time (sm121LowerPlace sm121LowerWaitBit))
def sm121LowerWaitSpan : Nat = (compile-time (sm121LowerPlace 6))
def sm121LowerReusePlace : Nat = (compile-time (sm121LowerPlace 58))
def sm121LowerReuseSpan : Nat = (compile-time (sm121LowerPlace 4))

-- the longest stall one instruction carries
def sm121LowerStallLongest : Nat = (nat-subtract sm121LowerStallSpan 1)

-- the scoreboard field's value for none
def sm121LowerNoBarrier : Nat = (nat-subtract sm121LowerBarrierSpan 1)

-- scoreboards a program may name (the sixth, 5, is the last resort for the
-- free one)
def sm121LowerBarrierCount : Nat = 6

-- the guard field of an instruction that always executes (PT, not negated)
def sm121LowerAlways : Nat = 7

-- RZ as a register operand
def sm121LowerZeroRegister : Nat = 255

-- ---------------------------------------------------------------------------
-- sm_121's words, as ptxas -arch=sm_121 emits them on the DGX Spark

def sm121LowerOpcodeConstantLoad : Nat = 0xb82
def sm121LowerOpcodeMultiplyAddRegister : Nat = 0x224
def sm121LowerOpcodeMultiplyAddWideRegister : Nat = 0x225
def sm121LowerOpcodeReduceGlobal : Nat = 0x9a6
-- REDG.E.ADD.F32.FTZ.RN.STRONG.GPU's modifier and descriptor bits (64..95)
def sm121LowerReduceGlobalModifiers : Nat = 0x0c12f304
def sm121LowerOpcodeNoOperation : Nat = 0x918
def sm121LowerOpcodeUniformMove : Nat = 0x882
-- LDC's width field: 32 and 64 bits; the registers a load writes are the
-- width's value less this
def sm121LowerConstantWord : Nat = 4
def sm121LowerConstantPair : Nat = 5
def sm121LowerConstantWidthBase : Nat = 3
-- the global-memory descriptor's uniform registers
def sm121LowerDescriptorLow : Nat = 4
def sm121LowerDescriptorHigh : Nat = 5
-- shared memory's base on sm_121
def sm121LowerSharedBase : Nat = 0x400

-- LDGSTS (cp.async.cg ... 16): the shared offset at 44 (20 bits); the high
-- word as ptxas -arch=sm_121a writes it against -arch=sm_86 (nvdisasm -hex
-- on the DGX Spark, 2026-09-26, research p4-hmma/cp_async*.cu): bits 70 and
-- 74 clear, 83 set -- desc[UR4] in place of sm_86's implicit descriptor
-- (0x...0b901c44 becomes 0x...0b981804; the byte offsets and registers are
-- the same)
def sm121LowerAsyncCopyOffsetPlace : Nat = (compile-time (sm121LowerPlace 44))
def sm121LowerAsyncCopyOffsetSpan : Nat = (compile-time (sm121LowerPlace 20))
def sm121LowerAsyncCopyOffsetLimit : Nat = (compile-time (sm121LowerPlace 20))
def sm121LowerAsyncCopyDescriptorEnable : Nat = (compile-time (sm121LowerPlace 6))
def sm121LowerAsyncCopyLegacyBit : Nat = (compile-time (sm121LowerPlace 10))
def sm121LowerAsyncCopyDescriptorBit : Nat = (compile-time (sm121LowerPlace 19))

-- The control of an instruction with no scoreboards: its stall and yield.
def sm121LowerControl =
  (lambda unrestricted stall : Nat .
    (lambda unrestricted yield : Nat .
      (nat-add
        (nat-add (nat-multiply stall sm121LowerStallPlace) (nat-multiply yield sm121LowerYieldPlace))
        (nat-add
          (nat-multiply sm121LowerNoBarrier sm121LowerWriteBarrierPlace)
          (nat-multiply sm121LowerNoBarrier sm121LowerReadBarrierPlace)))))

def sm121LowerAlwaysOpcode =
  (lambda unrestricted opcode : Nat .
    (nat-add opcode (nat-multiply sm121LowerAlways sm121LowerGuardPlace)))

def sm121LowerStall =
  (lambda unrestricted high : Nat .
    (sm121LowerField high sm121LowerStallPlace sm121LowerStallSpan))

def sm121LowerWithStall =
  (lambda unrestricted stall : Nat .
    (lambda unrestricted high : Nat .
      (sm121LowerWithField high sm121LowerStallPlace sm121LowerStallSpan stall)))

-- the operand-reuse flags describe the NEXT instruction; an instruction
-- placed after this one invalidates them
def sm121LowerWithoutReuse =
  (lambda unrestricted high : Nat .
    (sm121LowerWithField high sm121LowerReusePlace sm121LowerReuseSpan zero))

def sm121LowerWithWait =
  (lambda unrestricted waitPlace : Nat .
    (lambda unrestricted high : Nat .
      (sm121LowerWithField high waitPlace 2 1)))

def sm121LowerWithReadBarrier =
  (lambda unrestricted barrier : Nat .
    (lambda unrestricted high : Nat .
      (sm121LowerWithField high sm121LowerReadBarrierPlace sm121LowerBarrierSpan barrier)))

def sm121LowerSelect =
  (lambda erased value : Type 0 .
    (lambda unrestricted condition : Nat .
      (lambda unrestricted whenTrue : value .
        (lambda unrestricted whenFalse : value .
          (nat-eliminate
            (lambda unrestricted current : Nat . value)
            whenFalse
            (lambda unrestricted predecessor : Nat .
              (lambda unrestricted induction : value . whenTrue))
            condition)))))

-- A select whose arms may be expensive redexes. The machine evaluates a
-- function's arguments before the call, so a table call (or a union, a
-- drain, a rebuilt prefix) sitting in a plain select's arm runs for every
-- cell whether the arm is taken or not (measured: the read-after-write
-- table costs ~450 dispatch steps, paid per pending entry per operand
-- register -- two thirds of the schedule). Both arms thunked behind
-- lambdas are already values when passed; the taken one alone runs when
-- the result is applied. Plain `sm121LowerSelect` stays for cheap arms,
-- where the two closures would cost more than they save.
def sm121LowerSelectLazy =
  (lambda erased value : Type 0 .
    (lambda unrestricted condition : Nat .
      (lambda unrestricted whenTrue : (pi unrestricted u : Nat . value) .
        (lambda unrestricted whenFalse : (pi unrestricted u : Nat . value) .
          (app
            (nat-eliminate
              (lambda unrestricted current : Nat . (pi unrestricted u : Nat . value))
              whenFalse
              (lambda unrestricted predecessor : Nat .
                (lambda unrestricted induction : (pi unrestricted u : Nat . value) . whenTrue))
              condition)
            zero)))))

def sm121LowerMaximum =
  (lambda unrestricted left : Nat .
    (lambda unrestricted right : Nat .
      (nat-eliminate (lambda unrestricted current : Nat . Nat)
        left
        (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . right))
        (nat-less-than left right))))

def sm121LowerMinimum =
  (lambda unrestricted left : Nat .
    (lambda unrestricted right : Nat .
      (naturalSelect (nat-less-than right left) right left)))

-- ---------------------------------------------------------------------------
-- Latency classes (NAK's sm_120 categories)

-- Fixed-latency classes are timed by stalls; the others by scoreboards.
def sm121LowerFixed =
  (lambda unrestricted class : (family SM121LowerClass) .
    (eliminate
      SM121LowerClass
      (lambda unrestricted current : (family SM121LowerClass) . Nat)
      class
      (branch SM121LowerAlu . 1)
      (branch SM121LowerDualAlu . 1)
      (branch SM121LowerFma . 1)
      (branch SM121LowerWide . 1)
      (branch SM121LowerHalfToFloat . 1)
      (branch SM121LowerTensor . 1)
      (branch SM121LowerVariable . zero)
      (branch SM121LowerMemory . zero)
      (branch SM121LowerNoResult . zero)))

def sm121LowerIsTensor =
  (lambda unrestricted class : (family SM121LowerClass) .
    (eliminate
      SM121LowerClass
      (lambda unrestricted current : (family SM121LowerClass) . Nat)
      class
      (branch SM121LowerAlu . zero)
      (branch SM121LowerDualAlu . zero)
      (branch SM121LowerFma . zero)
      (branch SM121LowerWide . zero)
      (branch SM121LowerHalfToFloat . zero)
      (branch SM121LowerTensor . 1)
      (branch SM121LowerVariable . zero)
      (branch SM121LowerMemory . zero)
      (branch SM121LowerNoResult . zero)))

def sm121LowerHasResult =
  (lambda unrestricted class : (family SM121LowerClass) .
    (eliminate
      SM121LowerClass
      (lambda unrestricted current : (family SM121LowerClass) . Nat)
      class
      (branch SM121LowerAlu . 1)
      (branch SM121LowerDualAlu . 1)
      (branch SM121LowerFma . 1)
      (branch SM121LowerWide . 1)
      (branch SM121LowerHalfToFloat . 1)
      (branch SM121LowerTensor . 1)
      (branch SM121LowerVariable . 1)
      (branch SM121LowerMemory . 1)
      (branch SM121LowerNoResult . zero)))

-- The writer's column of NVIDIA's sm100 read-after-write table.
def sm121LowerWriterColumn =
  (lambda unrestricted class : (family SM121LowerClass) .
    (eliminate
      SM121LowerClass
      (lambda unrestricted current : (family SM121LowerClass) . Nat)
      class
      (branch SM121LowerAlu . zero)
      (branch SM121LowerDualAlu . 1)
      (branch SM121LowerFma . 2)
      (branch SM121LowerWide . 3)
      (branch SM121LowerHalfToFloat . 4)
      (branch SM121LowerTensor . 5)
      (branch SM121LowerVariable . 5)
      (branch SM121LowerMemory . 5)
      (branch SM121LowerNoResult . 5)))

def sm121LowerRow =
  (lambda unrestricted alu : Nat .
    (lambda unrestricted dualAlu : Nat .
      (lambda unrestricted fma : Nat .
        (lambda unrestricted wide : Nat .
          (lambda unrestricted halfToFloat : Nat .
            (lambda unrestricted other : Nat .
              (lambda unrestricted column : Nat .
                (naturalSelect (naturalEqual column zero) alu
                  (naturalSelect (naturalEqual column 1) dualAlu
                    (naturalSelect (naturalEqual column 2) fma
                      (naturalSelect (naturalEqual column 3) wide
                        (naturalSelect (naturalEqual column 4) halfToFloat other))))))))))))

-- The reader's row.
def sm121LowerReaderRow =
  (lambda unrestricted class : (family SM121LowerClass) .
    (lambda unrestricted column : Nat .
      (eliminate
        SM121LowerClass
        (lambda unrestricted current : (family SM121LowerClass) . Nat)
        class
        (branch SM121LowerAlu . (sm121LowerRow 4 4 5 5 5 19 column))
        (branch SM121LowerDualAlu . (sm121LowerRow 4 4 5 5 5 19 column))
        (branch SM121LowerFma . (sm121LowerRow 5 5 4 4 5 19 column))
        (branch SM121LowerWide . (sm121LowerRow 5 5 4 6 5 19 column))
        (branch SM121LowerHalfToFloat . (sm121LowerRow 5 5 5 5 5 19 column))
        (branch SM121LowerTensor . (sm121LowerRow 7 7 7 7 7 20 column))
        (branch SM121LowerVariable . (sm121LowerRow 4 4 4 4 4 19 column))
        (branch SM121LowerMemory . (sm121LowerRow 5 5 5 5 5 19 column))
        (branch SM121LowerNoResult . (sm121LowerRow 4 4 5 5 5 19 column)))))

-- NAK's padding: the tensor pipe's measured latency is longer than the
-- table's.
def sm121LowerPadding : Nat = 1
def sm121LowerTensorPadding : Nat = 9

-- Cycles from a fixed-latency writer's issue to a reader's.
def sm121LowerReadAfterWrite =
  (lambda unrestricted writer : (family SM121LowerClass) .
    (lambda unrestricted reader : (family SM121LowerClass) .
      (nat-add
        (sm121LowerReaderRow reader (sm121LowerWriterColumn writer))
        (naturalSelect
          (naturalOr (sm121LowerIsTensor writer) (sm121LowerIsTensor reader))
          sm121LowerTensorPadding
          sm121LowerPadding))))

-- The largest of `measure` over every class.
def sm121LowerOverClasses =
  (lambda unrestricted measure : (pi unrestricted class : (family SM121LowerClass) . Nat) .
      (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerAlu))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerDualAlu))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerFma))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerWide))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerHalfToFloat))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerTensor))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerVariable))
        (sm121LowerMaximum (measure (constructor SM121LowerClass SM121LowerMemory))
        (measure (constructor SM121LowerClass SM121LowerNoResult)))))))))))

-- The longest of them: a result pending longer ago than this is ready for
-- every reader, so the schedule forgets it (Proof.SM121LoweringContract:
-- it is the table's largest entry).
def sm121LowerLongestReadAfterWrite : Nat =
  (compile-time (sm121LowerReadAfterWrite
    (constructor SM121LowerClass SM121LowerTensor) (constructor SM121LowerClass SM121LowerTensor)))

def sm121LowerTableLongest : Nat =
  (sm121LowerOverClasses
    (lambda unrestricted writer : (family SM121LowerClass) .
      (sm121LowerOverClasses
        (lambda unrestricted reader : (family SM121LowerClass) .
          (sm121LowerReadAfterWrite writer reader)))))

-- ---------------------------------------------------------------------------
-- The latency rows the scheduler reads (and nothing else).
--
-- The machine evaluates a function's arguments before the call, so a table
-- call sitting in a per-cell select arm runs for every cell whether the arm
-- is taken or not (measured: ~450 dispatch steps per cell, ~2/3 of the
-- schedule). Each row below holds one reader's nine writer latencies and is
-- computed once per build from the table above (the proofs still read the
-- table itself); a scheduling step dispatches its reader's row once, and a
-- cell extracts only its writer's entry, and only on a register match.

-- A latency class as its row and column number.
def sm121LowerClassTag =
  (lambda unrestricted class : (family SM121LowerClass) .
    (eliminate
      SM121LowerClass
      (lambda unrestricted current : (family SM121LowerClass) . Nat)
      class
      (branch SM121LowerAlu . zero)
      (branch SM121LowerDualAlu . 1)
      (branch SM121LowerFma . 2)
      (branch SM121LowerWide . 3)
      (branch SM121LowerHalfToFloat . 4)
      (branch SM121LowerTensor . 5)
      (branch SM121LowerVariable . 6)
      (branch SM121LowerMemory . 7)
      (branch SM121LowerNoResult . 8)))

-- 64 to the powers 0 through 8, the places of a row's nine entries.
def sm121LowerLatencyPlace =
  (lambda unrestricted writerTag : Nat .
    (naturalSelect (naturalEqual writerTag zero) 1
    (naturalSelect (naturalEqual writerTag 1) 64
    (naturalSelect (naturalEqual writerTag 2) 4096
    (naturalSelect (naturalEqual writerTag 3) 262144
    (naturalSelect (naturalEqual writerTag 4) 16777216
    (naturalSelect (naturalEqual writerTag 5) 1073741824
    (naturalSelect (naturalEqual writerTag 6) 68719476736
    (naturalSelect (naturalEqual writerTag 7) 4398046511104
    281474976710656)))))))))

-- One reader's row: the entry for writer tag `w` holds the table's latency
-- for that writer and this reader, in 6 bits.
def sm121LowerLatencyRow =
  (lambda unrestricted reader : (family SM121LowerClass) .
    (nat-add (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerAlu) reader)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerDualAlu) reader) 64)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerFma) reader) 4096)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerWide) reader) 262144)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerHalfToFloat) reader) 16777216)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerTensor) reader) 1073741824)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerVariable) reader) 68719476736)
    (nat-add (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerMemory) reader) 4398046511104)
    (nat-multiply (sm121LowerReadAfterWrite (constructor SM121LowerClass SM121LowerNoResult) reader) 281474976710656))))))))))

-- Each reader's row, computed once at compile time and baked in as a
-- literal: the machine shares no evaluation across references, so a row
-- built per step would rerun its nine table calls for every instruction.
def sm121LowerRowAlu : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerAlu)))
def sm121LowerRowDualAlu : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerDualAlu)))
def sm121LowerRowFma : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerFma)))
def sm121LowerRowWide : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerWide)))
def sm121LowerRowHalfToFloat : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerHalfToFloat)))
def sm121LowerRowTensor : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerTensor)))
def sm121LowerRowVariable : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerVariable)))
def sm121LowerRowMemory : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerMemory)))
def sm121LowerRowNoResult : Nat = (compile-time (sm121LowerLatencyRow (constructor SM121LowerClass SM121LowerNoResult)))

-- A reader's row, dispatched once per scheduling step.
def sm121LowerReaderRowOf =
  (lambda unrestricted reader : (family SM121LowerClass) .
    (eliminate
      SM121LowerClass
      (lambda unrestricted current : (family SM121LowerClass) . Nat)
      reader
      (branch SM121LowerAlu . sm121LowerRowAlu)
      (branch SM121LowerDualAlu . sm121LowerRowDualAlu)
      (branch SM121LowerFma . sm121LowerRowFma)
      (branch SM121LowerWide . sm121LowerRowWide)
      (branch SM121LowerHalfToFloat . sm121LowerRowHalfToFloat)
      (branch SM121LowerTensor . sm121LowerRowTensor)
      (branch SM121LowerVariable . sm121LowerRowVariable)
      (branch SM121LowerMemory . sm121LowerRowMemory)
      (branch SM121LowerNoResult . sm121LowerRowNoResult)))

-- A writer's entry of a row.
def sm121LowerLatencyAt =
  (lambda unrestricted row : Nat .
    (lambda unrestricted writerTag : Nat .
      (nat-modulo (nat-divide row (sm121LowerLatencyPlace writerTag)) 64)))

-- A predicate written by ISETP read as a guard (NVIDIA's sm100 predicate
-- table, dual ALU to guard, plus NAK's padding).
def sm121LowerPredicateLatency : Nat = (nat-add 13 sm121LowerPadding)

-- ---------------------------------------------------------------------------
-- Registers

def sm121LowerNone : (family SM121LowerRegisters) =
  (constructor SM121LowerRegisters SM121LowerRegistersEnd)

def sm121LowerCons =
  (lambda unrestricted register : Nat .
    (lambda unrestricted rest : (family SM121LowerRegisters) .
      (constructor SM121LowerRegisters SM121LowerRegistersNext register rest)))

def sm121LowerAppend =
  (lambda unrestricted left : (family SM121LowerRegisters) .
    (lambda unrestricted right : (family SM121LowerRegisters) .
      (eliminate
        SM121LowerRegisters
        (lambda unrestricted current : (family SM121LowerRegisters) . (family SM121LowerRegisters))
        left
        (branch SM121LowerRegistersEnd . right)
        (branch SM121LowerRegistersNext head tail induction . (sm121LowerCons head induction)))))

-- A register operand as a list: RZ names none.
def sm121LowerNamed =
  (lambda unrestricted register : Nat .
    (nat-add (nat-less-than register sm121LowerZeroRegister) (nat-less-than sm121LowerZeroRegister register)))

def sm121LowerOne =
  (lambda unrestricted register : Nat .
    (nat-eliminate
      (lambda unrestricted current : Nat . (family SM121LowerRegisters))
      sm121LowerNone
      (lambda unrestricted predecessor : Nat .
        (lambda unrestricted induction : (family SM121LowerRegisters) .
          (sm121LowerCons register sm121LowerNone)))
      (sm121LowerNamed register)))

def sm121LowerTwo =
  (lambda unrestricted first : Nat .
    (lambda unrestricted second : Nat .
      (sm121LowerAppend (sm121LowerOne first) (sm121LowerOne second))))

def sm121LowerThree =
  (lambda unrestricted first : Nat .
    (lambda unrestricted second : Nat .
      (lambda unrestricted third : Nat .
        (sm121LowerAppend (sm121LowerOne first) (sm121LowerTwo second third)))))

-- `count` consecutive registers from `base`; none from RZ.
def sm121LowerSpan =
  (lambda unrestricted base : Nat .
    (lambda unrestricted count : Nat .
      (nat-eliminate
        (lambda unrestricted current : Nat . (family SM121LowerRegisters))
        sm121LowerNone
        (lambda unrestricted predecessor : Nat .
          (lambda unrestricted induction : (family SM121LowerRegisters) .
            (nat-eliminate
              (lambda unrestricted current : Nat . (family SM121LowerRegisters))
              sm121LowerNone
              (lambda unrestricted unused : Nat .
                (lambda unrestricted inner : (family SM121LowerRegisters) .
                  (sm121LowerAppend induction (sm121LowerOne (nat-add base predecessor)))))
              (sm121LowerNamed base))))
        count)))

def sm121LowerHas =
  (lambda unrestricted registers : (family SM121LowerRegisters) .
    (lambda unrestricted register : Nat .
      (eliminate
        SM121LowerRegisters
        (lambda unrestricted current : (family SM121LowerRegisters) . Nat)
        registers
        (branch SM121LowerRegistersEnd . zero)
        (branch SM121LowerRegistersNext head tail induction .
          (nat-eliminate (lambda unrestricted current : Nat . Nat)
            induction
            (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (succ zero)))
            (naturalEqual head register))))))

-- `left` with those of `right` it lacks
def sm121LowerUnion =
  (lambda unrestricted left : (family SM121LowerRegisters) .
    (lambda unrestricted right : (family SM121LowerRegisters) .
      (eliminate
        SM121LowerRegisters
        (lambda unrestricted current : (family SM121LowerRegisters) . (family SM121LowerRegisters))
        right
        (branch SM121LowerRegistersEnd . left)
        (branch SM121LowerRegistersNext head tail induction .
          (sm121LowerSelect (family SM121LowerRegisters)
            (sm121LowerHas induction head)
            induction
            (sm121LowerCons head induction))))))

-- whether any of `registers` is in `set`
def sm121LowerAny =
  (lambda unrestricted registers : (family SM121LowerRegisters) .
    (lambda unrestricted set : (family SM121LowerRegisters) .
      (eliminate
        SM121LowerRegisters
        (lambda unrestricted current : (family SM121LowerRegisters) . Nat)
        registers
        (branch SM121LowerRegistersEnd . zero)
        (branch SM121LowerRegistersNext head tail induction .
          (nat-eliminate (lambda unrestricted current : Nat . Nat)
            induction
            (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (succ zero)))
            (sm121LowerHas set head))))))

def sm121LowerIsEmpty =
  (lambda unrestricted registers : (family SM121LowerRegisters) .
    (eliminate
      SM121LowerRegisters
      (lambda unrestricted current : (family SM121LowerRegisters) . Nat)
      registers
      (branch SM121LowerRegistersEnd . 1)
      (branch SM121LowerRegistersNext head tail induction . zero)))

def sm121LowerHighest =
  (lambda unrestricted registers : (family SM121LowerRegisters) .
    (lambda unrestricted start : Nat .
      (eliminate
        SM121LowerRegisters
        (lambda unrestricted current : (family SM121LowerRegisters) . Nat)
        registers
        (branch SM121LowerRegistersEnd . start)
        (branch SM121LowerRegistersNext head tail induction .
          (sm121LowerMaximum (succ head) induction)))))

def sm121LowerMaskWordBits : Nat = 64

def sm121LowerMaskEmpty : (family SM121LowerMask) =
  (constructor SM121LowerMask SM121LowerMaskValue zero zero zero zero)

-- 2^bit for a bit below 64 without the doubling loop: 2^(bit mod 8) times
-- 256^(bit div 8), each small power one byte shift and 256^q = ((2^q)^2)^2)^2.
-- The loop costs a step per bit (~3.5k dispatch steps for high registers);
-- mask insert and test call this per register, so the loop would dominate
-- the once-per-program context folds that build the free scoreboard and
-- scratch registers.  (A chain of eight comparisons against bit div 8 took
-- 93 steps a call, 8% of realizing a tiled product.)
def sm121LowerPlaceOf =
  (lambda unrestricted bit : Nat .
    (let unrestricted low = (byte-to-nat (byte-shift-left (byte 1) (nat-to-byte (nat-modulo bit 8)))) in
    (let unrestricted q = (byte-to-nat (byte-shift-left (byte 1) (nat-to-byte (nat-divide bit 8)))) in
    (let unrestricted q2 = (nat-multiply q q) in
    (let unrestricted q4 = (nat-multiply q2 q2) in
      (nat-multiply low (nat-multiply q4 q4)))))))

def sm121LowerMaskSet =
  (lambda unrestricted word : Nat .
    (lambda unrestricted bit : Nat .
      (sm121LowerWithField word (sm121LowerPlaceOf bit) 2 1)))

def sm121LowerMaskInsert =
  (lambda unrestricted mask : (family SM121LowerMask) .
    (lambda unrestricted register : Nat .
      (eliminate
        SM121LowerMask
        (lambda unrestricted current : (family SM121LowerMask) . (family SM121LowerMask))
        mask
        (branch SM121LowerMaskValue w0 w1 w2 w3 .
          (let unrestricted index = (nat-divide register sm121LowerMaskWordBits) in
          (let unrestricted bit = (nat-modulo register sm121LowerMaskWordBits) in
            (constructor SM121LowerMask SM121LowerMaskValue
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                w0
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (sm121LowerMaskSet w0 bit)))
                (naturalEqual index zero))
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                w1
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (sm121LowerMaskSet w1 bit)))
                (naturalEqual index 1))
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                w2
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (sm121LowerMaskSet w2 bit)))
                (naturalEqual index 2))
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                w3
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (sm121LowerMaskSet w3 bit)))
                (naturalEqual index 3)))))))))

def sm121LowerMaskHas =
  (lambda unrestricted mask : (family SM121LowerMask) .
    (lambda unrestricted register : Nat .
      (eliminate
        SM121LowerMask
        (lambda unrestricted current : (family SM121LowerMask) . Nat)
        mask
        (branch SM121LowerMaskValue w0 w1 w2 w3 .
          (let unrestricted index = (nat-divide register sm121LowerMaskWordBits) in
          (let unrestricted place = (sm121LowerPlaceOf (nat-modulo register sm121LowerMaskWordBits)) in
            (sm121LowerField
              (naturalSelect (naturalEqual index zero) w0
                (naturalSelect (naturalEqual index 1) w1
                  (naturalSelect (naturalEqual index 2) w2 w3)))
              place
              2)))))))

def sm121LowerMaskAll =
  (lambda unrestricted mask : (family SM121LowerMask) .
    (lambda unrestricted registers : (family SM121LowerRegisters) .
      (app
        (eliminate
          SM121LowerRegisters
          (lambda unrestricted current : (family SM121LowerRegisters) .
            (pi unrestricted accumulated : (family SM121LowerMask) . (family SM121LowerMask)))
          registers
          (branch SM121LowerRegistersEnd .
            (lambda unrestricted accumulated : (family SM121LowerMask) . accumulated))
          (branch SM121LowerRegistersNext head tail induction .
            (lambda unrestricted accumulated : (family SM121LowerMask) .
              (induction (sm121LowerMaskInsert accumulated head)))))
        mask)))

-- The least register below `count` for which `free` holds, or `count`.
def sm121LowerFirst =
  (lambda unrestricted count : Nat .
    (lambda unrestricted free : (pi unrestricted register : Nat . Nat) .
      (nat-eliminate
        (lambda unrestricted current : Nat . Nat)
        count
        (lambda unrestricted predecessor : Nat .
          (lambda unrestricted induction : Nat .
            (let unrestricted register = (nat-subtract count (succ predecessor)) in
              (naturalSelect (free register) register induction))))
        count)))

-- ---------------------------------------------------------------------------
-- Lowered instructions

def sm121LowerRefusalText =
  (lambda unrestricted refusal : (family SM121LowerRefusal) .
    (eliminate
      SM121LowerRefusal
      (lambda unrestricted current : (family SM121LowerRefusal) . Bytes)
      refusal
      (branch SM121LowerRefusedEncoding . b"sm_121 lowering: the SM86 encoder refused an instruction")
      (branch SM121LowerRefusedBranchTarget . b"sm_121 lowering: a branch's target lies outside its program")
      (branch SM121LowerRefusedRegisterBeyondCount . b"sm_121 lowering: the program names a register beyond its declared count")
      (branch SM121LowerRefusedNoFreeRegister . b"sm_121 lowering: no register is free for a constant operand")
      (branch SM121LowerRefusedNoFreePair . b"sm_121 lowering: no register pair is free for a constant operand")
      (branch SM121LowerRefusedSharedOffset . b"sm_121 lowering: a shared-memory offset leaves the 24-bit field")
      (branch SM121LowerRefusedNoStall . b"sm_121 lowering: an instruction has no stall")
      (branch SM121LowerRefusedHazard . b"sm_121 lowering: the static scoreboard check found a hazard in the schedule")))

def sm121LowerRegister =
  (lambda unrestricted register : (family SM86Register) .
    (eliminate
      SM86Register
      (lambda unrestricted current : (family SM86Register) . Nat)
      register
      (branch SM86RegisterValue value . (byte-to-nat value))))

def sm121LowerShape =
  (lambda unrestricted rule : (family SM121LowerRule) .
    (lambda unrestricted class : (family SM121LowerClass) .
      (lambda unrestricted reads : (family SM121LowerRegisters) .
        (lambda unrestricted writes : (family SM121LowerRegisters) .
          (constructor SM121LowerShape SM121LowerShapeValue rule class reads writes sm121LowerNone)))))

def sm121LowerSame =
  (lambda unrestricted class : (family SM121LowerClass) .
    (lambda unrestricted reads : (family SM121LowerRegisters) .
      (lambda unrestricted writes : (family SM121LowerRegisters) .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleSame) class reads writes))))

def sm121LowerMatrixCount =
  (lambda unrestricted count : (family SM86SharedMatrixCount) .
    (eliminate
      SM86SharedMatrixCount
      (lambda unrestricted current : (family SM86SharedMatrixCount) . Nat)
      count
      (branch SM86SharedMatrix1 . 1)
      (branch SM86SharedMatrix2 . 2)
      (branch SM86SharedMatrix4 . 4)))

-- The shape of each SM86 instruction.
def sm121LowerShapeOf =
  (lambda unrestricted body : (family SM86InstructionBody) .
    (eliminate
      SM86InstructionBody
      (lambda unrestricted current : (family SM86InstructionBody) . (family SM121LowerShape))
      body
      (branch SM86MoveConstant d bank offset control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleConstantMove) (constructor SM121LowerClass SM121LowerVariable) sm121LowerNone
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86SpecialToRegister d special control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerVariable) sm121LowerNone (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86MoveImmediate d value control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerDualAlu) sm121LowerNone (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86IntegerMultiplyAddConstant d a bank offset c control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleConstantMultiplyAdd) (constructor SM121LowerClass SM121LowerFma)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister c))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86IntegerMultiplyAddImmediate d a value c control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerFma)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister c))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86IntegerMultiplyAddWideConstant d a b bank offset control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleConstantMultiplyAddWide) (constructor SM121LowerClass SM121LowerWide)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerSpan (sm121LowerRegister d) 2)))
      (branch SM86IntegerAddThreeImmediate d a value control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86IntegerAddThreeRegister d a b control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86ShiftRightImmediate d a amount control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86LogicThreeInputTruthTable d a b table control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86IntegerToFloat d a control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatAdd d a b control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerFma)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatMultiply d a b control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerFma)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatFusedMultiplyAdd d a b c control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerFma)
          (sm121LowerThree (sm121LowerRegister a) (sm121LowerRegister b) (sm121LowerRegister c))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86MultiFunctionUnitApproximation d a operation control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerVariable) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatMinimumOrMaximum d a b mode control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerDualAlu)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatNegate d a control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerFma) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatPairToPackedHalfPair d a b control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86FloatPairToPackedBFloat16Pair d a b control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister b))
          (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86HalfToFloat d a selector control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerHalfToFloat) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86BFloat16ToFloat d a selector control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerAlu) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32 d a b c control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerTensor)
          (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister a) 4)
            (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister b) 2) (sm121LowerSpan (sm121LowerRegister c) 4)))
          (sm121LowerSpan (sm121LowerRegister d) 4)))
      (branch SM86TensorCoreBFloat16MatrixMultiplyAccumulate16x8x16Float32 d a b c control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerTensor)
          (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister a) 4)
            (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister b) 2) (sm121LowerSpan (sm121LowerRegister c) 4)))
          (sm121LowerSpan (sm121LowerRegister d) 4)))
      (branch SM86LoadGlobal d a offset control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerMemory) (sm121LowerSpan (sm121LowerRegister a) 2) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86LoadGlobalWide d a offset control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerMemory) (sm121LowerSpan (sm121LowerRegister a) 2) (sm121LowerSpan (sm121LowerRegister d) 4)))
      (branch SM86WarpShuffle d a lane segment mode control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerMemory) (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86LoadShared d a offset control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleShared) (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerOne (sm121LowerRegister a)) (sm121LowerOne (sm121LowerRegister d))))
      (branch SM86LoadSharedMatrix d a offset count transpose control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleSharedMatrix) (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerOne (sm121LowerRegister a))
          (sm121LowerSpan (sm121LowerRegister d) (sm121LowerMatrixCount count))))
      (branch SM86StoreShared a v offset control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleShared) (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerTwo (sm121LowerRegister a) (sm121LowerRegister v)) sm121LowerNone))
      (branch SM86LoadGlobalToShared a offset source sourceOffset control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleAsyncCopy) (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerAppend (sm121LowerOne (sm121LowerRegister a)) (sm121LowerSpan (sm121LowerRegister source) 2))
          sm121LowerNone))
      (branch SM86CommitAsyncGroup control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone))
      (branch SM86WaitAsyncGroups count control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone))
      (branch SM86BarrierSynchronize control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone))
      (branch SM86StoreGlobal a v offset control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister a) 2) (sm121LowerOne (sm121LowerRegister v)))
          sm121LowerNone))
      (branch SM86StoreGlobalWide a v offset control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister a) 2) (sm121LowerSpan (sm121LowerRegister v) 4))
          sm121LowerNone))
      (branch SM86StoreGlobal64 a v offset control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister a) 2) (sm121LowerSpan (sm121LowerRegister v) 2))
          sm121LowerNone))
      (branch SM86ReduceGlobalAddFloat32 a v offset control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleReduction) (constructor SM121LowerClass SM121LowerMemory)
          (sm121LowerAppend (sm121LowerSpan (sm121LowerRegister a) 2) (sm121LowerOne (sm121LowerRegister v)))
          sm121LowerNone))
      (branch SM86PredicateGreaterThanImmediate p a value control .
        (constructor SM121LowerShape SM121LowerShapeValue (constructor SM121LowerRule SM121LowerRuleSame) (constructor SM121LowerClass SM121LowerDualAlu)
          (sm121LowerOne (sm121LowerRegister a)) sm121LowerNone
          (sm121LowerCons (byte-to-nat (sm86PredicateNumber p)) sm121LowerNone)))
      (branch SM86Branch offset descriptor control .
        (sm121LowerShape (constructor SM121LowerRule SM121LowerRuleBranch) (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone))
      (branch SM86Exit control .
        (sm121LowerSame (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone))))

-- The guard predicate an instruction reads (PT is none).
def sm121LowerGuardPredicate =
  (lambda unrestricted predicate : (family SM86Predicate) .
    (let unrestricted number = (byte-to-nat (sm86PredicateNumber predicate)) in
      (naturalSelect (naturalEqual number sm121LowerAlways) zero 1)))

def sm121LowerGuardReads =
  (lambda unrestricted guard : (family SM86InstructionGuard) .
    (eliminate
      SM86InstructionGuard
      (lambda unrestricted current : (family SM86InstructionGuard) . (family SM121LowerRegisters))
      guard
      (branch SM86InstructionAlways . sm121LowerNone)
      (branch SM86InstructionWhen predicate .
        (nat-eliminate
          (lambda unrestricted current : Nat . (family SM121LowerRegisters))
          sm121LowerNone
          (lambda unrestricted unused : Nat .
            (lambda unrestricted induction : (family SM121LowerRegisters) .
              (sm121LowerCons (byte-to-nat (sm86PredicateNumber predicate)) sm121LowerNone)))
          (sm121LowerGuardPredicate predicate)))
      (branch SM86InstructionWhenNot predicate .
        (nat-eliminate
          (lambda unrestricted current : Nat . (family SM121LowerRegisters))
          sm121LowerNone
          (lambda unrestricted unused : Nat .
            (lambda unrestricted induction : (family SM121LowerRegisters) .
              (sm121LowerCons (byte-to-nat (sm86PredicateNumber predicate)) sm121LowerNone)))
          (sm121LowerGuardPredicate predicate)))))

-- A little-endian word of `count` bytes.
def sm121LowerLittleEndian =
  (lambda unrestricted count : Nat .
    (nat-eliminate
      (lambda unrestricted current : Nat . (pi unrestricted bytes : Bytes . Nat))
      (lambda unrestricted bytes : Bytes . zero)
      (lambda unrestricted predecessor : Nat .
        (lambda unrestricted induction : (pi unrestricted bytes : Bytes . Nat) .
          (lambda unrestricted bytes : Bytes .
            (nat-add
              (byte-to-nat (bytes-head bytes))
              (nat-multiply sm121LowerRegisterSpan (induction (bytes-tail bytes)))))))
      count))

def sm121LowerDrop =
  (lambda unrestricted count : Nat .
    (nat-eliminate
      (lambda unrestricted current : Nat . (pi unrestricted bytes : Bytes . Bytes))
      (lambda unrestricted bytes : Bytes . bytes)
      (lambda unrestricted predecessor : Nat .
        (lambda unrestricted induction : (pi unrestricted bytes : Bytes . Bytes) .
          (lambda unrestricted bytes : Bytes . (induction (bytes-tail bytes)))))
      count))

def sm121LowerWordBytes : Nat = 8

def sm121LowerInstructionBytes : Nat = (nat-add sm121LowerWordBytes sm121LowerWordBytes)

-- An SM86 instruction, decoded: its two words (the first sixteen of
-- `words`) and its shape.
def sm121LowerDecode =
  (lambda unrestricted instruction : (family SM86Instruction) .
    (lambda unrestricted words : Bytes .
      (eliminate
        SM86Instruction
        (lambda unrestricted current : (family SM86Instruction) . (family SM121LowerDecoded))
        instruction
        (branch SM86InstructionValue guard body .
          (constructor SM121LowerDecoded SM121LowerDecodedValue
            (sm121LowerLittleEndian sm121LowerWordBytes words)
            (sm121LowerLittleEndian sm121LowerWordBytes (sm121LowerDrop sm121LowerWordBytes words))
            (sm121LowerGuardReads guard)
            (sm121LowerShapeOf body))))))

-- A program's instructions decoded against its SM86 encoding
-- (Accelerator.SM86's sm86EncodeProgram, the words its canonical placement
-- holds), sixteen bytes each.
def sm121LowerDecodeProgram =
  (lambda unrestricted program : (family SM86Program) .
    (eliminate
      SM86ProgramEncodingResult
      (lambda unrestricted current : (family SM86ProgramEncodingResult) . (family SM121LowerDecodedList))
      (sm86EncodeProgram program)
      (branch SM86ProgramEncodingSucceeded image telemetry .
        (app
          (eliminate
            SM86Program
            (lambda unrestricted current : (family SM86Program) .
              (pi unrestricted words : Bytes . (family SM121LowerDecodedList)))
            program
            (branch SM86ProgramEnd .
              (lambda unrestricted words : Bytes . (constructor SM121LowerDecodedList SM121LowerDecodedEnd)))
            (branch SM86ProgramNext head tail induction .
              (lambda unrestricted words : Bytes .
                (constructor SM121LowerDecodedList SM121LowerDecodedNext
                  (sm121LowerDecode head words)
                  (induction (sm121LowerDrop sm121LowerInstructionBytes words))))))
          image))
      (branch SM86ProgramEncodingFailed ordinal failure telemetry .
        (constructor SM121LowerDecodedList SM121LowerDecodedRefused
          (constructor SM121LowerRefusal SM121LowerRefusedEncoding)))))

-- ---------------------------------------------------------------------------
-- What the whole program fixes: the free scoreboard and scratch registers

def sm121LowerDecodedFold =
  (lambda erased result : Type 0 .
    (lambda unrestricted decoded : (family SM121LowerDecodedList) .
      (lambda unrestricted start : result .
        (lambda unrestricted step :
            (pi unrestricted accumulated : result .
              (pi unrestricted item : (family SM121LowerDecoded) . result)) .
          (app
            (eliminate
              SM121LowerDecodedList
              (lambda unrestricted current : (family SM121LowerDecodedList) .
                (pi unrestricted accumulated : result . result))
              decoded
              (branch SM121LowerDecodedEnd . (lambda unrestricted accumulated : result . accumulated))
              (branch SM121LowerDecodedNext item tail induction .
                (lambda unrestricted accumulated : result . (induction (step accumulated item))))
              (branch SM121LowerDecodedRefused refusal . (lambda unrestricted accumulated : result . accumulated)))
            start)))))

def sm121LowerDecodedRegisters =
  (lambda unrestricted item : (family SM121LowerDecoded) .
    (eliminate
      SM121LowerDecoded
      (lambda unrestricted current : (family SM121LowerDecoded) . (family SM121LowerRegisters))
      item
      (branch SM121LowerDecodedValue low high guardReads shape .
        (eliminate
          SM121LowerShape
          (lambda unrestricted current : (family SM121LowerShape) . (family SM121LowerRegisters))
          shape
          (branch SM121LowerShapeValue rule class reads writes predicateWrites .
            (sm121LowerAppend reads writes))))))

-- the scoreboards an instruction's control names, as a bit mask
def sm121LowerDecodedBarriers =
  (lambda unrestricted item : (family SM121LowerDecoded) .
    (eliminate
      SM121LowerDecoded
      (lambda unrestricted current : (family SM121LowerDecoded) . (family SM121LowerRegisters))
      item
      (branch SM121LowerDecodedValue low high guardReads shape .
        (sm121LowerTwo
          (sm121LowerField high sm121LowerWriteBarrierPlace sm121LowerBarrierSpan)
          (sm121LowerField high sm121LowerReadBarrierPlace sm121LowerBarrierSpan)))))

-- The wait mask naming each scoreboard `barriers` holds, and `freeBarrier`.
def sm121LowerWaitAllOf =
  (lambda unrestricted barriers : (family SM121LowerMask) .
    (lambda unrestricted freeBarrier : Nat .
      (nat-eliminate
        (lambda unrestricted current : Nat . Nat)
        zero
        (lambda unrestricted barrier : Nat .
          (lambda unrestricted induction : Nat .
            (nat-add induction
              (naturalSelect
                (naturalOr (sm121LowerMaskHas barriers barrier) (naturalEqual barrier freeBarrier))
                (sm121LowerPlace barrier)
                zero))))
        sm121LowerBarrierCount)))

def sm121LowerContextOf =
  (lambda unrestricted addressing : (family SM121SharedAddressing) .
  (lambda unrestricted decoded : (family SM121LowerDecodedList) .
    (lambda unrestricted count : Nat .
      (let unrestricted named =
        (sm121LowerDecodedFold (family SM121LowerMask) decoded sm121LowerMaskEmpty
          (lambda unrestricted mask : (family SM121LowerMask) .
            (lambda unrestricted item : (family SM121LowerDecoded) .
              (sm121LowerMaskAll mask (sm121LowerDecodedRegisters item))))) in
      (let unrestricted barriers =
        (sm121LowerDecodedFold (family SM121LowerMask) decoded sm121LowerMaskEmpty
          (lambda unrestricted mask : (family SM121LowerMask) .
            (lambda unrestricted item : (family SM121LowerDecoded) .
              (sm121LowerMaskAll mask (sm121LowerDecodedBarriers item))))) in
      -- the scoreboards of the instructions that read or write registers
      (let unrestricted bearing =
        (sm121LowerDecodedFold (family SM121LowerMask) decoded sm121LowerMaskEmpty
          (lambda unrestricted mask : (family SM121LowerMask) .
            (lambda unrestricted item : (family SM121LowerDecoded) .
              (sm121LowerSelect (family SM121LowerMask)
                (sm121LowerIsEmpty (sm121LowerDecodedRegisters item))
                mask
                (sm121LowerMaskAll mask (sm121LowerDecodedBarriers item)))))) in
      (let unrestricted unused =
        (sm121LowerFirst sm121LowerBarrierCount
          (lambda unrestricted barrier : Nat .
            (naturalIsZero (sm121LowerMaskHas barriers barrier)))) in
      (let unrestricted free =
        (lambda unrestricted register : Nat .
          (naturalIsZero (sm121LowerMaskHas named register))) in
      (let unrestricted freeBarrier =
        (naturalSelect (naturalEqual unused sm121LowerBarrierCount)
          (nat-subtract sm121LowerBarrierCount 1)
          unused) in
        (constructor SM121LowerContext SM121LowerContextValue
          freeBarrier
          (sm121LowerFirst count free)
          (sm121LowerFirst (nat-subtract count 1)
            (lambda unrestricted register : Nat .
              (naturalAnd
                (naturalIsZero (nat-modulo register 2))
                (naturalAnd (free register) (free (succ register))))))
          count
          addressing
          (sm121LowerWaitAllOf bearing freeBarrier)))))))))))

-- the highest register the program names, plus one
def sm121LowerNamedBound =
  (lambda unrestricted decoded : (family SM121LowerDecodedList) .
    (sm121LowerDecodedFold Nat decoded zero
      (lambda unrestricted bound : Nat .
        (lambda unrestricted item : (family SM121LowerDecoded) .
          (sm121LowerHighest (sm121LowerDecodedRegisters item) bound)))))

-- ---------------------------------------------------------------------------
-- The four rewrites

def sm121LowerOnly =
  (lambda unrestricted op : (family SM121LowerOp) .
    (constructor SM121LowerExpansion SM121LowerExpanded
      (constructor SM121LowerOps SM121LowerOpsNext op (constructor SM121LowerOps SM121LowerOpsEnd))))

def sm121LowerPair =
  (lambda unrestricted first : (family SM121LowerOp) .
    (lambda unrestricted second : (family SM121LowerOp) .
      (constructor SM121LowerExpansion SM121LowerExpanded
        (constructor SM121LowerOps SM121LowerOpsNext first
          (constructor SM121LowerOps SM121LowerOpsNext second (constructor SM121LowerOps SM121LowerOpsEnd))))))

def sm121LowerRefuse =
  (lambda unrestricted refusal : (family SM121LowerRefusal) .
    (constructor SM121LowerExpansion SM121LowerExpansionRefused refusal))

-- The low word's opcode replaced, the half above bit 32 cleared.
def sm121LowerRegisterForm =
  (lambda unrestricted low : Nat .
    (lambda unrestricted opcode : Nat .
      (sm121LowerWithField (nat-modulo low sm121LowerHalfWordSpan) 1 sm121LowerOpcodeSpan opcode)))

-- An LDS, STS or LDSM with its offset moved up past the kilobyte sm_121
-- reserves, or refused when it leaves the field.
def sm121LowerRebaseShared =
  (lambda unrestricted low : Nat .
    (lambda unrestricted high : Nat .
      (lambda unrestricted plain :
          (pi unrestricted opLow : Nat . (pi unrestricted opHigh : Nat .
            (pi unrestricted opReads : (family SM121LowerRegisters) . (family SM121LowerOp)))) .
        (lambda unrestricted reads : (family SM121LowerRegisters) .
          (let unrestricted offset =
            (nat-add sm121LowerSharedBase
              (sm121LowerField low sm121LowerSharedOffsetPlace sm121LowerSharedOffsetSpan)) in
            (sm121LowerSelect (family SM121LowerExpansion)
              (nat-less-than offset sm121LowerSharedOffsetLimit)
              (sm121LowerOnly
                (plain
                  (sm121LowerWithField low sm121LowerSharedOffsetPlace sm121LowerSharedOffsetSpan offset)
                  high
                  reads))
              (sm121LowerRefuse (constructor SM121LowerRefusal SM121LowerRefusedSharedOffset))))))))

-- A branch's displacement.  SM86's BRA holds the byte displacement from the
-- next instruction, 50 bits in two's complement: bits 32..63 of the low word,
-- then bits 0..17 of the high word (Accelerator.SM86's encoder writes the
-- offset word and its sign-extending descriptor there).  sm_121 moves the
-- same displacement, counted in words of 4 bytes (E = D / 4, 56 bits):
-- E's low 8 bits to bits 16..23, its next 30 to bits 34..63 and the rest to
-- bits 0..17 of the high word -- read off ptxas -arch=sm_121a's BRA against
-- nvdisasm's label addresses on the DGX Spark, 2026-09-26.  An SM86 word
-- passed through unchanged branches 256 times too far on sm_121.
def sm121LowerBranchOffsetPlace : Nat = sm121LowerHalfWordSpan
def sm121LowerBranchHighSpan : Nat = (compile-time (sm121LowerPlace 18))
def sm121LowerBranchSpan : Nat = (compile-time (sm121LowerPlace 50))
def sm121LowerBranchHalf : Nat = (compile-time (sm121LowerPlace 49))
def sm121LowerBranchWordSpan : Nat = (compile-time (sm121LowerPlace 56))
def sm121LowerBranchLowPlace : Nat = (compile-time (sm121LowerPlace 16))
def sm121LowerBranchMiddlePlace : Nat = (compile-time (sm121LowerPlace 34))
def sm121LowerBranchMiddleSpan : Nat = (compile-time (sm121LowerPlace 30))
-- bytes per word the sm_121 displacement counts
def sm121LowerBranchUnit : Nat = 4

-- The SM86 index + 1 of the instruction the branch at `index` jumps to (its
-- SM86 displacement counts whole instructions), or 0 when that would precede
-- the program; a target past its end has no label and is refused when the
-- schedule is relocated.
def sm121LowerBranchTarget =
  (lambda unrestricted low : Nat .
    (lambda unrestricted high : Nat .
      (lambda unrestricted index : Nat .
        (let unrestricted displacement =
          (nat-add
            (sm121LowerField low sm121LowerBranchOffsetPlace sm121LowerHalfWordSpan)
            (nat-multiply (sm121LowerField high 1 sm121LowerBranchHighSpan) sm121LowerHalfWordSpan)) in
        (let unrestricted next = (succ index) in
          (naturalSelect (nat-less-than displacement sm121LowerBranchHalf)
            (succ (nat-add next (nat-divide displacement sm121LowerInstructionBytes)))
            (let unrestricted back =
              (nat-divide (nat-subtract sm121LowerBranchSpan displacement) sm121LowerInstructionBytes) in
              (naturalSelect (nat-less-than next back) zero (succ (nat-subtract next back))))))))))

-- `low` and `high` of a branch with its displacement set to `words` words of
-- 4 bytes from the next instruction, given as `forward` words ahead or
-- `backward` words back (one of them zero), in sm_121's placement.
def sm121LowerWithDisplacement =
  (lambda unrestricted forward : Nat .
    (lambda unrestricted backward : Nat .
      (lambda unrestricted low : Nat .
        (lambda unrestricted high : Nat .
          (let unrestricted encoded =
            (nat-modulo (nat-subtract (nat-add forward sm121LowerBranchWordSpan) backward) sm121LowerBranchWordSpan) in
          (let unrestricted cleared =
            (sm121LowerWithField
              (sm121LowerWithField low sm121LowerBranchLowPlace sm121LowerRegisterSpan
                (nat-modulo encoded sm121LowerRegisterSpan))
              sm121LowerBranchOffsetPlace sm121LowerHalfWordSpan zero) in
            (constructor SM121LowerWords SM121LowerWordsValue
              (sm121LowerWithField cleared sm121LowerBranchMiddlePlace sm121LowerBranchMiddleSpan
                (nat-divide encoded sm121LowerRegisterSpan))
              (sm121LowerWithField high 1 sm121LowerBranchHighSpan
                (nat-divide encoded (nat-multiply sm121LowerRegisterSpan sm121LowerBranchMiddleSpan))))))))))

def sm121LowerExpand =
  (lambda unrestricted context : (family SM121LowerContext) .
  (lambda unrestricted index : Nat .
  (lambda unrestricted targets : (family SM121LowerRegisters) .
    (lambda unrestricted item : (family SM121LowerDecoded) .
      (eliminate
        SM121LowerContext
        (lambda unrestricted current : (family SM121LowerContext) . (family SM121LowerExpansion))
        context
        (branch SM121LowerContextValue barrier free freePair count addressing waitAll .
          (eliminate
            SM121LowerDecoded
            (lambda unrestricted current : (family SM121LowerDecoded) . (family SM121LowerExpansion))
            item
            (branch SM121LowerDecodedValue low high guardReads shape .
              (eliminate
                SM121LowerShape
                (lambda unrestricted current : (family SM121LowerShape) . (family SM121LowerExpansion))
                shape
                (branch SM121LowerShapeValue rule class reads writes predicateWrites .
                  -- the first op an instruction lowers to carries its label, and
                  -- the join mark when a branch targets it
                  (let unrestricted label = (succ index) in
                  (let unrestricted join = (sm121LowerHas targets index) in
                  (let unrestricted plainWith =
                    (lambda unrestricted opLabel : Nat .
                    (lambda unrestricted opJoin : Nat .
                    (lambda unrestricted opTarget : Nat .
                    (lambda unrestricted opLow : Nat .
                      (lambda unrestricted opHigh : Nat .
                        (lambda unrestricted opReads : (family SM121LowerRegisters) .
                          (constructor SM121LowerOp SM121LowerOpValue opLow opHigh class opReads writes
                            guardReads predicateWrites zero opLabel opJoin opTarget))))))) in
                  (let unrestricted plain = (plainWith label join zero) in
                  (let unrestricted plainAfter = (plainWith zero zero zero) in
                  -- LDC / LDC.64 Rd, c[bank][offset], under the instruction's
                  -- guard, on the free scoreboard; it keeps the instruction's
                  -- waits: it writes the destination before the instruction did
                  (let unrestricted constantLoad =
                    (lambda unrestricted width : Nat .
                      (lambda unrestricted target : Nat .
                        (constructor SM121LowerOp SM121LowerOpValue
                          (nat-add
                            (nat-add
                              (nat-add sm121LowerOpcodeConstantLoad
                                (nat-multiply (sm121LowerField low sm121LowerGuardPlace sm121LowerGuardSpan) sm121LowerGuardPlace))
                              (nat-add (nat-multiply target sm121LowerDestinationPlace)
                                (nat-multiply sm121LowerZeroRegister sm121LowerSourcePlace)))
                            (nat-add
                              (nat-multiply
                                (sm121LowerField low sm121LowerConstantOffsetPlace sm121LowerConstantOffsetSpan)
                                sm121LowerConstantOffsetPlace)
                              (nat-multiply
                                (sm121LowerField low sm121LowerConstantBankPlace sm121LowerConstantBankSpan)
                                sm121LowerConstantBankPlace)))
                          (nat-add
                            (nat-add
                              (nat-add (nat-multiply width sm121LowerConstantWidthPlace)
                                (nat-multiply (sm121LowerMaximum 1 (sm121LowerStall high)) sm121LowerStallPlace))
                              (nat-add
                                (nat-multiply (sm121LowerField high sm121LowerYieldPlace sm121LowerYieldSpan)
                                  sm121LowerYieldPlace)
                                (nat-multiply barrier sm121LowerWriteBarrierPlace)))
                            (nat-add
                              (nat-multiply sm121LowerNoBarrier sm121LowerReadBarrierPlace)
                              (nat-multiply (sm121LowerField high sm121LowerWaitPlace sm121LowerWaitSpan)
                                sm121LowerWaitPlace)))
                          (constructor SM121LowerClass SM121LowerVariable)
                          sm121LowerNone
                          (sm121LowerSpan target (nat-subtract width sm121LowerConstantWidthBase))
                          guardReads
                          sm121LowerNone
                          1
                          label join zero))) in
                  (let unrestricted destination =
                    (sm121LowerField low sm121LowerDestinationPlace sm121LowerRegisterSpan) in
                    (eliminate
                      SM121LowerRule
                      (lambda unrestricted current : (family SM121LowerRule) . (family SM121LowerExpansion))
                      rule
                      (branch SM121LowerRuleSame . (sm121LowerOnly (plain low high reads)))
                      (branch SM121LowerRuleConstantMove .
                        (sm121LowerOnly (constantLoad sm121LowerConstantWord destination)))
                      (branch SM121LowerRuleConstantMultiplyAdd .
                        (let unrestricted clash = (sm121LowerHas reads destination) in
                        (let unrestricted target = (naturalSelect clash free destination) in
                          (sm121LowerSelectLazy (family SM121LowerExpansion)
                            (naturalAnd clash (naturalEqual free count))
                            (lambda unrestricted u : Nat . (sm121LowerRefuse (constructor SM121LowerRefusal SM121LowerRefusedNoFreeRegister)))
                            (lambda unrestricted u : Nat . (sm121LowerPair
                              (constantLoad sm121LowerConstantWord target)
                              (plainAfter
                                (sm121LowerWithField
                                  (sm121LowerRegisterForm low sm121LowerOpcodeMultiplyAddRegister)
                                  sm121LowerSecondSourcePlace sm121LowerRegisterSpan target)
                                high
                                (sm121LowerCons target reads))))))))
                      (branch SM121LowerRuleConstantMultiplyAddWide .
                        (let unrestricted clash =
                          (naturalOr (sm121LowerHas reads destination) (sm121LowerHas reads (succ destination))) in
                        (let unrestricted target = (naturalSelect clash freePair destination) in
                          (sm121LowerSelectLazy (family SM121LowerExpansion)
                            (naturalAnd clash (naturalEqual freePair (nat-subtract count 1)))
                            (lambda unrestricted u : Nat . (sm121LowerRefuse (constructor SM121LowerRefusal SM121LowerRefusedNoFreePair)))
                            (lambda unrestricted u : Nat . (sm121LowerPair
                              (constantLoad sm121LowerConstantPair target)
                              (plainAfter
                                (sm121LowerWithField
                                  (sm121LowerRegisterForm low sm121LowerOpcodeMultiplyAddWideRegister)
                                  sm121LowerSecondSourcePlace sm121LowerRegisterSpan
                                  (sm121LowerField high sm121LowerThirdSourcePlace sm121LowerRegisterSpan))
                                (sm121LowerWithField high sm121LowerThirdSourcePlace sm121LowerRegisterSpan target)
                                (sm121LowerCons target (sm121LowerCons (succ target) reads)))))))))
                      (branch SM121LowerRuleReduction .
                        (sm121LowerOnly
                          (plain
                            (sm121LowerWithField low 1 sm121LowerOpcodeSpan sm121LowerOpcodeReduceGlobal)
                            (sm121LowerWithField high 1 sm121LowerHalfWordSpan sm121LowerReduceGlobalModifiers)
                            reads)))
                      (branch SM121LowerRuleShared .
                        (sm121LowerRebaseShared low
                          (eliminate
                            SM121SharedAddressing
                            (lambda unrestricted current : (family SM121SharedAddressing) . Nat)
                            addressing
                            (branch SM121SharedAddressScaled . high)
                            (branch SM121SharedAddressBytes .
                              (sm121LowerWithField high sm121LowerSharedScalePlace 2 zero)))
                          plain reads))
                      (branch SM121LowerRuleSharedMatrix . (sm121LowerRebaseShared low high plain reads))
                      (branch SM121LowerRuleAsyncCopy .
                        (let unrestricted offset =
                          (nat-add sm121LowerSharedBase
                            (sm121LowerField low sm121LowerAsyncCopyOffsetPlace sm121LowerAsyncCopyOffsetSpan)) in
                          (sm121LowerSelect (family SM121LowerExpansion)
                            (nat-less-than offset sm121LowerAsyncCopyOffsetLimit)
                            (sm121LowerOnly
                              (plain
                                (sm121LowerWithField low sm121LowerAsyncCopyOffsetPlace sm121LowerAsyncCopyOffsetSpan offset)
                                (sm121LowerWithField
                                  (sm121LowerWithField
                                    (sm121LowerWithField high sm121LowerAsyncCopyDescriptorEnable 2 zero)
                                    sm121LowerAsyncCopyLegacyBit 2 zero)
                                  sm121LowerAsyncCopyDescriptorBit 2 1)
                                reads))
                            (sm121LowerRefuse (constructor SM121LowerRefusal SM121LowerRefusedSharedOffset)))))
                      -- a branch keeps its SM86 word until the schedule has placed
                      -- its target (sm121LowerRelocate re-encodes the
                      -- displacement); it is a join
                      (branch SM121LowerRuleBranch .
                        (let unrestricted targetMark = (sm121LowerBranchTarget low high index) in
                          (sm121LowerSelect (family SM121LowerExpansion)
                            (naturalIsZero targetMark)
                            (sm121LowerRefuse (constructor SM121LowerRefusal SM121LowerRefusedBranchTarget))
                            (sm121LowerOnly (plainWith label 1 targetMark low high reads)))))))))))))))))))))))

-- ---------------------------------------------------------------------------
-- The schedule

def sm121LowerPendingNone : (family SM121LowerPending) =
  (constructor SM121LowerPending SM121LowerPendingEnd)

def sm121LowerPendingRemove =
  (lambda unrestricted pending : (family SM121LowerPending) .
    (lambda unrestricted register : Nat .
      (eliminate
        SM121LowerPending
        (lambda unrestricted current : (family SM121LowerPending) . (family SM121LowerPending))
        pending
        (branch SM121LowerPendingEnd . sm121LowerPendingNone)
        (branch SM121LowerPendingNext entry time class tail induction .
          (nat-eliminate (lambda unrestricted current : Nat . (family SM121LowerPending))
            (constructor SM121LowerPending SM121LowerPendingNext entry time class induction)
            (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : (family SM121LowerPending) . induction))
            (nat-eliminate (lambda unrestricted current : Nat . Nat)
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                (succ zero)
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . zero))
                (nat-less-than register entry))
              (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . zero))
              (nat-less-than entry register)))))))

-- entries that can still delay an instruction issued at `now` or later
def sm121LowerPendingPrune =
  (lambda unrestricted pending : (family SM121LowerPending) .
    (lambda unrestricted now : Nat .
      (lambda unrestricted latency : Nat .
        (eliminate
          SM121LowerPending
          (lambda unrestricted current : (family SM121LowerPending) . (family SM121LowerPending))
          pending
          (branch SM121LowerPendingEnd . sm121LowerPendingNone)
          (branch SM121LowerPendingNext entry time class tail induction .
            (sm121LowerSelect (family SM121LowerPending)
              (nat-less-than now (nat-add time latency))
              (constructor SM121LowerPending SM121LowerPendingNext entry time class induction)
              induction))))))

-- the cycle `register` is ready at for the reader whose row `row` holds (0
-- when nothing is pending). The row is dispatched once per scheduling step
-- (sm121LowerReaderRowOf); the taken arm alone extracts its writer's entry.
def sm121LowerReady =
  (lambda unrestricted pending : (family SM121LowerPending) .
    (lambda unrestricted row : Nat .
      (lambda unrestricted register : Nat .
        (eliminate
          SM121LowerPending
          (lambda unrestricted current : (family SM121LowerPending) . Nat)
          pending
          (branch SM121LowerPendingEnd . zero)
          (branch SM121LowerPendingNext entry time class tail induction .
            (nat-eliminate (lambda unrestricted current : Nat . Nat)
              induction
              (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (nat-add time (sm121LowerLatencyAt row (sm121LowerClassTag class)))))
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                (nat-eliminate (lambda unrestricted current : Nat . Nat)
                  (succ zero)
                  (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . zero))
                  (nat-less-than register entry))
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . zero))
                (nat-less-than entry register))))))))

def sm121LowerPredicateReady =
  (lambda unrestricted pending : (family SM121LowerPending) .
    (lambda unrestricted predicate : Nat .
      (eliminate
        SM121LowerPending
        (lambda unrestricted current : (family SM121LowerPending) . Nat)
        pending
        (branch SM121LowerPendingEnd . zero)
        (branch SM121LowerPendingNext entry time class tail induction .
          (nat-eliminate (lambda unrestricted current : Nat . Nat)
            induction
            (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (nat-add time sm121LowerPredicateLatency)))
            (nat-eliminate (lambda unrestricted current : Nat . Nat)
              (nat-eliminate (lambda unrestricted current : Nat . Nat)
                (succ zero)
                (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . zero))
                (nat-less-than predicate entry))
              (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . zero))
              (nat-less-than entry predicate)))))))

-- the latest of `ready` over `registers`, and `start`
def sm121LowerLatest =
  (lambda unrestricted registers : (family SM121LowerRegisters) .
    (lambda unrestricted ready : (pi unrestricted register : Nat . Nat) .
      (lambda unrestricted start : Nat .
        (eliminate
          SM121LowerRegisters
          (lambda unrestricted current : (family SM121LowerRegisters) . Nat)
          registers
          (branch SM121LowerRegistersEnd . start)
          (branch SM121LowerRegistersNext head tail induction .
            (sm121LowerMaximum (ready head) induction))))))

-- each register of `registers` pending from `time` with `class`, or (for a
-- variable-latency writer, timed by a scoreboard) no longer pending
def sm121LowerPendingWrite =
  (lambda unrestricted pending : (family SM121LowerPending) .
    (lambda unrestricted registers : (family SM121LowerRegisters) .
      (lambda unrestricted time : Nat .
        (lambda unrestricted class : (family SM121LowerClass) .
          (lambda unrestricted fixed : Nat .
            (eliminate
              SM121LowerRegisters
              (lambda unrestricted current : (family SM121LowerRegisters) . (family SM121LowerPending))
              registers
              (branch SM121LowerRegistersEnd . pending)
              (branch SM121LowerRegistersNext head tail induction .
                (sm121LowerSelect (family SM121LowerPending)
                  fixed
                  (constructor SM121LowerPending SM121LowerPendingNext head time class
                    (sm121LowerPendingRemove induction head))
                  (sm121LowerPendingRemove induction head)))))))))

def sm121LowerPush =
  (lambda unrestricted op : (family SM121LowerOp) .
    (lambda unrestricted rest : (family SM121LowerOps) .
      (constructor SM121LowerOps SM121LowerOpsNext op rest)))

-- `op` with its high word changed
def sm121LowerWithHigh =
  (lambda unrestricted change : (pi unrestricted high : Nat . Nat) .
    (lambda unrestricted op : (family SM121LowerOp) .
      (eliminate
        SM121LowerOp
        (lambda unrestricted current : (family SM121LowerOp) . (family SM121LowerOp))
        op
        (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
          (constructor SM121LowerOp SM121LowerOpValue low (change high) class reads writes
            predicateReads predicateWrites constantLoad label join target)))))

def sm121LowerHighOf =
  (lambda unrestricted op : (family SM121LowerOp) .
    (eliminate
      SM121LowerOp
      (lambda unrestricted current : (family SM121LowerOp) . Nat)
      op
      (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target . high)))

def sm121LowerNoOperation =
  (lambda unrestricted stall : Nat .
    (constructor SM121LowerOp SM121LowerOpValue
      (sm121LowerAlwaysOpcode sm121LowerOpcodeNoOperation)
      (sm121LowerControl stall zero)
      (constructor SM121LowerClass SM121LowerNoResult)
      sm121LowerNone sm121LowerNone sm121LowerNone sm121LowerNone zero zero zero zero))

-- Delay the next instruction by `deficit` cycles: the last one placed stalls
-- longer, up to fifteen cycles, and NOPs carry the rest.
def sm121LowerDelay =
  (lambda unrestricted deficit : Nat .
    (lambda unrestricted placed : (family SM121LowerOps) .
      (eliminate
        SM121LowerOps
        (lambda unrestricted current : (family SM121LowerOps) . (family SM121LowerOps))
        placed
        (branch SM121LowerOpsEnd . placed)
        (branch SM121LowerOpsNext last earlier induction .
          (let unrestricted stall = (sm121LowerStall (sm121LowerHighOf last)) in
          (let unrestricted grown = (sm121LowerMinimum sm121LowerStallLongest (nat-add stall deficit)) in
          (let unrestricted remaining = (nat-subtract deficit (nat-subtract grown stall)) in
          (let unrestricted previous =
            (sm121LowerPush
              (sm121LowerWithHigh
                (lambda unrestricted high : Nat .
                  (naturalSelect
                    (nat-less-than zero remaining)
                    (sm121LowerWithoutReuse (sm121LowerWithStall grown high))
                    (sm121LowerWithStall grown high)))
                last)
              earlier) in
            (nat-eliminate
              (lambda unrestricted current : Nat . (family SM121LowerOps))
              previous
              (lambda unrestricted index : Nat .
                (lambda unrestricted induction : (family SM121LowerOps) .
                  (sm121LowerPush
                    (sm121LowerNoOperation
                      (sm121LowerMinimum sm121LowerStallLongest
                        (nat-subtract remaining (nat-multiply index sm121LowerStallLongest))))
                    induction)))
              (nat-divide
                (nat-add remaining (nat-subtract sm121LowerStallLongest 1))
                sm121LowerStallLongest))))))))))

-- the operand-reuse flags of the last instruction placed cleared: an LDC
-- follows it
def sm121LowerDropReuse =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (eliminate
      SM121LowerOps
      (lambda unrestricted current : (family SM121LowerOps) . (family SM121LowerOps))
      placed
      (branch SM121LowerOpsEnd . placed)
      (branch SM121LowerOpsNext last earlier induction .
        (sm121LowerPush (sm121LowerWithHigh sm121LowerWithoutReuse last) earlier))))

def sm121LowerHasPlaced =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (eliminate
      SM121LowerOps
      (lambda unrestricted current : (family SM121LowerOps) . Nat)
      placed
      (branch SM121LowerOpsEnd . zero)
      (branch SM121LowerOpsNext last earlier induction . 1)))

-- The latest cycle any of `pending` is ready at for its slowest reader
-- (`latency` of the writer's class after it issued), and `start`.
def sm121LowerDrained =
  (lambda unrestricted pending : (family SM121LowerPending) .
    (lambda unrestricted latency : (pi unrestricted writer : (family SM121LowerClass) . Nat) .
      (lambda unrestricted start : Nat .
        (eliminate
          SM121LowerPending
          (lambda unrestricted current : (family SM121LowerPending) . Nat)
          pending
          (branch SM121LowerPendingEnd . start)
          (branch SM121LowerPendingNext entry issued class tail induction .
            (sm121LowerMaximum (nat-add issued (latency class)) induction))))))

-- the cycles until a result of `writer`'s class is ready for every reader
def sm121LowerWriterLongest =
  (lambda unrestricted writer : (family SM121LowerClass) .
    (sm121LowerOverClasses
      (lambda unrestricted reader : (family SM121LowerClass) . (sm121LowerReadAfterWrite writer reader))))

-- a predicate is ready for its guard after the predicate latency
def sm121LowerPredicateWait =
  (lambda unrestricted writer : (family SM121LowerClass) . sm121LowerPredicateLatency)

-- `high` waiting on every scoreboard `mask` names as well as those it waits on
def sm121LowerWithWaits =
  (lambda unrestricted mask : Nat .
    (lambda unrestricted high : Nat .
      (nat-eliminate
        (lambda unrestricted current : Nat . Nat)
        high
        (lambda unrestricted barrier : Nat .
          (lambda unrestricted induction : Nat .
            (naturalSelect (sm121LowerField mask (sm121LowerPlace barrier) 2)
              (sm121LowerWithWait (sm121LowerPlace (nat-add sm121LowerWaitBit barrier)) induction)
              induction)))
        sm121LowerBarrierCount)))

-- One instruction placed.
def sm121LowerStep =
  (lambda unrestricted barrier : Nat .
  (lambda unrestricted waitAll : Nat .
    (lambda unrestricted waitPlace : Nat .
      (lambda unrestricted schedule : (family SM121LowerSchedule) .
        (lambda unrestricted op : (family SM121LowerOp) .
          (eliminate
            SM121LowerSchedule
            (lambda unrestricted current : (family SM121LowerSchedule) . (family SM121LowerSchedule))
            schedule
            (branch SM121LowerScheduleValue time registers predicates loads guardedReads placed .
              (eliminate
                SM121LowerOp
                (lambda unrestricted current : (family SM121LowerOp) . (family SM121LowerSchedule))
                op
                (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
                  (let unrestricted before =
                    (sm121LowerSelectLazy (family SM121LowerOps) constantLoad
                      (lambda unrestricted u : Nat . (sm121LowerDropReuse placed))
                      (lambda unrestricted u : Nat . placed)) in
                  -- this step's reader row, dispatched once for every cell below
                  (let unrestricted row = (sm121LowerReaderRowOf class) in
                  (let unrestricted operandsReady =
                    (sm121LowerLatest reads (sm121LowerReady registers row)
                      (sm121LowerLatest writes (sm121LowerReady registers sm121LowerRowAlu)
                        (sm121LowerLatest
                          (sm121LowerAppend predicateReads predicateWrites)
                          (sm121LowerPredicateReady predicates)
                          time))) in
                  (let unrestricted required =
                    (sm121LowerSelectLazy Nat join
                      (lambda unrestricted u : Nat . (sm121LowerMaximum operandsReady
                        (sm121LowerDrained registers sm121LowerWriterLongest
                          (sm121LowerDrained predicates sm121LowerPredicateWait time))))
                      (lambda unrestricted u : Nat . operandsReady)) in
                  (let unrestricted deficit = (nat-subtract required time) in
                  (let unrestricted issue =
                    (naturalSelect (sm121LowerHasPlaced before) deficit zero) in
                  (let unrestricted waits =
                    (naturalOr join
                    (naturalOr
                      (sm121LowerAny (sm121LowerAppend reads writes) loads)
                      (sm121LowerAny writes guardedReads))) in
                  (let unrestricted waited =
                    (sm121LowerSelectLazy Nat join
                      (lambda unrestricted u : Nat . (sm121LowerWithWaits waitAll high))
                      (lambda unrestricted u : Nat . (naturalSelect waits (sm121LowerWithWait waitPlace high) high))) in
                  (let unrestricted fixed = (sm121LowerFixed class) in
                  (let unrestricted unguarded =
                    (naturalAnd
                      (naturalAnd (naturalIsZero fixed) (sm121LowerHasResult class))
                      (naturalAnd
                        (naturalIsZero constantLoad)
                        (naturalAnd
                          (naturalEqual
                            (sm121LowerField waited sm121LowerReadBarrierPlace sm121LowerBarrierSpan)
                            sm121LowerNoBarrier)
                          (naturalIsZero (sm121LowerIsEmpty reads))))) in
                  (let unrestricted placedHigh =
                    (naturalSelect unguarded (sm121LowerWithReadBarrier barrier waited) waited) in
                  (let unrestricted loadsKept =
                    (sm121LowerSelect (family SM121LowerRegisters) waits sm121LowerNone loads) in
                  (let unrestricted readsKept =
                    (sm121LowerSelect (family SM121LowerRegisters) waits sm121LowerNone guardedReads) in
                  (let unrestricted issued = (nat-add time issue) in
                  (let unrestricted next = (nat-add issued (sm121LowerStall placedHigh)) in
                    (sm121LowerSelect (family SM121LowerSchedule)
                      (naturalIsZero (sm121LowerStall placedHigh))
                      (constructor SM121LowerSchedule SM121LowerScheduleRefused (constructor SM121LowerRefusal SM121LowerRefusedNoStall))
                      (constructor SM121LowerSchedule SM121LowerScheduleValue
                        next
                        (sm121LowerPendingPrune
                          (sm121LowerPendingWrite registers writes issued class fixed)
                          next
                          sm121LowerLongestReadAfterWrite)
                        (sm121LowerPendingPrune
                          (sm121LowerPendingWrite predicates predicateWrites issued (constructor SM121LowerClass SM121LowerDualAlu) 1)
                          next
                          sm121LowerPredicateLatency)
                        (sm121LowerSelectLazy (family SM121LowerRegisters) constantLoad
                          (lambda unrestricted u : Nat . (sm121LowerUnion loadsKept writes))
                          (lambda unrestricted u : Nat . loadsKept))
                        (sm121LowerSelectLazy (family SM121LowerRegisters) unguarded
                          (lambda unrestricted u : Nat . (sm121LowerUnion readsKept reads))
                          (lambda unrestricted u : Nat . readsKept))
                        (sm121LowerPush
                          (constructor SM121LowerOp SM121LowerOpValue low placedHigh class reads writes
                            predicateReads predicateWrites constantLoad label join target)
                          (sm121LowerDelay deficit before))))))))))))))))))))))
            (branch SM121LowerScheduleRefused refusal . schedule)))))))

def sm121LowerStepAll =
  (lambda unrestricted barrier : Nat .
  (lambda unrestricted waitAll : Nat .
    (lambda unrestricted waitPlace : Nat .
      (lambda unrestricted schedule : (family SM121LowerSchedule) .
        (lambda unrestricted ops : (family SM121LowerOps) .
          (app
            (eliminate
              SM121LowerOps
              (lambda unrestricted current : (family SM121LowerOps) .
                (pi unrestricted carried : (family SM121LowerSchedule) . (family SM121LowerSchedule)))
              ops
              (branch SM121LowerOpsEnd .
                (lambda unrestricted carried : (family SM121LowerSchedule) . carried))
              (branch SM121LowerOpsNext op tail induction .
                (lambda unrestricted carried : (family SM121LowerSchedule) .
                  (induction (sm121LowerStep barrier waitAll waitPlace carried op)))))
            schedule))))))

-- UR4:UR5 = 0, the global-memory descriptor, before the first access
-- (uniform move to a memory reader: two cycles, plus padding).
def sm121LowerPrologue : (family SM121LowerOps) =
  (constructor SM121LowerOps SM121LowerOpsNext
    (constructor SM121LowerOp SM121LowerOpValue
      (nat-add (sm121LowerAlwaysOpcode sm121LowerOpcodeUniformMove)
        (nat-multiply sm121LowerDescriptorLow sm121LowerDestinationPlace))
      (sm121LowerControl 1 1)
      (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone sm121LowerNone sm121LowerNone zero zero zero zero)
    (constructor SM121LowerOps SM121LowerOpsNext
      (constructor SM121LowerOp SM121LowerOpValue
        (nat-add (sm121LowerAlwaysOpcode sm121LowerOpcodeUniformMove)
          (nat-multiply sm121LowerDescriptorHigh sm121LowerDestinationPlace))
        (sm121LowerControl 3 1)
        (constructor SM121LowerClass SM121LowerNoResult) sm121LowerNone sm121LowerNone sm121LowerNone sm121LowerNone zero zero zero zero)
      (constructor SM121LowerOps SM121LowerOpsEnd)))

def sm121LowerScheduleStart : (family SM121LowerSchedule) =
  (constructor SM121LowerSchedule SM121LowerScheduleValue
    zero sm121LowerPendingNone sm121LowerPendingNone sm121LowerNone sm121LowerNone
    (constructor SM121LowerOps SM121LowerOpsEnd))

-- ---------------------------------------------------------------------------
-- The program

def sm121LowerLittleEndianBytes =
  (lambda unrestricted count : Nat .
    (nat-eliminate
      (lambda unrestricted current : Nat . (pi unrestricted value : Nat . Bytes))
      (lambda unrestricted value : Nat . b"")
      (lambda unrestricted predecessor : Nat .
        (lambda unrestricted induction : (pi unrestricted value : Nat . Bytes) .
          (lambda unrestricted value : Nat .
            (bytes-cons
              (nat-to-byte (nat-modulo value sm121LowerRegisterSpan))
              (induction (nat-divide value sm121LowerRegisterSpan))))))
      count))

def sm121LowerWordsBytes =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (bytes-builder-build
      (eliminate
        SM121LowerOps
        (lambda unrestricted current : (family SM121LowerOps) . BytesBuilder)
        placed
        (branch SM121LowerOpsEnd . (bytes-builder-empty))
        (branch SM121LowerOpsNext last earlier induction .
          (eliminate
            SM121LowerOp
            (lambda unrestricted current : (family SM121LowerOp) . BytesBuilder)
            last
            (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
          (bytes-builder-append
            induction
            (bytes-builder-chunk
              (bytes-append
                (sm121LowerLittleEndianBytes sm121LowerWordBytes low)
                (sm121LowerLittleEndianBytes sm121LowerWordBytes high))))))))))

-- `decoded` folded with each item's SM86 index.
def sm121LowerDecodedFoldIndexed =
  (lambda erased result : Type 0 .
    (lambda unrestricted decoded : (family SM121LowerDecodedList) .
      (lambda unrestricted start : result .
        (lambda unrestricted step :
            (pi unrestricted accumulated : result .
              (pi unrestricted index : Nat .
                (pi unrestricted item : (family SM121LowerDecoded) . result))) .
          (app (app
            (eliminate
              SM121LowerDecodedList
              (lambda unrestricted current : (family SM121LowerDecodedList) .
                (pi unrestricted accumulated : result . (pi unrestricted index : Nat . result)))
              decoded
              (branch SM121LowerDecodedEnd .
                (lambda unrestricted accumulated : result . (lambda unrestricted index : Nat . accumulated)))
              (branch SM121LowerDecodedNext item tail induction .
                (lambda unrestricted accumulated : result . (lambda unrestricted index : Nat .
                  (induction (step accumulated index item) (succ index)))))
              (branch SM121LowerDecodedRefused refusal .
                (lambda unrestricted accumulated : result . (lambda unrestricted index : Nat . accumulated))))
            start) zero)))))

def sm121LowerIsBranchRule =
  (lambda unrestricted rule : (family SM121LowerRule) .
    (eliminate
      SM121LowerRule
      (lambda unrestricted current : (family SM121LowerRule) . Nat)
      rule
      (branch SM121LowerRuleSame . zero)
      (branch SM121LowerRuleConstantMove . zero)
      (branch SM121LowerRuleConstantMultiplyAdd . zero)
      (branch SM121LowerRuleConstantMultiplyAddWide . zero)
      (branch SM121LowerRuleReduction . zero)
      (branch SM121LowerRuleShared . zero)
      (branch SM121LowerRuleSharedMatrix . zero)
      (branch SM121LowerRuleAsyncCopy . zero)
      (branch SM121LowerRuleBranch . 1)))

-- The SM86 indices branches in `decoded` jump to.
def sm121LowerBranchTargets =
  (lambda unrestricted decoded : (family SM121LowerDecodedList) .
    (sm121LowerDecodedFoldIndexed (family SM121LowerRegisters) decoded sm121LowerNone
      (lambda unrestricted targets : (family SM121LowerRegisters) .
        (lambda unrestricted index : Nat .
          (lambda unrestricted item : (family SM121LowerDecoded) .
            (eliminate
              SM121LowerDecoded
              (lambda unrestricted current : (family SM121LowerDecoded) . (family SM121LowerRegisters))
              item
              (branch SM121LowerDecodedValue low high guardReads shape .
                (eliminate
                  SM121LowerShape
                  (lambda unrestricted current : (family SM121LowerShape) . (family SM121LowerRegisters))
                  shape
                  (branch SM121LowerShapeValue rule class reads writes predicateWrites .
                    (let unrestricted mark = (sm121LowerBranchTarget low high index) in
                      (sm121LowerSelect (family SM121LowerRegisters)
                        (naturalAnd (sm121LowerIsBranchRule rule) (naturalIsZero (naturalIsZero mark)))
                        (sm121LowerCons (nat-subtract mark 1) targets)
                        targets)))))))))))

def sm121LowerScheduleProgram =
  (lambda unrestricted context : (family SM121LowerContext) .
    (lambda unrestricted decoded : (family SM121LowerDecodedList) .
      (eliminate
        SM121LowerContext
        (lambda unrestricted current : (family SM121LowerContext) . (family SM121LowerSchedule))
        context
        (branch SM121LowerContextValue barrier free freePair count addressing waitAll .
          (let unrestricted waitPlace = (sm121LowerPlace (nat-add sm121LowerWaitBit barrier)) in
          (let unrestricted targets = (sm121LowerBranchTargets decoded) in
            (sm121LowerDecodedFoldIndexed (family SM121LowerSchedule) decoded
              (sm121LowerStepAll barrier waitAll waitPlace sm121LowerScheduleStart sm121LowerPrologue)
              (lambda unrestricted schedule : (family SM121LowerSchedule) .
                (lambda unrestricted index : Nat .
                  (lambda unrestricted item : (family SM121LowerDecoded) .
                    (eliminate
                      SM121LowerExpansion
                      (lambda unrestricted current : (family SM121LowerExpansion) . (family SM121LowerSchedule))
                      (sm121LowerExpand context index targets item)
                      (branch SM121LowerExpanded ops . (sm121LowerStepAll barrier waitAll waitPlace schedule ops))
                      (branch SM121LowerExpansionRefused refusal .
                        (constructor SM121LowerSchedule SM121LowerScheduleRefused refusal)))))))))))))

-- ---------------------------------------------------------------------------
-- The static scoreboard check
--
-- A lowered program is walked again from its final words alone -- the issue
-- cycle of each instruction from the stalls before it, the scoreboards from
-- its control fields -- and every hazard counted:
--   * a register read (or written) before a fixed-latency writer's result is
--     ready for it, or a guard before its predicate is;
--   * a register read or written while a variable-latency result for it is
--     on a scoreboard the instruction has not waited on;
--   * a register written while an instruction may still be reading it (a
--     read barrier not waited on, or an unguarded late read: never safe);
--   * a variable-latency result signalling no scoreboard.
-- (scripts/alpha/sbcheck.py over nvdisasm's text was the first form of it.)

-- an asynchronous read no read barrier guards
def sm121LowerUnguarded : Nat = 9

def sm121LowerBoardsNone : (family SM121LowerBoards) =
  (constructor SM121LowerBoards SM121LowerBoardsEnd)

def sm121LowerBoardsHas =
  (lambda unrestricted boards : (family SM121LowerBoards) .
    (lambda unrestricted register : Nat .
      (eliminate
        SM121LowerBoards
        (lambda unrestricted current : (family SM121LowerBoards) . Nat)
        boards
        (branch SM121LowerBoardsEnd . zero)
        (branch SM121LowerBoardsNext entry barrier tail induction .
          (nat-eliminate (lambda unrestricted current : Nat . Nat)
            induction
            (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (succ zero)))
            (naturalEqual entry register))))))

-- entries whose barrier the wait mask names retire
def sm121LowerBoardsWait =
  (lambda unrestricted boards : (family SM121LowerBoards) .
    (lambda unrestricted mask : Nat .
      (eliminate
        SM121LowerBoards
        (lambda unrestricted current : (family SM121LowerBoards) . (family SM121LowerBoards))
        boards
        (branch SM121LowerBoardsEnd . sm121LowerBoardsNone)
        (branch SM121LowerBoardsNext entry barrier tail induction .
          (sm121LowerSelect (family SM121LowerBoards)
            (naturalAnd
              (nat-less-than barrier sm121LowerBarrierCount)
              (sm121LowerField mask (sm121LowerPlace barrier) 2))
            induction
            (constructor SM121LowerBoards SM121LowerBoardsNext entry barrier induction))))))

def sm121LowerBoardsRemove =
  (lambda unrestricted boards : (family SM121LowerBoards) .
    (lambda unrestricted register : Nat .
      (eliminate
        SM121LowerBoards
        (lambda unrestricted current : (family SM121LowerBoards) . (family SM121LowerBoards))
        boards
        (branch SM121LowerBoardsEnd . sm121LowerBoardsNone)
        (branch SM121LowerBoardsNext entry barrier tail induction .
          (nat-eliminate (lambda unrestricted current : Nat . (family SM121LowerBoards))
            (constructor SM121LowerBoards SM121LowerBoardsNext entry barrier induction)
            (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : (family SM121LowerBoards) . induction))
            (naturalEqual entry register))))))

-- each of `registers` now on `barrier` (its earlier entry replaced)
def sm121LowerBoardsAdd =
  (lambda unrestricted boards : (family SM121LowerBoards) .
    (lambda unrestricted registers : (family SM121LowerRegisters) .
      (lambda unrestricted barrier : Nat .
        (eliminate
          SM121LowerRegisters
          (lambda unrestricted current : (family SM121LowerRegisters) . (family SM121LowerBoards))
          registers
          (branch SM121LowerRegistersEnd . boards)
          (branch SM121LowerRegistersNext head tail induction .
            (constructor SM121LowerBoards SM121LowerBoardsNext head barrier
              (sm121LowerBoardsRemove induction head)))))))

-- how many of `registers` `hazard` holds for
def sm121LowerCount =
  (lambda unrestricted registers : (family SM121LowerRegisters) .
    (lambda unrestricted hazard : (pi unrestricted register : Nat . Nat) .
      (eliminate
        SM121LowerRegisters
        (lambda unrestricted current : (family SM121LowerRegisters) . Nat)
        registers
        (branch SM121LowerRegistersEnd . zero)
        (branch SM121LowerRegistersNext head tail induction .
          (nat-add (nat-eliminate (lambda unrestricted current : Nat . Nat)
                     zero
                     (lambda unrestricted flagPredecessor : Nat . (lambda unrestricted flagRest : Nat . (succ zero)))
                     (hazard head)) induction)))))

def sm121LowerTraceStart : (family SM121LowerTrace) =
  (constructor SM121LowerTrace SM121LowerTraceValue
    zero sm121LowerPendingNone sm121LowerPendingNone sm121LowerBoardsNone sm121LowerBoardsNone zero zero)

def sm121LowerBoardsEmpty =
  (lambda unrestricted boards : (family SM121LowerBoards) .
    (eliminate
      SM121LowerBoards
      (lambda unrestricted current : (family SM121LowerBoards) . Nat)
      boards
      (branch SM121LowerBoardsEnd . 1)
      (branch SM121LowerBoardsNext entry barrier tail induction . zero)))

def sm121LowerCheckStep =
  (lambda unrestricted joins : (family SM121LowerRegisters) .
  (lambda unrestricted trace : (family SM121LowerTrace) .
    (lambda unrestricted op : (family SM121LowerOp) .
      (eliminate
        SM121LowerTrace
        (lambda unrestricted current : (family SM121LowerTrace) . (family SM121LowerTrace))
        trace
        (branch SM121LowerTraceValue time fixedResults predicates written reading hazards position .
          (eliminate
            SM121LowerOp
            (lambda unrestricted current : (family SM121LowerOp) . (family SM121LowerTrace))
            op
            (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
              (let unrestricted mask = (sm121LowerField high sm121LowerWaitPlace sm121LowerWaitSpan) in
              (let unrestricted writtenNow = (sm121LowerBoardsWait written mask) in
              (let unrestricted readingNow = (sm121LowerBoardsWait reading mask) in
              (let unrestricted fixed = (sm121LowerFixed class) in
              (let unrestricted row = (sm121LowerReaderRowOf class) in
              (let unrestricted early =
                (nat-add
                  (sm121LowerCount reads
                    (lambda unrestricted register : Nat .
                      (nat-less-than time (sm121LowerReady fixedResults row register))))
                  (nat-add
                    (sm121LowerCount writes
                      (lambda unrestricted register : Nat .
                        (nat-less-than time
                          (sm121LowerReady fixedResults sm121LowerRowAlu register))))
                    (sm121LowerCount (sm121LowerAppend predicateReads predicateWrites)
                      (lambda unrestricted predicate : Nat .
                        (nat-less-than time (sm121LowerPredicateReady predicates predicate)))))) in
              (let unrestricted unwaited =
                (nat-add
                  (sm121LowerCount (sm121LowerAppend reads writes) (sm121LowerBoardsHas writtenNow))
                  (sm121LowerCount writes (sm121LowerBoardsHas readingNow))) in
              (let unrestricted variableWrite =
                (naturalAnd (naturalIsZero fixed) (naturalIsZero (sm121LowerIsEmpty writes))) in
              (let unrestricted writeBarrier =
                (sm121LowerField high sm121LowerWriteBarrierPlace sm121LowerBarrierSpan) in
              (let unrestricted readBarrier =
                (sm121LowerField high sm121LowerReadBarrierPlace sm121LowerBarrierSpan) in
              (let unrestricted lateRead =
                (naturalAnd
                  (naturalAnd (naturalIsZero fixed) (sm121LowerHasResult class))
                  (naturalAnd (naturalIsZero constantLoad) (naturalIsZero (sm121LowerIsEmpty reads)))) in
                (constructor SM121LowerTrace SM121LowerTraceValue
                  (nat-add time (sm121LowerStall high))
                  (sm121LowerPendingPrune
                    (sm121LowerPendingWrite fixedResults writes time class fixed)
                    (nat-add time (sm121LowerStall high))
                    sm121LowerLongestReadAfterWrite)
                  (sm121LowerPendingPrune
                    (sm121LowerPendingWrite predicates predicateWrites time
                      (constructor SM121LowerClass SM121LowerDualAlu) 1)
                    (nat-add time (sm121LowerStall high))
                    sm121LowerPredicateLatency)
                  (sm121LowerSelect (family SM121LowerBoards) variableWrite
                    (sm121LowerBoardsAdd writtenNow writes writeBarrier)
                    writtenNow)
                  (sm121LowerSelect (family SM121LowerBoards) lateRead
                    (sm121LowerBoardsAdd readingNow reads
                      (naturalSelect (naturalEqual readBarrier sm121LowerNoBarrier)
                        sm121LowerUnguarded readBarrier))
                    readingNow)
                  (nat-add hazards
                    (nat-add
                      (nat-add (nat-add early unwaited)
                        (naturalSelect
                          (naturalAnd variableWrite (naturalEqual writeBarrier sm121LowerNoBarrier))
                          1 zero))
                      -- a join: nothing in flight once its waits retire
                      (sm121LowerSelectLazy Nat (sm121LowerHas joins position)
                        (lambda unrestricted u : Nat .
                          (nat-add
                            (nat-add
                              (naturalSelect (nat-less-than time (sm121LowerDrained fixedResults sm121LowerWriterLongest zero)) 1 zero)
                              (naturalSelect (nat-less-than time (sm121LowerDrained predicates sm121LowerPredicateWait zero)) 1 zero))
                            (nat-add
                              (naturalSelect (sm121LowerBoardsEmpty writtenNow) zero 1)
                              (naturalSelect (sm121LowerBoardsEmpty readingNow) zero 1))))
                        (lambda unrestricted u : Nat . zero))))
                  (succ position))))))))))))))))))))

def sm121LowerOpcodeBranch : Nat = 0x947

-- The joins of `placed` (the last op first), from its words alone: a BRA
-- and the position its sm_121 displacement lands on.
def sm121LowerJoinsOf =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (eliminate
      SM121LowerOps
      (lambda unrestricted current : (family SM121LowerOps) . (family SM121LowerJoins))
      placed
      (branch SM121LowerOpsEnd .
        (constructor SM121LowerJoins SM121LowerJoinsValue zero sm121LowerNone zero))
      (branch SM121LowerOpsNext last earlier induction .
        (eliminate
          SM121LowerJoins
          (lambda unrestricted current : (family SM121LowerJoins) . (family SM121LowerJoins))
          induction
          (branch SM121LowerJoinsValue position joins bad .
            (eliminate
              SM121LowerOp
              (lambda unrestricted current : (family SM121LowerOp) . (family SM121LowerJoins))
              last
              (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
                (let unrestricted encoded =
                  (nat-add
                    (nat-add
                      (sm121LowerField low sm121LowerBranchLowPlace sm121LowerRegisterSpan)
                      (nat-multiply (sm121LowerField low sm121LowerBranchMiddlePlace sm121LowerBranchMiddleSpan)
                        sm121LowerRegisterSpan))
                    (nat-multiply (sm121LowerField high 1 sm121LowerBranchHighSpan)
                      (nat-multiply sm121LowerRegisterSpan sm121LowerBranchMiddleSpan))) in
                (let unrestricted next = (succ position) in
                (let unrestricted forward = (nat-less-than encoded (nat-divide sm121LowerBranchWordSpan 2)) in
                (let unrestricted words =
                  (naturalSelect forward encoded (nat-subtract sm121LowerBranchWordSpan encoded)) in
                (let unrestricted aligned =
                  (naturalIsZero (nat-modulo words sm121LowerBranchUnit)) in
                (let unrestricted distance = (nat-divide words sm121LowerBranchUnit) in
                (let unrestricted inside = (naturalOr forward (naturalIsZero (nat-less-than next distance))) in
                (let unrestricted landing =
                  (naturalSelect forward (nat-add next distance) (nat-subtract next distance)) in
                  (sm121LowerSelect (family SM121LowerJoins)
                    (naturalEqual (nat-modulo low sm121LowerOpcodeSpan) sm121LowerOpcodeBranch)
                    (constructor SM121LowerJoins SM121LowerJoinsValue
                      next
                      (sm121LowerCons position (sm121LowerCons landing joins))
                      (naturalSelect (naturalAnd aligned inside) bad (succ bad)))
                    (constructor SM121LowerJoins SM121LowerJoinsValue next joins bad)))))))))))))))))

-- how many of `positions` lie at or past `count`
def sm121LowerBeyond =
  (lambda unrestricted positions : (family SM121LowerRegisters) .
    (lambda unrestricted count : Nat .
      (sm121LowerCount positions
        (lambda unrestricted position : Nat . (naturalIsZero (nat-less-than position count))))))

-- The hazards in `placed` (the last instruction first).
def sm121LowerHazards =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (eliminate
      SM121LowerJoins
      (lambda unrestricted current : (family SM121LowerJoins) . Nat)
      (sm121LowerJoinsOf placed)
      (branch SM121LowerJoinsValue count joins bad .
        (nat-add (nat-add bad (sm121LowerBeyond joins count))
          (eliminate
            SM121LowerTrace
            (lambda unrestricted current : (family SM121LowerTrace) . Nat)
            (app
              (eliminate
                SM121LowerOps
                (lambda unrestricted current : (family SM121LowerOps) .
                  (pi unrestricted trace : (family SM121LowerTrace) . (family SM121LowerTrace)))
                placed
                (branch SM121LowerOpsEnd . (lambda unrestricted trace : (family SM121LowerTrace) . trace))
                (branch SM121LowerOpsNext last earlier induction .
                  (lambda unrestricted trace : (family SM121LowerTrace) .
                    (sm121LowerCheckStep joins (induction trace) last))))
              sm121LowerTraceStart)
            (branch SM121LowerTraceValue time fixedResults predicates written reading hazards position . hazards))))))

-- ---------------------------------------------------------------------------
-- Relocation: branches re-encoded against where their targets were placed

-- the placed program's length and where each op a branch lands on was
-- placed (`placed` holds the last op first, so an op's position is the count
-- of the ops before it).  Only those ops are recorded: the table stays as
-- small as the program's loops, and relocation stays linear in its length.
def sm121LowerLabelsOf =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (eliminate
      SM121LowerOps
      (lambda unrestricted current : (family SM121LowerOps) . (family SM121LowerCounted))
      placed
      (branch SM121LowerOpsEnd .
        (constructor SM121LowerCounted SM121LowerCountedValue zero (constructor SM121LowerLabels SM121LowerLabelsEnd)))
      (branch SM121LowerOpsNext last earlier induction .
        (eliminate
          SM121LowerCounted
          (lambda unrestricted current : (family SM121LowerCounted) . (family SM121LowerCounted))
          induction
          (branch SM121LowerCountedValue count labels .
            (eliminate
              SM121LowerOp
              (lambda unrestricted current : (family SM121LowerOp) . (family SM121LowerCounted))
              last
              (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
                (constructor SM121LowerCounted SM121LowerCountedValue
                  (succ count)
                  (sm121LowerSelect (family SM121LowerLabels)
                    (naturalAnd join (naturalIsZero target))
                    (constructor SM121LowerLabels SM121LowerLabelsNext label count labels)
                    labels)))))))))

-- the position + 1 of the op labelled `label` (0 when none is)
def sm121LowerFind =
  (lambda unrestricted labels : (family SM121LowerLabels) .
    (lambda unrestricted label : Nat .
      (eliminate
        SM121LowerLabels
        (lambda unrestricted current : (family SM121LowerLabels) . Nat)
        labels
        (branch SM121LowerLabelsEnd . zero)
        (branch SM121LowerLabelsNext entry position tail induction .
          (naturalSelect (naturalEqual entry label) (succ position) induction)))))

def sm121LowerRelocateWith =
  (lambda unrestricted labels : (family SM121LowerLabels) .
    (lambda unrestricted placed : (family SM121LowerOps) .
      (eliminate
        SM121LowerOps
        (lambda unrestricted current : (family SM121LowerOps) . (family SM121LowerRelocation))
        placed
        (branch SM121LowerOpsEnd .
          (constructor SM121LowerRelocation SM121LowerRelocated zero (constructor SM121LowerOps SM121LowerOpsEnd)))
        (branch SM121LowerOpsNext last earlier induction .
          (eliminate
            SM121LowerRelocation
            (lambda unrestricted current : (family SM121LowerRelocation) . (family SM121LowerRelocation))
            induction
            (branch SM121LowerRelocated count ops .
              (eliminate
                SM121LowerOp
                (lambda unrestricted current : (family SM121LowerOp) . (family SM121LowerRelocation))
                last
                (branch SM121LowerOpValue low high class reads writes predicateReads predicateWrites constantLoad label join target .
                  (let unrestricted found = (sm121LowerFind labels target) in
                  (let unrestricted place = (nat-subtract found 1) in
                  (let unrestricted ahead = (nat-less-than count place) in
                    (sm121LowerSelect (family SM121LowerRelocation)
                      (naturalIsZero target)
                      (constructor SM121LowerRelocation SM121LowerRelocated (succ count) (sm121LowerPush last ops))
                      (sm121LowerSelect (family SM121LowerRelocation)
                        (naturalIsZero found)
                        (constructor SM121LowerRelocation SM121LowerRelocationRefused)
                        (eliminate
                          SM121LowerWords
                          (lambda unrestricted current : (family SM121LowerWords) . (family SM121LowerRelocation))
                          (sm121LowerWithDisplacement
                            (naturalSelect ahead (nat-multiply sm121LowerBranchUnit (nat-subtract place (succ count))) zero)
                            (naturalSelect ahead zero (nat-multiply sm121LowerBranchUnit (nat-subtract (succ count) place)))
                            low high)
                          (branch SM121LowerWordsValue newLow newHigh .
                            (constructor SM121LowerRelocation SM121LowerRelocated (succ count)
                              (sm121LowerPush
                                (constructor SM121LowerOp SM121LowerOpValue newLow newHigh class reads writes
                                  predicateReads predicateWrites constantLoad label join target)
                                ops))))))))))))
            (branch SM121LowerRelocationRefused . induction))))))

-- every branch re-encoded for sm_121 against its target's placement, or
-- refused when a target was not placed
def sm121LowerRelocate =
  (lambda unrestricted placed : (family SM121LowerOps) .
    (eliminate
      SM121LowerCounted
      (lambda unrestricted current : (family SM121LowerCounted) . (family SM121LowerRelocation))
      (sm121LowerLabelsOf placed)
      (branch SM121LowerCountedValue count labels . (sm121LowerRelocateWith labels placed))))

def sm121LowerFinish =
  (lambda unrestricted addressing : (family SM121SharedAddressing) .
  (lambda unrestricted decoded : (family SM121LowerDecodedList) .
    (lambda unrestricted registers : Nat .
      (sm121LowerSelect (family SM121LowerResult)
        (nat-less-than registers (sm121LowerNamedBound decoded))
        (constructor SM121LowerResult SM121LowerRefused (constructor SM121LowerRefusal SM121LowerRefusedRegisterBeyondCount))
        (eliminate
          SM121LowerSchedule
          (lambda unrestricted current : (family SM121LowerSchedule) . (family SM121LowerResult))
          (sm121LowerScheduleProgram (sm121LowerContextOf addressing decoded registers) decoded)
          (branch SM121LowerScheduleValue time pending predicates loads guardedReads scheduled .
            (eliminate
              SM121LowerRelocation
              (lambda unrestricted current : (family SM121LowerRelocation) . (family SM121LowerResult))
              (sm121LowerRelocate scheduled)
              (branch SM121LowerRelocated count placed .
                (sm121LowerSelect (family SM121LowerResult)
                  (naturalIsZero (sm121LowerHazards placed))
                  (constructor SM121LowerResult SM121LowerLowered (sm121LowerWordsBytes placed))
                  (constructor SM121LowerResult SM121LowerRefused
                    (constructor SM121LowerRefusal SM121LowerRefusedHazard))))
              (branch SM121LowerRelocationRefused .
                (constructor SM121LowerResult SM121LowerRefused
                  (constructor SM121LowerRefusal SM121LowerRefusedBranchTarget)))))
          (branch SM121LowerScheduleRefused refusal .
            (constructor SM121LowerResult SM121LowerRefused refusal)))))))

-- The lowering of `program`, whose launch declares `registers` registers
-- and whose plain shared accesses address memory as `addressing` says.
def sm121LowerWith =
  (lambda unrestricted addressing : (family SM121SharedAddressing) .
  (lambda unrestricted program : (family SM86Program) .
    (lambda unrestricted registers : Nat .
      (let unrestricted decoded = (sm121LowerDecodeProgram program) in
        (eliminate
          SM121LowerDecodedList
          (lambda unrestricted current : (family SM121LowerDecodedList) . (family SM121LowerResult))
          decoded
          (branch SM121LowerDecodedEnd . (sm121LowerFinish addressing decoded registers))
          (branch SM121LowerDecodedNext item tail induction . (sm121LowerFinish addressing decoded registers))
          (branch SM121LowerDecodedRefused refusal .
            (constructor SM121LowerResult SM121LowerRefused refusal)))))))

def sm121LowerScaled : (family SM121SharedAddressing) =
  (constructor SM121SharedAddressing SM121SharedAddressScaled)

def sm121LowerBytes : (family SM121SharedAddressing) =
  (constructor SM121SharedAddressing SM121SharedAddressBytes)

-- The lowered machine code alone, for a region given as sm_121 code
-- (NvidiaDeviceRegionSM121): none when the lowering refuses, and the backend
-- refuses a region with no machine code.  A region given as its SM86 program
-- (sm121LowerRealization) keeps the refusal's reason instead.
def sm121LowerMachineCode =
  (lambda unrestricted addressing : (family SM121SharedAddressing) .
  (lambda unrestricted program : (family SM86Program) .
    (lambda unrestricted registers : Nat .
      (eliminate
        SM121LowerResult
        (lambda unrestricted current : (family SM121LowerResult) . Bytes)
        (sm121LowerWith addressing program registers)
        (branch SM121LowerLowered bytes . bytes)
        (branch SM121LowerRefused refusal . b"")))))

def sm121LowerProgram = (sm121LowerMachineCode sm121LowerScaled)

def sm121LowerByteAddressedSharedProgram = (sm121LowerMachineCode sm121LowerBytes)

-- The shared kilobyte the lowering reserves in front of a program's own:
-- a region given as sm_121 code declares its source extent plus this.
def sm121LoweredSharedOrigin : Nat = sm121LowerSharedBase
def sm121LoweredSharedBytes =
  (lambda unrestricted sourceBytes : Nat . (naturalAdd sm121LoweredSharedOrigin sourceBytes))

-- The realization a device region carries for sm_121.
def sm121LowerRealization =
  (lambda unrestricted program : (family SM86Program) .
    (lambda unrestricted registers : Nat .
      (eliminate
        SM121LowerResult
        (lambda unrestricted current : (family SM121LowerResult) . (family NvidiaDeviceRealization))
        (sm121LowerWith sm121LowerScaled program registers)
        (branch SM121LowerLowered bytes .
          (constructor NvidiaDeviceRealization NvidiaDeviceRealized bytes))
        (branch SM121LowerRefused refusal .
          (constructor NvidiaDeviceRealization NvidiaDeviceRealizationRefused (sm121LowerRefusalText refusal))))))

-- The static scoreboard check's count over `program`'s schedule as the
-- scheduler placed it, before the lowering refuses on it (a program the
-- lowering refuses for another reason counts once).
def sm121LowerFaults =
  (lambda unrestricted program : (family SM86Program) .
    (lambda unrestricted registers : Nat .
      (let unrestricted decoded = (sm121LowerDecodeProgram program) in
      (let unrestricted scheduled =
        (eliminate
          SM121LowerSchedule
          (lambda unrestricted current : (family SM121LowerSchedule) . Nat)
          (sm121LowerScheduleProgram (sm121LowerContextOf sm121LowerScaled decoded registers) decoded)
          (branch SM121LowerScheduleValue time pending predicates loads guardedReads scheduled .
            (eliminate
              SM121LowerRelocation
              (lambda unrestricted current : (family SM121LowerRelocation) . Nat)
              (sm121LowerRelocate scheduled)
              (branch SM121LowerRelocated count placed . (sm121LowerHazards placed))
              (branch SM121LowerRelocationRefused . 1)))
          (branch SM121LowerScheduleRefused refusal . 1)) in
        (eliminate
          SM121LowerDecodedList
          (lambda unrestricted current : (family SM121LowerDecodedList) . Nat)
          decoded
          (branch SM121LowerDecodedEnd . scheduled)
          (branch SM121LowerDecodedNext item tail induction .
            (naturalSelect (nat-less-than registers (sm121LowerNamedBound decoded)) 1 scheduled))
          (branch SM121LowerDecodedRefused refusal . 1))))))

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.