Source/Packages

Realization.Nvidia.SM86.IndexedRowScatterSM86

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

926 lines148 declarations38.7 KiBSHA-256 c86946ffe8c6

Complete file · line 53

IndexedRowScatterSM86.alpha

Definition view
1module Realization.Nvidia.SM86.IndexedRowScatterSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.InstructionEncoding
7import Accelerator.SM86.Program
8import Accelerator.SM86.Types
9import Data.SHA256Digest
10import Std.Natural
11import Training.NativeObligation
12
13family IndexedRowScatterSM86Geometry : Type 0
14constructor IndexedRowScatterPositiveClassGradientRows
15constructor IndexedRowScatterNegativeClassGradientRows
16constructor IndexedRowScatterTokenEmbeddingGradientRows
17constructor IndexedRowScatterTokenEmbeddingGradientRows512x1024
18
19end-family
20
21family IndexedRowScatterSM86ReductionSemantics : Type 0
22constructor IndexedRowScatterAtomicAddF32FTZRNStrongGPU
23
24end-family
25
26family IndexedRowScatterSM86FailureCode : Type 0
27constructor IndexedRowScatterSourceRowOutOfBounds
28constructor IndexedRowScatterSourceComponentOutOfBounds
29constructor IndexedRowScatterClassIdOutOfBounds
30constructor IndexedRowScatterDestinationComponentOutOfBounds
31constructor IndexedRowScatterInstructionCountMismatch
32constructor IndexedRowScatterEncodedByteCountMismatch
33constructor IndexedRowScatterEncodingFailed
34constructor IndexedRowScatterIdentityFailed
35constructor IndexedRowScatterIdentityLengthInvalid
36
37end-family
38
39family IndexedRowScatterSM86IndexResult : Type 0
40constructor IndexedRowScatterIndexValid
41field unrestricted indexedRowScatterIndexValue : Nat
42constructor IndexedRowScatterIndexInvalid
43field unrestricted indexedRowScatterIndexFailure : (family IndexedRowScatterSM86FailureCode)
44field unrestricted indexedRowScatterIndexRejectedMajor : Nat
45field unrestricted indexedRowScatterIndexRejectedMinor : Nat
46
47end-family
48
49family IndexedRowScatterSM86ABI : Type 0
50constructor IndexedRowScatterSM86ABIValue
51field unrestricted indexedRowScatterABIDestinationPointer : Nat
52field unrestricted indexedRowScatterABISourcePointer : Nat
53field unrestricted indexedRowScatterABIClassIdPointer : Nat
54
55end-family
56
57family IndexedRowScatterSM86Extents : Type 0
58constructor IndexedRowScatterSM86ExtentsValue
59field unrestricted indexedRowScatterExtentSourceRows : Nat
60field unrestricted indexedRowScatterExtentSourceWidth : Nat
61field unrestricted indexedRowScatterExtentClassIds : Nat
62field unrestricted indexedRowScatterExtentDestinationRows : Nat
63field unrestricted indexedRowScatterExtentDestinationWidth : Nat
64
65end-family
66
67family IndexedRowScatterSM86Manifest : Type 0
68constructor IndexedRowScatterSM86ManifestValue
69field unrestricted indexedRowScatterManifestObligation : (family NativeObligation)
70field unrestricted indexedRowScatterManifestGeometry : (family IndexedRowScatterSM86Geometry)
71field unrestricted indexedRowScatterManifestExpectedInstructions : Nat
72field unrestricted indexedRowScatterManifestExpectedEncodedBytes : Nat
73field unrestricted indexedRowScatterManifestRegisters : Nat
74field unrestricted indexedRowScatterManifestSharedBytes : Nat
75field unrestricted indexedRowScatterManifestGridX : Nat
76field unrestricted indexedRowScatterManifestGridY : Nat
77field unrestricted indexedRowScatterManifestBlockX : Nat
78field unrestricted indexedRowScatterManifestABI : (family IndexedRowScatterSM86ABI)
79field unrestricted indexedRowScatterManifestExtents : (family IndexedRowScatterSM86Extents)
80field unrestricted indexedRowScatterManifestReductionSemantics : (family IndexedRowScatterSM86ReductionSemantics)
81field unrestricted indexedRowScatterManifestAtomicReductionInstructions : Nat
82field unrestricted indexedRowScatterManifestHostFallbackOperations : Nat
83
84end-family
85
86family IndexedRowScatterSM86Telemetry : Type 0
87constructor IndexedRowScatterSM86TelemetryValue
88field unrestricted indexedRowScatterTelemetryManifest : (family IndexedRowScatterSM86Manifest)
89field unrestricted indexedRowScatterTelemetryObservedInstructions : Nat
90field unrestricted indexedRowScatterTelemetryObservedEncodedBytes : Nat
91field unrestricted indexedRowScatterTelemetryAtomicReductionInstructions : Nat
92
93end-family
94
95family IndexedRowScatterSM86BuildResult : Type 0
96constructor IndexedRowScatterSM86BuildSucceeded
97field unrestricted indexedRowScatterEncodedBytes : Bytes
98field unrestricted indexedRowScatterImageIdentity : Bytes
99field unrestricted indexedRowScatterProgramEncodingTelemetry : (family SM86ProgramEncodingTelemetry)
100field unrestricted indexedRowScatterIdentityTelemetry : (family SHA256DigestTelemetry)
101field unrestricted indexedRowScatterBuildTelemetry : (family IndexedRowScatterSM86Telemetry)
102constructor IndexedRowScatterSM86ContractFailed
103field unrestricted indexedRowScatterContractFailure : (family IndexedRowScatterSM86FailureCode)
104field unrestricted indexedRowScatterContractFailureTelemetry : (family IndexedRowScatterSM86Telemetry)
105constructor IndexedRowScatterSM86ImageEncodingFailed
106field unrestricted indexedRowScatterEncodingFailure : (family IndexedRowScatterSM86FailureCode)
107field unrestricted indexedRowScatterFailedEncoding : (family SM86ProgramEncodingResult)
108field unrestricted indexedRowScatterEncodingFailureTelemetry : (family IndexedRowScatterSM86Telemetry)
109constructor IndexedRowScatterSM86ImageIdentityFailed
110field unrestricted indexedRowScatterIdentityFailure : (family IndexedRowScatterSM86FailureCode)
111field unrestricted indexedRowScatterFailedIdentity : (family SHA256HexResult)
112field unrestricted indexedRowScatterIdentityFailureTelemetry : (family IndexedRowScatterSM86Telemetry)
113
114end-family
115
116-- Field projection (fields do not create definitions).
117def indexedRowScatterTelemetryManifest =
118  (lambda unrestricted value : (family IndexedRowScatterSM86Telemetry) .
119    (eliminate
120      IndexedRowScatterSM86Telemetry
121      (lambda unrestricted current : (family IndexedRowScatterSM86Telemetry) .
122        (family IndexedRowScatterSM86Manifest))
123      value
124      (branch
125        IndexedRowScatterSM86TelemetryValue
126        indexedRowScatterTelemetryManifestField
127        indexedRowScatterTelemetryObservedInstructions
128        indexedRowScatterTelemetryObservedEncodedBytes
129        indexedRowScatterTelemetryAtomicReductionInstructions
130        .
131        indexedRowScatterTelemetryManifestField)))
132
133-- Field projection (fields do not create definitions).
134def indexedRowScatterManifestSharedBytes =
135  (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) .
136    (eliminate
137      IndexedRowScatterSM86Manifest
138      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
139      value
140      (branch
141        IndexedRowScatterSM86ManifestValue
142        indexedRowScatterManifestObligation
143        indexedRowScatterManifestGeometry
144        indexedRowScatterManifestExpectedInstructions
145        indexedRowScatterManifestExpectedEncodedBytes
146        indexedRowScatterManifestRegisters
147        indexedRowScatterManifestSharedBytesField
148        indexedRowScatterManifestGridX
149        indexedRowScatterManifestGridY
150        indexedRowScatterManifestBlockX
151        indexedRowScatterManifestABI
152        indexedRowScatterManifestExtents
153        indexedRowScatterManifestReductionSemantics
154        indexedRowScatterManifestAtomicReductionInstructions
155        indexedRowScatterManifestHostFallbackOperations
156        .
157        indexedRowScatterManifestSharedBytesField)))
158
159-- Field projection (fields do not create definitions).
160def indexedRowScatterManifestHostFallbackOperations =
161  (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) .
162    (eliminate
163      IndexedRowScatterSM86Manifest
164      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
165      value
166      (branch
167        IndexedRowScatterSM86ManifestValue
168        indexedRowScatterManifestObligation
169        indexedRowScatterManifestGeometry
170        indexedRowScatterManifestExpectedInstructions
171        indexedRowScatterManifestExpectedEncodedBytes
172        indexedRowScatterManifestRegisters
173        indexedRowScatterManifestSharedBytes
174        indexedRowScatterManifestGridX
175        indexedRowScatterManifestGridY
176        indexedRowScatterManifestBlockX
177        indexedRowScatterManifestABI
178        indexedRowScatterManifestExtents
179        indexedRowScatterManifestReductionSemantics
180        indexedRowScatterManifestAtomicReductionInstructions
181        indexedRowScatterManifestHostFallbackOperationsField
182        .
183        indexedRowScatterManifestHostFallbackOperationsField)))
184
185-- Field projection (fields do not create definitions).
186def indexedRowScatterManifestRegisters =
187  (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) .
188    (eliminate
189      IndexedRowScatterSM86Manifest
190      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
191      value
192      (branch
193        IndexedRowScatterSM86ManifestValue
194        indexedRowScatterManifestObligation
195        indexedRowScatterManifestGeometry
196        indexedRowScatterManifestExpectedInstructions
197        indexedRowScatterManifestExpectedEncodedBytes
198        indexedRowScatterManifestRegistersField
199        indexedRowScatterManifestSharedBytes
200        indexedRowScatterManifestGridX
201        indexedRowScatterManifestGridY
202        indexedRowScatterManifestBlockX
203        indexedRowScatterManifestABI
204        indexedRowScatterManifestExtents
205        indexedRowScatterManifestReductionSemantics
206        indexedRowScatterManifestAtomicReductionInstructions
207        indexedRowScatterManifestHostFallbackOperations
208        .
209        indexedRowScatterManifestRegistersField)))
210
211-- Field projection (fields do not create definitions).
212def indexedRowScatterManifestBlockX =
213  (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) .
214    (eliminate
215      IndexedRowScatterSM86Manifest
216      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
217      value
218      (branch
219        IndexedRowScatterSM86ManifestValue
220        indexedRowScatterManifestObligation
221        indexedRowScatterManifestGeometry
222        indexedRowScatterManifestExpectedInstructions
223        indexedRowScatterManifestExpectedEncodedBytes
224        indexedRowScatterManifestRegisters
225        indexedRowScatterManifestSharedBytes
226        indexedRowScatterManifestGridX
227        indexedRowScatterManifestGridY
228        indexedRowScatterManifestBlockXField
229        indexedRowScatterManifestABI
230        indexedRowScatterManifestExtents
231        indexedRowScatterManifestReductionSemantics
232        indexedRowScatterManifestAtomicReductionInstructions
233        indexedRowScatterManifestHostFallbackOperations
234        .
235        indexedRowScatterManifestBlockXField)))
236
237-- Field projection (fields do not create definitions).
238def indexedRowScatterManifestGridX =
239  (lambda unrestricted value : (family IndexedRowScatterSM86Manifest) .
240    (eliminate
241      IndexedRowScatterSM86Manifest
242      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
243      value
244      (branch
245        IndexedRowScatterSM86ManifestValue
246        indexedRowScatterManifestObligation
247        indexedRowScatterManifestGeometry
248        indexedRowScatterManifestExpectedInstructions
249        indexedRowScatterManifestExpectedEncodedBytes
250        indexedRowScatterManifestRegisters
251        indexedRowScatterManifestSharedBytes
252        indexedRowScatterManifestGridXField
253        indexedRowScatterManifestGridY
254        indexedRowScatterManifestBlockX
255        indexedRowScatterManifestABI
256        indexedRowScatterManifestExtents
257        indexedRowScatterManifestReductionSemantics
258        indexedRowScatterManifestAtomicReductionInstructions
259        indexedRowScatterManifestHostFallbackOperations
260        .
261        indexedRowScatterManifestGridXField)))
262
263def indexedRowScatterSM86FailureCodeBytes =
264  (lambda unrestricted code : (family IndexedRowScatterSM86FailureCode) .
265    (eliminate
266      IndexedRowScatterSM86FailureCode
267      (lambda unrestricted current : (family IndexedRowScatterSM86FailureCode) . Bytes)
268      code
269      (branch IndexedRowScatterSourceRowOutOfBounds . b"ALPHA-SM86-IDXSCAT-001")
270      (branch IndexedRowScatterSourceComponentOutOfBounds . b"ALPHA-SM86-IDXSCAT-002")
271      (branch IndexedRowScatterClassIdOutOfBounds . b"ALPHA-SM86-IDXSCAT-003")
272      (branch IndexedRowScatterDestinationComponentOutOfBounds . b"ALPHA-SM86-IDXSCAT-004")
273      (branch IndexedRowScatterInstructionCountMismatch . b"ALPHA-SM86-IDXSCAT-005")
274      (branch IndexedRowScatterEncodedByteCountMismatch . b"ALPHA-SM86-IDXSCAT-006")
275      (branch IndexedRowScatterEncodingFailed . b"ALPHA-SM86-IDXSCAT-007")
276      (branch IndexedRowScatterIdentityFailed . b"ALPHA-SM86-IDXSCAT-008")
277      (branch IndexedRowScatterIdentityLengthInvalid . b"ALPHA-SM86-IDXSCAT-009")))
278
279def indexedRowScatterSM86N4 =
280  (byte-to-nat (byte 4))
281
282def indexedRowScatterSM86N12 =
283  (byte-to-nat (byte 12))
284
285def indexedRowScatterSM86N16 =
286  (byte-to-nat (byte 16))
287
288def indexedRowScatterSM86N24 =
289  (byte-to-nat (byte 24))
290
291def indexedRowScatterSM86N48 =
292  (byte-to-nat (byte 48))
293
294def indexedRowScatterSM86N64 =
295  (byte-to-nat (byte 64))
296
297def indexedRowScatterSM86N96 =
298  (byte-to-nat (byte 96))
299
300def indexedRowScatterSM86N104 =
301  (byte-to-nat (byte 104))
302
303def indexedRowScatterSM86N112 =
304  (byte-to-nat (byte 112))
305
306def indexedRowScatterSM86N256 =
307  (succ (byte-to-nat (byte 255)))
308
309def indexedRowScatterSM86N512 =
310  512
311
312def indexedRowScatterSM86N1024 =
313  (naturalMultiply indexedRowScatterSM86N4 indexedRowScatterSM86N256)
314
315def indexedRowScatterSM86N4096 =
316  (naturalPowerOfTwo indexedRowScatterSM86N12)
317
318def indexedRowScatterSM86N6144 =
319  (naturalMultiply indexedRowScatterSM86N24 indexedRowScatterSM86N256)
320
321def indexedRowScatterSM86N12288 =
322  (naturalMultiply indexedRowScatterSM86N48 indexedRowScatterSM86N256)
323
324def indexedRowScatterSM86N192 =
325  (naturalMultiply indexedRowScatterSM86N12 indexedRowScatterSM86N16)
326
327def indexedRowScatterSM86DestinationPointerNatural =
328  (naturalAdd indexedRowScatterSM86N256 indexedRowScatterSM86N96)
329
330def indexedRowScatterSM86SourcePointerNatural =
331  (naturalAdd indexedRowScatterSM86N256 indexedRowScatterSM86N104)
332
333def indexedRowScatterSM86ClassIdPointerNatural =
334  (naturalAdd indexedRowScatterSM86N256 indexedRowScatterSM86N112)
335
336-- Both the atomic and ordered realizations share this three-pointer ABI.
337def indexedRowScatterSM86DestinationArgument : Nat = 0
338def indexedRowScatterSM86SourceArgument : Nat = 1
339def indexedRowScatterSM86ClassIdArgument : Nat = 2
340def indexedRowScatterSM86ArgumentCount : Nat = 3
341
342def indexedRowScatterSM86U0 =
343  sm86Unsigned32Zero
344
345def indexedRowScatterSM86U4 =
346  (sm86Unsigned32 (byte 4) (byte 0) (byte 0) (byte 0))
347
348def indexedRowScatterSM86U256 =
349  (sm86Unsigned32 (byte 0) (byte 1) (byte 0) (byte 0))
350
351def indexedRowScatterSM86U512 =
352  (sm86Unsigned32 (byte 0) (byte 2) (byte 0) (byte 0))
353
354def indexedRowScatterSM86U1024 =
355  (sm86Unsigned32 (byte 0) (byte 4) (byte 0) (byte 0))
356
357def indexedRowScatterSM86DestinationPointer =
358  (sm86Unsigned32 (byte 96) (byte 1) (byte 0) (byte 0))
359
360def indexedRowScatterSM86SourcePointer =
361  (sm86Unsigned32 (byte 104) (byte 1) (byte 0) (byte 0))
362
363def indexedRowScatterSM86ClassIdPointer =
364  (sm86Unsigned32 (byte 112) (byte 1) (byte 0) (byte 0))
365
366def indexedRowScatterSM86R0 =
367  (sm86Register (byte 0))
368
369def indexedRowScatterSM86R1 =
370  (sm86Register (byte 1))
371
372def indexedRowScatterSM86R2 =
373  (sm86Register (byte 2))
374
375def indexedRowScatterSM86R3 =
376  (sm86Register (byte 3))
377
378def indexedRowScatterSM86R5 =
379  (sm86Register (byte 5))
380
381def indexedRowScatterSM86R6 =
382  (sm86Register (byte 6))
383
384def indexedRowScatterSM86R10 =
385  (sm86Register (byte 10))
386
387def indexedRowScatterSM86R11 =
388  (sm86Register (byte 11))
389
390def indexedRowScatterSM86R14 =
391  (sm86Register (byte 14))
392
393def indexedRowScatterSM86Set0 =
394  (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0))
395
396def indexedRowScatterSM86Set3 =
397  (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier3))
398
399def indexedRowScatterSM86Set4 =
400  (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier4))
401
402def indexedRowScatterSM86WaitSpecials =
403  (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0))
404
405def indexedRowScatterSM86WaitClassId =
406  (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier3))
407
408def indexedRowScatterSM86WaitSource =
409  (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier4))
410
411def indexedRowScatterSM86WidthNatural =
412  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
413    (eliminate
414      IndexedRowScatterSM86Geometry
415      (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) . Nat)
416      geometry
417      (branch IndexedRowScatterPositiveClassGradientRows . indexedRowScatterSM86N256)
418      (branch IndexedRowScatterNegativeClassGradientRows . indexedRowScatterSM86N256)
419      (branch IndexedRowScatterTokenEmbeddingGradientRows . indexedRowScatterSM86N1024)
420      (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . indexedRowScatterSM86N512)))
421
422def indexedRowScatterSM86RowsNatural =
423  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
424    (eliminate
425      IndexedRowScatterSM86Geometry
426      (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) . Nat)
427      geometry
428      (branch IndexedRowScatterPositiveClassGradientRows . indexedRowScatterSM86N6144)
429      (branch IndexedRowScatterNegativeClassGradientRows . indexedRowScatterSM86N4096)
430      (branch IndexedRowScatterTokenEmbeddingGradientRows . indexedRowScatterSM86N6144)
431      (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . indexedRowScatterSM86N1024)))
432
433def indexedRowScatterSM86WidthImmediate =
434  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
435    (eliminate
436      IndexedRowScatterSM86Geometry
437      (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) .
438        (family SM86Unsigned32))
439      geometry
440      (branch IndexedRowScatterPositiveClassGradientRows . indexedRowScatterSM86U256)
441      (branch IndexedRowScatterNegativeClassGradientRows . indexedRowScatterSM86U256)
442      (branch IndexedRowScatterTokenEmbeddingGradientRows . indexedRowScatterSM86U1024)
443      (branch IndexedRowScatterTokenEmbeddingGradientRows512x1024 . indexedRowScatterSM86U512)))
444
445def indexedRowScatterSM86SourceIndex =
446  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
447    (lambda unrestricted row : Nat .
448      (lambda unrestricted component : Nat .
449        (nat-eliminate
450          (lambda unrestricted rowInBounds : Nat . (family IndexedRowScatterSM86IndexResult))
451          (constructor
452            IndexedRowScatterSM86IndexResult
453            IndexedRowScatterIndexInvalid
454            (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterSourceRowOutOfBounds)
455            row
456            component)
457          (lambda unrestricted rowPredecessor : Nat .
458            (lambda unrestricted rowInduction : (family IndexedRowScatterSM86IndexResult) .
459              (nat-eliminate
460                (lambda unrestricted componentInBounds : Nat .
461                  (family IndexedRowScatterSM86IndexResult))
462                (constructor
463                  IndexedRowScatterSM86IndexResult
464                  IndexedRowScatterIndexInvalid
465                  (constructor
466                    IndexedRowScatterSM86FailureCode
467                    IndexedRowScatterSourceComponentOutOfBounds)
468                  row
469                  component)
470                (lambda unrestricted componentPredecessor : Nat .
471                  (lambda unrestricted componentInduction : (family IndexedRowScatterSM86IndexResult) .
472                    (constructor
473                      IndexedRowScatterSM86IndexResult
474                      IndexedRowScatterIndexValid
475                      (naturalAdd
476                        (naturalMultiply row (indexedRowScatterSM86WidthNatural geometry))
477                        component))))
478                (naturalLess component (indexedRowScatterSM86WidthNatural geometry)))))
479          (naturalLess row (indexedRowScatterSM86RowsNatural geometry))))))
480
481def indexedRowScatterSM86DestinationIndex =
482  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
483    (lambda unrestricted classId : Nat .
484      (lambda unrestricted component : Nat .
485        (nat-eliminate
486          (lambda unrestricted classIdInBounds : Nat . (family IndexedRowScatterSM86IndexResult))
487          (constructor
488            IndexedRowScatterSM86IndexResult
489            IndexedRowScatterIndexInvalid
490            (constructor IndexedRowScatterSM86FailureCode IndexedRowScatterClassIdOutOfBounds)
491            classId
492            component)
493          (lambda unrestricted classIdPredecessor : Nat .
494            (lambda unrestricted classIdInduction : (family IndexedRowScatterSM86IndexResult) .
495              (nat-eliminate
496                (lambda unrestricted componentInBounds : Nat .
497                  (family IndexedRowScatterSM86IndexResult))
498                (constructor
499                  IndexedRowScatterSM86IndexResult
500                  IndexedRowScatterIndexInvalid
501                  (constructor
502                    IndexedRowScatterSM86FailureCode
503                    IndexedRowScatterDestinationComponentOutOfBounds)
504                  classId
505                  component)
506                (lambda unrestricted componentPredecessor : Nat .
507                  (lambda unrestricted componentInduction : (family IndexedRowScatterSM86IndexResult) .
508                    (constructor
509                      IndexedRowScatterSM86IndexResult
510                      IndexedRowScatterIndexValid
511                      (naturalAdd
512                        (naturalMultiply classId (indexedRowScatterSM86WidthNatural geometry))
513                        component))))
514                (naturalLess component (indexedRowScatterSM86WidthNatural geometry)))))
515          (naturalLess classId indexedRowScatterSM86N12288)))))
516
517def indexedRowScatterSM86ProgramWithWidth =
518  (lambda unrestricted width : (family SM86Unsigned32) .
519    (constructor
520      SM86Program
521      SM86ProgramNext
522      (sm86Instruction
523        (constructor
524          SM86InstructionBody
525          SM86SpecialToRegister
526          indexedRowScatterSM86R0
527          (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
528          indexedRowScatterSM86Set0))
529      (constructor
530        SM86Program
531        SM86ProgramNext
532        (sm86Instruction
533          (constructor
534            SM86InstructionBody
535            SM86SpecialToRegister
536            indexedRowScatterSM86R1
537            (constructor SM86SpecialRegister SM86ThreadIdX)
538            indexedRowScatterSM86Set0))
539        (constructor
540          SM86Program
541          SM86ProgramNext
542          (sm86Instruction
543            (constructor
544              SM86InstructionBody
545              SM86MoveImmediate
546              indexedRowScatterSM86R5
547              indexedRowScatterSM86U4
548              sm86SafeControl))
549          (constructor
550            SM86Program
551            SM86ProgramNext
552            (sm86Instruction
553              (constructor
554                SM86InstructionBody
555                SM86IntegerMultiplyAddWideConstant
556                indexedRowScatterSM86R6
557                indexedRowScatterSM86R0
558                indexedRowScatterSM86R5
559                (byte 0)
560                indexedRowScatterSM86ClassIdPointer
561                indexedRowScatterSM86WaitSpecials))
562            (constructor
563              SM86Program
564              SM86ProgramNext
565              (sm86Instruction
566                (constructor
567                  SM86InstructionBody
568                  SM86LoadGlobal
569                  indexedRowScatterSM86R11
570                  indexedRowScatterSM86R6
571                  indexedRowScatterSM86U0
572                  indexedRowScatterSM86Set3))
573              (constructor
574                SM86Program
575                SM86ProgramNext
576                (sm86Instruction
577                  (constructor
578                    SM86InstructionBody
579                    SM86IntegerMultiplyAddImmediate
580                    indexedRowScatterSM86R2
581                    indexedRowScatterSM86R0
582                    width
583                    indexedRowScatterSM86R1
584                    sm86SafeControl))
585                (constructor
586                  SM86Program
587                  SM86ProgramNext
588                  (sm86Instruction
589                    (constructor
590                      SM86InstructionBody
591                      SM86IntegerMultiplyAddWideConstant
592                      indexedRowScatterSM86R14
593                      indexedRowScatterSM86R2
594                      indexedRowScatterSM86R5
595                      (byte 0)
596                      indexedRowScatterSM86SourcePointer
597                      sm86SafeControl))
598                  (constructor
599                    SM86Program
600                    SM86ProgramNext
601                    (sm86Instruction
602                      (constructor
603                        SM86InstructionBody
604                        SM86LoadGlobal
605                        indexedRowScatterSM86R10
606                        indexedRowScatterSM86R14
607                        indexedRowScatterSM86U0
608                        indexedRowScatterSM86Set4))
609                    (constructor
610                      SM86Program
611                      SM86ProgramNext
612                      (sm86Instruction
613                        (constructor
614                          SM86InstructionBody
615                          SM86IntegerMultiplyAddImmediate
616                          indexedRowScatterSM86R3
617                          indexedRowScatterSM86R11
618                          width
619                          indexedRowScatterSM86R1
620                          indexedRowScatterSM86WaitClassId))
621                      (constructor
622                        SM86Program
623                        SM86ProgramNext
624                        (sm86Instruction
625                          (constructor
626                            SM86InstructionBody
627                            SM86IntegerMultiplyAddWideConstant
628                            indexedRowScatterSM86R6
629                            indexedRowScatterSM86R3
630                            indexedRowScatterSM86R5
631                            (byte 0)
632                            indexedRowScatterSM86DestinationPointer
633                            sm86SafeControl))
634                        (constructor
635                          SM86Program
636                          SM86ProgramNext
637                          (sm86Instruction
638                            (constructor
639                              SM86InstructionBody
640                              SM86ReduceGlobalAddFloat32
641                              indexedRowScatterSM86R6
642                              indexedRowScatterSM86R10
643                              indexedRowScatterSM86U0
644                              indexedRowScatterSM86WaitSource))
645                          (constructor
646                            SM86Program
647                            SM86ProgramNext
648                            (sm86Instruction
649                              (constructor SM86InstructionBody SM86Exit sm86SafeControl))
650                            (constructor SM86Program SM86ProgramEnd))))))))))))))
651
652def indexedRowScatterSM86Program =
653  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
654    (indexedRowScatterSM86ProgramWithWidth (indexedRowScatterSM86WidthImmediate geometry)))
655
656def indexedRowScatterSM86ABIValue =
657  (constructor
658    IndexedRowScatterSM86ABI
659    IndexedRowScatterSM86ABIValue
660    indexedRowScatterSM86DestinationPointerNatural
661    indexedRowScatterSM86SourcePointerNatural
662    indexedRowScatterSM86ClassIdPointerNatural)
663
664def indexedRowScatterSM86ExtentsFor =
665  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
666    (constructor
667      IndexedRowScatterSM86Extents
668      IndexedRowScatterSM86ExtentsValue
669      (indexedRowScatterSM86RowsNatural geometry)
670      (indexedRowScatterSM86WidthNatural geometry)
671      (indexedRowScatterSM86RowsNatural geometry)
672      indexedRowScatterSM86N12288
673      (indexedRowScatterSM86WidthNatural geometry)))
674
675def indexedRowScatterSM86ObligationFor =
676  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
677    (eliminate
678      IndexedRowScatterSM86Geometry
679      (lambda unrestricted current : (family IndexedRowScatterSM86Geometry) .
680        (family NativeObligation))
681      geometry
682      (branch
683        IndexedRowScatterPositiveClassGradientRows
684        .
685        (constructor NativeObligation NativeClassRowScatter256))
686      (branch
687        IndexedRowScatterNegativeClassGradientRows
688        .
689        (constructor NativeObligation NativeClassRowScatter256))
690      (branch
691        IndexedRowScatterTokenEmbeddingGradientRows
692        .
693        (constructor NativeObligation NativeEmbeddingScatter1024))
694      (branch
695        IndexedRowScatterTokenEmbeddingGradientRows512x1024
696        .
697        (constructor NativeObligation NativeEmbeddingScatter512))))
698
699def indexedRowScatterSM86ManifestFor =
700  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
701    (constructor
702      IndexedRowScatterSM86Manifest
703      IndexedRowScatterSM86ManifestValue
704      (indexedRowScatterSM86ObligationFor geometry)
705      geometry
706      indexedRowScatterSM86N12
707      indexedRowScatterSM86N192
708      indexedRowScatterSM86N24
709      zero
710      (indexedRowScatterSM86RowsNatural geometry)
711      (succ zero)
712      (indexedRowScatterSM86WidthNatural geometry)
713      indexedRowScatterSM86ABIValue
714      (indexedRowScatterSM86ExtentsFor geometry)
715      (constructor
716        IndexedRowScatterSM86ReductionSemantics
717        IndexedRowScatterAtomicAddF32FTZRNStrongGPU)
718      (succ zero)
719      zero))
720
721def indexedRowScatterSM86ManifestExpectedInstructions =
722  (lambda unrestricted manifest : (family IndexedRowScatterSM86Manifest) .
723    (eliminate
724      IndexedRowScatterSM86Manifest
725      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
726      manifest
727      (branch
728        IndexedRowScatterSM86ManifestValue
729        obligation
730        geometry
731        instructions
732        encodedBytes
733        registers
734        sharedBytes
735        gridX
736        gridY
737        blockX
738        abi
739        extents
740        reduction
741        atomic
742        hostFallback
743        .
744        instructions)))
745
746def indexedRowScatterSM86ManifestExpectedBytes =
747  (lambda unrestricted manifest : (family IndexedRowScatterSM86Manifest) .
748    (eliminate
749      IndexedRowScatterSM86Manifest
750      (lambda unrestricted current : (family IndexedRowScatterSM86Manifest) . Nat)
751      manifest
752      (branch
753        IndexedRowScatterSM86ManifestValue
754        obligation
755        geometry
756        instructions
757        encodedBytes
758        registers
759        sharedBytes
760        gridX
761        gridY
762        blockX
763        abi
764        extents
765        reduction
766        atomic
767        hostFallback
768        .
769        encodedBytes)))
770
771def indexedRowScatterSM86TelemetryFor =
772  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
773    (lambda unrestricted observedInstructions : Nat .
774      (lambda unrestricted observedEncodedBytes : Nat .
775        (constructor
776          IndexedRowScatterSM86Telemetry
777          IndexedRowScatterSM86TelemetryValue
778          (indexedRowScatterSM86ManifestFor geometry)
779          observedInstructions
780          observedEncodedBytes
781          (succ zero)))))
782
783def indexedRowScatterSM86Build =
784  (lambda unrestricted geometry : (family IndexedRowScatterSM86Geometry) .
785    (app
786      (lambda unrestricted program : (family SM86Program) .
787        (app
788          (lambda unrestricted observedInstructions : Nat .
789            (app
790              (lambda unrestricted manifest : (family IndexedRowScatterSM86Manifest) .
791                (nat-eliminate
792                  (lambda unrestricted countMatched : Nat .
793                    (family IndexedRowScatterSM86BuildResult))
794                  (constructor
795                    IndexedRowScatterSM86BuildResult
796                    IndexedRowScatterSM86ContractFailed
797                    (constructor
798                      IndexedRowScatterSM86FailureCode
799                      IndexedRowScatterInstructionCountMismatch)
800                    (indexedRowScatterSM86TelemetryFor geometry observedInstructions zero))
801                  (lambda unrestricted countPredecessor : Nat .
802                    (lambda unrestricted countInduction : (family IndexedRowScatterSM86BuildResult) .
803                      (app
804                        (lambda unrestricted encoding : (family SM86ProgramEncodingResult) .
805                          (eliminate
806                            SM86ProgramEncodingResult
807                            (lambda unrestricted current : (family SM86ProgramEncodingResult) .
808                              (family IndexedRowScatterSM86BuildResult))
809                            encoding
810                            (branch
811                              SM86ProgramEncodingSucceeded
812                              bytes
813                              encodingTelemetry
814                              .
815                              (app
816                                (lambda unrestricted telemetry : (family IndexedRowScatterSM86Telemetry) .
817                                  (nat-eliminate
818                                    (lambda unrestricted bytesMatched : Nat .
819                                      (family IndexedRowScatterSM86BuildResult))
820                                    (constructor
821                                      IndexedRowScatterSM86BuildResult
822                                      IndexedRowScatterSM86ContractFailed
823                                      (constructor
824                                        IndexedRowScatterSM86FailureCode
825                                        IndexedRowScatterEncodedByteCountMismatch)
826                                      telemetry)
827                                    (lambda unrestricted bytesPredecessor : Nat .
828                                      (lambda unrestricted bytesInduction : (family IndexedRowScatterSM86BuildResult) .
829                                        (app
830                                        (lambda unrestricted identityResult : (family SHA256HexResult) .
831                                        (eliminate
832                                        SHA256HexResult
833                                        (lambda unrestricted current : (family SHA256HexResult) .
834                                        (family IndexedRowScatterSM86BuildResult))
835                                        identityResult
836                                        (branch
837                                        SHA256HexSucceeded
838                                        identity
839                                        identityTelemetry
840                                        .
841                                        (nat-eliminate
842                                        (lambda unrestricted identityLengthMatched : Nat .
843                                        (family IndexedRowScatterSM86BuildResult))
844                                        (constructor
845                                        IndexedRowScatterSM86BuildResult
846                                        IndexedRowScatterSM86ImageIdentityFailed
847                                        (constructor
848                                        IndexedRowScatterSM86FailureCode
849                                        IndexedRowScatterIdentityLengthInvalid)
850                                        identityResult
851                                        telemetry)
852                                        (lambda unrestricted identityLengthPredecessor : Nat .
853                                        (lambda unrestricted identityLengthInduction : (family IndexedRowScatterSM86BuildResult) .
854                                        (constructor
855                                        IndexedRowScatterSM86BuildResult
856                                        IndexedRowScatterSM86BuildSucceeded
857                                        bytes
858                                        identity
859                                        encodingTelemetry
860                                        identityTelemetry
861                                        telemetry)))
862                                        (naturalEqual
863                                        (bytes-length identity)
864                                        indexedRowScatterSM86N64)))
865                                        (branch
866                                        SHA256HexFailed
867                                        error
868                                        ordinal
869                                        identityTelemetry
870                                        .
871                                        (constructor
872                                        IndexedRowScatterSM86BuildResult
873                                        IndexedRowScatterSM86ImageIdentityFailed
874                                        (constructor
875                                        IndexedRowScatterSM86FailureCode
876                                        IndexedRowScatterIdentityFailed)
877                                        identityResult
878                                        telemetry))))
879                                        (sha256Hex bytes))))
880                                    (naturalEqual
881                                      (bytes-length bytes)
882                                      (indexedRowScatterSM86ManifestExpectedBytes manifest))))
883                                (indexedRowScatterSM86TelemetryFor
884                                  geometry
885                                  observedInstructions
886                                  (bytes-length bytes))))
887                            (branch
888                              SM86ProgramEncodingFailed
889                              instructionIndex
890                              failure
891                              encodingTelemetry
892                              .
893                              (constructor
894                                IndexedRowScatterSM86BuildResult
895                                IndexedRowScatterSM86ImageEncodingFailed
896                                (constructor
897                                  IndexedRowScatterSM86FailureCode
898                                  IndexedRowScatterEncodingFailed)
899                                encoding
900                                (indexedRowScatterSM86TelemetryFor
901                                  geometry
902                                  observedInstructions
903                                  zero)))))
904                        (sm86EncodeProgram program))))
905                  (naturalEqual
906                    observedInstructions
907                    (indexedRowScatterSM86ManifestExpectedInstructions manifest))))
908              (indexedRowScatterSM86ManifestFor geometry)))
909          (sm86ProgramCount program)))
910      (indexedRowScatterSM86Program geometry)))
911
912def indexedRowScatterSM86BuildPositiveClassGradientRows =
913  (indexedRowScatterSM86Build
914    (constructor IndexedRowScatterSM86Geometry IndexedRowScatterPositiveClassGradientRows))
915
916def indexedRowScatterSM86BuildNegativeClassGradientRows =
917  (indexedRowScatterSM86Build
918    (constructor IndexedRowScatterSM86Geometry IndexedRowScatterNegativeClassGradientRows))
919
920def indexedRowScatterSM86BuildTokenEmbeddingGradientRows =
921  (indexedRowScatterSM86Build
922    (constructor IndexedRowScatterSM86Geometry IndexedRowScatterTokenEmbeddingGradientRows))
923
924def indexedRowScatterSM86BuildTokenEmbeddingGradientRows512x1024 =
925  (indexedRowScatterSM86Build
926    (constructor IndexedRowScatterSM86Geometry IndexedRowScatterTokenEmbeddingGradientRows512x1024))

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.