Source/Packages

Representation.Schema

packages/representations/src/Representation/Schema.alpha

701 lines99 declarations24.8 KiBSHA-256 209c8ba5712e

Complete file · line 367

Schema.alpha

Definition view
1module Representation.Schema
2
3import Std.Foundation
4import Std.Natural
5import Std.List
6
7-- A stable parameter identity is its structural component path. Names remain
8-- separate path components, so concatenation cannot create accidental aliases.
9family ParameterIdentity : Type 0
10constructor ParameterIdentityOf
11field unrestricted parameterIdentityComponents : (family StdList Bytes)
12
13end-family
14
15family ParameterShape : Type 0
16constructor ParameterShapeOf
17field unrestricted parameterShapeDimensions : (family StdList Nat)
18
19end-family
20
21family LogicalOwnership : Type 0
22constructor LogicalOwned
23constructor LogicalShared
24constructor LogicalBorrowed
25
26end-family
27
28family LogicalLifetime : Type 0
29constructor LogicalPersistent
30constructor LogicalStep
31constructor LogicalInvocation
32
33end-family
34
35-- Logical storage has an extent and alignment but deliberately no address.
36family LogicalStorage : Type 0
37constructor LogicalStorageOf
38field unrestricted logicalStorageExtent : Nat
39field unrestricted logicalStorageAlignment : Nat
40field unrestricted logicalStorageOwnership : (family LogicalOwnership)
41field unrestricted logicalStorageLifetime : (family LogicalLifetime)
42
43end-family
44
45family ParameterSchema : Type 0
46constructor ParameterSchemaOf
47field unrestricted parameterSchemaIdentity : (family ParameterIdentity)
48field unrestricted parameterSchemaShape : (family ParameterShape)
49field unrestricted parameterSchemaRepresentation : Bytes
50field unrestricted parameterSchemaStorage : (family LogicalStorage)
51
52end-family
53
54family LearnerContract : Type 0
55constructor LearnerSGD
56constructor LearnerMomentum
57constructor LearnerAdamW
58
59end-family
60
61family LearnerStateKind : Type 0
62constructor LearnerVelocity
63constructor LearnerFirstMoment
64constructor LearnerSecondMoment
65constructor LearnerStepCounter
66
67end-family
68
69family LearnerStateSlot : Type 0
70constructor LearnerParameterState
71field unrestricted learnerStateParameter : (family ParameterIdentity)
72field unrestricted learnerStateKind : (family LearnerStateKind)
73constructor LearnerGlobalState
74field unrestricted learnerGlobalStateKind : (family LearnerStateKind)
75
76end-family
77
78family LearnerStateSchema : Type 0
79constructor LearnerStateSchemaOf
80field unrestricted learnerStateSlots : (family StdList (family LearnerStateSlot))
81
82end-family
83
84-- PRD 19 L24: a system is its structural parameters and its learner; everything
85-- the continuation and the update plan need is DERIVED from this value.
86family SystemDefinition : Type 0
87constructor SystemDefinitionOf
88field unrestricted systemName : Bytes
89field unrestricted systemParameters : (family StdList (family ParameterSchema))
90field unrestricted systemLearner : (family LearnerContract)
91
92end-family
93
94-- PRD 18 K2: the categories a continuation must carry.
95family ContinuationCategory : Type 0
96constructor ContinuationParameters
97constructor ContinuationComponentState
98constructor ContinuationLearnerState
99constructor ContinuationRng
100constructor ContinuationDataCursor
101constructor ContinuationCounters
102constructor ContinuationSchedule
103constructor ContinuationIdentities
104
105end-family
106
107family ContinuationEntry : Type 0
108constructor ContinuationEntryOf
109field unrestricted continuationEntryCategory : (family ContinuationCategory)
110field unrestricted continuationEntryKey : (family StdList Bytes)
111
112end-family
113
114family ContinuationSchema : Type 0
115constructor ContinuationSchemaOf
116field unrestricted continuationSchemaSystem : Bytes
117field unrestricted continuationSchemaEntries : (family StdList (family ContinuationEntry))
118
119end-family
120
121-- One step of the derived update plan: a parameter's update, or one learner-state slot's update.
122family UpdateStep : Type 0
123constructor UpdateParameter
124field unrestricted updateStepParameter : (family ParameterIdentity)
125constructor UpdateLearnerState
126field unrestricted updateStepState : (family LearnerStateSlot)
127
128end-family
129
130def parameterIdentityEqual =
131  (lambda unrestricted left : (family ParameterIdentity) .
132    (lambda unrestricted right : (family ParameterIdentity) .
133      (eliminate
134        ParameterIdentity
135        (lambda unrestricted current : (family ParameterIdentity) . (family StdBool))
136        left
137        (branch
138          ParameterIdentityOf
139          leftComponents
140          .
141          (eliminate
142            ParameterIdentity
143            (lambda unrestricted current : (family ParameterIdentity) . (family StdBool))
144            right
145            (branch
146              ParameterIdentityOf
147              rightComponents
148              .
149              (app
150                (eliminate
151                  StdList
152                  (lambda unrestricted current : (family StdList Bytes) .
153                    (pi unrestricted other : (family StdList Bytes) . (family StdBool)))
154                  leftComponents
155                  (branch
156                    StdListEmpty
157                    .
158                    (lambda unrestricted other : (family StdList Bytes) .
159                      (stdListIsEmpty Bytes other)))
160                  (branch
161                    StdListCons
162                    head
163                    tail
164                    induction
165                    .
166                    (lambda unrestricted other : (family StdList Bytes) .
167                      (eliminate
168                        StdList
169                        (lambda unrestricted current : (family StdList Bytes) . (family StdBool))
170                        other
171                        (branch StdListEmpty . (constructor StdBool StdFalse))
172                        (branch
173                          StdListCons
174                          otherHead
175                          otherTail
176                          otherInduction
177                          .
178                          (stdBoolAnd
179                            (stdBoolFromNatural (bytes-equal head otherHead))
180                            (induction otherTail)))))))
181                rightComponents)))))))
182
183def parameterSchemaIdentityOf =
184  (lambda unrestricted schema : (family ParameterSchema) .
185    (eliminate
186      ParameterSchema
187      (lambda unrestricted current : (family ParameterSchema) . (family ParameterIdentity))
188      schema
189      (branch ParameterSchemaOf identity shape representation storage . identity)))
190
191def parameterSchemaContainsIdentity =
192  (lambda unrestricted wanted : (family ParameterIdentity) .
193    (lambda unrestricted schemas : (family StdList (family ParameterSchema)) .
194      (eliminate
195        StdList
196        (lambda unrestricted current : (family StdList (family ParameterSchema)) . (family StdBool))
197        schemas
198        (branch StdListEmpty . (constructor StdBool StdFalse))
199        (branch
200          StdListCons
201          head
202          tail
203          induction
204          .
205          (stdBoolOr (parameterIdentityEqual wanted (parameterSchemaIdentityOf head)) induction)))))
206
207-- True exactly when two schema entries claim the same structural path.
208def parameterSchemaHasCollision =
209  (lambda unrestricted schemas : (family StdList (family ParameterSchema)) .
210    (eliminate
211      StdList
212      (lambda unrestricted current : (family StdList (family ParameterSchema)) . (family StdBool))
213      schemas
214      (branch StdListEmpty . (constructor StdBool StdFalse))
215      (branch
216        StdListCons
217        head
218        tail
219        induction
220        .
221        (stdBoolOr
222          (parameterSchemaContainsIdentity (parameterSchemaIdentityOf head) tail)
223          induction))))
224
225def learnerSlotsForKind =
226  (lambda unrestricted kind : (family LearnerStateKind) .
227    (lambda unrestricted parameters : (family StdList (family ParameterSchema)) .
228      (stdListMap
229        (family ParameterSchema)
230        (family LearnerStateSlot)
231        (lambda unrestricted parameter : (family ParameterSchema) .
232          (constructor
233            LearnerStateSlot
234            LearnerParameterState
235            (parameterSchemaIdentityOf parameter)
236            kind))
237        parameters)))
238
239-- Optimizer state is derived from the optimizer contract and parameter list:
240-- SGD has none; momentum has one velocity per parameter; AdamW has first and
241-- second moments per parameter plus one global step counter.
242def deriveLearnerStateSchema =
243  (lambda unrestricted contract : (family LearnerContract) .
244    (lambda unrestricted parameters : (family StdList (family ParameterSchema)) .
245      (eliminate
246        LearnerContract
247        (lambda unrestricted current : (family LearnerContract) . (family LearnerStateSchema))
248        contract
249        (branch
250          LearnerSGD
251          .
252          (constructor
253            LearnerStateSchema
254            LearnerStateSchemaOf
255            (constructor StdList StdListEmpty (family LearnerStateSlot))))
256        (branch
257          LearnerMomentum
258          .
259          (constructor
260            LearnerStateSchema
261            LearnerStateSchemaOf
262            (learnerSlotsForKind (constructor LearnerStateKind LearnerVelocity) parameters)))
263        (branch
264          LearnerAdamW
265          .
266          (constructor
267            LearnerStateSchema
268            LearnerStateSchemaOf
269            (stdListAppend
270              (family LearnerStateSlot)
271              (learnerSlotsForKind (constructor LearnerStateKind LearnerFirstMoment) parameters)
272              (stdListAppend
273                (family LearnerStateSlot)
274                (learnerSlotsForKind (constructor LearnerStateKind LearnerSecondMoment) parameters)
275                (constructor
276                  StdList
277                  StdListCons
278                  (family LearnerStateSlot)
279                  (constructor
280                    LearnerStateSlot
281                    LearnerGlobalState
282                    (constructor LearnerStateKind LearnerStepCounter))
283                  (constructor StdList StdListEmpty (family LearnerStateSlot))))))))))
284
285def learnerStateSchemaSize =
286  (lambda unrestricted schema : (family LearnerStateSchema) .
287    (eliminate
288      LearnerStateSchema
289      (lambda unrestricted current : (family LearnerStateSchema) . Nat)
290      schema
291      (branch LearnerStateSchemaOf slots . (stdListLength (family LearnerStateSlot) slots))))
292
293def parameterIdentityComponentsOf =
294  (lambda unrestricted identity : (family ParameterIdentity) .
295    (eliminate
296      ParameterIdentity
297      (lambda unrestricted current : (family ParameterIdentity) . (family StdList Bytes))
298      identity
299      (branch ParameterIdentityOf components . components)))
300
301def learnerStateKindName =
302  (lambda unrestricted kind : (family LearnerStateKind) .
303    (eliminate
304      LearnerStateKind
305      (lambda unrestricted current : (family LearnerStateKind) . Bytes)
306      kind
307      (branch LearnerVelocity . b"velocity")
308      (branch LearnerFirstMoment . b"first-moment")
309      (branch LearnerSecondMoment . b"second-moment")
310      (branch LearnerStepCounter . b"step-counter")))
311
312-- a learner-state slot's key: the parameter's path then the state kind, or the kind alone
313def learnerStateSlotKey =
314  (lambda unrestricted slot : (family LearnerStateSlot) .
315    (eliminate
316      LearnerStateSlot
317      (lambda unrestricted current : (family LearnerStateSlot) . (family StdList Bytes))
318      slot
319      (branch
320        LearnerParameterState
321        parameter
322        kind
323        .
324        (stdListAppend
325          Bytes
326          (parameterIdentityComponentsOf parameter)
327          (constructor
328            StdList
329            StdListCons
330            Bytes
331            (learnerStateKindName kind)
332            (constructor StdList StdListEmpty Bytes))))
333      (branch
334        LearnerGlobalState
335        kind
336        .
337        (constructor
338          StdList
339          StdListCons
340          Bytes
341          (learnerStateKindName kind)
342          (constructor StdList StdListEmpty Bytes)))))
343
344def continuationCategoryCode =
345  (lambda unrestricted category : (family ContinuationCategory) .
346    (eliminate
347      ContinuationCategory
348      (lambda unrestricted current : (family ContinuationCategory) . Nat)
349      category
350      (branch ContinuationParameters . 1)
351      (branch ContinuationComponentState . 2)
352      (branch ContinuationLearnerState . 3)
353      (branch ContinuationRng . 4)
354      (branch ContinuationDataCursor . 5)
355      (branch ContinuationCounters . 6)
356      (branch ContinuationSchedule . 7)
357      (branch ContinuationIdentities . 8)))
358
359def continuationEntryCategoryOf =
360  (lambda unrestricted entry : (family ContinuationEntry) .
361    (eliminate
362      ContinuationEntry
363      (lambda unrestricted current : (family ContinuationEntry) . (family ContinuationCategory))
364      entry
365      (branch ContinuationEntryOf category key . category)))
366
367def continuationSystemEntry =
368  (lambda unrestricted category : (family ContinuationCategory) .
369    (lambda unrestricted name : Bytes .
370      (constructor
371        ContinuationEntry
372        ContinuationEntryOf
373        category
374        (constructor StdList StdListCons Bytes name (constructor StdList StdListEmpty Bytes)))))
375
376def continuationEntriesOfSix =
377  (lambda unrestricted a : (family ContinuationEntry) .
378    (lambda unrestricted b : (family ContinuationEntry) .
379      (lambda unrestricted c : (family ContinuationEntry) .
380        (lambda unrestricted d : (family ContinuationEntry) .
381          (lambda unrestricted e : (family ContinuationEntry) .
382            (lambda unrestricted f : (family ContinuationEntry) .
383              (constructor
384                StdList
385                StdListCons
386                (family ContinuationEntry)
387                a
388                (constructor
389                  StdList
390                  StdListCons
391                  (family ContinuationEntry)
392                  b
393                  (constructor
394                    StdList
395                    StdListCons
396                    (family ContinuationEntry)
397                    c
398                    (constructor
399                      StdList
400                      StdListCons
401                      (family ContinuationEntry)
402                      d
403                      (constructor
404                        StdList
405                        StdListCons
406                        (family ContinuationEntry)
407                        e
408                        (constructor
409                          StdList
410                          StdListCons
411                          (family ContinuationEntry)
412                          f
413                          (constructor StdList StdListEmpty (family ContinuationEntry))))))))))))))
414
415-- The continuation schema of a system: one Parameters entry per parameter, one LearnerState
416-- entry per derived learner-state slot, and the system-wide categories.
417def deriveContinuationSchema =
418  (lambda unrestricted system : (family SystemDefinition) .
419    (eliminate
420      SystemDefinition
421      (lambda unrestricted current : (family SystemDefinition) . (family ContinuationSchema))
422      system
423      (branch
424        SystemDefinitionOf
425        name
426        parameters
427        learner
428        .
429        (constructor
430          ContinuationSchema
431          ContinuationSchemaOf
432          name
433          (stdListAppend
434            (family ContinuationEntry)
435            (stdListMap
436              (family ParameterSchema)
437              (family ContinuationEntry)
438              (lambda unrestricted parameter : (family ParameterSchema) .
439                (constructor
440                  ContinuationEntry
441                  ContinuationEntryOf
442                  (constructor ContinuationCategory ContinuationParameters)
443                  (parameterIdentityComponentsOf (parameterSchemaIdentityOf parameter))))
444              parameters)
445            (stdListAppend
446              (family ContinuationEntry)
447              (eliminate
448                LearnerStateSchema
449                (lambda unrestricted current : (family LearnerStateSchema) .
450                  (family StdList (family ContinuationEntry)))
451                (deriveLearnerStateSchema learner parameters)
452                (branch
453                  LearnerStateSchemaOf
454                  slots
455                  .
456                  (stdListMap
457                    (family LearnerStateSlot)
458                    (family ContinuationEntry)
459                    (lambda unrestricted slot : (family LearnerStateSlot) .
460                      (constructor
461                        ContinuationEntry
462                        ContinuationEntryOf
463                        (constructor ContinuationCategory ContinuationLearnerState)
464                        (learnerStateSlotKey slot)))
465                    slots)))
466              (continuationEntriesOfSix
467                (continuationSystemEntry
468                  (constructor ContinuationCategory ContinuationComponentState)
469                  b"component-state")
470                (continuationSystemEntry (constructor ContinuationCategory ContinuationRng) b"rng")
471                (continuationSystemEntry
472                  (constructor ContinuationCategory ContinuationDataCursor)
473                  b"data-cursor")
474                (continuationSystemEntry
475                  (constructor ContinuationCategory ContinuationCounters)
476                  b"counters")
477                (continuationSystemEntry
478                  (constructor ContinuationCategory ContinuationSchedule)
479                  b"schedule")
480                (continuationSystemEntry
481                  (constructor ContinuationCategory ContinuationIdentities)
482                  b"identities"))))))))
483
484def continuationSchemaEntriesOf =
485  (lambda unrestricted schema : (family ContinuationSchema) .
486    (eliminate
487      ContinuationSchema
488      (lambda unrestricted current : (family ContinuationSchema) .
489        (family StdList (family ContinuationEntry)))
490      schema
491      (branch ContinuationSchemaOf system entries . entries)))
492
493def continuationSchemaSize =
494  (lambda unrestricted schema : (family ContinuationSchema) .
495    (stdListLength (family ContinuationEntry) (continuationSchemaEntriesOf schema)))
496
497-- does the schema carry at least one entry of the category?
498def continuationSchemaHasCategory =
499  (lambda unrestricted category : (family ContinuationCategory) .
500    (lambda unrestricted schema : (family ContinuationSchema) .
501      (eliminate
502        StdList
503        (lambda unrestricted current : (family StdList (family ContinuationEntry)) .
504          (family StdBool))
505        (continuationSchemaEntriesOf schema)
506        (branch StdListEmpty . (constructor StdBool StdFalse))
507        (branch
508          StdListCons
509          head
510          tail
511          induction
512          .
513          (stdBoolOr
514            (stdBoolFromNatural
515              (naturalEqual
516                (continuationCategoryCode (continuationEntryCategoryOf head))
517                (continuationCategoryCode category)))
518            induction)))))
519
520-- CKP-003: a comparison is complete only when every K2 category is present
521def continuationSchemaCoversAll =
522  (lambda unrestricted schema : (family ContinuationSchema) .
523    (stdBoolAnd
524      (continuationSchemaHasCategory
525        (constructor ContinuationCategory ContinuationParameters)
526        schema)
527      (stdBoolAnd
528        (continuationSchemaHasCategory
529          (constructor ContinuationCategory ContinuationComponentState)
530          schema)
531        (stdBoolAnd
532          (continuationSchemaHasCategory
533            (constructor ContinuationCategory ContinuationLearnerState)
534            schema)
535          (stdBoolAnd
536            (continuationSchemaHasCategory
537              (constructor ContinuationCategory ContinuationRng)
538              schema)
539            (stdBoolAnd
540              (continuationSchemaHasCategory
541                (constructor ContinuationCategory ContinuationDataCursor)
542                schema)
543              (stdBoolAnd
544                (continuationSchemaHasCategory
545                  (constructor ContinuationCategory ContinuationCounters)
546                  schema)
547                (stdBoolAnd
548                  (continuationSchemaHasCategory
549                    (constructor ContinuationCategory ContinuationSchedule)
550                    schema)
551                  (continuationSchemaHasCategory
552                    (constructor ContinuationCategory ContinuationIdentities)
553                    schema)))))))))
554
555-- The update plan of a system: every parameter's update, then every learner-state slot's update.
556def deriveUpdatePlan =
557  (lambda unrestricted system : (family SystemDefinition) .
558    (eliminate
559      SystemDefinition
560      (lambda unrestricted current : (family SystemDefinition) .
561        (family StdList (family UpdateStep)))
562      system
563      (branch
564        SystemDefinitionOf
565        name
566        parameters
567        learner
568        .
569        (stdListAppend
570          (family UpdateStep)
571          (stdListMap
572            (family ParameterSchema)
573            (family UpdateStep)
574            (lambda unrestricted parameter : (family ParameterSchema) .
575              (constructor UpdateStep UpdateParameter (parameterSchemaIdentityOf parameter)))
576            parameters)
577          (eliminate
578            LearnerStateSchema
579            (lambda unrestricted current : (family LearnerStateSchema) .
580              (family StdList (family UpdateStep)))
581            (deriveLearnerStateSchema learner parameters)
582            (branch
583              LearnerStateSchemaOf
584              slots
585              .
586              (stdListMap
587                (family LearnerStateSlot)
588                (family UpdateStep)
589                (lambda unrestricted slot : (family LearnerStateSlot) .
590                  (constructor UpdateStep UpdateLearnerState slot))
591                slots)))))))
592
593def updatePathBytes =
594  (lambda unrestricted components : (family StdList Bytes) .
595    (stdListFold
596      Bytes
597      Bytes
598      (lambda unrestricted component : Bytes .
599        (lambda unrestricted rest : Bytes . (bytes-append component (bytes-append b"/" rest))))
600      b""
601      components))
602
603-- The state digest a plan's transition produces: the ordered record of every state it writes.
604-- Dropping or duplicating one step changes it.
605def updatePlanDigest =
606  (lambda unrestricted plan : (family StdList (family UpdateStep)) .
607    (stdListFold
608      (family UpdateStep)
609      Bytes
610      (lambda unrestricted step : (family UpdateStep) .
611        (lambda unrestricted rest : Bytes .
612          (bytes-append
613            (eliminate
614              UpdateStep
615              (lambda unrestricted current : (family UpdateStep) . Bytes)
616              step
617              (branch
618                UpdateParameter
619                parameter
620                .
621                (bytes-append b"p:" (updatePathBytes (parameterIdentityComponentsOf parameter))))
622              (branch
623                UpdateLearnerState
624                slot
625                .
626                (bytes-append b"s:" (updatePathBytes (learnerStateSlotKey slot)))))
627            (bytes-append b";" rest))))
628      b""
629      plan))
630
631def systemNameOf =
632  (lambda unrestricted system : (family SystemDefinition) .
633    (eliminate
634      SystemDefinition
635      (lambda unrestricted current : (family SystemDefinition) . Bytes)
636      system
637      (branch SystemDefinitionOf name parameters learner . name)))
638
639def systemParametersOf =
640  (lambda unrestricted system : (family SystemDefinition) .
641    (eliminate
642      SystemDefinition
643      (lambda unrestricted current : (family SystemDefinition) .
644        (family StdList (family ParameterSchema)))
645      system
646      (branch SystemDefinitionOf name parameters learner . parameters)))
647
648def systemLearnerStateSchema =
649  (lambda unrestricted system : (family SystemDefinition) .
650    (eliminate
651      SystemDefinition
652      (lambda unrestricted current : (family SystemDefinition) . (family LearnerStateSchema))
653      system
654      (branch
655        SystemDefinitionOf
656        name
657        parameters
658        learner
659        .
660        (deriveLearnerStateSchema learner parameters))))
661
662def continuationCategoryName =
663  (lambda unrestricted category : (family ContinuationCategory) .
664    (eliminate
665      ContinuationCategory
666      (lambda unrestricted current : (family ContinuationCategory) . Bytes)
667      category
668      (branch ContinuationParameters . b"parameters")
669      (branch ContinuationComponentState . b"component-state")
670      (branch ContinuationLearnerState . b"learner-state")
671      (branch ContinuationRng . b"rng")
672      (branch ContinuationDataCursor . b"data-cursor")
673      (branch ContinuationCounters . b"counters")
674      (branch ContinuationSchedule . b"schedule")
675      (branch ContinuationIdentities . b"identities")))
676
677def continuationEntryKeyOf =
678  (lambda unrestricted entry : (family ContinuationEntry) .
679    (eliminate
680      ContinuationEntry
681      (lambda unrestricted current : (family ContinuationEntry) . (family StdList Bytes))
682      entry
683      (branch ContinuationEntryOf category key . key)))
684
685-- one line per entry: the category, then the structural key, for rendering
686def continuationSchemaKeys =
687  (lambda unrestricted schema : (family ContinuationSchema) .
688    (stdListFold
689      (family ContinuationEntry)
690      Bytes
691      (lambda unrestricted entry : (family ContinuationEntry) .
692        (lambda unrestricted rest : Bytes .
693          (bytes-append
694            (continuationCategoryName (continuationEntryCategoryOf entry))
695            (bytes-append
696              b" "
697              (bytes-append
698                (updatePathBytes (continuationEntryKeyOf entry))
699                (bytes-append b"\n" rest))))))
700      b""
701      (continuationSchemaEntriesOf schema)))

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.