Source/Systems

Coppelius.Build.DeviceImages

systems/coppelius/src/Coppelius/Build/DeviceImages.alpha

726 lines107 declarations39.5 KiBSHA-256 aca105a01057

Complete file · line 319

DeviceImages.alpha

Definition view
1module Coppelius.Build.DeviceImages
2
3import Coppelius.ArenaPlan
4import Coppelius.SM86Capability
5import Accelerator.SM86.InstructionEncoding
6import Accelerator.SM86.Instruction
7import Accelerator.SM86.Immediate
8import Accelerator.SM86.Operands
9import Accelerator.SM86.Program
10import Accelerator.SM86.HalfFormat
11import Accelerator.SM121.Lowering
12import Coppelius.Learner
13import Coppelius.Model
14import Coppelius.TrainingRun
15import Accelerator.SM86.RNG.RandomNormal
16import Data.Bytes
17import Hardware.Nvidia.SM86.Command.WholeProgramPlan
18import Realization.Nvidia.SM86.AdamWHalfSM86
19import Realization.Nvidia.SM86.AttentionHeadLayoutSM86
20import Realization.Nvidia.SM86.Cast.SM86
21import Realization.Nvidia.SM86.RotaryTableSM86
22import Realization.Nvidia.SM86.Elementwise.GeneralSM86
23import Realization.Nvidia.SM86.ElementwiseVectorSM86
24import Realization.Nvidia.SM86.EmbeddingGatherF32SM86
25import Realization.Nvidia.SM86.ExactCrossEntropyBackwardSM86
26import Realization.Nvidia.SM86.ExactCrossEntropySM86
27import Realization.Nvidia.SM86.GELUSM86
28import Realization.Nvidia.SM86.GatedGELUSM86
29import Realization.Nvidia.SM86.GEMM.HMMA.RegisterTileSM86
30import Realization.Nvidia.SM86.HMMAProductionSM86
31import Realization.Nvidia.SM86.GreedyArgmaxSM86
32import Realization.Nvidia.SM86.IndexedRowScatterOrderedSM86
33import Realization.Nvidia.SM86.IndexedRowScatterSM86
34import Realization.Nvidia.SM86.LayerNormSM86
35import Realization.Nvidia.SM86.SoftmaxSM86
36import Realization.Nvidia.SM86.StallCompaction
37import Realization.Nvidia.SM86.StreamingAttentionSM86
38import Realization.Nvidia.SM86.TargetGatherSM86
39import Realization.Nvidia.SM86.TiledProductSM86
40import Coppelius.Build.TiledChoices
41import SM86.Scoreboard
42import Realization.Nvidia.SM86.Transpose.FP16SM86
43import Std.List
44import Std.Foundation
45import Std.Natural
46import Std.Word
47
48-- A device region remains an Alpha value until the final physical artifact is
49-- assembled.  The resource fields travel with the encoded program so QMD
50-- construction cannot accidentally select a program address independently of
51-- its register, block and shared-memory requirements.
52family CoppeliusDeviceImage : Type 0
53constructor CoppeliusDeviceImageValue
54field unrestricted coppeliusDeviceImageIdentity : Bytes
55field unrestricted coppeliusDeviceImageProgram : (family SM86Program)
56field unrestricted coppeliusDeviceImageMaterial : Bytes
57field unrestricted coppeliusDeviceImageRegisters : Nat
58field unrestricted coppeliusDeviceImageBlockX : Nat
59field unrestricted coppeliusDeviceImageSharedBytes : Nat
60
61end-family
62
63family CoppeliusDeviceImages : Type 0
64constructor CoppeliusDeviceImagesEnd
65constructor CoppeliusDeviceImagesNext
66field unrestricted coppeliusDeviceImagesHead : (family CoppeliusDeviceImage)
67recursive unrestricted coppeliusDeviceImagesTail
68
69end-family
70
71family CoppeliusDeviceImagesBuildResult : Type 0
72constructor CoppeliusDeviceImagesBuildReady
73field unrestricted coppeliusDeviceImagesBuildBytes : BytesBuilder
74field unrestricted coppeliusDeviceImagesBuildCursor : Nat
75field unrestricted coppeliusDeviceImagesBuildCount : Nat
76constructor CoppeliusDeviceImagesBuildFailed
77field unrestricted coppeliusDeviceImagesBuildFailureIdentity : Bytes
78
79end-family
80
81-- a product the model launches (the tiled products below)
82family CoppeliusTiledShape : Type 0
83constructor CoppeliusTiledShapeValue
84field unrestricted coppeliusTiledShapeA : (family TiledProductLayout)
85field unrestricted coppeliusTiledShapeB : (family TiledProductLayout)
86field unrestricted coppeliusTiledShapeM : Nat
87field unrestricted coppeliusTiledShapeN : Nat
88field unrestricted coppeliusTiledShapeK : Nat
89-- 1: C = A B + R (a residual added in the epilogue)
90field unrestricted coppeliusTiledShapeResidual : Nat
91end-family
92
93def coppeliusDeviceImagesEnd : (family CoppeliusDeviceImages) =
94  (constructor CoppeliusDeviceImages CoppeliusDeviceImagesEnd)
95
96def coppeliusDeviceImagesNext =
97  (lambda unrestricted image : (family CoppeliusDeviceImage) .
98    (lambda unrestricted tail : (family CoppeliusDeviceImages) .
99      (constructor CoppeliusDeviceImages CoppeliusDeviceImagesNext image tail)))
100
101def coppeliusDeviceImagesAppend =
102  (lambda unrestricted left : (family CoppeliusDeviceImages) .
103    (lambda unrestricted right : (family CoppeliusDeviceImages) .
104      (eliminate
105        CoppeliusDeviceImages
106        (lambda unrestricted current : (family CoppeliusDeviceImages) .
107          (family CoppeliusDeviceImages))
108        left
109        (branch CoppeliusDeviceImagesEnd . right)
110        (branch
111          CoppeliusDeviceImagesNext
112          image
113          tail
114          induction
115          .
116          (coppeliusDeviceImagesNext image induction)))))
117
118-- Preserve the typed identity and launch-resource facts until the NVIDIA
119-- backend has seen the complete launch graph.  The old byte-only fold remains
120-- as an oracle helper, but release realization consumes this target-owned
121-- region vocabulary so unreachable images can be removed before QMD emission.
122def coppeliusNvidiaDeviceRegions =
123  (lambda unrestricted images : (family CoppeliusDeviceImages) .
124    (eliminate
125      CoppeliusDeviceImages
126      (lambda unrestricted current : (family CoppeliusDeviceImages) .
127        (family NvidiaDeviceRegions))
128      images
129      (branch
130        CoppeliusDeviceImagesEnd
131        .
132        (constructor NvidiaDeviceRegions NvidiaDeviceRegionsEnd))
133      (branch
134        CoppeliusDeviceImagesNext
135        image
136        tail
137        induction
138        .
139        (eliminate
140          CoppeliusDeviceImage
141          (lambda unrestricted current : (family CoppeliusDeviceImage) .
142            (family NvidiaDeviceRegions))
143          image
144          (branch
145            CoppeliusDeviceImageValue
146            identity
147            program
148            material
149            registers
150            blockX
151            sharedBytes
152            .
153            (constructor
154              NvidiaDeviceRegions
155              NvidiaDeviceRegionsNext
156              (constructor
157                NvidiaDeviceRegion
158                NvidiaDeviceProgramRegion
159                identity
160                program
161                registers
162                blockX
163                sharedBytes
164                (sm121LowerRealization program registers))
165              induction))))))
166
167def coppeliusDeviceZeroBytes =
168  (lambda unrestricted count : Nat .
169    (nat-eliminate
170      (lambda unrestricted current : Nat . Bytes)
171      b""
172      (lambda unrestricted predecessor : Nat .
173        (lambda unrestricted induction : Bytes .
174          (bytes-cons (byte 0) induction)))
175      count))
176
177def coppeliusDeviceAlignment : Nat =
178  256
179
180def coppeliusDevicePadding =
181  (lambda unrestricted cursor : Nat .
182    (naturalModuloUnchecked
183      (naturalSaturatingSubtract
184        coppeliusDeviceAlignment
185        (naturalModuloUnchecked cursor coppeliusDeviceAlignment))
186      coppeliusDeviceAlignment))
187
188def coppeliusBuildDeviceImagesFrom =
189  (lambda unrestricted images : (family CoppeliusDeviceImages) .
190    (eliminate
191      CoppeliusDeviceImages
192      (lambda unrestricted current : (family CoppeliusDeviceImages) .
193        (pi unrestricted builder : BytesBuilder .
194          (pi unrestricted cursor : Nat .
195            (pi unrestricted count : Nat .
196              (family CoppeliusDeviceImagesBuildResult)))))
197      images
198      (branch
199        CoppeliusDeviceImagesEnd
200        .
201        (lambda unrestricted builder : BytesBuilder .
202          (lambda unrestricted cursor : Nat .
203            (lambda unrestricted count : Nat .
204              (constructor
205                CoppeliusDeviceImagesBuildResult
206                CoppeliusDeviceImagesBuildReady
207                builder
208                cursor
209                count)))))
210      (branch
211        CoppeliusDeviceImagesNext
212        image
213        tail
214        induction
215        .
216        (lambda unrestricted builder : BytesBuilder .
217          (lambda unrestricted cursor : Nat .
218            (lambda unrestricted count : Nat .
219              (eliminate
220                CoppeliusDeviceImage
221                (lambda unrestricted current : (family CoppeliusDeviceImage) .
222                  (family CoppeliusDeviceImagesBuildResult))
223                image
224                (branch
225                  CoppeliusDeviceImageValue
226                  identity
227                  program
228                  material
229                  registers
230                  blockX
231                  sharedBytes
232                  .
233                  (nat-eliminate
234                    (lambda unrestricted nonempty : Nat .
235                      (family CoppeliusDeviceImagesBuildResult))
236                    (constructor
237                      CoppeliusDeviceImagesBuildResult
238                      CoppeliusDeviceImagesBuildFailed
239                      identity)
240                    (lambda unrestricted materialPredecessor : Nat .
241                      (lambda unrestricted materialInduction : (family CoppeliusDeviceImagesBuildResult) .
242                        (let unrestricted padding = (coppeliusDevicePadding cursor)
243                        in
244                          (induction
245                            (bytes-builder-append
246                              builder
247                              (bytes-builder-append
248                                (bytes-builder-chunk (coppeliusDeviceZeroBytes padding))
249                                (bytes-builder-chunk material)))
250                            (naturalAdd cursor (naturalAdd padding (bytes-length material)))
251                            (succ count)))))
252                    (naturalNonzero (bytes-length material)))))))))))
253
254def coppeliusBuildDeviceImages =
255  (lambda unrestricted images : (family CoppeliusDeviceImages) .
256    (coppeliusBuildDeviceImagesFrom images (bytes-builder-empty) zero zero))
257
258-- Build-result projections deliberately return empty bytes on refusal.  The
259-- list fold above converts that sentinel into a named failed region, and the
260-- final artifact projection therefore fails closed rather than publishing a
261-- truncated program table.
262def coppeliusProgramImage =
263  (lambda unrestricted program : (family SM86Program) .
264    (eliminate
265      SM86ProgramEncodingResult
266      (lambda unrestricted current : (family SM86ProgramEncodingResult) . Bytes)
267      (sm86EncodeProgram program)
268      (branch SM86ProgramEncodingSucceeded image telemetry . image)
269      (branch SM86ProgramEncodingFailed ordinal failure telemetry . b"")))
270
271def coppeliusTransposeHeadTile : Nat = 64
272
273-- Preserve already accepted schedules byte for byte. The RTX 3090 gather
274-- launch completed with its original controls but stalled when every LDG
275-- was given a read barrier. Repair only programs the scoreboard refuses,
276-- then admit the placed program through the same checker below.
277def coppeliusGuardLateReadsIfRequired =
278  (lambda unrestricted program : (family SM86Program) .
279    (nat-eliminate (lambda unrestricted accepted : Nat . (family SM86Program))
280      (sm86GuardLateReads program)
281      (lambda unrestricted predecessor : Nat .
282        (lambda unrestricted induction : (family SM86Program) . program))
283      (naturalIsZero (sm86Scoreboard program))))
284
285-- An image: its typed program, and the program's SM86 encoding, which
286-- places it (Coppelius.Build.Graph names launches by that placement); the
287-- backend realizes the region in the target's machine code.
288-- Its program's halves are in the plan's format (Coppelius.Learner).
289def coppeliusDeviceImage =
290  (lambda unrestricted identity : Bytes .
291    (lambda unrestricted program : (family SM86Program) .
292      (lambda unrestricted registers : Nat .
293        (lambda unrestricted blockX : Nat .
294          (lambda unrestricted sharedBytes : Nat .
295            (let unrestricted halves = (sm86ProgramWithHalfFormat coppeliusHalfFormat program) in
296            (let unrestricted guarded = (coppeliusGuardLateReadsIfRequired halves) in
297            (constructor
298              CoppeliusDeviceImage
299              CoppeliusDeviceImageValue
300              identity
301              guarded
302              (coppeliusProgramImage guarded)
303              registers
304              blockX
305              sharedBytes))))))))
306
307def coppeliusElementwiseImage =
308  (lambda unrestricted operation : (family Elementwise256) .
309    (emitElementwise256SM86 operation))
310
311def coppeliusArgmaxImage =
312  (lambda unrestricted artifact : (family GreedyArgmaxSM86Artifact) .
313    (greedyArgmaxSM86ProgramFor greedyArgmaxSM86PromotedVocabulary))
314
315def coppeliusCastImage =
316  (lambda unrestricted direction : (family CastDirection) .
317    (emitCastSM86 direction))
318
319def coppeliusVectorImage =
320  (lambda unrestricted kind : (family ElementwiseVectorSM86Kind) .
321    (elementwiseVectorProgram kind))
322
323def coppeliusGatherImage : (family SM86Program) =
324  (embeddingGatherF32ProgramFor
325    (constructor EmbeddingGatherF32Width EmbeddingGatherF32Width512))
326
327def coppeliusLayerNormImage =
328  (lambda unrestricted variant : (family LayerNormSM86Variant) .
329    (layerNormSM86Program variant 1024))
330
331def coppeliusGELUImage =
332  (lambda unrestricted kind : (family GELUSM86ProgramKind) .
333    (geluSM86Program kind))
334
335def coppeliusAttentionImage =
336  (lambda unrestricted kind : (family AttentionHeadLayoutSM86Kind) .
337    (attentionHeadLayoutSM86Program kind))
338
339def coppeliusSoftmaxImage =
340  (lambda unrestricted variant : (family SoftmaxSM86Variant) .
341    (softmaxSM86ProgramFor variant))
342
343def coppeliusCrossEntropyImage : (family SM86Program) =
344  (exactCrossEntropySM86ProgramFor
345    (constructor ExactCrossEntropySM86Variant ExactCrossEntropyRows))
346
347def coppeliusCrossEntropyBackwardImage : (family SM86Program) =
348  exactCrossEntropyBackwardProgram
349
350-- the embedding's gradient collected in a fixed order (the rows of the
351-- positions naming a token added in position order, by the block of the
352-- first: Realization.Nvidia.SM86.IndexedRowScatterOrderedSM86), so training
353-- takes the same bits every run; the geometry is the atomic form's
354def coppeliusScatterGeometry : (family IndexedRowScatterSM86Geometry) =
355  (constructor IndexedRowScatterSM86Geometry IndexedRowScatterTokenEmbeddingGradientRows512x1024)
356
357def coppeliusScatterImage : (family SM86Program) =
358  (indexedRowScatterOrderedSM86Program
359    (indexedRowScatterSM86WidthNatural coppeliusScatterGeometry)
360    (indexedRowScatterSM86RowsNatural coppeliusScatterGeometry))
361
362-- AdamW over the whole flat bank in one launch, the
363-- parameters' half copies written with it (Realization.Nvidia.SM86.
364-- AdamWHalfSM86, its stalls compacted): a grid of coppeliusAdamWBlocks
365-- blocks walks the bank's element quads with a grid stride.  The bank's
366-- elements are the checkpoint's (it is P, M and V, each the parameter
367-- bank); the plan holds its own bank to the same count (Coppelius.Build.
368-- Graph.cgAdamWCoversTheBank).
369def coppeliusAdamWBlocks : Nat = 1024
370def coppeliusAdamWElements : Nat =
371  (naturalDivideUnchecked coppeliusCheckpointPayloadBytesNatural 12)
372def coppeliusAdamWQuads : Nat = (naturalDivideUnchecked coppeliusAdamWElements adamWHalfSM86Elements)
373def coppeliusAdamWStride : Nat = (naturalMultiply coppeliusAdamWBlocks adamWHalfSM86Threads)
374def coppeliusAdamWIterations : Nat =
375  (adamWHalfSM86Iterations coppeliusAdamWQuads coppeliusAdamWStride)
376def coppeliusAdamWImage : (family SM86Program) =
377  (adamWHalfSM86 sm86CompactStallsDrained coppeliusAdamWQuads coppeliusAdamWStride)
378-- what the schedule checks read (the body twice, no branch)
379def coppeliusAdamWChecked : (family SM86Program) =
380  (adamWHalfSM86Checked sm86CompactStallsDrained coppeliusAdamWQuads coppeliusAdamWStride)
381
382-- the bank is whole quads, the grid stride reaches every quad (the
383-- image's bound leaves the rest of the last pass alone), and the bound is
384-- a 32-bit immediate
385def coppeliusAdamWWholeQuads :
386  (equal Nat (naturalMultiply adamWHalfSM86Elements coppeliusAdamWQuads) coppeliusAdamWElements) =
387  (refl Nat coppeliusAdamWElements)
388def coppeliusAdamWReachesEveryElement :
389  (equal Nat (naturalLessOrEqual coppeliusAdamWQuads (naturalMultiply coppeliusAdamWIterations coppeliusAdamWStride)) 1) =
390  (refl Nat 1)
391def coppeliusAdamWBoundFits :
392  (equal Nat (naturalLess coppeliusAdamWQuads 4294967296) 1) =
393  (refl Nat 1)
394
395-- the gated activation's forward and backward
396-- (Realization.Nvidia.SM86.GatedGELUSM86), their stalls compacted
397def coppeliusGatedGELUForwardImage : (family SM86Program) = (sm86CompactStalls gatedGELUForwardSM86)
398def coppeliusGatedGELUBackwardImage : (family SM86Program) = (sm86CompactStalls gatedGELUBackwardSM86)
399
400def coppeliusCoreDeviceImages : (family CoppeliusDeviceImages) =
401  (coppeliusDeviceImagesNext
402    (coppeliusDeviceImage b"random-normal" randomNormalSM86Program 24 256 0)
403    (coppeliusDeviceImagesNext
404    (coppeliusDeviceImage b"rotary-table" rotaryTableSM86Program rotaryTableSM86RegisterCount 256 0)
405    (coppeliusDeviceImagesNext
406    (coppeliusDeviceImage b"fill" (coppeliusVectorImage (constructor ElementwiseVectorSM86Kind ElementwiseVectorSM86Fill)) 16 256 0)
407    (coppeliusDeviceImagesNext
408    (coppeliusDeviceImage b"copy" (coppeliusVectorImage (constructor ElementwiseVectorSM86Kind ElementwiseVectorSM86Copy)) 24 256 0)
409    (coppeliusDeviceImagesNext
410    (coppeliusDeviceImage b"cast-down" (coppeliusCastImage (constructor CastDirection Float32ToFloat16)) 24 256 0)
411    (coppeliusDeviceImagesNext
412    (coppeliusDeviceImage b"cast-up" (coppeliusCastImage (constructor CastDirection Float16ToFloat32)) 24 256 0)
413    (coppeliusDeviceImagesNext
414    (coppeliusDeviceImage b"add" (coppeliusElementwiseImage (constructor Elementwise256 Add256)) 24 256 0)
415    (coppeliusDeviceImagesNext
416    (coppeliusDeviceImage b"scale" (coppeliusElementwiseImage (constructor Elementwise256 Scale256)) 24 256 0)
417    (coppeliusDeviceImagesNext
418    (coppeliusDeviceImage b"argmax" (coppeliusArgmaxImage greedyArgmaxSM86BuildPromoted) 64 256 96)
419    (coppeliusDeviceImagesNext
420    (coppeliusDeviceImage b"gather" coppeliusGatherImage 40 128 0)
421    (coppeliusDeviceImagesNext
422    (coppeliusDeviceImage b"layernorm-forward" (coppeliusLayerNormImage (constructor LayerNormSM86Variant LayerNormSM86Forward512)) 40 (layerNormSM86BlockX (constructor LayerNormSM86Variant LayerNormSM86Forward512)) 128)
423    (coppeliusDeviceImagesNext
424    (coppeliusDeviceImage b"qk-rope" (coppeliusAttentionImage (constructor AttentionHeadLayoutSM86Kind AttentionHeadLayoutQKRotary)) 64 128 0)
425    (coppeliusDeviceImagesNext
426    (coppeliusDeviceImage b"value-transpose" (coppeliusAttentionImage (constructor AttentionHeadLayoutSM86Kind AttentionHeadLayoutValueTranspose)) 24 512 0)
427    (coppeliusDeviceImagesNext
428    (coppeliusDeviceImage b"head-to-token" (coppeliusAttentionImage (constructor AttentionHeadLayoutSM86Kind AttentionHeadLayoutHeadToToken)) 32 256 0)
429    (coppeliusDeviceImagesNext
430    (coppeliusDeviceImage b"token-to-head" (coppeliusAttentionImage (constructor AttentionHeadLayoutSM86Kind AttentionHeadLayoutTokenToHead)) 32 256 0)
431    (coppeliusDeviceImagesNext
432    (coppeliusDeviceImage b"inverse-rope-qkv" (coppeliusAttentionImage (constructor AttentionHeadLayoutSM86Kind AttentionHeadLayoutInverseRoPEQKVMerge)) 64 128 0)
433    (coppeliusDeviceImagesNext
434    (coppeliusDeviceImage b"softmax-forward" (coppeliusSoftmaxImage (constructor SoftmaxSM86Variant SoftmaxSM86CausalForward1024)) 32 256 128)
435    (coppeliusDeviceImagesNext
436    (coppeliusDeviceImage b"softmax-backward" (coppeliusSoftmaxImage (constructor SoftmaxSM86Variant SoftmaxSM86CausalBackward1024)) 32 256 128)
437    (coppeliusDeviceImagesNext
438    (coppeliusDeviceImage b"layernorm-backward" (coppeliusLayerNormImage (constructor LayerNormSM86Variant LayerNormSM86InputBackward512)) 48 (layerNormSM86BlockX (constructor LayerNormSM86Variant LayerNormSM86InputBackward512)) 128)
439    (coppeliusDeviceImagesNext
440    (coppeliusDeviceImage b"layernorm-parameters" (coppeliusLayerNormImage (constructor LayerNormSM86Variant LayerNormSM86ParameterGradient512)) 40 (layerNormSM86BlockX (constructor LayerNormSM86Variant LayerNormSM86ParameterGradient512)) 128)
441    (coppeliusDeviceImagesNext
442    (coppeliusDeviceImage b"cross-entropy" coppeliusCrossEntropyImage 48 256 128)
443    (coppeliusDeviceImagesNext
444    (coppeliusDeviceImage b"cross-entropy-backward" coppeliusCrossEntropyBackwardImage 40 256 0)
445    (coppeliusDeviceImagesNext
446    (coppeliusDeviceImage b"embedding-scatter" coppeliusScatterImage 40 512 0)
447    (coppeliusDeviceImagesNext
448    (coppeliusDeviceImage b"adamw" coppeliusAdamWImage adamWHalfSM86Registers adamWHalfSM86Threads 0)
449    (coppeliusDeviceImagesNext
450    (coppeliusDeviceImage b"gated-gelu-forward" coppeliusGatedGELUForwardImage gatedGELUForwardSM86Registers gatedGELUSM86Threads 0)
451    (coppeliusDeviceImagesNext
452    (coppeliusDeviceImage b"gated-gelu-backward" coppeliusGatedGELUBackwardImage gatedGELUBackwardSM86Registers gatedGELUSM86Threads 0)
453    coppeliusDeviceImagesEnd))))))))))))))))))))))))))
454
455-- the attention's head layouts (the products read their operands in place)
456def coppeliusTransposeDeviceImages : (family CoppeliusDeviceImages) =
457  (coppeliusDeviceImagesNext
458    (coppeliusDeviceImage b"transpose-g8-r1024-c64" (coppeliusFP16TransposeTiledProgram coppeliusTransposeHeadTile 8 1024 64) 40 (coppeliusFP16TransposeBlockThreads coppeliusTransposeHeadTile) 0)
459    (coppeliusDeviceImagesNext
460      (coppeliusDeviceImage b"transpose-g8-r64-c1024" (coppeliusFP16TransposeTiledProgram coppeliusTransposeHeadTile 8 64 1024) 40 (coppeliusFP16TransposeBlockThreads coppeliusTransposeHeadTile) 0)
461      coppeliusDeviceImagesEnd))
462
463-- the images every launch but the loss's uses, in the order the backend
464-- places them
465def coppeliusImportedDeviceImages : (family CoppeliusDeviceImages) =
466  (coppeliusDeviceImagesAppend coppeliusCoreDeviceImages coppeliusTransposeDeviceImages)
467
468-- the training loss's target-logit gather (Realization.Nvidia.SM86.
469-- TargetGatherSM86), placed after them: it fills `correct` with each row's
470-- target logit (docs/observability PRD item 1)
471def coppeliusTargetGatherImage : (family SM86Program) = targetGatherSM86Coppelius
472def coppeliusLossDeviceImages : (family CoppeliusDeviceImages) =
473  (coppeliusDeviceImagesNext
474    (coppeliusDeviceImage b"loss-target-gather" coppeliusTargetGatherImage (sm86RegisterDemand targetGatherSM86Coppelius) 256 0)
475    coppeliusDeviceImagesEnd)
476
477-- the heads' attention as streaming launches
478-- (Realization.Nvidia.SM86.StreamingAttentionSM86) at the model's sequence
479-- and heads, the forward and the backward's two, their stalls compacted
480-- (Realization.Nvidia.SM86.StallCompaction)
481def coppeliusStreamingImage =
482  (lambda unrestricted kernel : (pi unrestricted seq : Nat . (pi unrestricted heads : Nat . (family SM86Program))) .
483    (sm86CompactStalls (kernel (stdU32ToNatural modelCoppeliusBlock) (stdU32ToNatural modelCoppeliusHeads))))
484def coppeliusStreamingDeviceImages : (family CoppeliusDeviceImages) =
485  (coppeliusDeviceImagesNext
486    (coppeliusDeviceImage b"attention-forward" (coppeliusStreamingImage streamingAttentionForwardSM86)
487      streamingAttentionForwardSM86Registers streamingAttentionForwardSM86Threads 0)
488    (coppeliusDeviceImagesNext
489      (coppeliusDeviceImage b"attention-query" (coppeliusStreamingImage streamingAttentionQuerySM86)
490        streamingAttentionQuerySM86Registers streamingAttentionQuerySM86Threads 0)
491      (coppeliusDeviceImagesNext
492        (coppeliusDeviceImage b"attention-key" (coppeliusStreamingImage streamingAttentionKeySM86)
493          streamingAttentionKeySM86Registers streamingAttentionKeySM86Threads 0)
494        coppeliusDeviceImagesEnd)))
495
496-- ---- the tiled products ----
497def stdListFromShapes = (lambda unrestricted head : (family CoppeliusTiledShape) .
498  (lambda unrestricted tail : (family StdList (family CoppeliusTiledShape)) .
499    (constructor StdList StdListCons (family CoppeliusTiledShape) head tail)))
500
501-- C (m x n) = A B (Realization.Nvidia.SM86.TiledProductSM86), each operand
502-- read where it is stored: A as rows (m x k) or as columns (A^T, k x m), B
503-- as rows (B^T, n x k) or as columns (k x n).  The model's products are
504-- three patterns: the forward's X W^T (rows, rows), a data gradient's dY W
505-- (rows, columns) and a weight gradient's dY^T X (columns, columns).  Each
506-- product's tile and stages are a candidate measured on the card and
507-- recorded (Coppelius.Build.TiledChoices, scripts/alpha/tiled-choices.py)
508-- among the candidates below.
509
510def coppeliusTiledRows : (family TiledProductLayout) = (constructor TiledProductLayout TiledProductRows)
511def coppeliusTiledColumns : (family TiledProductLayout) = (constructor TiledProductLayout TiledProductColumns)
512def coppeliusTiledIsColumns = (lambda unrestricted layout : (family TiledProductLayout) .
513  (eliminate TiledProductLayout (lambda unrestricted c : (family TiledProductLayout) . Nat) layout
514    (branch TiledProductRows . 0) (branch TiledProductColumns . 1)))
515-- the candidates, as scripts/alpha/tiled-choices.py numbers them (block M x
516-- block N x stages, warps along M x N, each warp's rows x columns):
517--   0  128 x 128 x 4   2 x 4   64 x 32
518--   1  128 x 128 x 2   2 x 4   64 x 32
519--   2   64 x  64 x 4   2 x 2   32 x 32
520--   3   64 x  64 x 2   2 x 2   32 x 32
521--   4  128 x  64 x 4   4 x 2   32 x 32   (48 KiB: two blocks an SM)
522--   5   64 x 128 x 4   2 x 4   32 x 32
523--   6  256 x 128 x 2   4 x 2   64 x 64   (the most registers the SM86 file allows)
524--   7  128 x 256 x 2   2 x 4   64 x 64
525-- Candidates over a card's shared memory or registers exist but never serve
526-- (coppeliusTiledTileServes).  The cuBLAS tiles measured on the GB10
527-- (research gemm/cublas-shapes.md: 128 x 96 .. 256 x 128, three stages) are
528-- where 4 .. 7 came from.
529def coppeliusTiledCandidateCount : Nat = 8
530-- candidate i's entry of a column of eight
531def coppeliusTiledPick = (lambda unrestricted i : Nat .
532  (lambda unrestricted v0 : Nat . (lambda unrestricted v1 : Nat . (lambda unrestricted v2 : Nat . (lambda unrestricted v3 : Nat .
533  (lambda unrestricted v4 : Nat . (lambda unrestricted v5 : Nat . (lambda unrestricted v6 : Nat . (lambda unrestricted v7 : Nat .
534    (naturalSelect (naturalEqual i 0) v0 (naturalSelect (naturalEqual i 1) v1 (naturalSelect (naturalEqual i 2) v2
535    (naturalSelect (naturalEqual i 3) v3 (naturalSelect (naturalEqual i 4) v4 (naturalSelect (naturalEqual i 5) v5
536    (naturalSelect (naturalEqual i 6) v6 v7))))))))))))))))
537def coppeliusTiledCandidate = (lambda unrestricted i : Nat .
538  (tiledProductTile
539    (coppeliusTiledPick i 2 2 2 2 4 2 4 2)
540    (coppeliusTiledPick i 4 4 2 2 2 4 2 4)
541    (coppeliusTiledPick i 4 4 2 2 2 2 4 4)
542    (coppeliusTiledPick i 4 4 4 4 4 4 8 8)
543    (coppeliusTiledPick i 4 2 4 2 4 4 2 2)))
544-- a product's recorded candidate
545def coppeliusTiledChoiceFor = (lambda unrestricted residual : Nat .
546  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
547  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
548    (coppeliusTiledChoice residual (coppeliusTiledIsColumns a) (coppeliusTiledIsColumns b) m n k)))))))
549def coppeliusTiledTile = (lambda unrestricted residual : Nat .
550  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
551  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
552    (coppeliusTiledCandidate (coppeliusTiledChoiceFor residual a b m n k))))))))
553-- 1 when a tile serves the product: it divides C, its stages divide K's
554-- steps, its threads copy its operand tiles evenly, its registers fit the
555-- SM86 file, and its stages fit the shared memory -- every compatible
556-- card's (Coppelius.SM86Capability), not only the GB10's
557def coppeliusTiledTileServes = (lambda unrestricted t : (family TiledProductTile) .
558  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
559    (tiledProductSM86ShapeAdmitted t m n k coppeliusSM86CompatSharedBytesPerBlock)))))
560def coppeliusTiledLayoutName = (lambda unrestricted layout : (family TiledProductLayout) .
561  (eliminate TiledProductLayout (lambda unrestricted c : (family TiledProductLayout) . Bytes) layout
562    (branch TiledProductRows . b"r") (branch TiledProductColumns . b"c")))
563-- the stored rows' lengths: A's (k as rows, m as columns), B's (k, n)
564def coppeliusTiledStrideA = (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted m : Nat . (lambda unrestricted k : Nat .
565  (naturalSelect (coppeliusTiledIsColumns layout) m k))))
566def coppeliusTiledStrideB = (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
567  (naturalSelect (coppeliusTiledIsColumns layout) n k))))
568def coppeliusTiledName = (lambda unrestricted residual : Nat .
569  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
570  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
571    (bytes-append b"tiled-" (bytes-append (coppeliusTiledLayoutName a) (bytes-append (coppeliusTiledLayoutName b)
572      (bytes-append b"-m" (bytes-append (naturalDecimalBytesWithin 8 m) (bytes-append b"-n" (bytes-append (naturalDecimalBytesWithin 8 n)
573        (bytes-append b"-k" (bytes-append (naturalDecimalBytesWithin 8 k)
574          (nat-eliminate (lambda unrestricted c : Nat . Bytes) b"" (lambda unrestricted p : Nat . (lambda unrestricted ignored : Bytes . b"-add")) residual))))))))))))))))
575
576-- a product's program with tile t: looped 1 the loop, 0 the checked
577-- unrolling (Proof.GB10TiledChoiceProbe times every candidate this way)
578def coppeliusTiledProgramWith = (lambda unrestricted t : (family TiledProductTile) .
579  (lambda unrestricted looped : Nat . (lambda unrestricted residual : Nat .
580  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
581  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
582    (let unrestricted sa = (coppeliusTiledStrideA a m k) in (let unrestricted sb = (coppeliusTiledStrideB b n k) in
583    (eliminate StdBool (lambda unrestricted c : (family StdBool) . (family SM86Program)) (stdBoolFromNatural looped)
584      (branch StdTrue .
585        (eliminate StdBool (lambda unrestricted c : (family StdBool) . (family SM86Program)) (stdBoolFromNatural residual)
586          (branch StdTrue . (tiledProductSM86Residual sm86CompactStallsDrained t a b k sa sb n))
587          (branch StdFalse . (tiledProductSM86 sm86CompactStallsDrained t a b k sa sb n))))
588      (branch StdFalse .
589        (eliminate StdBool (lambda unrestricted c : (family StdBool) . (family SM86Program)) (stdBoolFromNatural residual)
590          (branch StdTrue . (tiledProductSM86ResidualChecked sm86CompactStallsDrained t a b k sa sb n))
591          (branch StdFalse . (tiledProductSM86Checked sm86CompactStallsDrained t a b k sa sb n)))))))))))))))
592
593-- a product's program with its recorded tile
594def coppeliusTiledProgramAs = (lambda unrestricted looped : Nat . (lambda unrestricted residual : Nat .
595  (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
596  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
597    (coppeliusTiledProgramWith (coppeliusTiledTile residual a b m n k) looped residual a b m n k))))))))
598
599def coppeliusTiledShapeImage = (lambda unrestricted shape : (family CoppeliusTiledShape) .
600  (eliminate CoppeliusTiledShape (lambda unrestricted c : (family CoppeliusTiledShape) . (family CoppeliusDeviceImage)) shape
601    (branch CoppeliusTiledShapeValue a b m n k residual .
602      (let unrestricted tile = (coppeliusTiledTile residual a b m n k) in
603      (coppeliusDeviceImage (coppeliusTiledName residual a b m n k) (coppeliusTiledProgramAs 1 residual a b m n k)
604        (tiledProductSM86Registers tile) (tiledProductSM86Threads tile) (tiledProductSM86SharedBytes tile))))))
605-- 1 when the product has a recorded choice, its tile serves it, and the
606-- checked schedule holds on both of SM86.Scoreboard's passes
607def coppeliusTiledShapeAdmitted = (lambda unrestricted shape : (family CoppeliusTiledShape) .
608  (eliminate CoppeliusTiledShape (lambda unrestricted c : (family CoppeliusTiledShape) . Nat) shape
609    (branch CoppeliusTiledShapeValue a b m n k residual .
610      (let unrestricted checked = (coppeliusGuardLateReadsIfRequired (coppeliusTiledProgramAs 0 residual a b m n k)) in
611      -- a recorded choice, a candidate that serves the product
612      (naturalAnd (naturalLess (coppeliusTiledChoiceFor residual a b m n k) coppeliusTiledCandidateCount)
613      (naturalAnd (coppeliusTiledTileServes (coppeliusTiledTile residual a b m n k) m n k)
614        (naturalAnd (naturalIsZero (sm86Scoreboard checked)) (naturalIsZero (sm86FixedLatencyHazard checked)))))))))
615def coppeliusTiledShapeWith = (lambda unrestricted residual : Nat . (lambda unrestricted a : (family TiledProductLayout) . (lambda unrestricted b : (family TiledProductLayout) .
616  (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat .
617    (constructor CoppeliusTiledShape CoppeliusTiledShapeValue a b m n k residual)))))))
618def coppeliusTiledShape = (coppeliusTiledShapeWith 0)
619-- C = A B + R
620def coppeliusTiledShapeAdding = (coppeliusTiledShapeWith 1)
621-- every product Coppelius.Build.Graph launches (sequence 1024, width 512,
622-- the query/key/value 1536, the feed-forward 1408, the vocabulary 12288)
623def coppeliusTiledShapes : (family StdList (family CoppeliusTiledShape)) =
624  (stdListFromShapes
625    (coppeliusTiledShape coppeliusTiledRows coppeliusTiledRows 1024 1536 512)
626    (stdListFromShapes (coppeliusTiledShapeAdding coppeliusTiledRows coppeliusTiledRows 1024 512 512)
627    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledRows 1024 1408 512)
628    (stdListFromShapes (coppeliusTiledShapeAdding coppeliusTiledRows coppeliusTiledRows 1024 512 1408)
629    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledRows 1024 12288 512)
630    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledColumns 1024 512 1536)
631    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledColumns 1024 512 512)
632    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledColumns 1024 512 1408)
633    (stdListFromShapes (coppeliusTiledShapeAdding coppeliusTiledRows coppeliusTiledColumns 1024 512 1408)
634    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledColumns 1024 1408 512)
635    (stdListFromShapes (coppeliusTiledShape coppeliusTiledRows coppeliusTiledColumns 1024 512 12288)
636    (stdListFromShapes (coppeliusTiledShape coppeliusTiledColumns coppeliusTiledColumns 1536 512 1024)
637    (stdListFromShapes (coppeliusTiledShape coppeliusTiledColumns coppeliusTiledColumns 512 512 1024)
638    (stdListFromShapes (coppeliusTiledShape coppeliusTiledColumns coppeliusTiledColumns 1408 512 1024)
639    (stdListFromShapes (coppeliusTiledShape coppeliusTiledColumns coppeliusTiledColumns 512 1408 1024)
640    (stdListFromShapes (coppeliusTiledShape coppeliusTiledColumns coppeliusTiledColumns 12288 512 1024)
641      (constructor StdList StdListEmpty (family CoppeliusTiledShape))))))))))))))))))
642-- every product Coppelius launches has a recorded choice, and it serves the
643-- product (decided before any program is built)
644def coppeliusTiledChoicesAdmitted : Nat =
645  (eliminate StdList (lambda unrestricted c : (family StdList (family CoppeliusTiledShape)) . Nat) coppeliusTiledShapes
646    (branch StdListEmpty . 1)
647    (branch StdListCons head tail induction .
648      (naturalAnd induction
649        (eliminate CoppeliusTiledShape (lambda unrestricted c : (family CoppeliusTiledShape) . Nat) head
650          (branch CoppeliusTiledShapeValue a b m n k residual .
651            (naturalAnd (naturalLess (coppeliusTiledChoiceFor residual a b m n k) coppeliusTiledCandidateCount)
652              (coppeliusTiledTileServes (coppeliusTiledTile residual a b m n k) m n k)))))))
653def coppeliusTiledChoicesServe : (equal Nat coppeliusTiledChoicesAdmitted 1) = (refl Nat 1)
654def coppeliusTiledDeviceImages : (family CoppeliusDeviceImages) =
655  (eliminate StdList (lambda unrestricted c : (family StdList (family CoppeliusTiledShape)) . (family CoppeliusDeviceImages)) coppeliusTiledShapes
656    (branch StdListEmpty . coppeliusDeviceImagesEnd)
657    (branch StdListCons head tail induction . (coppeliusDeviceImagesNext (coppeliusTiledShapeImage head) induction)))
658def coppeliusTiledSchedulesAdmitted : Nat =
659  (eliminate StdList (lambda unrestricted c : (family StdList (family CoppeliusTiledShape)) . Nat) coppeliusTiledShapes
660    (branch StdListEmpty . 1)
661    (branch StdListCons head tail induction . (naturalAnd (coppeliusTiledShapeAdmitted head) induction)))
662
663-- 1 when a program's schedule holds on both of SM86.Scoreboard's passes
664def coppeliusScheduleHolds = (lambda unrestricted program : (family SM86Program) .
665  (let unrestricted guarded = (coppeliusGuardLateReadsIfRequired program) in
666    (naturalAnd (naturalIsZero (sm86Scoreboard guarded)) (naturalIsZero (sm86FixedLatencyHazard guarded)))))
667def coppeliusDeviceImages : (family CoppeliusDeviceImages) =
668  (coppeliusDeviceImagesAppend coppeliusImportedDeviceImages
669    (coppeliusDeviceImagesAppend coppeliusLossDeviceImages
670      (coppeliusDeviceImagesAppend coppeliusStreamingDeviceImages coppeliusTiledDeviceImages)))
671
672-- Every placed SM86 image is checked after half-format selection and the
673-- late-read guard. This covers the streaming and imported images that the
674-- older sampled admission did not inspect. The fixed-latency pass remains on
675-- the previously qualified unrolled paths: older imported kernels, beginning
676-- with random-normal, have pre-existing fixed-latency model mismatches and
677-- require a separate rescheduling migration before that pass can gate all.
678def coppeliusDeviceImageScheduleRefusal : Bytes =
679  (eliminate CoppeliusDeviceImages
680    (lambda unrestricted current : (family CoppeliusDeviceImages) . Bytes)
681    coppeliusDeviceImages
682    (branch CoppeliusDeviceImagesEnd . b"")
683    (branch CoppeliusDeviceImagesNext image tail induction .
684      (eliminate CoppeliusDeviceImage
685        (lambda unrestricted current : (family CoppeliusDeviceImage) . Bytes) image
686        (branch CoppeliusDeviceImageValue identity program material registers blockX sharedBytes .
687          (nat-eliminate (lambda unrestricted accepted : Nat . Bytes)
688            identity
689            (lambda unrestricted predecessor : Nat . (lambda unrestricted ignored : Bytes . induction))
690            (naturalIsZero (sm86Scoreboard program)))))))
691
692def coppeliusDeviceImageSchedulesAdmitted : Nat =
693  (naturalIsZero (bytes-length coppeliusDeviceImageScheduleRefusal))
694
695def coppeliusSchedulesAdmitted : Nat =
696  (naturalAnd coppeliusDeviceImageSchedulesAdmitted
697    (naturalAnd coppeliusTiledSchedulesAdmitted
698      (naturalAnd (coppeliusScheduleHolds coppeliusAdamWChecked)
699        (naturalAnd (coppeliusScheduleHolds coppeliusGatedGELUForwardImage)
700          (coppeliusScheduleHolds coppeliusGatedGELUBackwardImage)))))
701
702-- the device address of the image named `identity`: where the backend
703-- places its region, by the same fold (0, which no launch can name, when
704-- there is none)
705def coppeliusDeviceImageAddress =
706  (lambda unrestricted identity : Bytes .
707    (app
708      (eliminate CoppeliusDeviceImages
709        (lambda unrestricted current : (family CoppeliusDeviceImages) . (pi unrestricted cursor : Nat . Nat))
710        coppeliusDeviceImages
711        (branch CoppeliusDeviceImagesEnd . (lambda unrestricted cursor : Nat . 0))
712        (branch CoppeliusDeviceImagesNext image tail induction .
713          (lambda unrestricted cursor : Nat .
714            (eliminate CoppeliusDeviceImage (lambda unrestricted current : (family CoppeliusDeviceImage) . Nat) image
715              (branch CoppeliusDeviceImageValue name program material registers blockX sharedBytes .
716                (let unrestricted offset = (naturalAdd cursor (coppeliusDevicePadding cursor)) in
717                  (naturalSelect (bytes-equal name identity)
718                    (naturalAdd coppeliusProgramBase offset)
719                    (induction (naturalAdd offset (bytes-length material))))))))))
720      0))
721
722def coppeliusWholeProgramDeviceRegions : (family NvidiaDeviceRegions) =
723  (coppeliusNvidiaDeviceRegions coppeliusDeviceImages)
724
725def coppeliusDeviceImagesBuild : (family CoppeliusDeviceImagesBuildResult) =
726  (coppeliusBuildDeviceImages coppeliusDeviceImages)

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.