Configuration¶
Configuration dataclasses for training and datasets.
- class taktiny.trainer.TrainingConfig(max_steps=None, learning_rate=0.001, schedule=None, optimizer=None, weight_decay=0.0, log_interval=10, seed=42, jit_compile=True, output_dir=None, save_steps=None, save_total_limit=None, save_at_end=False, save_optimizer_state=True, save_async=False, eval_strategy='no', eval_steps=None, metric_for_best_model='eval_loss', greater_is_better=None, load_best_model_at_end=False, gradient_accumulation_steps=1, max_grad_norm=None, compute_grad_norm=True, skip_non_finite=True, ema_decay=None, loss_scale=None, initial_loss_scale=32768.0, loss_scale_growth_interval=2000)[source]¶
Bases:
objectTraining controls and native Orbax checkpoint policy.
save_async uses Orbax’s asynchronous checkpointer; train() waits for pending writes before returning. save_optimizer_state=False creates a weights/state snapshot that cannot provide exact training resume. Retention keeps the best and latest checkpoints, which can exceed save_total_limit=1 by one directory.
- Parameters:
max_steps (int | None)
learning_rate (float)
optimizer (Any)
weight_decay (float)
log_interval (int)
seed (int)
jit_compile (bool)
output_dir (str | PathLike | None)
save_steps (int | None)
save_total_limit (int | None)
save_at_end (bool)
save_optimizer_state (bool)
save_async (bool)
eval_strategy (str)
eval_steps (int | None)
metric_for_best_model (str)
greater_is_better (bool | None)
load_best_model_at_end (bool)
gradient_accumulation_steps (int)
max_grad_norm (float | None)
compute_grad_norm (bool)
skip_non_finite (bool)
ema_decay (float | None)
initial_loss_scale (float)
loss_scale_growth_interval (int)
- class taktiny.trainer.DatasetConfig(train_dataloader=None, validation_dataloader=None, batch_sharding=None, prefetch_size=2)[source]¶
Bases:
objectCaller-provided iterables of batches and optional device placement.
Loading, preprocessing, batching, and shuffling belong to the caller. Trainer never downloads datasets or changes a loader’s sampling policy.