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: object

Training 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)

  • schedule (Callable[[Any], Any] | None)

  • 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)

  • loss_scale (float | str | 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: object

Caller-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.

Parameters:
  • train_dataloader (Iterable[Batch] | None)

  • validation_dataloader (Iterable[Batch] | None)

  • batch_sharding (PyTree | None)

  • prefetch_size (int)