///|
const CHECKPOINT_VERSION : Int = 10
///|
const CHECKPOINT_MAGIC : Int = 1296649799
///|
priv struct MiniGPTCheckpoint {
model_kind : String
tokenizer_kind : String
vocabulary : Array[String]
bpe_merges : Array[@tokenizer.BpeMerge]
n_embd : Int
n_head : Int
n_layer : Int
block_size : Int
parameters : MiniGPTParameterData
training_config : TrainingConfig?
optimizer : @optim.AdamWCheckpoint?
iter_num : Int
best_val_loss : Double
}
///|
pub struct TrainingCheckpoint {
priv model : MiniGPT
priv tokenizer : @tokenizer.Tokenizer
priv optimizer : @optim.AdamW
priv iter_num : Int
priv best_val_loss : Double
priv config : TrainingConfig
}
///|
pub fn TrainingCheckpoint::TrainingCheckpoint(
model : MiniGPT,
tokenizer : @tokenizer.Tokenizer,
optimizer : @optim.AdamW,
iter_num : Int,
best_val_loss : Double,
config : TrainingConfig,
) -> TrainingCheckpoint {
if iter_num < 0 {
abort("iter_num must not be negative")
}
{ model, tokenizer, optimizer, iter_num, best_val_loss, config }
}
///|
pub fn TrainingCheckpoint::from_state(
state : TrainingState,
tokenizer : @tokenizer.Tokenizer,
) -> TrainingCheckpoint {
TrainingCheckpoint(
state.model(),
tokenizer,
state.optimizer(),
state.iter_num(),
state.best_val_loss(),
state.config(),
)
}
///|
pub fn TrainingCheckpoint::iter_num(self : TrainingCheckpoint) -> Int {
self.iter_num
}
///|
pub fn TrainingCheckpoint::best_val_loss(self : TrainingCheckpoint) -> Double {
self.best_val_loss
}
///|
fn MiniGPTCheckpoint::from_model(
model : MiniGPT,
tokenizer : @tokenizer.Tokenizer,
) -> MiniGPTCheckpoint {
{
model_kind: model.kind_name(),
tokenizer_kind: tokenizer.kind_name(),
vocabulary: tokenizer.vocabulary(),
bpe_merges: tokenizer.bpe_merges(),
n_embd: model.n_embd(),
n_head: model.n_head(),
n_layer: model.n_layer(),
block_size: model.block_size(),
parameters: model.parameter_data(),
training_config: None,
optimizer: None,
iter_num: 0,
best_val_loss: 1.0e9,
}
}
///|
fn MiniGPTCheckpoint::to_model_and_tokenizer(
self : MiniGPTCheckpoint,
) -> (MiniGPT, @tokenizer.Tokenizer) {
if self.model_kind != "gpt-transformer" {
abort("unsupported MiniGPT checkpoint model kind")
}
let tokenizer = tokenizer_from_checkpoint_data(
self.tokenizer_kind,
self.vocabulary,
self.bpe_merges,
)
let model = MiniGPT::from_parameter_data(
tokenizer.vocab_size(),
self.n_embd,
self.n_head,
self.n_layer,
self.block_size,
self.parameters,
)
(model, tokenizer)
}
///|
fn tokenizer_from_checkpoint_data(
kind : String,
vocabulary : Array[String],
merges : Array[@tokenizer.BpeMerge],
) -> @tokenizer.Tokenizer {
match kind {
"char" => {
let chars : Array[Char] = []
for token in vocabulary {
chars.push(single_char_checkpoint_token(token))
}
@tokenizer.Tokenizer::from_char(
@tokenizer.CharTokenizer::from_chars(chars),
)
}
"word" =>
@tokenizer.Tokenizer::from_word(
@tokenizer.WordTokenizer::from_tokens(vocabulary),
)
"bpe" =>
@tokenizer.Tokenizer::from_bpe(
@tokenizer.BpeTokenizer::from_vocabulary_and_merges(vocabulary, merges),
)
_ => abort("unsupported tokenizer kind: \{kind}")
}
}
///|
fn single_char_checkpoint_token(text : String) -> Char {
let chars : Array[Char] = []
for ch in text {
chars.push(ch)
}
if chars.length() != 1 {
abort("character tokenizer vocabulary item must contain one character")
}
chars[0]
}
///|
pub fn encode_checkpoint(
model : MiniGPT,
tokenizer : @tokenizer.Tokenizer,
) -> Bytes {
let checkpoint = MiniGPTCheckpoint::from_model(model, tokenizer)
encode_checkpoint_data(checkpoint)
}
///|
pub fn encode_training_checkpoint(checkpoint : TrainingCheckpoint) -> Bytes {
let checkpoint_data = MiniGPTCheckpoint::from_model(
checkpoint.model,
checkpoint.tokenizer,
)
encode_checkpoint_data({
..checkpoint_data,
training_config: Some(checkpoint.config),
optimizer: Some(checkpoint.optimizer.checkpoint()),
iter_num: checkpoint.iter_num,
best_val_loss: checkpoint.best_val_loss,
})
}
///|
fn encode_checkpoint_data(checkpoint : MiniGPTCheckpoint) -> Bytes {
let params = checkpoint.parameters
let buffer = Buffer::Buffer()
buffer.write_int_le(CHECKPOINT_MAGIC)
buffer.write_int_le(CHECKPOINT_VERSION)
write_string(buffer, checkpoint.model_kind)
write_string(buffer, checkpoint.tokenizer_kind)
buffer.write_int_le(checkpoint.vocabulary.length())
for token in checkpoint.vocabulary {
write_string(buffer, token)
}
buffer.write_int_le(checkpoint.bpe_merges.length())
for merge in checkpoint.bpe_merges {
write_string(buffer, merge.left())
write_string(buffer, merge.right())
write_string(buffer, merge.merged())
}
buffer.write_int_le(checkpoint.n_embd)
buffer.write_int_le(checkpoint.n_head)
buffer.write_int_le(checkpoint.n_layer)
buffer.write_int_le(checkpoint.block_size)
write_doubles(buffer, params.token_embedding_table)
write_doubles(buffer, params.position_embedding_table)
write_double_arrays(buffer, params.ln1_weight)
write_double_arrays(buffer, params.ln1_bias)
write_double_arrays(buffer, params.wq)
write_double_arrays(buffer, params.bq)
write_double_arrays(buffer, params.wk)
write_double_arrays(buffer, params.bk)
write_double_arrays(buffer, params.wv)
write_double_arrays(buffer, params.bv)
write_double_arrays(buffer, params.wo)
write_double_arrays(buffer, params.bo)
write_double_arrays(buffer, params.ln2_weight)
write_double_arrays(buffer, params.ln2_bias)
write_double_arrays(buffer, params.w_fc)
write_double_arrays(buffer, params.b_fc)
write_double_arrays(buffer, params.w_proj)
write_double_arrays(buffer, params.b_proj)
write_doubles(buffer, params.ln_f_weight)
write_doubles(buffer, params.ln_f_bias)
buffer.write_int_le(checkpoint.iter_num)
buffer.write_double_le(checkpoint.best_val_loss)
match (checkpoint.training_config, checkpoint.optimizer) {
(Some(config), Some(optimizer)) => {
buffer.write_int_le(1)
write_training_config(buffer, config)
write_adamw_checkpoint(buffer, optimizer)
}
(None, None) => buffer.write_int_le(0)
_ => abort("training checkpoint requires both config and optimizer")
}
buffer.to_bytes()
}
///|
pub fn decode_checkpoint(bytes : Bytes) -> (MiniGPT, @tokenizer.Tokenizer) {
let reader = BytesReader(bytes)
let magic = reader.read_int_le()
if magic != CHECKPOINT_MAGIC {
abort("invalid MiniGPT checkpoint magic")
}
let version = reader.read_int_le()
if version != CHECKPOINT_VERSION {
abort("unsupported MiniGPT checkpoint version")
}
let model_kind = reader.read_string()
let tokenizer_kind = reader.read_string()
let vocab_size = reader.read_int_le()
let vocabulary : Array[String] = []
for _ in 0.. Unit {
buffer.write_int_le(config.batch_size)
buffer.write_int_le(config.block_size)
buffer.write_int_le(config.steps)
buffer.write_double_le(config.learning_rate)
buffer.write_double_le(config.min_lr)
buffer.write_int_le(config.warmup_iters)
buffer.write_int_le(config.eval_interval)
buffer.write_int_le(config.eval_iters)
buffer.write_int_le(config.log_interval)
buffer.write_double_le(config.weight_decay)
buffer.write_double_le(config.beta1)
buffer.write_double_le(config.beta2)
buffer.write_double_le(config.grad_clip)
buffer.write_int_le(if config.always_save_checkpoint { 1 } else { 0 })
}
///|
fn write_adamw_checkpoint(
buffer : Buffer,
optimizer : @optim.AdamWCheckpoint,
) -> Unit {
write_double_arrays(buffer, optimizer.m())
write_double_arrays(buffer, optimizer.v())
buffer.write_double_le(optimizer.lr())
buffer.write_double_le(optimizer.beta1())
buffer.write_double_le(optimizer.beta2())
buffer.write_double_le(optimizer.eps())
write_doubles(buffer, optimizer.weight_decays())
buffer.write_int_le(optimizer.step())
buffer.write_double_le(optimizer.beta1_power())
buffer.write_double_le(optimizer.beta2_power())
}
///|
fn write_string(buffer : Buffer, text : String) -> Unit {
let bytes = @utf8.encode(text[:])
write_bytes(buffer, bytes)
}
///|
fn write_bytes(buffer : Buffer, bytes : Bytes) -> Unit {
buffer.write_int_le(bytes.length())
buffer.write_bytes(bytes[:])
}
///|
fn write_doubles(buffer : Buffer, values : Array[Double]) -> Unit {
buffer.write_int_le(values.length())
for value in values {
buffer.write_double_le(value)
}
}
///|
fn write_double_arrays(buffer : Buffer, values : Array[Array[Double]]) -> Unit {
buffer.write_int_le(values.length())
for value in values {
write_doubles(buffer, value)
}
}
///|
priv struct BytesReader {
bytes : Bytes
mut offset : Int
}
///|
fn BytesReader::BytesReader(bytes : Bytes) -> BytesReader {
{ bytes, offset: 0 }
}
///|
fn BytesReader::is_done(self : BytesReader) -> Bool {
self.offset == self.bytes.length()
}
///|
fn BytesReader::read_byte(self : BytesReader) -> Int {
if self.offset >= self.bytes.length() {
abort("checkpoint ended unexpectedly")
}
let value = self.bytes[self.offset].to_int()
self.offset += 1
value
}
///|
fn BytesReader::read_int_le(self : BytesReader) -> Int {
let b0 = self.read_byte()
let b1 = self.read_byte()
let b2 = self.read_byte()
let b3 = self.read_byte()
b0 | (b1 << 8) | (b2 << 16) | (b3 << 24)
}
///|
fn BytesReader::read_string(self : BytesReader) -> String {
let bytes = self.read_bytes()
@utf8.decode(bytes[:]) catch {
_ => abort("invalid checkpoint UTF-8 string")
}
}
///|
fn BytesReader::read_bytes(self : BytesReader) -> Bytes {
let length = self.read_int_le()
if length < 0 || self.offset + length > self.bytes.length() {
abort("invalid checkpoint bytes length")
}
let view = self.bytes.view(start=self.offset, end=self.offset + length)
self.offset += length
view.to_owned()
}
///|
fn BytesReader::read_double_le(self : BytesReader) -> Double {
let b0 = self.read_byte().to_int64()
let b1 = self.read_byte().to_int64()
let b2 = self.read_byte().to_int64()
let b3 = self.read_byte().to_int64()
let b4 = self.read_byte().to_int64()
let b5 = self.read_byte().to_int64()
let b6 = self.read_byte().to_int64()
let b7 = self.read_byte().to_int64()
let bits = b0 |
(b1 << 8) |
(b2 << 16) |
(b3 << 24) |
(b4 << 32) |
(b5 << 40) |
(b6 << 48) |
(b7 << 56)
bits.reinterpret_as_double()
}
///|
fn BytesReader::read_doubles(self : BytesReader) -> Array[Double] {
let length = self.read_int_le()
if length < 0 {
abort("invalid checkpoint array length")
}
let values : Array[Double] = []
for _ in 0.. Array[Array[Double]] {
let length = self.read_int_le()
if length < 0 {
abort("invalid checkpoint nested array length")
}
let values : Array[Array[Double]] = []
for _ in 0.. Array[@tokenizer.BpeMerge] {
let length = self.read_int_le()
if length < 0 {
abort("invalid checkpoint BPE merge count")
}
let merges : Array[@tokenizer.BpeMerge] = []
for _ in 0.. TrainingConfig {
TrainingConfig(
batch_size=self.read_int_le(),
block_size=self.read_int_le(),
steps=self.read_int_le(),
learning_rate=self.read_double_le(),
min_lr=self.read_double_le(),
warmup_iters=self.read_int_le(),
eval_interval=self.read_int_le(),
eval_iters=self.read_int_le(),
log_interval=self.read_int_le(),
weight_decay=self.read_double_le(),
beta1=self.read_double_le(),
beta2=self.read_double_le(),
grad_clip=self.read_double_le(),
always_save_checkpoint=self.read_bool(),
)
}
///|
fn BytesReader::read_bool(self : BytesReader) -> Bool {
match self.read_int_le() {
0 => false
1 => true
_ => abort("invalid checkpoint boolean")
}
}
///|
fn BytesReader::skip_adamw_checkpoint(self : BytesReader) -> Unit {
ignore(self.read_double_arrays())
ignore(self.read_double_arrays())
ignore(self.read_double_le())
ignore(self.read_double_le())
ignore(self.read_double_le())
ignore(self.read_double_le())
ignore(self.read_doubles())
ignore(self.read_int_le())
ignore(self.read_double_le())
ignore(self.read_double_le())
}