Source/Packages

Realization.Nvidia.SM86.GEMM.HMMA.RegisterTileSM86

packages/realizations/cooperative/nvidia-sm86/src/Realization/Nvidia/SM86/GEMM/HMMA/RegisterTileSM86.alpha

416 lines83 declarations15.7 KiBSHA-256 bfab13489a67

Complete file · line 58

RegisterTileSM86.alpha

Definition view
1module Realization.Nvidia.SM86.GEMM.HMMA.RegisterTileSM86
2
3import Compiler.ApplicationBuilder
4import Realization.Nvidia.SM86.HMMAProductionSM86
5import Std.Natural
6
7family RegisterTileHMMAGeometries : Type 0
8constructor RegisterTileHMMAGeometriesEnd
9constructor RegisterTileHMMAGeometriesNext
10field unrestricted registerTileHMMAHead : (family HMMAProductionSM86Geometry)
11recursive unrestricted registerTileHMMATail
12
13end-family
14
15family RegisterTileHMMAErrorCode : Type 0
16constructor RegisterTileHMMANoError
17constructor RegisterTileHMMAGeometryInvalid
18constructor RegisterTileHMMAInstructionCountMismatch
19constructor RegisterTileHMMAEncodedByteCountMismatch
20constructor RegisterTileHMMAEncodingFailed
21constructor RegisterTileHMMAIdentityFailed
22constructor RegisterTileHMMAHostFallbackForbidden
23constructor RegisterTileHMMAApplicationBuilderContractInvalid
24
25end-family
26
27family RegisterTileHMMATelemetry : Type 0
28constructor RegisterTileHMMATelemetryValue
29field unrestricted registerTileHMMATelemetryM : Nat
30field unrestricted registerTileHMMATelemetryN : Nat
31field unrestricted registerTileHMMATelemetryK : Nat
32field unrestricted registerTileHMMATelemetryKStages : Nat
33field unrestricted registerTileHMMATelemetryGridX : Nat
34field unrestricted registerTileHMMATelemetryGridY : Nat
35field unrestricted registerTileHMMATelemetryGridZ : Nat
36field unrestricted registerTileHMMATelemetryBlockX : Nat
37field unrestricted registerTileHMMATelemetryInstructions : Nat
38field unrestricted registerTileHMMATelemetryEncodedBytes : Nat
39field unrestricted registerTileHMMATelemetryRegisters : Nat
40field unrestricted registerTileHMMATelemetrySharedBytes : Nat
41field unrestricted registerTileHMMATelemetryHostFallbackOperations : Nat
42field unrestricted registerTileHMMATelemetryApplicationBuilderContract : Nat
43
44end-family
45
46family RegisterTileHMMABuildResult : Type 0
47constructor RegisterTileHMMABuildSucceeded
48field unrestricted registerTileHMMAEncodedBytes : Bytes
49field unrestricted registerTileHMMAImageSHA256 : Bytes
50field unrestricted registerTileHMMANativeResult : (family HMMAProductionNativeBuildResult)
51field unrestricted registerTileHMMABuildTelemetry : (family RegisterTileHMMATelemetry)
52constructor RegisterTileHMMABuildFailed
53field unrestricted registerTileHMMAFailureCode : (family RegisterTileHMMAErrorCode)
54field unrestricted registerTileHMMAFailureTelemetry : (family RegisterTileHMMATelemetry)
55
56end-family
57
58family RegisterTileHMMAGridResult : Type 0
59constructor RegisterTileHMMAGridSucceeded
60field unrestricted registerTileHMMAGridX : Nat
61field unrestricted registerTileHMMAGridY : Nat
62field unrestricted registerTileHMMAGridZ : Nat
63field unrestricted registerTileHMMAGridTelemetry : (family RegisterTileHMMATelemetry)
64constructor RegisterTileHMMAGridFailed
65field unrestricted registerTileHMMAGridFailure : (family RegisterTileHMMAErrorCode)
66field unrestricted registerTileHMMAGridFailureTelemetry : (family RegisterTileHMMATelemetry)
67
68end-family
69
70def registerTileHMMAN8 =
71  (byte-to-nat (byte 8))
72
73def registerTileHMMAN16 =
74  (byte-to-nat (byte 16))
75
76def registerTileHMMAN21 =
77  (byte-to-nat (byte 21))
78
79def registerTileHMMAN32 =
80  (byte-to-nat (byte 32))
81
82def registerTileHMMAN40 =
83  (byte-to-nat (byte 40))
84
85def registerTileHMMAN64 =
86  (byte-to-nat (byte 64))
87
88def registerTileHMMAN69 =
89  (byte-to-nat (byte 69))
90
91def registerTileHMMAN128 =
92  (naturalMultiply registerTileHMMAN16 registerTileHMMAN8)
93
94def registerTileHMMAN256 =
95  (naturalMultiply registerTileHMMAN16 registerTileHMMAN16)
96
97def registerTileHMMAN384 =
98  (naturalMultiply registerTileHMMAN16 (byte-to-nat (byte 24)))
99
100def registerTileHMMAN640 =
101  (naturalMultiply registerTileHMMAN64 (byte-to-nat (byte 10)))
102
103def registerTileHMMAN1024 =
104  (naturalMultiply registerTileHMMAN16 registerTileHMMAN64)
105
106def registerTileHMMAN3072 =
107  (naturalMultiply (byte-to-nat (byte 48)) registerTileHMMAN64)
108
109def registerTileHMMAN4096 =
110  (naturalMultiply registerTileHMMAN64 registerTileHMMAN64)
111
112def registerTileHMMAN6144 =
113  (naturalMultiply (byte-to-nat (byte 96)) registerTileHMMAN64)
114
115def registerTileHMMAN8192 =
116  (naturalMultiply registerTileHMMAN128 registerTileHMMAN64)
117
118def registerTileHMMAHostFallbackOperations : Nat =
119  zero
120
121def registerTileHMMAErrorStableCode =
122  (lambda unrestricted code : (family RegisterTileHMMAErrorCode) .
123    (eliminate
124      RegisterTileHMMAErrorCode
125      (lambda unrestricted current : (family RegisterTileHMMAErrorCode) . Bytes)
126      code
127      (branch RegisterTileHMMANoError . b"HMMA-RT-000")
128      (branch RegisterTileHMMAGeometryInvalid . b"HMMA-RT-001")
129      (branch RegisterTileHMMAInstructionCountMismatch . b"HMMA-RT-002")
130      (branch RegisterTileHMMAEncodedByteCountMismatch . b"HMMA-RT-003")
131      (branch RegisterTileHMMAEncodingFailed . b"HMMA-RT-004")
132      (branch RegisterTileHMMAIdentityFailed . b"HMMA-RT-005")
133      (branch RegisterTileHMMAHostFallbackForbidden . b"HMMA-RT-006")
134      (branch
135        RegisterTileHMMAApplicationBuilderContractInvalid
136        .
137        b"HMMA-RT-007")))
138
139def registerTileHMMAGeometry =
140  (lambda unrestricted m : Nat .
141    (lambda unrestricted n : Nat .
142      (lambda unrestricted k : Nat .
143        (constructor
144          HMMAProductionSM86Geometry
145          HMMAProductionSM86GeometryValue
146          m
147          n
148          k
149          (naturalDivideUnchecked k registerTileHMMAN32)))))
150
151def registerTileHMMAGeometryEqual =
152  (lambda unrestricted left : (family HMMAProductionSM86Geometry) .
153    (lambda unrestricted right : (family HMMAProductionSM86Geometry) .
154      (naturalAnd
155        (naturalEqual (hmmaNativeGeometryM left) (hmmaNativeGeometryM right))
156        (naturalAnd
157          (naturalEqual (hmmaNativeGeometryN left) (hmmaNativeGeometryN right))
158          (naturalAnd
159            (naturalEqual (hmmaNativeGeometryK left) (hmmaNativeGeometryK right))
160            (naturalEqual (hmmaNativeGeometryStages left) (hmmaNativeGeometryStages right)))))))
161
162def registerTileHMMAInstructionCount =
163  (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
164    (naturalAdd
165      registerTileHMMAN69
166      (naturalMultiply registerTileHMMAN21 (hmmaNativeGeometryStages geometry))))
167
168def registerTileHMMATelemetryFor =
169  (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
170    (app
171      (lambda unrestricted instructions : Nat .
172        (constructor
173          RegisterTileHMMATelemetry
174          RegisterTileHMMATelemetryValue
175          (hmmaNativeGeometryM geometry)
176          (hmmaNativeGeometryN geometry)
177          (hmmaNativeGeometryK geometry)
178          (hmmaNativeGeometryStages geometry)
179          (naturalDivideUnchecked (hmmaNativeGeometryM geometry) registerTileHMMAN32)
180          (naturalDivideUnchecked (hmmaNativeGeometryN geometry) registerTileHMMAN32)
181          (succ zero)
182          registerTileHMMAN128
183          instructions
184          (naturalMultiply instructions registerTileHMMAN16)
185          registerTileHMMAN32
186          registerTileHMMAN8192
187          registerTileHMMAHostFallbackOperations
188          applicationBuilderNativeOnly))
189      (registerTileHMMAInstructionCount geometry)))
190
191def registerTileHMMAFail =
192  (lambda unrestricted code : (family RegisterTileHMMAErrorCode) .
193    (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
194      (constructor
195        RegisterTileHMMABuildResult
196        RegisterTileHMMABuildFailed
197        code
198        (registerTileHMMATelemetryFor geometry))))
199
200def registerTileHMMAMapNativeResult =
201  (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
202    (lambda unrestricted result : (family HMMAProductionNativeBuildResult) .
203      (eliminate
204        HMMAProductionNativeBuildResult
205        (lambda unrestricted current : (family HMMAProductionNativeBuildResult) .
206          (family RegisterTileHMMABuildResult))
207        result
208        (branch
209          HMMAProductionNativeBuildSucceeded
210          encoded
211          identity
212          encodingTelemetry
213          identityTelemetry
214          nativeTelemetry
215          .
216          (constructor
217            RegisterTileHMMABuildResult
218            RegisterTileHMMABuildSucceeded
219            encoded
220            identity
221            result
222            (registerTileHMMATelemetryFor geometry)))
223        (branch
224          HMMAProductionNativeContractFailed
225          failure
226          nativeTelemetry
227          .
228          (eliminate
229            HMMAProductionSM86FailureCode
230            (lambda unrestricted current : (family HMMAProductionSM86FailureCode) .
231              (family RegisterTileHMMABuildResult))
232            failure
233            (branch
234              HMMAProductionInvalidGeometry
235              .
236              (registerTileHMMAFail
237                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAGeometryInvalid)
238                geometry))
239            (branch
240              HMMAProductionInvalidGrouping
241              .
242              (registerTileHMMAFail
243                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAGeometryInvalid)
244                geometry))
245            (branch
246              HMMAProductionInstructionCountMismatch
247              .
248              (registerTileHMMAFail
249                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAInstructionCountMismatch)
250                geometry))
251            (branch
252              HMMAProductionEncodedByteCountMismatch
253              .
254              (registerTileHMMAFail
255                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAEncodedByteCountMismatch)
256                geometry))
257            (branch
258              HMMAProductionEncodingFailed
259              .
260              (registerTileHMMAFail
261                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAEncodingFailed)
262                geometry))
263            (branch
264              HMMAProductionIdentityFailed
265              .
266              (registerTileHMMAFail
267                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAIdentityFailed)
268                geometry))
269            (branch
270              HMMAProductionIdentityLengthInvalid
271              .
272              (registerTileHMMAFail
273                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAIdentityFailed)
274                geometry))))
275        (branch
276          HMMAProductionNativeEncodingFailed
277          failure
278          nativeTelemetry
279          .
280          (registerTileHMMAFail
281            (constructor RegisterTileHMMAErrorCode RegisterTileHMMAEncodingFailed)
282            geometry))
283        (branch
284          HMMAProductionNativeIdentityFailed
285          failure
286          nativeTelemetry
287          .
288          (registerTileHMMAFail
289            (constructor RegisterTileHMMAErrorCode RegisterTileHMMAIdentityFailed)
290            geometry)))))
291
292def registerTileHMMAFinalize =
293  (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
294    (lambda unrestricted result : (family RegisterTileHMMABuildResult) .
295      (nat-eliminate
296        (lambda unrestricted builderValid : Nat . (family RegisterTileHMMABuildResult))
297        (registerTileHMMAFail
298          (constructor RegisterTileHMMAErrorCode RegisterTileHMMAApplicationBuilderContractInvalid)
299          geometry)
300        (lambda unrestricted builderPredecessor : Nat .
301          (lambda unrestricted builderInduction : (family RegisterTileHMMABuildResult) .
302            (nat-eliminate
303              (lambda unrestricted fallbackFree : Nat . (family RegisterTileHMMABuildResult))
304              (registerTileHMMAFail
305                (constructor RegisterTileHMMAErrorCode RegisterTileHMMAHostFallbackForbidden)
306                geometry)
307              (lambda unrestricted fallbackPredecessor : Nat .
308                (lambda unrestricted fallbackInduction : (family RegisterTileHMMABuildResult) .
309                  result))
310              (naturalEqual registerTileHMMAHostFallbackOperations zero))))
311        applicationBuilderNativeOnly)))
312
313def emitRegisterTileHMMASM86 =
314  (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
315    (registerTileHMMAFinalize
316      geometry
317      (registerTileHMMAMapNativeResult
318        geometry
319        (hmmaNativeBuild
320          (constructor HMMAProductionSM86Variant HMMAProductionDenseForward)
321          (constructor HMMAProductionSM86Grouping HMMAProductionDense)
322          (constructor HMMAProductionSM86Orientation HMMAProductionABTransposed)
323          geometry))))
324
325def registerTileHMMAGrid =
326  (lambda unrestricted geometry : (family HMMAProductionSM86Geometry) .
327    (nat-eliminate
328      (lambda unrestricted valid : Nat . (family RegisterTileHMMAGridResult))
329      (constructor
330        RegisterTileHMMAGridResult
331        RegisterTileHMMAGridFailed
332        (constructor RegisterTileHMMAErrorCode RegisterTileHMMAGeometryInvalid)
333        (registerTileHMMATelemetryFor geometry))
334      (lambda unrestricted predecessor : Nat .
335        (lambda unrestricted induction : (family RegisterTileHMMAGridResult) .
336          (constructor
337            RegisterTileHMMAGridResult
338            RegisterTileHMMAGridSucceeded
339            (naturalDivideUnchecked (hmmaNativeGeometryM geometry) registerTileHMMAN32)
340            (naturalDivideUnchecked (hmmaNativeGeometryN geometry) registerTileHMMAN32)
341            (succ zero)
342            (registerTileHMMATelemetryFor geometry))))
343      (hmmaNativeGeometryValid geometry)))
344
345def registerTileHMMASupportedGeometries : (family RegisterTileHMMAGeometries) =
346  (constructor
347    RegisterTileHMMAGeometries
348    RegisterTileHMMAGeometriesNext
349    (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN64 registerTileHMMAN1024)
350    (constructor
351      RegisterTileHMMAGeometries
352      RegisterTileHMMAGeometriesNext
353      (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN3072 registerTileHMMAN64)
354      (constructor
355        RegisterTileHMMAGeometries
356        RegisterTileHMMAGeometriesNext
357        (registerTileHMMAGeometry registerTileHMMAN6144 registerTileHMMAN1024 registerTileHMMAN64)
358        (constructor
359          RegisterTileHMMAGeometries
360          RegisterTileHMMAGeometriesNext
361          (registerTileHMMAGeometry
362            registerTileHMMAN6144
363            registerTileHMMAN256
364            registerTileHMMAN1024)
365          (constructor
366            RegisterTileHMMAGeometries
367            RegisterTileHMMAGeometriesNext
368            (registerTileHMMAGeometry
369              registerTileHMMAN6144
370              registerTileHMMAN4096
371              registerTileHMMAN256)
372            (constructor
373              RegisterTileHMMAGeometries
374              RegisterTileHMMAGeometriesNext
375              (registerTileHMMAGeometry
376                registerTileHMMAN256
377                registerTileHMMAN256
378                registerTileHMMAN256)
379              (constructor
380                RegisterTileHMMAGeometries
381                RegisterTileHMMAGeometriesNext
382                (registerTileHMMAGeometry
383                  registerTileHMMAN384
384                  registerTileHMMAN640
385                  registerTileHMMAN1024)
386                (constructor
387                  RegisterTileHMMAGeometries
388                  RegisterTileHMMAGeometriesNext
389                  (registerTileHMMAGeometry
390                    registerTileHMMAN384
391                    registerTileHMMAN1024
392                    registerTileHMMAN640)
393                  (constructor RegisterTileHMMAGeometries RegisterTileHMMAGeometriesEnd)))))))))
394
395def emitRegisterTileHMMAM256N256K256SM86 =
396  (emitRegisterTileHMMASM86
397    (registerTileHMMAGeometry registerTileHMMAN256 registerTileHMMAN256 registerTileHMMAN256))
398
399def registerTileHMMAM256N256K256InstructionCount =
400  (registerTileHMMAInstructionCount
401    (registerTileHMMAGeometry registerTileHMMAN256 registerTileHMMAN256 registerTileHMMAN256))
402
403def registerTileHMMARegisterCount =
404  registerTileHMMAN32
405
406def registerTileHMMASharedBytes =
407  registerTileHMMAN8192
408
409def registerTileHMMAM256N256K256RegisterCount =
410  registerTileHMMARegisterCount
411
412def registerTileHMMAM256N256K256SharedBytes =
413  registerTileHMMASharedBytes
414
415def registerTileHMMAApplicationBuilderContract =
416  applicationBuilderNativeOnly

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.