device | str | "cuda" | The device to run the model on. Defaults to “cuda”. |
epochs | int | 30 | The total number of passes through the fine-tuning data. Defaults to 30. |
time_limit | int | None | None | Time limit in seconds for fine-tuning. If None, no time limit is applied. Defaults to None. |
learning_rate | float | 1e-05 | The learning rate for the AdamW optimizer. A small value is crucial for stable fine-tuning. Defaults to 1e-5. |
weight_decay | float | 0.01 | The weight decay for the AdamW optimizer. Defaults to 0.01. |
validation_split_ratio | float | None | 0.1 | Fraction of the original training data reserved as a validation set for early stopping and monitoring. Set to 0 or None to disable validation: all data is then used for fine-tuning, per-epoch evaluation is skipped, and early stopping is disabled. Ignored when explicit validation data is passed to fit. Defaults to 0.1. |
n_finetune_ctx_plus_query_samples | int | 50000 | The total number of samples per meta-dataset during fine-tuning (context plus query) before applying the finetune_ctx_query_split_ratio. Defaults to 50_000. |
finetune_ctx_query_split_ratio | float | 0.2 | The proportion of each fine-tuning meta-dataset to use as query samples for calculating the loss. The remainder is used as context. Defaults to 0.2. |
n_inference_subsample_samples | int | None | None | The total number of subsampled training samples per estimator during validation and final inference. If None, no subsampling is applied and the full training set is used as context. Defaults to None. |
random_state | int | 0 | Seed for reproducibility of data splitting and model initialization. Defaults to 0. |
early_stopping | bool | True | Whether to use early stopping based on validation performance. When enabled, the best-performing weights are restored at the end of training and a best checkpoint is saved alongside the interval checkpoints. When disabled, training runs all epochs and the last-epoch weights are kept (no best checkpoint is saved). Defaults to True. |
early_stopping_patience | int | 8 | Number of validation checks to wait for improvement before early stopping. Defaults to 8. |
validation_frequency | int | 1 | Number of epochs between validation checks. A value of 1 (default) validates after every epoch, preserving the existing behavior. The initial evaluation of the unfine-tuned model still runs whenever validation data is available. With a value greater than 1, early_stopping_patience counts validation checks rather than epochs. Must be a positive integer. |
min_delta | float | 0.0001 | Minimum change in metric to be considered as an improvement. Defaults to 1e-4. |
grad_clip_value | float | None | 1.0 | Maximum norm for gradient clipping. If None, gradient clipping is disabled. Gradient clipping helps stabilize training by preventing exploding gradients. Defaults to 1.0. |
use_lr_scheduler | bool | True | Whether to use a learning rate scheduler (linear warmup with optional cosine decay) during fine-tuning. Defaults to True. |
lr_warmup_only | bool | False | If True, only performs linear warmup to the base learning rate and then keeps it constant. If False, applies cosine decay after warmup. Defaults to False. |
n_estimators_finetune | int | 2 | If set, overrides n_estimators of the underlying estimator only during fine-tuning to control the number of estimators (ensemble size) used in the training loop. If None, the value from kwargs or the estimator default is used. Defaults to 2. |
n_estimators_validation | int | 2 | If set, overrides n_estimators only for validation-time evaluation during fine-tuning (early-stopping / monitoring). If None, the value from kwargs or the estimator default is used. Defaults to 2. |
n_estimators_final_inference | int | 2 | If set, overrides n_estimators only for the final fitted inference model that is used after fine-tuning. If None, the value from kwargs or the estimator default is used. Defaults to 2. |
use_activation_checkpointing | bool | True | Whether to use activation checkpointing to reduce memory usage. Defaults to True. |
shard_estimators_across_gpus | bool | False | When True under DDP, every rank processes the same data chunk but only a shard of the fine-tuning estimators. DDP then averages gradients across estimator shards. This reduces per-rank activation memory instead of only distributing data chunks. Defaults to False. |
save_checkpoint_interval | int | None | 10 | Number of epochs between checkpoint saves. This only has an effect if output_dir is provided during the fit() call. If None, no intermediate checkpoints are saved. The best model checkpoint is always saved regardless of this setting. Defaults to 10. |
use_fixed_preprocessing_seed | bool | True | Whether to use a fixed preprocessing seed. If True, the preprocessing will always use the same random seed throughout data batches. This is helpful in most cases because, e.g., the column order will stay the same across batches. If False, the preprocessing will use a different random seed for each batch. |
experiment_logger | FinetuningLogger | None | None | An optional logger implementing the FinetuningLogger protocol (e.g., WandbLogger) for experiment tracking. If None, a no-op NullLogger is used. Defaults to None. |
model_version | ModelVersion | None | None | Which TabPFN model version to fine-tune. If None (default), uses the package default version (settings.tabpfn.model_version) — the same version a default TabPFNClassifier/TabPFNRegressor loads — so fine-tuning tracks the current default model rather than a hardcoded one. Defaults to None. |