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.