module Realization.Nvidia.SM86.IndexedRowScatterSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.InstructionEncoding import Accelerator.SM86.Program import Accelerator.SM86.Types import Data.SHA256Digest import Std.Natural import Training.NativeObligation family IndexedRowScatterSM86Geometry : Type 0 constructor IndexedRowScatterPositiveClassGradientRows constructor IndexedRowScatterNegativeClassGradientRows constructor IndexedRowScatterTokenEmbeddingGradientRows constructor IndexedRowScatterTokenEmbeddingGradientRows512x1024 end-family family IndexedRowScatterSM86ReductionSemantics : Type 0 constructor IndexedRowScatterAtomicAddF32FTZRNStrongGPU end-family family IndexedRowScatterSM86FailureCode : Type 0 constructor IndexedRowScatterSourceRowOutOfBounds constructor IndexedRowScatterSourceComponentOutOfBounds constructor IndexedRowScatterClassIdOutOfBounds constructor IndexedRowScatterDestinationComponentOutOfBounds constructor IndexedRowScatterInstructionCountMismatch constructor IndexedRowScatterEncodedByteCountMismatch constructor IndexedRowScatterEncodingFailed constructor IndexedRowScatterIdentityFailed constructor IndexedRowScatterIdentityLengthInvalid end-family family IndexedRowScatterSM86IndexResult : Type 0 constructor IndexedRowScatterIndexValid field unrestricted indexedRowScatterIndexValue : Nat constructor IndexedRowScatterIndexInvalid field unrestricted indexedRowScatterIndexFailure : (family IndexedRowScatterSM86FailureCode) field unrestricted indexedRowScatterIndexRejectedMajor : Nat field unrestricted indexedRowScatterIndexRejectedMinor : Nat end-family family IndexedRowScatterSM86ABI : Type 0 constructor IndexedRowScatterSM86ABIValue field unrestricted indexedRowScatterABIDestinationPointer : Nat field unrestricted indexedRowScatterABISourcePointer : Nat field unrestricted indexedRowScatterABIClassIdPointer : Nat end-family family IndexedRowScatterSM86Extents : Type 0 constructor IndexedRowScatterSM86ExtentsValue field unrestricted indexedRowScatterExtentSourceRows : Nat field unrestricted indexedRowScatterExtentSourceWidth : Nat field unrestricted indexedRowScatterExtentClassIds : Nat field unrestricted indexedRowScatterExtentDestinationRows : Nat field unrestricted indexedRowScatterExtentDestinationWidth : Nat end-family family IndexedRowScatterSM86Manifest : Type 0 constructor IndexedRowScatterSM86ManifestValue field unrestricted indexedRowScatterManifestObligation : (family NativeObligation) field unrestricted indexedRowScatterManifestGeometry : (family IndexedRowScatterSM86Geometry) field unrestricted indexedRowScatterManifestExpectedInstructions : Nat field unrestricted indexedRowScatterManifestExpectedEncodedBytes : Nat field unrestricted indexedRowScatterManifestRegisters : Nat field unrestricted indexedRowScatterManifestSharedBytes : Nat field unrestricted indexedRowScatterManifestGridX : Nat field unrestricted indexedRowScatterManifestGridY : Nat field unrestricted indexedRowScatterManifestBlockX : Nat field unrestricted indexedRowScatterManifestABI : (family IndexedRowScatterSM86ABI) field unrestricted indexedRowScatterManifestExtents : (family IndexedRowScatterSM86Extents) field unrestricted indexedRowScatterManifestReductionSemantics : (family IndexedRowScatterSM86ReductionSemantics) field unrestricted indexedRowScatterManifestAtomicReductionInstructions : Nat field unrestricted indexedRowScatterManifestHostFallbackOperations : Nat end-family family IndexedRowScatterSM86Telemetry : Type 0 constructor IndexedRowScatterSM86TelemetryValue field unrestricted indexedRowScatterTelemetryManifest : (family IndexedRowScatterSM86Manifest) field unrestricted indexedRowScatterTelemetryObservedInstructions : Nat field unrestricted indexedRowScatterTelemetryObservedEncodedBytes : Nat field unrestricted indexedRowScatterTelemetryAtomicReductionInstructions : Nat end-family family IndexedRowScatterSM86BuildResult : Type 0 constructor IndexedRowScatterSM86BuildSucceeded field unrestricted indexedRowScatterEncodedBytes : Bytes field unrestricted indexedRowScatterImageIdentity : Bytes field unrestricted indexedRowScatterProgramEncodingTelemetry : (family SM86ProgramEncodingTelemetry) field unrestricted indexedRowScatterIdentityTelemetry : (family SHA256DigestTelemetry) field unrestricted indexedRowScatterBuildTelemetry : (family IndexedRowScatterSM86Telemetry) constructor IndexedRowScatterSM86ContractFailed field unrestricted indexedRowScatterContractFailure : (family IndexedRowScatterSM86FailureCode) field unrestricted indexedRowScatterContractFailureTelemetry : (family IndexedRowScatterSM86Telemetry) constructor IndexedRowScatterSM86ImageEncodingFailed field unrestricted indexedRowScatterEncodingFailure : (family IndexedRowScatterSM86FailureCode) field unrestricted indexedRowScatterFailedEncoding : (family SM86ProgramEncodingResult) field unrestricted indexedRowScatterEncodingFailureTelemetry : (family IndexedRowScatterSM86Telemetry) constructor IndexedRowScatterSM86ImageIdentityFailed field unrestricted indexedRowScatterIdentityFailure : (family IndexedRowScatterSM86FailureCode) field unrestricted indexedRowScatterFailedIdentity : (family SHA256HexResult) field unrestricted indexedRowScatterIdentityFailureTelemetry : (family IndexedRowScatterSM86Telemetry) end-family -- Field projection (fields do not create definitions). def indexedRowScatterTelemetryManifest = (lambda unrestricted value : (family IndexedRowScatterSM86Telemetry) . (eliminate IndexedRowScatterSM86Telemetry (lambda unrestricted current : (family IndexedRowScatterSM86Telemetry) . (family IndexedRowScatterSM86Manifest)) value (branch IndexedRowScatterSM86TelemetryValue indexedRowScatterTelemetryManifestField indexedRowScatterTelemetryObservedInstructions indexedRowScatterTelemetryObservedEncodedBytes indexedRowScatterTelemetryAtomicReductionInstructions . indexedRowScatterTelemetryManifestField))) -- Field projection (fields do not create definitions). def indexedRowScatterManifestSharedBytes = (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) value (branch IndexedRowScatterSM86ManifestValue indexedRowScatterManifestObligation indexedRowScatterManifestGeometry indexedRowScatterManifestExpectedInstructions indexedRowScatterManifestExpectedEncodedBytes indexedRowScatterManifestRegisters indexedRowScatterManifestSharedBytesField indexedRowScatterManifestGridX indexedRowScatterManifestGridY indexedRowScatterManifestBlockX indexedRowScatterManifestABI indexedRowScatterManifestExtents indexedRowScatterManifestReductionSemantics indexedRowScatterManifestAtomicReductionInstructions indexedRowScatterManifestHostFallbackOperations . indexedRowScatterManifestSharedBytesField))) -- Field projection (fields do not create definitions). def indexedRowScatterManifestHostFallbackOperations = (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) value (branch IndexedRowScatterSM86ManifestValue indexedRowScatterManifestObligation indexedRowScatterManifestGeometry indexedRowScatterManifestExpectedInstructions indexedRowScatterManifestExpectedEncodedBytes indexedRowScatterManifestRegisters indexedRowScatterManifestSharedBytes indexedRowScatterManifestGridX indexedRowScatterManifestGridY indexedRowScatterManifestBlockX indexedRowScatterManifestABI indexedRowScatterManifestExtents indexedRowScatterManifestReductionSemantics indexedRowScatterManifestAtomicReductionInstructions indexedRowScatterManifestHostFallbackOperationsField . indexedRowScatterManifestHostFallbackOperationsField))) -- Field projection (fields do not create definitions). def indexedRowScatterManifestRegisters = (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) value (branch IndexedRowScatterSM86ManifestValue indexedRowScatterManifestObligation indexedRowScatterManifestGeometry indexedRowScatterManifestExpectedInstructions indexedRowScatterManifestExpectedEncodedBytes indexedRowScatterManifestRegistersField indexedRowScatterManifestSharedBytes indexedRowScatterManifestGridX indexedRowScatterManifestGridY indexedRowScatterManifestBlockX indexedRowScatterManifestABI indexedRowScatterManifestExtents indexedRowScatterManifestReductionSemantics indexedRowScatterManifestAtomicReductionInstructions indexedRowScatterManifestHostFallbackOperations . indexedRowScatterManifestRegistersField))) -- Field projection (fields do not create definitions). def indexedRowScatterManifestBlockX = (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) value (branch IndexedRowScatterSM86ManifestValue indexedRowScatterManifestObligation indexedRowScatterManifestGeometry indexedRowScatterManifestExpectedInstructions indexedRowScatterManifestExpectedEncodedBytes indexedRowScatterManifestRegisters indexedRowScatterManifestSharedBytes indexedRowScatterManifestGridX indexedRowScatterManifestGridY indexedRowScatterManifestBlockXField indexedRowScatterManifestABI indexedRowScatterManifestExtents indexedRowScatterManifestReductionSemantics indexedRowScatterManifestAtomicReductionInstructions indexedRowScatterManifestHostFallbackOperations . indexedRowScatterManifestBlockXField))) -- Field projection (fields do not create definitions). def indexedRowScatterManifestGridX = (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) value (branch IndexedRowScatterSM86ManifestValue indexedRowScatterManifestObligation indexedRowScatterManifestGeometry indexedRowScatterManifestExpectedInstructions indexedRowScatterManifestExpectedEncodedBytes indexedRowScatterManifestRegisters indexedRowScatterManifestSharedBytes indexedRowScatterManifestGridXField indexedRowScatterManifestGridY indexedRowScatterManifestBlockX indexedRowScatterManifestABI indexedRowScatterManifestExtents indexedRowScatterManifestReductionSemantics indexedRowScatterManifestAtomicReductionInstructions indexedRowScatterManifestHostFallbackOperations . indexedRowScatterManifestGridXField))) def indexedRowScatterSM86FailureCodeBytes = (lambda unrestricted code : (family IndexedRowScatterSM86FailureCode) . (eliminate IndexedRowScatterSM86FailureCode (lambda unrestricted current : (family IndexedRowScatterSM86FailureCode) . Bytes) code (branch IndexedRowScatterSourceRowOutOfBounds . b"ALPHA-SM86-IDXSCAT-001") (branch IndexedRowScatterSourceComponentOutOfBounds . b"ALPHA-SM86-IDXSCAT-002") (branch IndexedRowScatterClassIdOutOfBounds . b"ALPHA-SM86-IDXSCAT-003") (branch IndexedRowScatterDestinationComponentOutOfBounds . b"ALPHA-SM86-IDXSCAT-004") (branch IndexedRowScatterInstructionCountMismatch . b"ALPHA-SM86-IDXSCAT-005") (branch IndexedRowScatterEncodedByteCountMismatch . b"ALPHA-SM86-IDXSCAT-006") (branch IndexedRowScatterEncodingFailed . b"ALPHA-SM86-IDXSCAT-007") (branch IndexedRowScatterIdentityFailed . b"ALPHA-SM86-IDXSCAT-008") (branch IndexedRowScatterIdentityLengthInvalid . b"ALPHA-SM86-IDXSCAT-009"))) def indexedRowScatterSM86N4 = (byte-to-nat (byte 4)) def indexedRowScatterSM86N12 = (byte-to-nat (byte 12)) def indexedRowScatterSM86N16 = (byte-to-nat (byte 16)) def indexedRowScatterSM86N24 = (byte-to-nat (byte 24)) def indexedRowScatterSM86N48 = (byte-to-nat (byte 48)) def indexedRowScatterSM86N64 = (byte-to-nat (byte 64)) def indexedRowScatterSM86N96 = (byte-to-nat (byte 96)) def indexedRowScatterSM86N104 = (byte-to-nat (byte 104)) def indexedRowScatterSM86N112 = (byte-to-nat (byte 112)) def indexedRowScatterSM86N256 = (succ (byte-to-nat (byte 255))) def indexedRowScatterSM86N512 = 512 def indexedRowScatterSM86N1024 = (naturalMultiply indexedRowScatterSM86N4 indexedRowScatterSM86N256) def indexedRowScatterSM86N4096 = (naturalPowerOfTwo indexedRowScatterSM86N12) def indexedRowScatterSM86N6144 = (naturalMultiply indexedRowScatterSM86N24 indexedRowScatterSM86N256) def indexedRowScatterSM86N12288 = (naturalMultiply indexedRowScatterSM86N48 indexedRowScatterSM86N256) def indexedRowScatterSM86N192 = (naturalMultiply indexedRowScatterSM86N12 indexedRowScatterSM86N16) def indexedRowScatterSM86DestinationPointerNatural = (naturalAdd indexedRowScatterSM86N256 indexedRowScatterSM86N96) def indexedRowScatterSM86SourcePointerNatural = (naturalAdd indexedRowScatterSM86N256 indexedRowScatterSM86N104) def indexedRowScatterSM86ClassIdPointerNatural = (naturalAdd indexedRowScatterSM86N256 indexedRowScatterSM86N112) -- Both the atomic and ordered realizations share this three-pointer ABI. def indexedRowScatterSM86DestinationArgument : Nat = 0 def indexedRowScatterSM86SourceArgument : Nat = 1 def indexedRowScatterSM86ClassIdArgument : Nat = 2 def indexedRowScatterSM86ArgumentCount : Nat = 3 def indexedRowScatterSM86U0 = sm86Unsigned32Zero def indexedRowScatterSM86U4 = (sm86Unsigned32 (byte 4) (byte 0) (byte 0) (byte 0)) def indexedRowScatterSM86U256 = (sm86Unsigned32 (byte 0) (byte 1) (byte 0) (byte 0)) def indexedRowScatterSM86U512 = (sm86Unsigned32 (byte 0) (byte 2) (byte 0) (byte 0)) def indexedRowScatterSM86U1024 = (sm86Unsigned32 (byte 0) (byte 4) (byte 0) (byte 0)) def indexedRowScatterSM86DestinationPointer = (sm86Unsigned32 (byte 96) (byte 1) (byte 0) (byte 0)) def indexedRowScatterSM86SourcePointer = (sm86Unsigned32 (byte 104) (byte 1) (byte 0) (byte 0)) def indexedRowScatterSM86ClassIdPointer = (sm86Unsigned32 (byte 112) (byte 1) (byte 0) (byte 0)) def indexedRowScatterSM86R0 = (sm86Register (byte 0)) def indexedRowScatterSM86R1 = (sm86Register (byte 1)) def indexedRowScatterSM86R2 = (sm86Register (byte 2)) def indexedRowScatterSM86R3 = (sm86Register (byte 3)) def indexedRowScatterSM86R5 = (sm86Register (byte 5)) def indexedRowScatterSM86R6 = (sm86Register (byte 6)) def indexedRowScatterSM86R10 = (sm86Register (byte 10)) def indexedRowScatterSM86R11 = (sm86Register (byte 11)) def indexedRowScatterSM86R14 = (sm86Register (byte 14)) def indexedRowScatterSM86Set0 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0)) def indexedRowScatterSM86Set3 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier3)) def indexedRowScatterSM86Set4 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier4)) def indexedRowScatterSM86WaitSpecials = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0)) def indexedRowScatterSM86WaitClassId = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier3)) def indexedRowScatterSM86WaitSource = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier4)) def indexedRowScatterSM86WidthNatural = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (eliminate IndexedRowScatterSM86Geometry (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) . Nat) geometry (branch IndexedRowScatterPositiveClassGradientRows . indexedRowScatterSM86N256) (branch IndexedRowScatterNegativeClassGradientRows . indexedRowScatterSM86N256) (branch IndexedRowScatterTokenEmbeddingGradientRows . indexedRowScatterSM86N1024) (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . indexedRowScatterSM86N512))) def indexedRowScatterSM86RowsNatural = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (eliminate IndexedRowScatterSM86Geometry (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) . Nat) geometry (branch IndexedRowScatterPositiveClassGradientRows . indexedRowScatterSM86N6144) (branch IndexedRowScatterNegativeClassGradientRows . indexedRowScatterSM86N4096) (branch IndexedRowScatterTokenEmbeddingGradientRows . indexedRowScatterSM86N6144) (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . indexedRowScatterSM86N1024))) def indexedRowScatterSM86WidthImmediate = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (eliminate IndexedRowScatterSM86Geometry (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) . (family SM86Unsigned32)) geometry (branch IndexedRowScatterPositiveClassGradientRows . indexedRowScatterSM86U256) (branch IndexedRowScatterNegativeClassGradientRows . indexedRowScatterSM86U256) (branch IndexedRowScatterTokenEmbeddingGradientRows . indexedRowScatterSM86U1024) (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . indexedRowScatterSM86U512))) def indexedRowScatterSM86SourceIndex = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (lambda unrestricted row : Nat . (lambda unrestricted component : Nat . (nat-eliminate (lambda unrestricted rowInBounds : Nat . (family IndexedRowScatterSM86IndexResult)) (constructor IndexedRowScatterSM86IndexResult IndexedRowScatterIndexInvalid (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterSourceRowOutOfBounds) row component) (lambda unrestricted rowPredecessor : Nat . (lambda unrestricted rowInduction : (family IndexedRowScatterSM86IndexResult) . (nat-eliminate (lambda unrestricted componentInBounds : Nat . (family IndexedRowScatterSM86IndexResult)) (constructor IndexedRowScatterSM86IndexResult IndexedRowScatterIndexInvalid (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterSourceComponentOutOfBounds) row component) (lambda unrestricted componentPredecessor : Nat . (lambda unrestricted componentInduction : (family IndexedRowScatterSM86IndexResult) . (constructor IndexedRowScatterSM86IndexResult IndexedRowScatterIndexValid (naturalAdd (naturalMultiply row (indexedRowScatterSM86WidthNatural geometry)) component)))) (naturalLess component (indexedRowScatterSM86WidthNatural geometry))))) (naturalLess row (indexedRowScatterSM86RowsNatural geometry)))))) def indexedRowScatterSM86DestinationIndex = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (lambda unrestricted classId : Nat . (lambda unrestricted component : Nat . (nat-eliminate (lambda unrestricted classIdInBounds : Nat . (family IndexedRowScatterSM86IndexResult)) (constructor IndexedRowScatterSM86IndexResult IndexedRowScatterIndexInvalid (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterClassIdOutOfBounds) classId component) (lambda unrestricted classIdPredecessor : Nat . (lambda unrestricted classIdInduction : (family IndexedRowScatterSM86IndexResult) . (nat-eliminate (lambda unrestricted componentInBounds : Nat . (family IndexedRowScatterSM86IndexResult)) (constructor IndexedRowScatterSM86IndexResult IndexedRowScatterIndexInvalid (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterDestinationComponentOutOfBounds) classId component) (lambda unrestricted componentPredecessor : Nat . (lambda unrestricted componentInduction : (family IndexedRowScatterSM86IndexResult) . (constructor IndexedRowScatterSM86IndexResult IndexedRowScatterIndexValid (naturalAdd (naturalMultiply classId (indexedRowScatterSM86WidthNatural geometry)) component)))) (naturalLess component (indexedRowScatterSM86WidthNatural geometry))))) (naturalLess classId indexedRowScatterSM86N12288))))) def indexedRowScatterSM86ProgramWithWidth = (lambda unrestricted width : (family SM86Unsigned32) . (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86SpecialToRegister indexedRowScatterSM86R0 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) indexedRowScatterSM86Set0)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86SpecialToRegister indexedRowScatterSM86R1 (constructor SM86SpecialRegister SM86ThreadIdX) indexedRowScatterSM86Set0)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86MoveImmediate indexedRowScatterSM86R5 indexedRowScatterSM86U4 sm86SafeControl)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant indexedRowScatterSM86R6 indexedRowScatterSM86R0 indexedRowScatterSM86R5 (byte 0) indexedRowScatterSM86ClassIdPointer indexedRowScatterSM86WaitSpecials)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86LoadGlobal indexedRowScatterSM86R11 indexedRowScatterSM86R6 indexedRowScatterSM86U0 indexedRowScatterSM86Set3)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate indexedRowScatterSM86R2 indexedRowScatterSM86R0 width indexedRowScatterSM86R1 sm86SafeControl)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant indexedRowScatterSM86R14 indexedRowScatterSM86R2 indexedRowScatterSM86R5 (byte 0) indexedRowScatterSM86SourcePointer sm86SafeControl)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86LoadGlobal indexedRowScatterSM86R10 indexedRowScatterSM86R14 indexedRowScatterSM86U0 indexedRowScatterSM86Set4)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate indexedRowScatterSM86R3 indexedRowScatterSM86R11 width indexedRowScatterSM86R1 indexedRowScatterSM86WaitClassId)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant indexedRowScatterSM86R6 indexedRowScatterSM86R3 indexedRowScatterSM86R5 (byte 0) indexedRowScatterSM86DestinationPointer sm86SafeControl)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86ReduceGlobalAddFloat32 indexedRowScatterSM86R6 indexedRowScatterSM86R10 indexedRowScatterSM86U0 indexedRowScatterSM86WaitSource)) (constructor SM86Program SM86ProgramNext (sm86Instruction (constructor SM86InstructionBody SM86Exit sm86SafeControl)) (constructor SM86Program SM86ProgramEnd)))))))))))))) def indexedRowScatterSM86Program = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (indexedRowScatterSM86ProgramWithWidth (indexedRowScatterSM86WidthImmediate geometry))) def indexedRowScatterSM86ABIValue = (constructor IndexedRowScatterSM86ABI IndexedRowScatterSM86ABIValue indexedRowScatterSM86DestinationPointerNatural indexedRowScatterSM86SourcePointerNatural indexedRowScatterSM86ClassIdPointerNatural) def indexedRowScatterSM86ExtentsFor = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (constructor IndexedRowScatterSM86Extents IndexedRowScatterSM86ExtentsValue (indexedRowScatterSM86RowsNatural geometry) (indexedRowScatterSM86WidthNatural geometry) (indexedRowScatterSM86RowsNatural geometry) indexedRowScatterSM86N12288 (indexedRowScatterSM86WidthNatural geometry))) def indexedRowScatterSM86ObligationFor = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (eliminate IndexedRowScatterSM86Geometry (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) . (family NativeObligation)) geometry (branch IndexedRowScatterPositiveClassGradientRows . (constructor NativeObligation NativeClassRowScatter256)) (branch IndexedRowScatterNegativeClassGradientRows . (constructor NativeObligation NativeClassRowScatter256)) (branch IndexedRowScatterTokenEmbeddingGradientRows . (constructor NativeObligation NativeEmbeddingScatter1024)) (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . (constructor NativeObligation NativeEmbeddingScatter512)))) def indexedRowScatterSM86ManifestFor = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (constructor IndexedRowScatterSM86Manifest IndexedRowScatterSM86ManifestValue (indexedRowScatterSM86ObligationFor geometry) geometry indexedRowScatterSM86N12 indexedRowScatterSM86N192 indexedRowScatterSM86N24 zero (indexedRowScatterSM86RowsNatural geometry) (succ zero) (indexedRowScatterSM86WidthNatural geometry) indexedRowScatterSM86ABIValue (indexedRowScatterSM86ExtentsFor geometry) (constructor IndexedRowScatterSM86ReductionSemantics IndexedRowScatterAtomicAddF32FTZRNStrongGPU) (succ zero) zero)) def indexedRowScatterSM86ManifestExpectedInstructions = (lambda unrestricted manifest : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) manifest (branch IndexedRowScatterSM86ManifestValue obligation geometry instructions encodedBytes registers sharedBytes gridX gridY blockX abi extents reduction atomic hostFallback . instructions))) def indexedRowScatterSM86ManifestExpectedBytes = (lambda unrestricted manifest : (family IndexedRowScatterSM86Manifest) . (eliminate IndexedRowScatterSM86Manifest (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat) manifest (branch IndexedRowScatterSM86ManifestValue obligation geometry instructions encodedBytes registers sharedBytes gridX gridY blockX abi extents reduction atomic hostFallback . encodedBytes))) def indexedRowScatterSM86TelemetryFor = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (lambda unrestricted observedInstructions : Nat . (lambda unrestricted observedEncodedBytes : Nat . (constructor IndexedRowScatterSM86Telemetry IndexedRowScatterSM86TelemetryValue (indexedRowScatterSM86ManifestFor geometry) observedInstructions observedEncodedBytes (succ zero))))) def indexedRowScatterSM86Build = (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) . (app (lambda unrestricted program : (family SM86Program) . (app (lambda unrestricted observedInstructions : Nat . (app (lambda unrestricted manifest : (family IndexedRowScatterSM86Manifest) . (nat-eliminate (lambda unrestricted countMatched : Nat . (family IndexedRowScatterSM86BuildResult)) (constructor IndexedRowScatterSM86BuildResult IndexedRowScatterSM86ContractFailed (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterInstructionCountMismatch) (indexedRowScatterSM86TelemetryFor geometry observedInstructions zero)) (lambda unrestricted countPredecessor : Nat . (lambda unrestricted countInduction : (family IndexedRowScatterSM86BuildResult) . (app (lambda unrestricted encoding : (family SM86ProgramEncodingResult) . (eliminate SM86ProgramEncodingResult (lambda unrestricted current : (family SM86ProgramEncodingResult) . (family IndexedRowScatterSM86BuildResult)) encoding (branch SM86ProgramEncodingSucceeded bytes encodingTelemetry . (app (lambda unrestricted telemetry : (family IndexedRowScatterSM86Telemetry) . (nat-eliminate (lambda unrestricted bytesMatched : Nat . (family IndexedRowScatterSM86BuildResult)) (constructor IndexedRowScatterSM86BuildResult IndexedRowScatterSM86ContractFailed (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterEncodedByteCountMismatch) telemetry) (lambda unrestricted bytesPredecessor : Nat . (lambda unrestricted bytesInduction : (family IndexedRowScatterSM86BuildResult) . (app (lambda unrestricted identityResult : (family SHA256HexResult) . (eliminate SHA256HexResult (lambda unrestricted current : (family SHA256HexResult) . (family IndexedRowScatterSM86BuildResult)) identityResult (branch SHA256HexSucceeded identity identityTelemetry . (nat-eliminate (lambda unrestricted identityLengthMatched : Nat . (family IndexedRowScatterSM86BuildResult)) (constructor IndexedRowScatterSM86BuildResult IndexedRowScatterSM86ImageIdentityFailed (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterIdentityLengthInvalid) identityResult telemetry) (lambda unrestricted identityLengthPredecessor : Nat . (lambda unrestricted identityLengthInduction : (family IndexedRowScatterSM86BuildResult) . (constructor IndexedRowScatterSM86BuildResult IndexedRowScatterSM86BuildSucceeded bytes identity encodingTelemetry identityTelemetry telemetry))) (naturalEqual (bytes-length identity) indexedRowScatterSM86N64))) (branch SHA256HexFailed error ordinal identityTelemetry . (constructor IndexedRowScatterSM86BuildResult IndexedRowScatterSM86ImageIdentityFailed (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterIdentityFailed) identityResult telemetry)))) (sha256Hex bytes)))) (naturalEqual (bytes-length bytes) (indexedRowScatterSM86ManifestExpectedBytes manifest)))) (indexedRowScatterSM86TelemetryFor geometry observedInstructions (bytes-length bytes)))) (branch SM86ProgramEncodingFailed instructionIndex failure encodingTelemetry . (constructor IndexedRowScatterSM86BuildResult IndexedRowScatterSM86ImageEncodingFailed (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterEncodingFailed) encoding (indexedRowScatterSM86TelemetryFor geometry observedInstructions zero))))) (sm86EncodeProgram program)))) (naturalEqual observedInstructions (indexedRowScatterSM86ManifestExpectedInstructions manifest)))) (indexedRowScatterSM86ManifestFor geometry))) (sm86ProgramCount program))) (indexedRowScatterSM86Program geometry))) def indexedRowScatterSM86BuildPositiveClassGradientRows = (indexedRowScatterSM86Build (constructor IndexedRowScatterSM86Geometry IndexedRowScatterPositiveClassGradientRows)) def indexedRowScatterSM86BuildNegativeClassGradientRows = (indexedRowScatterSM86Build (constructor IndexedRowScatterSM86Geometry IndexedRowScatterNegativeClassGradientRows)) def indexedRowScatterSM86BuildTokenEmbeddingGradientRows = (indexedRowScatterSM86Build (constructor IndexedRowScatterSM86Geometry IndexedRowScatterTokenEmbeddingGradientRows)) def indexedRowScatterSM86BuildTokenEmbeddingGradientRows512x1024 = (indexedRowScatterSM86Build (constructor IndexedRowScatterSM86Geometry IndexedRowScatterTokenEmbeddingGradientRows512x1024))