///|
pub const BATCH_SIZE : Int = 64
///|
pub const BLOCK_SIZE : Int = 256
///|
pub const TRAINING_STEPS : Int = 5000
///|
pub const LEARNING_RATE : Double = 1.0e-3
///|
pub const MIN_LR : Double = 1.0e-4
///|
pub const WARMUP_ITERS : Int = 100
///|
pub const EVAL_INTERVAL : Int = 250
///|
pub const EVAL_ITERS : Int = 200
///|
pub const LOG_INTERVAL : Int = 10
///|
pub const WEIGHT_DECAY : Double = 1.0e-1
///|
pub const BETA1 : Double = 0.9
///|
pub const BETA2 : Double = 0.99
///|
pub const GRAD_CLIP : Double = 1.0
///|
pub struct TrainingConfig {
priv batch_size : Int
priv block_size : Int
priv steps : Int
priv learning_rate : Double
priv min_lr : Double
priv warmup_iters : Int
priv eval_interval : Int
priv eval_iters : Int
priv log_interval : Int
priv weight_decay : Double
priv beta1 : Double
priv beta2 : Double
priv grad_clip : Double
priv always_save_checkpoint : Bool
}
///|
pub fn TrainingConfig::TrainingConfig(
batch_size~ : Int,
block_size~ : Int,
steps~ : Int,
learning_rate~ : Double,
min_lr? : Double = MIN_LR,
warmup_iters? : Int = WARMUP_ITERS,
eval_interval? : Int = EVAL_INTERVAL,
eval_iters? : Int = EVAL_ITERS,
log_interval? : Int = LOG_INTERVAL,
weight_decay? : Double = WEIGHT_DECAY,
beta1? : Double = BETA1,
beta2? : Double = BETA2,
grad_clip? : Double = GRAD_CLIP,
always_save_checkpoint? : Bool = false,
) -> TrainingConfig {
if batch_size <= 0 {
abort("batch_size must be positive")
}
if block_size <= 0 {
abort("block_size must be positive")
}
if steps < 0 {
abort("steps must not be negative")
}
if learning_rate <= 0.0 {
abort("learning_rate must be positive")
}
if min_lr <= 0.0 {
abort("min_lr must be positive")
}
if warmup_iters < 0 {
abort("warmup_iters must not be negative")
}
if eval_interval <= 0 {
abort("eval_interval must be positive")
}
if eval_iters <= 0 {
abort("eval_iters must be positive")
}
if log_interval <= 0 {
abort("log_interval must be positive")
}
if weight_decay < 0.0 {
abort("weight_decay must not be negative")
}
if beta1 < 0.0 || beta1 >= 1.0 || beta2 < 0.0 || beta2 >= 1.0 {
abort("AdamW beta values must be in [0, 1)")
}
if grad_clip < 0.0 {
abort("grad_clip must not be negative")
}
{
batch_size,
block_size,
steps,
learning_rate,
min_lr,
warmup_iters,
eval_interval,
eval_iters,
log_interval,
weight_decay,
beta1,
beta2,
grad_clip,
always_save_checkpoint,
}
}
///|
pub fn TrainingConfig::recommended() -> TrainingConfig {
TrainingConfig(
batch_size=BATCH_SIZE,
block_size=BLOCK_SIZE,
steps=TRAINING_STEPS,
learning_rate=LEARNING_RATE,
)
}
///|
pub fn TrainingConfig::batch_size(self : TrainingConfig) -> Int {
self.batch_size
}
///|
pub fn TrainingConfig::block_size(self : TrainingConfig) -> Int {
self.block_size
}
///|
pub fn TrainingConfig::steps(self : TrainingConfig) -> Int {
self.steps
}
///|
pub fn TrainingConfig::learning_rate(self : TrainingConfig) -> Double {
self.learning_rate
}
///|
pub fn TrainingConfig::min_lr(self : TrainingConfig) -> Double {
self.min_lr
}
///|
pub fn TrainingConfig::warmup_iters(self : TrainingConfig) -> Int {
self.warmup_iters
}
///|
pub fn TrainingConfig::eval_interval(self : TrainingConfig) -> Int {
self.eval_interval
}
///|
pub fn TrainingConfig::eval_iters(self : TrainingConfig) -> Int {
self.eval_iters
}
///|
pub fn TrainingConfig::log_interval(self : TrainingConfig) -> Int {
self.log_interval
}
///|
pub fn TrainingConfig::weight_decay(self : TrainingConfig) -> Double {
self.weight_decay
}
///|
pub fn TrainingConfig::beta1(self : TrainingConfig) -> Double {
self.beta1
}
///|
pub fn TrainingConfig::beta2(self : TrainingConfig) -> Double {
self.beta2
}
///|
pub fn TrainingConfig::grad_clip(self : TrainingConfig) -> Double {
self.grad_clip
}
///|
pub fn TrainingConfig::always_save_checkpoint(self : TrainingConfig) -> Bool {
self.always_save_checkpoint
}
///|
pub fn TrainingConfig::learning_rate_at(
self : TrainingConfig,
iter : Int,
) -> Double {
if iter < 0 {
abort("iter must not be negative")
}
if iter < self.warmup_iters {
return self.learning_rate *
(iter + 1).to_double() /
(self.warmup_iters + 1).to_double()
}
if iter > self.steps {
return self.min_lr
}
if self.steps <= self.warmup_iters {
return self.min_lr
}
let decay_ratio = (iter - self.warmup_iters).to_double() /
(self.steps - self.warmup_iters).to_double()
let coeff = 0.5 * (1.0 + @math.cos(@math.PI * decay_ratio))
self.min_lr + coeff * (self.learning_rate - self.min_lr)
}
///|
pub struct TrainingStats {
losses : Array[Double]
initial_eval_loss : Double
final_eval_loss : Double
best_val_loss : Double
saved_checkpoints : Int
}
///|
pub struct TrainingResult {
model : MiniGPT
tokenizer : @tokenizer.Tokenizer
stats : TrainingStats
}
///|
pub struct TrainingState {
priv model : MiniGPT
priv optimizer : @optim.AdamW
priv iter_num : Int
priv best_val_loss : Double
priv config : TrainingConfig
}
///|
fn TrainingState::TrainingState(
model : MiniGPT,
optimizer : @optim.AdamW,
iter_num : Int,
best_val_loss : Double,
config : TrainingConfig,
) -> TrainingState {
{ model, optimizer, iter_num, best_val_loss, config }
}
///|
pub fn TrainingState::model(self : TrainingState) -> MiniGPT {
self.model
}
///|
pub fn TrainingState::optimizer(self : TrainingState) -> @optim.AdamW {
self.optimizer
}
///|
pub fn TrainingState::iter_num(self : TrainingState) -> Int {
self.iter_num
}
///|
pub fn TrainingState::best_val_loss(self : TrainingState) -> Double {
self.best_val_loss
}
///|
pub fn TrainingState::config(self : TrainingState) -> TrainingConfig {
self.config
}
///|
pub struct EvalEvent {
iter_num : Int
train_loss : Double
val_loss : Double
}
///|
pub struct LogEvent {
iter_num : Int
loss : Double
learning_rate : Double
}
///|
pub struct TrainingCallbacks {
priv on_eval : (EvalEvent) -> Unit
priv on_log : (LogEvent) -> Unit
priv on_checkpoint : (TrainingState) -> Unit
}
///|
pub fn TrainingCallbacks::TrainingCallbacks(
on_eval? : (EvalEvent) -> Unit = fn(_event) { () },
on_log? : (LogEvent) -> Unit = fn(_event) { () },
on_checkpoint? : (TrainingState) -> Unit = fn(_state) { () },
) -> TrainingCallbacks {
{ on_eval, on_log, on_checkpoint }
}
///|
fn sample_batch(
token_ids : Array[Int],
batch_size : Int,
block_size : Int,
rng : @random.Rand,
) -> (@tensor.TokenIds, @tensor.TokenIds) {
if batch_size <= 0 {
abort("batch_size must be positive")
}
if block_size <= 0 {
abort("block_size must be positive")
}
if token_ids.length() <= block_size {
abort("token_ids must contain more items than block_size")
}
let inputs : Array[Int] = []
let targets : Array[Int] = []
let limit = token_ids.length() - block_size
for _ in 0.. Double {
let mut sum = 0.0
for _ in 0.. TrainingStats {
train_token_ids_internal(
model,
train_ids,
val_ids,
config,
rng,
TrainingCallbacks(),
)
}
///|
pub fn train_text(
text : String,
tokenizer_config : @tokenizer.TokenizerConfig,
config : TrainingConfig,
rng : @random.Rand,
) -> TrainingResult {
train_text_with_architecture(
text,
tokenizer_config,
config,
ArchitectureConfig(),
rng,
)
}
///|
pub fn train_text_with_architecture(
text : String,
tokenizer_config : @tokenizer.TokenizerConfig,
config : TrainingConfig,
architecture : ArchitectureConfig,
rng : @random.Rand,
) -> TrainingResult {
let dataset = @tokenizer.prepare_token_dataset(text, tokenizer_config)
let tokenizer = dataset.tokenizer()
let train_ids = dataset.train_ids()
let val_ids = dataset.val_ids()
let model = MiniGPT(
ModelConfig::from_architecture(
tokenizer.vocab_size(),
architecture,
block_size=config.block_size,
),
rng,
)
let stats = train_token_ids(model, train_ids, val_ids, config, rng)
{ model, tokenizer, stats }
}
///|
pub fn train_token_ids_with_callbacks(
model : MiniGPT,
train_ids : Array[Int],
val_ids : Array[Int],
config : TrainingConfig,
rng : @random.Rand,
callbacks : TrainingCallbacks,
) -> TrainingStats {
train_token_ids_internal(model, train_ids, val_ids, config, rng, callbacks)
}
///|
fn train_token_ids_internal(
model : MiniGPT,
train_ids : Array[Int],
val_ids : Array[Int],
config : TrainingConfig,
rng : @random.Rand,
callbacks : TrainingCallbacks,
) -> TrainingStats {
let losses : Array[Double] = []
let params = model.parameters()
let optimizer_config = @optim.AdamWConfig(
config.learning_rate,
beta1=config.beta1,
beta2=config.beta2,
eps=1.0e-8,
)
let optimizer = @optim.AdamW::with_parameter_weight_decays(
params,
optimizer_config,
model.parameter_weight_decays(config.weight_decay),
)
let mut iter_num = 0
let mut best_val_loss = 1.0e9
let mut saved_checkpoints = 0
let mut initial_eval_loss = @double.not_a_number
let mut final_eval_loss = @double.not_a_number
while true {
let lr = config.learning_rate_at(iter_num)
optimizer.set_learning_rate(lr)
if iter_num % config.eval_interval == 0 {
let train_loss = eval_token_ids(model, train_ids, config, rng)
let val_loss = eval_token_ids(model, val_ids, config, rng)
(callbacks.on_eval)({ iter_num, train_loss, val_loss })
if initial_eval_loss != initial_eval_loss {
initial_eval_loss = val_loss
}
final_eval_loss = val_loss
let improved = val_loss < best_val_loss
if improved {
best_val_loss = val_loss
}
if iter_num > 0 && (improved || config.always_save_checkpoint) {
(callbacks.on_checkpoint)(
TrainingState(model, optimizer, iter_num, best_val_loss, config),
)
saved_checkpoints += 1
}
}
let (inputs, targets) = sample_batch(
train_ids,
config.batch_size,
config.block_size,
rng,
)
let loss = model.loss_train(inputs, targets, rng)
let loss_value = loss.data()[0]
losses.push(loss_value)
model.zero_grad()
loss.backward()
model.clear_graph()
optimizer.step_with_grad_clip(config.grad_clip)
if iter_num % config.log_interval == 0 {
(callbacks.on_log)({ iter_num, loss: loss_value, learning_rate: lr })
}
iter_num += 1
if iter_num > config.steps {
break
}
}
{
losses,
initial_eval_loss,
final_eval_loss,
best_val_loss,
saved_checkpoints,
}
}