module Representation.Schema import Std.Foundation import Std.Natural import Std.List -- A stable parameter identity is its structural component path. Names remain -- separate path components, so concatenation cannot create accidental aliases. family ParameterIdentity : Type 0 constructor ParameterIdentityOf field unrestricted parameterIdentityComponents : (family StdList Bytes) end-family family ParameterShape : Type 0 constructor ParameterShapeOf field unrestricted parameterShapeDimensions : (family StdList Nat) end-family family LogicalOwnership : Type 0 constructor LogicalOwned constructor LogicalShared constructor LogicalBorrowed end-family family LogicalLifetime : Type 0 constructor LogicalPersistent constructor LogicalStep constructor LogicalInvocation end-family -- Logical storage has an extent and alignment but deliberately no address. family LogicalStorage : Type 0 constructor LogicalStorageOf field unrestricted logicalStorageExtent : Nat field unrestricted logicalStorageAlignment : Nat field unrestricted logicalStorageOwnership : (family LogicalOwnership) field unrestricted logicalStorageLifetime : (family LogicalLifetime) end-family family ParameterSchema : Type 0 constructor ParameterSchemaOf field unrestricted parameterSchemaIdentity : (family ParameterIdentity) field unrestricted parameterSchemaShape : (family ParameterShape) field unrestricted parameterSchemaRepresentation : Bytes field unrestricted parameterSchemaStorage : (family LogicalStorage) end-family family LearnerContract : Type 0 constructor LearnerSGD constructor LearnerMomentum constructor LearnerAdamW end-family family LearnerStateKind : Type 0 constructor LearnerVelocity constructor LearnerFirstMoment constructor LearnerSecondMoment constructor LearnerStepCounter end-family family LearnerStateSlot : Type 0 constructor LearnerParameterState field unrestricted learnerStateParameter : (family ParameterIdentity) field unrestricted learnerStateKind : (family LearnerStateKind) constructor LearnerGlobalState field unrestricted learnerGlobalStateKind : (family LearnerStateKind) end-family family LearnerStateSchema : Type 0 constructor LearnerStateSchemaOf field unrestricted learnerStateSlots : (family StdList (family LearnerStateSlot)) end-family -- PRD 19 L24: a system is its structural parameters and its learner; everything -- the continuation and the update plan need is DERIVED from this value. family SystemDefinition : Type 0 constructor SystemDefinitionOf field unrestricted systemName : Bytes field unrestricted systemParameters : (family StdList (family ParameterSchema)) field unrestricted systemLearner : (family LearnerContract) end-family -- PRD 18 K2: the categories a continuation must carry. family ContinuationCategory : Type 0 constructor ContinuationParameters constructor ContinuationComponentState constructor ContinuationLearnerState constructor ContinuationRng constructor ContinuationDataCursor constructor ContinuationCounters constructor ContinuationSchedule constructor ContinuationIdentities end-family family ContinuationEntry : Type 0 constructor ContinuationEntryOf field unrestricted continuationEntryCategory : (family ContinuationCategory) field unrestricted continuationEntryKey : (family StdList Bytes) end-family family ContinuationSchema : Type 0 constructor ContinuationSchemaOf field unrestricted continuationSchemaSystem : Bytes field unrestricted continuationSchemaEntries : (family StdList (family ContinuationEntry)) end-family -- One step of the derived update plan: a parameter's update, or one learner-state slot's update. family UpdateStep : Type 0 constructor UpdateParameter field unrestricted updateStepParameter : (family ParameterIdentity) constructor UpdateLearnerState field unrestricted updateStepState : (family LearnerStateSlot) end-family def parameterIdentityEqual = (lambda unrestricted left : (family ParameterIdentity) . (lambda unrestricted right : (family ParameterIdentity) . (eliminate ParameterIdentity (lambda unrestricted current : (family ParameterIdentity) . (family StdBool)) left (branch ParameterIdentityOf leftComponents . (eliminate ParameterIdentity (lambda unrestricted current : (family ParameterIdentity) . (family StdBool)) right (branch ParameterIdentityOf rightComponents . (app (eliminate StdList (lambda unrestricted current : (family StdList Bytes) . (pi unrestricted other : (family StdList Bytes) . (family StdBool))) leftComponents (branch StdListEmpty . (lambda unrestricted other : (family StdList Bytes) . (stdListIsEmpty Bytes other))) (branch StdListCons head tail induction . (lambda unrestricted other : (family StdList Bytes) . (eliminate StdList (lambda unrestricted current : (family StdList Bytes) . (family StdBool)) other (branch StdListEmpty . (constructor StdBool StdFalse)) (branch StdListCons otherHead otherTail otherInduction . (stdBoolAnd (stdBoolFromNatural (bytes-equal head otherHead)) (induction otherTail))))))) rightComponents))))))) def parameterSchemaIdentityOf = (lambda unrestricted schema : (family ParameterSchema) . (eliminate ParameterSchema (lambda unrestricted current : (family ParameterSchema) . (family ParameterIdentity)) schema (branch ParameterSchemaOf identity shape representation storage . identity))) def parameterSchemaContainsIdentity = (lambda unrestricted wanted : (family ParameterIdentity) . (lambda unrestricted schemas : (family StdList (family ParameterSchema)) . (eliminate StdList (lambda unrestricted current : (family StdList (family ParameterSchema)) . (family StdBool)) schemas (branch StdListEmpty . (constructor StdBool StdFalse)) (branch StdListCons head tail induction . (stdBoolOr (parameterIdentityEqual wanted (parameterSchemaIdentityOf head)) induction))))) -- True exactly when two schema entries claim the same structural path. def parameterSchemaHasCollision = (lambda unrestricted schemas : (family StdList (family ParameterSchema)) . (eliminate StdList (lambda unrestricted current : (family StdList (family ParameterSchema)) . (family StdBool)) schemas (branch StdListEmpty . (constructor StdBool StdFalse)) (branch StdListCons head tail induction . (stdBoolOr (parameterSchemaContainsIdentity (parameterSchemaIdentityOf head) tail) induction)))) def learnerSlotsForKind = (lambda unrestricted kind : (family LearnerStateKind) . (lambda unrestricted parameters : (family StdList (family ParameterSchema)) . (stdListMap (family ParameterSchema) (family LearnerStateSlot) (lambda unrestricted parameter : (family ParameterSchema) . (constructor LearnerStateSlot LearnerParameterState (parameterSchemaIdentityOf parameter) kind)) parameters))) -- Optimizer state is derived from the optimizer contract and parameter list: -- SGD has none; momentum has one velocity per parameter; AdamW has first and -- second moments per parameter plus one global step counter. def deriveLearnerStateSchema = (lambda unrestricted contract : (family LearnerContract) . (lambda unrestricted parameters : (family StdList (family ParameterSchema)) . (eliminate LearnerContract (lambda unrestricted current : (family LearnerContract) . (family LearnerStateSchema)) contract (branch LearnerSGD . (constructor LearnerStateSchema LearnerStateSchemaOf (constructor StdList StdListEmpty (family LearnerStateSlot)))) (branch LearnerMomentum . (constructor LearnerStateSchema LearnerStateSchemaOf (learnerSlotsForKind (constructor LearnerStateKind LearnerVelocity) parameters))) (branch LearnerAdamW . (constructor LearnerStateSchema LearnerStateSchemaOf (stdListAppend (family LearnerStateSlot) (learnerSlotsForKind (constructor LearnerStateKind LearnerFirstMoment) parameters) (stdListAppend (family LearnerStateSlot) (learnerSlotsForKind (constructor LearnerStateKind LearnerSecondMoment) parameters) (constructor StdList StdListCons (family LearnerStateSlot) (constructor LearnerStateSlot LearnerGlobalState (constructor LearnerStateKind LearnerStepCounter)) (constructor StdList StdListEmpty (family LearnerStateSlot)))))))))) def learnerStateSchemaSize = (lambda unrestricted schema : (family LearnerStateSchema) . (eliminate LearnerStateSchema (lambda unrestricted current : (family LearnerStateSchema) . Nat) schema (branch LearnerStateSchemaOf slots . (stdListLength (family LearnerStateSlot) slots)))) def parameterIdentityComponentsOf = (lambda unrestricted identity : (family ParameterIdentity) . (eliminate ParameterIdentity (lambda unrestricted current : (family ParameterIdentity) . (family StdList Bytes)) identity (branch ParameterIdentityOf components . components))) def learnerStateKindName = (lambda unrestricted kind : (family LearnerStateKind) . (eliminate LearnerStateKind (lambda unrestricted current : (family LearnerStateKind) . Bytes) kind (branch LearnerVelocity . b"velocity") (branch LearnerFirstMoment . b"first-moment") (branch LearnerSecondMoment . b"second-moment") (branch LearnerStepCounter . b"step-counter"))) -- a learner-state slot's key: the parameter's path then the state kind, or the kind alone def learnerStateSlotKey = (lambda unrestricted slot : (family LearnerStateSlot) . (eliminate LearnerStateSlot (lambda unrestricted current : (family LearnerStateSlot) . (family StdList Bytes)) slot (branch LearnerParameterState parameter kind . (stdListAppend Bytes (parameterIdentityComponentsOf parameter) (constructor StdList StdListCons Bytes (learnerStateKindName kind) (constructor StdList StdListEmpty Bytes)))) (branch LearnerGlobalState kind . (constructor StdList StdListCons Bytes (learnerStateKindName kind) (constructor StdList StdListEmpty Bytes))))) def continuationCategoryCode = (lambda unrestricted category : (family ContinuationCategory) . (eliminate ContinuationCategory (lambda unrestricted current : (family ContinuationCategory) . Nat) category (branch ContinuationParameters . 1) (branch ContinuationComponentState . 2) (branch ContinuationLearnerState . 3) (branch ContinuationRng . 4) (branch ContinuationDataCursor . 5) (branch ContinuationCounters . 6) (branch ContinuationSchedule . 7) (branch ContinuationIdentities . 8))) def continuationEntryCategoryOf = (lambda unrestricted entry : (family ContinuationEntry) . (eliminate ContinuationEntry (lambda unrestricted current : (family ContinuationEntry) . (family ContinuationCategory)) entry (branch ContinuationEntryOf category key . category))) def continuationSystemEntry = (lambda unrestricted category : (family ContinuationCategory) . (lambda unrestricted name : Bytes . (constructor ContinuationEntry ContinuationEntryOf category (constructor StdList StdListCons Bytes name (constructor StdList StdListEmpty Bytes))))) def continuationEntriesOfSix = (lambda unrestricted a : (family ContinuationEntry) . (lambda unrestricted b : (family ContinuationEntry) . (lambda unrestricted c : (family ContinuationEntry) . (lambda unrestricted d : (family ContinuationEntry) . (lambda unrestricted e : (family ContinuationEntry) . (lambda unrestricted f : (family ContinuationEntry) . (constructor StdList StdListCons (family ContinuationEntry) a (constructor StdList StdListCons (family ContinuationEntry) b (constructor StdList StdListCons (family ContinuationEntry) c (constructor StdList StdListCons (family ContinuationEntry) d (constructor StdList StdListCons (family ContinuationEntry) e (constructor StdList StdListCons (family ContinuationEntry) f (constructor StdList StdListEmpty (family ContinuationEntry)))))))))))))) -- The continuation schema of a system: one Parameters entry per parameter, one LearnerState -- entry per derived learner-state slot, and the system-wide categories. def deriveContinuationSchema = (lambda unrestricted system : (family SystemDefinition) . (eliminate SystemDefinition (lambda unrestricted current : (family SystemDefinition) . (family ContinuationSchema)) system (branch SystemDefinitionOf name parameters learner . (constructor ContinuationSchema ContinuationSchemaOf name (stdListAppend (family ContinuationEntry) (stdListMap (family ParameterSchema) (family ContinuationEntry) (lambda unrestricted parameter : (family ParameterSchema) . (constructor ContinuationEntry ContinuationEntryOf (constructor ContinuationCategory ContinuationParameters) (parameterIdentityComponentsOf (parameterSchemaIdentityOf parameter)))) parameters) (stdListAppend (family ContinuationEntry) (eliminate LearnerStateSchema (lambda unrestricted current : (family LearnerStateSchema) . (family StdList (family ContinuationEntry))) (deriveLearnerStateSchema learner parameters) (branch LearnerStateSchemaOf slots . (stdListMap (family LearnerStateSlot) (family ContinuationEntry) (lambda unrestricted slot : (family LearnerStateSlot) . (constructor ContinuationEntry ContinuationEntryOf (constructor ContinuationCategory ContinuationLearnerState) (learnerStateSlotKey slot))) slots))) (continuationEntriesOfSix (continuationSystemEntry (constructor ContinuationCategory ContinuationComponentState) b"component-state") (continuationSystemEntry (constructor ContinuationCategory ContinuationRng) b"rng") (continuationSystemEntry (constructor ContinuationCategory ContinuationDataCursor) b"data-cursor") (continuationSystemEntry (constructor ContinuationCategory ContinuationCounters) b"counters") (continuationSystemEntry (constructor ContinuationCategory ContinuationSchedule) b"schedule") (continuationSystemEntry (constructor ContinuationCategory ContinuationIdentities) b"identities")))))))) def continuationSchemaEntriesOf = (lambda unrestricted schema : (family ContinuationSchema) . (eliminate ContinuationSchema (lambda unrestricted current : (family ContinuationSchema) . (family StdList (family ContinuationEntry))) schema (branch ContinuationSchemaOf system entries . entries))) def continuationSchemaSize = (lambda unrestricted schema : (family ContinuationSchema) . (stdListLength (family ContinuationEntry) (continuationSchemaEntriesOf schema))) -- does the schema carry at least one entry of the category? def continuationSchemaHasCategory = (lambda unrestricted category : (family ContinuationCategory) . (lambda unrestricted schema : (family ContinuationSchema) . (eliminate StdList (lambda unrestricted current : (family StdList (family ContinuationEntry)) . (family StdBool)) (continuationSchemaEntriesOf schema) (branch StdListEmpty . (constructor StdBool StdFalse)) (branch StdListCons head tail induction . (stdBoolOr (stdBoolFromNatural (naturalEqual (continuationCategoryCode (continuationEntryCategoryOf head)) (continuationCategoryCode category))) induction))))) -- CKP-003: a comparison is complete only when every K2 category is present def continuationSchemaCoversAll = (lambda unrestricted schema : (family ContinuationSchema) . (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationParameters) schema) (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationComponentState) schema) (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationLearnerState) schema) (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationRng) schema) (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationDataCursor) schema) (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationCounters) schema) (stdBoolAnd (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationSchedule) schema) (continuationSchemaHasCategory (constructor ContinuationCategory ContinuationIdentities) schema))))))))) -- The update plan of a system: every parameter's update, then every learner-state slot's update. def deriveUpdatePlan = (lambda unrestricted system : (family SystemDefinition) . (eliminate SystemDefinition (lambda unrestricted current : (family SystemDefinition) . (family StdList (family UpdateStep))) system (branch SystemDefinitionOf name parameters learner . (stdListAppend (family UpdateStep) (stdListMap (family ParameterSchema) (family UpdateStep) (lambda unrestricted parameter : (family ParameterSchema) . (constructor UpdateStep UpdateParameter (parameterSchemaIdentityOf parameter))) parameters) (eliminate LearnerStateSchema (lambda unrestricted current : (family LearnerStateSchema) . (family StdList (family UpdateStep))) (deriveLearnerStateSchema learner parameters) (branch LearnerStateSchemaOf slots . (stdListMap (family LearnerStateSlot) (family UpdateStep) (lambda unrestricted slot : (family LearnerStateSlot) . (constructor UpdateStep UpdateLearnerState slot)) slots))))))) def updatePathBytes = (lambda unrestricted components : (family StdList Bytes) . (stdListFold Bytes Bytes (lambda unrestricted component : Bytes . (lambda unrestricted rest : Bytes . (bytes-append component (bytes-append b"/" rest)))) b"" components)) -- The state digest a plan's transition produces: the ordered record of every state it writes. -- Dropping or duplicating one step changes it. def updatePlanDigest = (lambda unrestricted plan : (family StdList (family UpdateStep)) . (stdListFold (family UpdateStep) Bytes (lambda unrestricted step : (family UpdateStep) . (lambda unrestricted rest : Bytes . (bytes-append (eliminate UpdateStep (lambda unrestricted current : (family UpdateStep) . Bytes) step (branch UpdateParameter parameter . (bytes-append b"p:" (updatePathBytes (parameterIdentityComponentsOf parameter)))) (branch UpdateLearnerState slot . (bytes-append b"s:" (updatePathBytes (learnerStateSlotKey slot))))) (bytes-append b";" rest)))) b"" plan)) def systemNameOf = (lambda unrestricted system : (family SystemDefinition) . (eliminate SystemDefinition (lambda unrestricted current : (family SystemDefinition) . Bytes) system (branch SystemDefinitionOf name parameters learner . name))) def systemParametersOf = (lambda unrestricted system : (family SystemDefinition) . (eliminate SystemDefinition (lambda unrestricted current : (family SystemDefinition) . (family StdList (family ParameterSchema))) system (branch SystemDefinitionOf name parameters learner . parameters))) def systemLearnerStateSchema = (lambda unrestricted system : (family SystemDefinition) . (eliminate SystemDefinition (lambda unrestricted current : (family SystemDefinition) . (family LearnerStateSchema)) system (branch SystemDefinitionOf name parameters learner . (deriveLearnerStateSchema learner parameters)))) def continuationCategoryName = (lambda unrestricted category : (family ContinuationCategory) . (eliminate ContinuationCategory (lambda unrestricted current : (family ContinuationCategory) . Bytes) category (branch ContinuationParameters . b"parameters") (branch ContinuationComponentState . b"component-state") (branch ContinuationLearnerState . b"learner-state") (branch ContinuationRng . b"rng") (branch ContinuationDataCursor . b"data-cursor") (branch ContinuationCounters . b"counters") (branch ContinuationSchedule . b"schedule") (branch ContinuationIdentities . b"identities"))) def continuationEntryKeyOf = (lambda unrestricted entry : (family ContinuationEntry) . (eliminate ContinuationEntry (lambda unrestricted current : (family ContinuationEntry) . (family StdList Bytes)) entry (branch ContinuationEntryOf category key . key))) -- one line per entry: the category, then the structural key, for rendering def continuationSchemaKeys = (lambda unrestricted schema : (family ContinuationSchema) . (stdListFold (family ContinuationEntry) Bytes (lambda unrestricted entry : (family ContinuationEntry) . (lambda unrestricted rest : Bytes . (bytes-append (continuationCategoryName (continuationEntryCategoryOf entry)) (bytes-append b" " (bytes-append (updatePathBytes (continuationEntryKeyOf entry)) (bytes-append b"\n" rest)))))) b"" (continuationSchemaEntriesOf schema)))