module Checkpoint.Envelope import Data.Bytes import Data.SHA256Digest import Std.Foundation import Std.Natural -- THE CHECKPOINT ENVELOPE: what a native host writes around a system's raw -- training state, and refuses to resume from unless it holds. -- -- A checkpoint file is a header of `headerBytes` (a whole number of pages) -- and then the payload, unmodified (for Bob and Coppelius: the P/M/V -- planes). The header, little-endian: -- 0 magic "ALPHACKP" -- 8 version (1) -- 16 header bytes -- 24 payload bytes -- 32 chunk bytes -- 40 full chunks (payload bytes / chunk bytes) -- 48 tail bytes (payload bytes mod chunk bytes) -- 56 completed updates (the learner's step count) -- 64 completed invocations (the sampler's position: batches consumed) -- 72 SHA-256 of the state schema's identity (model, state layout) -- 104 SHA-256 of the learner's identity (optimizer, hyperparameters) -- 136 SHA-256 of the data contract's identity (tokenizer, batch format) -- 168 SHA-256 of the initialization's identity (the RNG seed) -- 200 the SHA-256 of each chunk of the payload, in order (the tail last) -- then zeros to `headerBytes`. The file is complete when its size is -- exactly header + payload bytes. -- -- A system states its contract (below); the first 56 bytes and the four -- identity digests of any checkpoint it resumes from must be the ones the -- contract gives, and every chunk must hash to its digest. The first -- failure is named (`CheckpointEnvelopeVerdict`): a short or long file is -- truncated, another version or layout is refused, another system's, -- learner's, data's or seed's state is foreign, and a chunk whose digest -- differs -- a swapped or corrupted file -- is refused by its index. family CheckpointEnvelopeContract : Type 0 constructor CheckpointEnvelopeContractValue field unrestricted checkpointEnvelopeSchemaIdentity : Bytes field unrestricted checkpointEnvelopeLearnerIdentity : Bytes field unrestricted checkpointEnvelopeDataIdentity : Bytes field unrestricted checkpointEnvelopeSeedIdentity : Bytes field unrestricted checkpointEnvelopePayloadBytes : Nat field unrestricted checkpointEnvelopeChunkBytes : Nat field unrestricted checkpointEnvelopeUpdatesPerInvocation : Nat end-family family CheckpointEnvelopeVerdict : Type 0 constructor CheckpointEnvelopeAccepted constructor CheckpointEnvelopeTruncated constructor CheckpointEnvelopeLayout constructor CheckpointEnvelopeForeignSchema constructor CheckpointEnvelopeForeignLearner constructor CheckpointEnvelopeForeignData constructor CheckpointEnvelopeForeignSeed constructor CheckpointEnvelopeChunkMismatch field unrestricted checkpointEnvelopeMismatchedChunk : Nat end-family def checkpointEnvelopeMagic : Bytes = b"ALPHACKP" def checkpointEnvelopeVersion : Nat = 1 def checkpointEnvelopePageBytes : Nat = 4096 def checkpointEnvelopeDigestBytes : Nat = 32 -- the header's fields, as byte offsets def checkpointEnvelopeVersionAt : Nat = 8 def checkpointEnvelopeHeaderBytesAt : Nat = 16 def checkpointEnvelopePayloadBytesAt : Nat = 24 def checkpointEnvelopeChunkBytesAt : Nat = 32 def checkpointEnvelopeFullChunksAt : Nat = 40 def checkpointEnvelopeTailBytesAt : Nat = 48 def checkpointEnvelopeUpdatesAt : Nat = 56 def checkpointEnvelopeInvocationsAt : Nat = 64 def checkpointEnvelopeSchemaAt : Nat = 72 def checkpointEnvelopeLearnerAt : Nat = 104 def checkpointEnvelopeDataAt : Nat = 136 def checkpointEnvelopeSeedAt : Nat = 168 def checkpointEnvelopeDigestsAt : Nat = 200 -- the bytes every checkpoint of a contract begins with: magic to tail def checkpointEnvelopeFixedBytes : Nat = 56 def checkpointEnvelopeContractPayload = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (eliminate CheckpointEnvelopeContract (lambda unrestricted current : (family CheckpointEnvelopeContract) . Nat) contract (branch CheckpointEnvelopeContractValue schema learner data seed payload chunk updates . payload))) def checkpointEnvelopeContractChunk = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (eliminate CheckpointEnvelopeContract (lambda unrestricted current : (family CheckpointEnvelopeContract) . Nat) contract (branch CheckpointEnvelopeContractValue schema learner data seed payload chunk updates . chunk))) def checkpointEnvelopeContractUpdates = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (eliminate CheckpointEnvelopeContract (lambda unrestricted current : (family CheckpointEnvelopeContract) . Nat) contract (branch CheckpointEnvelopeContractValue schema learner data seed payload chunk updates . updates))) -- the four identities' digests, in header order def checkpointEnvelopeIdentityDigests = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (eliminate CheckpointEnvelopeContract (lambda unrestricted current : (family CheckpointEnvelopeContract) . Bytes) contract (branch CheckpointEnvelopeContractValue schema learner data seed payload chunk updates . (bytes-append (sha256RawDigestOrEmpty schema) (bytes-append (sha256RawDigestOrEmpty learner) (bytes-append (sha256RawDigestOrEmpty data) (sha256RawDigestOrEmpty seed))))))) def checkpointEnvelopeFullChunks = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (naturalDivideUnchecked (checkpointEnvelopeContractPayload contract) (checkpointEnvelopeContractChunk contract))) def checkpointEnvelopeTailBytes = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (naturalModuloUnchecked (checkpointEnvelopeContractPayload contract) (checkpointEnvelopeContractChunk contract))) -- the chunks a payload has: the full ones and, when there is one, the tail def checkpointEnvelopeChunkCount = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (naturalAdd (checkpointEnvelopeFullChunks contract) (naturalNonzero (checkpointEnvelopeTailBytes contract)))) -- The final digest covers only bytes present in the payload. The generic -- byte-prefix operation zero-pads, so passing the full chunk extent for the -- tail would disagree with native file hashing while still passing a codec -- round trip that made the same mistake on both sides. def checkpointEnvelopeChunkExtent = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (lambda unrestricted index : Nat . (naturalSelect (naturalLess index (checkpointEnvelopeFullChunks contract)) (checkpointEnvelopeContractChunk contract) (naturalSelect (naturalEqual index (checkpointEnvelopeFullChunks contract)) (checkpointEnvelopeTailBytes contract) 0)))) def checkpointEnvelopeHeaderBytes = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (naturalMultiply (naturalDivideUnchecked (naturalAdd (naturalAdd checkpointEnvelopeDigestsAt (naturalMultiply checkpointEnvelopeDigestBytes (checkpointEnvelopeChunkCount contract))) (naturalSaturatingSubtract checkpointEnvelopePageBytes 1)) checkpointEnvelopePageBytes) checkpointEnvelopePageBytes)) def checkpointEnvelopeFileBytes = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (naturalAdd (checkpointEnvelopeHeaderBytes contract) (checkpointEnvelopeContractPayload contract))) -- a natural as eight little-endian bytes (it is below 2^64) def checkpointEnvelopeWord = (lambda unrestricted value : Nat . (bytes-builder-build (app (nat-eliminate (lambda unrestricted current : Nat . (pi unrestricted rest : Nat . BytesBuilder)) (lambda unrestricted rest : Nat . (bytes-builder-chunk b"")) (lambda unrestricted p : Nat . (lambda unrestricted induction : (pi unrestricted rest : Nat . BytesBuilder) . (lambda unrestricted rest : Nat . (bytes-builder-append (bytes-builder-chunk (bytes (nat-to-byte (naturalModuloUnchecked rest 256)))) (induction (naturalDivideUnchecked rest 256)))))) 8) value))) -- the fixed bytes: magic, version, header, payload, chunk, full, tail def checkpointEnvelopeFixed = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (bytes-append checkpointEnvelopeMagic (bytes-append (checkpointEnvelopeWord checkpointEnvelopeVersion) (bytes-append (checkpointEnvelopeWord (checkpointEnvelopeHeaderBytes contract)) (bytes-append (checkpointEnvelopeWord (checkpointEnvelopeContractPayload contract)) (bytes-append (checkpointEnvelopeWord (checkpointEnvelopeContractChunk contract)) (bytes-append (checkpointEnvelopeWord (checkpointEnvelopeFullChunks contract)) (checkpointEnvelopeWord (checkpointEnvelopeTailBytes contract))))))))) -- a choice between verdicts on a 0/1 natural def checkpointEnvelopeIf = (lambda unrestricted condition : Nat . (lambda unrestricted whenTrue : (family CheckpointEnvelopeVerdict) . (lambda unrestricted whenFalse : (family CheckpointEnvelopeVerdict) . (nat-eliminate (lambda unrestricted current : Nat . (family CheckpointEnvelopeVerdict)) whenFalse (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family CheckpointEnvelopeVerdict) . whenTrue)) condition)))) def checkpointEnvelopeZeros = (lambda unrestricted count : Nat . (bytes-builder-build (nat-eliminate (lambda unrestricted current : Nat . BytesBuilder) (bytes-builder-chunk b"") (lambda unrestricted p : Nat . (lambda unrestricted induction : BytesBuilder . (bytes-builder-append (bytes-builder-chunk (bytes 0)) induction))) count))) -- ---- the model's reading of a checkpoint file ---- def checkpointEnvelopeSlice = (lambda unrestricted file : Bytes . (lambda unrestricted offset : Nat . (lambda unrestricted length : Nat . (dataBytesTakeValidated length (dataBytesDropValidated offset file))))) -- the index of the first chunk whose digest differs, or the chunk count def checkpointEnvelopeFirstBadChunk = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (lambda unrestricted file : Bytes . (let unrestricted count = (checkpointEnvelopeChunkCount contract) in (let unrestricted header = (checkpointEnvelopeHeaderBytes contract) in (let unrestricted chunk = (checkpointEnvelopeContractChunk contract) in (app (nat-eliminate (lambda unrestricted remaining : Nat . (pi unrestricted index : Nat . Nat)) (lambda unrestricted index : Nat . index) (lambda unrestricted p : Nat . (lambda unrestricted induction : (pi unrestricted index : Nat . Nat) . (lambda unrestricted index : Nat . (nat-eliminate (lambda unrestricted same : Nat . Nat) index (lambda unrestricted q : Nat . (lambda unrestricted ignored : Nat . (induction (succ index)))) (bytes-equal (checkpointEnvelopeSlice file (naturalAdd checkpointEnvelopeDigestsAt (naturalMultiply checkpointEnvelopeDigestBytes index)) checkpointEnvelopeDigestBytes) (sha256RawDigestOrEmpty (checkpointEnvelopeSlice file (naturalAdd header (naturalMultiply chunk index)) (checkpointEnvelopeChunkExtent contract index)))))))) count) 0)))))) def checkpointEnvelopeVerdictOf = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (lambda unrestricted file : Bytes . (let unrestricted digests = (checkpointEnvelopeIdentityDigests contract) in (let unrestricted identity = (lambda unrestricted index : Nat . (lambda unrestricted at : Nat . (bytes-equal (checkpointEnvelopeSlice file at checkpointEnvelopeDigestBytes) (checkpointEnvelopeSlice digests (naturalMultiply index checkpointEnvelopeDigestBytes) checkpointEnvelopeDigestBytes)))) in (let unrestricted bad = (checkpointEnvelopeFirstBadChunk contract file) in (checkpointEnvelopeIf (naturalEqual (bytes-length file) (checkpointEnvelopeFileBytes contract)) (checkpointEnvelopeIf (bytes-equal (checkpointEnvelopeSlice file 0 checkpointEnvelopeFixedBytes) (checkpointEnvelopeFixed contract)) (checkpointEnvelopeIf (identity 0 checkpointEnvelopeSchemaAt) (checkpointEnvelopeIf (identity 1 checkpointEnvelopeLearnerAt) (checkpointEnvelopeIf (identity 2 checkpointEnvelopeDataAt) (checkpointEnvelopeIf (identity 3 checkpointEnvelopeSeedAt) (checkpointEnvelopeIf (naturalEqual bad (checkpointEnvelopeChunkCount contract)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeAccepted) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeChunkMismatch bad)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeForeignSeed)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeForeignData)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeForeignLearner)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeForeignSchema)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeLayout)) (constructor CheckpointEnvelopeVerdict CheckpointEnvelopeTruncated))))))) -- the checkpoint a contract's host writes: the header -- fixed bytes, -- `updates` and `invocations`, the identity digests, each chunk's digest -- -- and the payload def checkpointEnvelopeWrite = (lambda unrestricted contract : (family CheckpointEnvelopeContract) . (lambda unrestricted updates : Nat . (lambda unrestricted invocations : Nat . (lambda unrestricted payload : Bytes . (let unrestricted chunk = (checkpointEnvelopeContractChunk contract) in (let unrestricted digests = (bytes-builder-build (app (nat-eliminate (lambda unrestricted remaining : Nat . (pi unrestricted index : Nat . BytesBuilder)) (lambda unrestricted index : Nat . (bytes-builder-chunk b"")) (lambda unrestricted p : Nat . (lambda unrestricted induction : (pi unrestricted index : Nat . BytesBuilder) . (lambda unrestricted index : Nat . (bytes-builder-append (bytes-builder-chunk (sha256RawDigestOrEmpty (checkpointEnvelopeSlice payload (naturalMultiply chunk index) (checkpointEnvelopeChunkExtent contract index)))) (induction (succ index)))))) (checkpointEnvelopeChunkCount contract)) 0)) in (let unrestricted body = (bytes-append (checkpointEnvelopeFixed contract) (bytes-append (checkpointEnvelopeWord updates) (bytes-append (checkpointEnvelopeWord invocations) (bytes-append (checkpointEnvelopeIdentityDigests contract) digests)))) in (bytes-append body (bytes-append (checkpointEnvelopeZeros (naturalSaturatingSubtract (checkpointEnvelopeHeaderBytes contract) (bytes-length body))) payload)))))))))