Source/Packages

Realization.Nvidia.SM86.Cast.SM86

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

629 lines93 declarations24.6 KiBSHA-256 bc3a43948136

Complete file · line 54

SM86.alpha

Definition view
1module Realization.Nvidia.SM86.Cast.SM86
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 Compiler.ApplicationBuilder
10import Data.SHA256Digest
11import Std.Natural
12import Std.Byte
13
14family CastDirection : Type 0
15constructor Float32ToFloat16
16constructor Float16ToFloat32
17
18end-family
19
20family CastGeometry : Type 0
21constructor CastGeometryValue
22field unrestricted castElementCount : Nat
23field unrestricted castBlockThreads : Nat
24
25end-family
26
27family CastSM86ErrorCode : Type 0
28constructor CastSM86ElementCountZero
29constructor CastSM86ElementCountOdd
30constructor CastSM86BlockThreadsZero
31constructor CastSM86BlockThreadsTooLarge
32constructor CastSM86PairCountNotDivisible
33constructor CastSM86InstructionCountMismatch
34constructor CastSM86EncodingFailed
35constructor CastSM86ByteCountMismatch
36constructor CastSM86IdentityInvalid
37constructor CastSM86HostFallbackObserved
38
39end-family
40
41family CastSM86Telemetry : Type 0
42constructor CastSM86TelemetryValue
43field unrestricted castTelemetryDirection : (family CastDirection)
44field unrestricted castTelemetryElements : Nat
45field unrestricted castTelemetryPairs : Nat
46field unrestricted castTelemetryBlockThreads : Nat
47field unrestricted castTelemetryGridX : Nat
48field unrestricted castTelemetryExpectedInstructions : Nat
49field unrestricted castTelemetryActualInstructions : Nat
50field unrestricted castTelemetryExpectedBytes : Nat
51field unrestricted castTelemetryActualBytes : Nat
52field unrestricted castTelemetryRegisters : Nat
53field unrestricted castTelemetryGlobalLoads : Nat
54field unrestricted castTelemetryGlobalStores : Nat
55field unrestricted castTelemetryConversions : Nat
56field unrestricted castTelemetryHostFallbackCalls : Nat
57field unrestricted castTelemetryApplicationBuilderContract : Nat
58
59end-family
60
61family CastSM86BuildResult : Type 0
62constructor CastSM86BuildReady
63field unrestricted castReadyDirection : (family CastDirection)
64field unrestricted castReadyGeometry : (family CastGeometry)
65field unrestricted castReadyGridX : Nat
66field unrestricted castReadyProgram : (family SM86Program)
67field unrestricted castReadyImage : Bytes
68field unrestricted castReadySHA256 : Bytes
69field unrestricted castReadyTelemetry : (family CastSM86Telemetry)
70constructor CastSM86BuildRejected
71field unrestricted castRejectedCode : (family CastSM86ErrorCode)
72field unrestricted castRejectedOrdinal : Nat
73field unrestricted castRejectedDetail : Bytes
74
75end-family
76
77family CastGeometryValidation : Type 0
78constructor CastGeometryAccepted
79field unrestricted castAcceptedPairs : Nat
80field unrestricted castAcceptedGridX : Nat
81constructor CastGeometryRejected
82field unrestricted castGeometryRejectedCode : (family CastSM86ErrorCode)
83
84end-family
85
86def castSM86N1 =
87  (succ zero)
88
89def castSM86N2 =
90  (byte-to-nat (byte 2))
91
92def castSM86N13 =
93  (byte-to-nat (byte 13))
94
95def castSM86N14 =
96  (byte-to-nat (byte 14))
97
98def castSM86N16 =
99  (byte-to-nat (byte 16))
100
101def castSM86N24 =
102  (byte-to-nat (byte 24))
103
104def castSM86N1024 =
105  (naturalMultiply byteNaturalTwoHundredFiftySix (byte-to-nat (byte 4)))
106
107-- The instruction stream reads the pair stride from the driver's CB0 block-X
108-- word and reads output/input pointers from the first two launch arguments.
109-- WholeProgramPlan owns that driver's word offset; callers patch it with the
110-- admitted block size before launching the image.
111def castSM86OutputArgument : Nat = 0
112def castSM86InputArgument : Nat = 1
113def castSM86ArgumentCount : Nat = 2
114def castSM86PairsPerThread : Nat = 1
115def castSM86BlockThreads : Nat = byteNaturalTwoHundredFiftySix
116
117def castSM86Word32 =
118  (lambda unrestricted b0 : Byte .
119    (lambda unrestricted b1 : Byte .
120      (lambda unrestricted b2 : Byte .
121        (lambda unrestricted b3 : Byte . (sm86Unsigned32 b0 b1 b2 b3)))))
122
123def castSM86Zero32 =
124  (castSM86Word32 (byte 0) (byte 0) (byte 0) (byte 0))
125
126def castSM86Two32 =
127  (castSM86Word32 (byte 2) (byte 0) (byte 0) (byte 0))
128
129def castSM86Four32 =
130  (castSM86Word32 (byte 4) (byte 0) (byte 0) (byte 0))
131
132def castSM86ConstantBase =
133  (castSM86Word32 (byte 40) (byte 0) (byte 0) (byte 0))
134
135-- Constant-bank word zero is the CTA stride in element pairs.  A 256-thread
136-- launch therefore binds 256 here; leaving it zero repeats the first tile.
137def castSM86OutputPointer =
138  (castSM86Word32 (byte 96) (byte 1) (byte 0) (byte 0))
139
140def castSM86InputPointer =
141  (castSM86Word32 (byte 104) (byte 1) (byte 0) (byte 0))
142
143def castSM86Register =
144  (lambda unrestricted index : Byte . (sm86Register index))
145
146def castSM86Instruction =
147  (lambda unrestricted body : (family SM86InstructionBody) .
148    (constructor
149      SM86Instruction
150      SM86InstructionValue
151      (constructor SM86InstructionGuard SM86InstructionAlways)
152      body))
153
154def castSM86Next =
155  (lambda unrestricted body : (family SM86InstructionBody) .
156    (lambda unrestricted tail : (family SM86Program) .
157      (constructor SM86Program SM86ProgramNext (castSM86Instruction body) tail)))
158
159def castSM86End =
160  (constructor SM86Program SM86ProgramEnd)
161
162def castSM86Set0 =
163  (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0))
164
165def castSM86Wait0 =
166  (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0))
167
168def castSM86Prefix =
169  (lambda unrestricted tail : (family SM86Program) .
170    (castSM86Next
171      (constructor
172        SM86InstructionBody
173        SM86MoveConstant
174        (castSM86Register (byte 1))
175        (byte 0)
176        castSM86ConstantBase
177        sm86SafeControl)
178      (castSM86Next
179        (constructor
180          SM86InstructionBody
181          SM86SpecialToRegister
182          (castSM86Register (byte 0))
183          (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
184          castSM86Set0)
185        (castSM86Next
186          (constructor
187            SM86InstructionBody
188            SM86MoveImmediate
189            (castSM86Register (byte 3))
190            castSM86Four32
191            castSM86Wait0)
192          (castSM86Next
193            (constructor
194              SM86InstructionBody
195              SM86SpecialToRegister
196              (castSM86Register (byte 2))
197              (constructor SM86SpecialRegister SM86ThreadIdX)
198              castSM86Set0)
199            (castSM86Next
200              (constructor
201                SM86InstructionBody
202                SM86IntegerMultiplyAddConstant
203                (castSM86Register (byte 0))
204                (castSM86Register (byte 0))
205                (byte 0)
206                castSM86Zero32
207                (castSM86Register (byte 2))
208                castSM86Wait0)
209              (castSM86Next
210                (constructor
211                  SM86InstructionBody
212                  SM86IntegerMultiplyAddImmediate
213                  (castSM86Register (byte 4))
214                  (castSM86Register (byte 0))
215                  castSM86Two32
216                  sm86ZeroRegister
217                  sm86SafeControl)
218                tail)))))))
219
220def castSM86F32ToF16Tail : (family SM86Program) =
221  (castSM86Next
222    (constructor
223      SM86InstructionBody
224      SM86IntegerMultiplyAddWideConstant
225      (castSM86Register (byte 6))
226      (castSM86Register (byte 4))
227      (castSM86Register (byte 3))
228      (byte 0)
229      castSM86InputPointer
230      sm86SafeControl)
231    (castSM86Next
232      (constructor
233        SM86InstructionBody
234        SM86LoadGlobal
235        (castSM86Register (byte 8))
236        (castSM86Register (byte 6))
237        castSM86Zero32
238        castSM86Set0)
239      (castSM86Next
240        (constructor
241          SM86InstructionBody
242          SM86LoadGlobal
243          (castSM86Register (byte 9))
244          (castSM86Register (byte 6))
245          castSM86Four32
246          castSM86Set0)
247        (castSM86Next
248          (constructor
249            SM86InstructionBody
250            SM86FloatPairToPackedHalfPair
251            (castSM86Register (byte 10))
252            (castSM86Register (byte 9))
253            (castSM86Register (byte 8))
254            castSM86Wait0)
255          (castSM86Next
256            (constructor
257              SM86InstructionBody
258              SM86IntegerMultiplyAddWideConstant
259              (castSM86Register (byte 12))
260              (castSM86Register (byte 0))
261              (castSM86Register (byte 3))
262              (byte 0)
263              castSM86OutputPointer
264              sm86SafeControl)
265            (castSM86Next
266              (constructor
267                SM86InstructionBody
268                SM86StoreGlobal
269                (castSM86Register (byte 12))
270                (castSM86Register (byte 10))
271                castSM86Zero32
272                sm86SafeControl)
273              (castSM86Next (constructor SM86InstructionBody SM86Exit sm86SafeControl) castSM86End)))))))
274
275def castSM86F16ToF32Tail : (family SM86Program) =
276  (castSM86Next
277    (constructor
278      SM86InstructionBody
279      SM86IntegerMultiplyAddWideConstant
280      (castSM86Register (byte 6))
281      (castSM86Register (byte 0))
282      (castSM86Register (byte 3))
283      (byte 0)
284      castSM86InputPointer
285      sm86SafeControl)
286    (castSM86Next
287      (constructor
288        SM86InstructionBody
289        SM86LoadGlobal
290        (castSM86Register (byte 10))
291        (castSM86Register (byte 6))
292        castSM86Zero32
293        castSM86Set0)
294      (castSM86Next
295        (constructor
296          SM86InstructionBody
297          SM86HalfToFloat
298          (castSM86Register (byte 14))
299          (castSM86Register (byte 10))
300          (constructor SM86HalfSelector SM86LowHalf)
301          castSM86Wait0)
302        (castSM86Next
303          (constructor
304            SM86InstructionBody
305            SM86HalfToFloat
306            (castSM86Register (byte 15))
307            (castSM86Register (byte 10))
308            (constructor SM86HalfSelector SM86HighHalf)
309            sm86SafeControl)
310          (castSM86Next
311            (constructor
312              SM86InstructionBody
313              SM86IntegerMultiplyAddWideConstant
314              (castSM86Register (byte 12))
315              (castSM86Register (byte 4))
316              (castSM86Register (byte 3))
317              (byte 0)
318              castSM86OutputPointer
319              sm86SafeControl)
320            (castSM86Next
321              (constructor
322                SM86InstructionBody
323                SM86StoreGlobal
324                (castSM86Register (byte 12))
325                (castSM86Register (byte 14))
326                castSM86Zero32
327                sm86SafeControl)
328              (castSM86Next
329                (constructor
330                  SM86InstructionBody
331                  SM86StoreGlobal
332                  (castSM86Register (byte 12))
333                  (castSM86Register (byte 15))
334                  castSM86Four32
335                  sm86SafeControl)
336                (castSM86Next
337                  (constructor SM86InstructionBody SM86Exit sm86SafeControl)
338                  castSM86End))))))))
339
340def emitCastSM86 =
341  (lambda unrestricted direction : (family CastDirection) .
342    (eliminate
343      CastDirection
344      (lambda unrestricted current : (family CastDirection) . (family SM86Program))
345      direction
346      (branch Float32ToFloat16 . (castSM86Prefix castSM86F32ToF16Tail))
347      (branch Float16ToFloat32 . (castSM86Prefix castSM86F16ToF32Tail))))
348
349def castInstructionCount =
350  (lambda unrestricted direction : (family CastDirection) .
351    (eliminate
352      CastDirection
353      (lambda unrestricted current : (family CastDirection) . Nat)
354      direction
355      (branch Float32ToFloat16 . castSM86N13)
356      (branch Float16ToFloat32 . castSM86N14)))
357
358def castRegisterCount : Nat =
359  castSM86N24
360
361def castSM86HostFallbackCalls : Nat =
362  zero
363
364def castSM86ApplicationBuilderContract : Nat =
365  Compiler.ApplicationBuilder/applicationBuilderNativeOnly
366
367def castSM86ErrorStableCode =
368  (lambda unrestricted code : (family CastSM86ErrorCode) .
369    (eliminate
370      CastSM86ErrorCode
371      (lambda unrestricted current : (family CastSM86ErrorCode) . Bytes)
372      code
373      (branch CastSM86ElementCountZero . b"CAST-001")
374      (branch CastSM86ElementCountOdd . b"CAST-002")
375      (branch CastSM86BlockThreadsZero . b"CAST-003")
376      (branch CastSM86BlockThreadsTooLarge . b"CAST-004")
377      (branch CastSM86PairCountNotDivisible . b"CAST-005")
378      (branch CastSM86InstructionCountMismatch . b"CAST-006")
379      (branch CastSM86EncodingFailed . b"CAST-007")
380      (branch CastSM86ByteCountMismatch . b"CAST-008")
381      (branch CastSM86IdentityInvalid . b"CAST-009")
382      (branch CastSM86HostFallbackObserved . b"CAST-010")))
383
384def validateCastGeometry =
385  (lambda unrestricted geometry : (family CastGeometry) .
386    (eliminate
387      CastGeometry
388      (lambda unrestricted current : (family CastGeometry) . (family CastGeometryValidation))
389      geometry
390      (branch
391        CastGeometryValue
392        elements
393        blockThreads
394        .
395        (nat-eliminate
396          (lambda unrestricted elementsNonzero : Nat . (family CastGeometryValidation))
397          (constructor
398            CastGeometryValidation
399            CastGeometryRejected
400            (constructor CastSM86ErrorCode CastSM86ElementCountZero))
401          (lambda unrestricted elementPredecessor : Nat .
402            (lambda unrestricted elementInduction : (family CastGeometryValidation) .
403              (nat-eliminate
404                (lambda unrestricted evenElements : Nat . (family CastGeometryValidation))
405                (constructor
406                  CastGeometryValidation
407                  CastGeometryRejected
408                  (constructor CastSM86ErrorCode CastSM86ElementCountOdd))
409                (lambda unrestricted evenPredecessor : Nat .
410                  (lambda unrestricted evenInduction : (family CastGeometryValidation) .
411                    (nat-eliminate
412                      (lambda unrestricted blockNonzero : Nat . (family CastGeometryValidation))
413                      (constructor
414                        CastGeometryValidation
415                        CastGeometryRejected
416                        (constructor CastSM86ErrorCode CastSM86BlockThreadsZero))
417                      (lambda unrestricted blockPredecessor : Nat .
418                        (lambda unrestricted blockInduction : (family CastGeometryValidation) .
419                          (nat-eliminate
420                            (lambda unrestricted blockFits : Nat . (family CastGeometryValidation))
421                            (constructor
422                              CastGeometryValidation
423                              CastGeometryRejected
424                              (constructor CastSM86ErrorCode CastSM86BlockThreadsTooLarge))
425                            (lambda unrestricted fitsPredecessor : Nat .
426                              (lambda unrestricted fitsInduction : (family CastGeometryValidation) .
427                                (app
428                                  (lambda unrestricted pairs : Nat .
429                                    (nat-eliminate
430                                      (lambda unrestricted divisible : Nat .
431                                        (family CastGeometryValidation))
432                                      (constructor
433                                        CastGeometryValidation
434                                        CastGeometryRejected
435                                        (constructor
436                                        CastSM86ErrorCode
437                                        CastSM86PairCountNotDivisible))
438                                      (lambda unrestricted divisiblePredecessor : Nat .
439                                        (lambda unrestricted divisibleInduction : (family CastGeometryValidation) .
440                                        (constructor
441                                        CastGeometryValidation
442                                        CastGeometryAccepted
443                                        pairs
444                                        (naturalDivideUnchecked pairs blockThreads))))
445                                      (naturalIsZero (naturalModuloUnchecked pairs blockThreads))))
446                                  (naturalDivideUnchecked elements castSM86N2))))
447                            (naturalLessOrEqual blockThreads castSM86N1024))))
448                      (naturalNonzero blockThreads))))
449                (naturalIsZero (naturalModuloUnchecked elements castSM86N2)))))
450          (naturalNonzero elements)))))
451
452def castGridX =
453  (lambda unrestricted geometry : (family CastGeometry) . (validateCastGeometry geometry))
454
455-- Only an accepted geometry supplies a launch grid. A rejected geometry
456-- evaluates to zero, so callers must prove a nonzero grid before emission.
457def castSM86AcceptedGridX =
458  (lambda unrestricted geometry : (family CastGeometry) .
459    (eliminate CastGeometryValidation
460      (lambda unrestricted current : (family CastGeometryValidation) . Nat)
461      (validateCastGeometry geometry)
462      (branch CastGeometryAccepted pairs gridX . gridX)
463      (branch CastGeometryRejected code . 0)))
464
465def castSM86TelemetryFor =
466  (lambda unrestricted direction : (family CastDirection) .
467    (lambda unrestricted geometry : (family CastGeometry) .
468      (lambda unrestricted pairs : Nat .
469        (lambda unrestricted gridX : Nat .
470          (lambda unrestricted actualInstructions : Nat .
471            (lambda unrestricted actualBytes : Nat .
472              (eliminate
473                CastGeometry
474                (lambda unrestricted current : (family CastGeometry) . (family CastSM86Telemetry))
475                geometry
476                (branch
477                  CastGeometryValue
478                  elements
479                  blockThreads
480                  .
481                  (constructor
482                    CastSM86Telemetry
483                    CastSM86TelemetryValue
484                    direction
485                    elements
486                    pairs
487                    blockThreads
488                    gridX
489                    (castInstructionCount direction)
490                    actualInstructions
491                    (naturalMultiply (castInstructionCount direction) castSM86N16)
492                    actualBytes
493                    castRegisterCount
494                    (eliminate
495                      CastDirection
496                      (lambda unrestricted current : (family CastDirection) . Nat)
497                      direction
498                      (branch Float32ToFloat16 . castSM86N2)
499                      (branch Float16ToFloat32 . castSM86N1))
500                    (eliminate
501                      CastDirection
502                      (lambda unrestricted current : (family CastDirection) . Nat)
503                      direction
504                      (branch Float32ToFloat16 . castSM86N1)
505                      (branch Float16ToFloat32 . castSM86N2))
506                    (eliminate
507                      CastDirection
508                      (lambda unrestricted current : (family CastDirection) . Nat)
509                      direction
510                      (branch Float32ToFloat16 . castSM86N1)
511                      (branch Float16ToFloat32 . castSM86N2))
512                    castSM86HostFallbackCalls
513                    castSM86ApplicationBuilderContract)))))))))
514
515def castImageSHA256 =
516  (lambda unrestricted direction : (family CastDirection) .
517    (lambda unrestricted geometry : (family CastGeometry) .
518      (eliminate
519        CastGeometryValidation
520        (lambda unrestricted current : (family CastGeometryValidation) .
521          (family CastSM86BuildResult))
522        (validateCastGeometry geometry)
523        (branch
524          CastGeometryAccepted
525          pairs
526          gridX
527          .
528          (app
529            (lambda unrestricted program : (family SM86Program) .
530              (app
531                (lambda unrestricted actualInstructions : Nat .
532                  (nat-eliminate
533                    (lambda unrestricted countMatches : Nat . (family CastSM86BuildResult))
534                    (constructor
535                      CastSM86BuildResult
536                      CastSM86BuildRejected
537                      (constructor CastSM86ErrorCode CastSM86InstructionCountMismatch)
538                      actualInstructions
539                      (castSM86ErrorStableCode
540                        (constructor CastSM86ErrorCode CastSM86InstructionCountMismatch)))
541                    (lambda unrestricted countPredecessor : Nat .
542                      (lambda unrestricted countInduction : (family CastSM86BuildResult) .
543                        (eliminate
544                          SM86ProgramEncodingResult
545                          (lambda unrestricted current : (family SM86ProgramEncodingResult) .
546                            (family CastSM86BuildResult))
547                          (sm86EncodeProgram program)
548                          (branch
549                            SM86ProgramEncodingSucceeded
550                            image
551                            encodingTelemetry
552                            .
553                            (app
554                              (lambda unrestricted actualBytes : Nat .
555                                (nat-eliminate
556                                  (lambda unrestricted bytesMatch : Nat .
557                                    (family CastSM86BuildResult))
558                                  (constructor
559                                    CastSM86BuildResult
560                                    CastSM86BuildRejected
561                                    (constructor CastSM86ErrorCode CastSM86ByteCountMismatch)
562                                    actualBytes
563                                    (castSM86ErrorStableCode
564                                      (constructor CastSM86ErrorCode CastSM86ByteCountMismatch)))
565                                  (lambda unrestricted bytePredecessor : Nat .
566                                    (lambda unrestricted byteInduction : (family CastSM86BuildResult) .
567                                      (app
568                                        (lambda unrestricted identity : Bytes .
569                                        (nat-eliminate
570                                        (lambda unrestricted identityValid : Nat .
571                                        (family CastSM86BuildResult))
572                                        (constructor
573                                        CastSM86BuildResult
574                                        CastSM86BuildRejected
575                                        (constructor CastSM86ErrorCode CastSM86IdentityInvalid)
576                                        zero
577                                        (castSM86ErrorStableCode
578                                        (constructor CastSM86ErrorCode CastSM86IdentityInvalid)))
579                                        (lambda unrestricted identityPredecessor : Nat .
580                                        (lambda unrestricted identityInduction : (family CastSM86BuildResult) .
581                                        (constructor
582                                        CastSM86BuildResult
583                                        CastSM86BuildReady
584                                        direction
585                                        geometry
586                                        gridX
587                                        program
588                                        image
589                                        identity
590                                        (castSM86TelemetryFor
591                                        direction
592                                        geometry
593                                        pairs
594                                        gridX
595                                        actualInstructions
596                                        actualBytes))))
597                                        (naturalEqual
598                                        (bytes-length identity)
599                                        (byte-to-nat (byte 64)))))
600                                        (sha256HexBytesOrEmpty (sha256Hex image)))))
601                                  (naturalEqual
602                                    actualBytes
603                                    (naturalMultiply (castInstructionCount direction) castSM86N16))))
604                              (bytes-length image)))
605                          (branch
606                            SM86ProgramEncodingFailed
607                            ordinal
608                            failure
609                            encodingTelemetry
610                            .
611                            (constructor
612                              CastSM86BuildResult
613                              CastSM86BuildRejected
614                              (constructor CastSM86ErrorCode CastSM86EncodingFailed)
615                              ordinal
616                              (sm86InstructionEncodingStableCode failure))))))
617                    (naturalEqual actualInstructions (castInstructionCount direction))))
618                (sm86ProgramCount program)))
619            (emitCastSM86 direction)))
620        (branch
621          CastGeometryRejected
622          code
623          .
624          (constructor
625            CastSM86BuildResult
626            CastSM86BuildRejected
627            code
628            zero
629            (castSM86ErrorStableCode code))))))

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.