Skip to main content
TensorFlow is deprecated. New TensorFlow model uploads are no longer accepted — use PyTorch for new models. The TensorFlow-specific optimizers, learning-rate schedules, loss functions, layer freezing, and augmentation flags documented below apply only to existing TensorFlow experiments, which remain readable.

Training Parameters

All parameters are set through the training_plan after linking your model with the dataset.
To see all current parameter settings, run training_plan.get_training_plan(). To run consecutive experiments, overwrite parameters and re-start training with training_plan.start().
You can refer to the TensorFlow Documentation for more information on TensorFlow augmentation parameters and the PyTorch Documentation for more information on PyTorch augmentation parameters.
Basic training configuration parameters that control the fundamental aspects of your training process.
The default validation split is often very small: 3 classes and 100 images give 0.03. Set it explicitly, for example training_plan.validation_split(0.15), and check it with training_plan.get_training_plan().
To try another batch size, you can copy a model file and change batch_size in the copy. For text models (causal language modeling included), the copy also needs its tokenizer. The SDK looks for it in this order: the tokenizer= argument to user.upload_model(...), then a <model-file-name>_tokenizer.json next to the model file, then a plain tokenizer.json next to it. Without one, the upload fails. So when you copy or rename a model-zoo file, copy and rename its _tokenizer.json too: distilgpt2.py and distilgpt2_tokenizer.json become distilgpt2_bs4.py and distilgpt2_bs4_tokenizer.json. The upload error also mentions tokenizer_id, but don’t declare one: the upload rejects hub references (see Model optimization).

Core Hyperparameters

1. Optimizer

Controls how the model’s parameters are updated during training. Supports different optimizers for TensorFlow and PyTorch. The default optimizer is SGD, except for text tasks: text classification, sentence pair classification, token classification, masked language modeling, causal language modeling, seq2seq and embeddings default to AdamW (SDK 1.2.38 and later). Supported Optimizers:
  • TensorFlow: adam, rmsprop, sgd, adadelta, adagrad, adamax, nadam, ftrl
  • PyTorch: adam, adamw, rmsprop, sgd, adadelta, adagrad, adamax

2. Learning Rate

Controls the rate at which the model learns. Supports three different types: Default: {'type': 'constant', 'value': 0.001}. Text tasks default to {'type': 'constant', 'value': 5e-05} (SDK 1.2.38 and later). The optimizer and learning rate defaults are a pair. If you set only one of them, the other keeps its default: on a text task, training_plan.learning_rate({'type': 'constant', 'value': 0.001}) alone trains AdamW at 0.001, which can be too high for a pretrained transformer, and training_plan.optimizer('sgd') alone trains SGD at 5e-05, which learns very slowly. Set both when you change either.
  • Custom for TensorFlow: Define a custom learning rate function, then pass it via learning_rate() with type: 'custom':

3. Loss Function

Defines how the model measures prediction errors. Pass a standard loss by name, or your own custom loss. The allowed standard losses and the default depend on the task: For object detection, the experiment’s settings show mse, but the platform does not use it: torchvision detectors compute their own losses, and yolo models always use the loss.py shipped with the model.
On PyTorch you can also pass a custom loss: a top-level function that takes (outputs, targets), or the path to a loss.py that defines Custom_loss. The SDK runs a short local training check with it before you start. scikit-learn models accept standard losses only.
An nn.Module class works too, but only when you import it from a .py file: the SDK reads the loss’s source code, and it cannot read the source of a class defined in a notebook cell. Alternatively, put the class in a loss.py as Custom_loss and pass that file’s path. TensorFlow experiments (deprecated) used binary_crossentropy, categorical_crossentropy and mse.

Training Control

Layer Freezing

Specify which layers should remain unchanged during training (TensorFlow only):
Layer freezing is currently supported for TensorFlow models only. PyTorch models will receive a “not supported” message.

Callbacks

Control training behavior with various callbacks: The metric must be one of accuracy, loss, val_accuracy or val_loss. Object detection does not measure training accuracy, so monitor a loss for object detection models.

Federated aggregation

At the end of each cycle, the platform combines the models trained in each secure environment. The default strategy is fedavg. Choose another one, with its optional hyperparameters, through aggregation_strategy:
Hyperparameters you leave out take the strategy’s defaults. An unknown strategy or hyperparameter name is refused.

Preprocessing (Tabular & Time Series)

Configure the platform’s built-in preprocessing per experiment. Available for tabular classification, tabular regression, time-to-event, and time-series classification; time-series forecasting supports the imputation pair (handle_missing_values / imputation_strategy). The calls below are illustrative alternatives, not a recommended combination — for example, handle_missing_values(False) and imputation_strategy('iterative') contradict each other (the first disables the step the second configures).
When to disable imputation: if your model embeds its own imputer (an sklearn Pipeline) or handles NaNs natively (XGBoost, LightGBM, HistGradientBoosting, CatBoost, EBM), set handle_missing_values(False) — otherwise the built-in step fills every gap first and your model never sees a raw NaN.
knn imputation and the QuantileTransformer scaler are deliberately unavailable: their fitted state memorizes raw training data, which would leave the secure environment inside the preprocessing artifact. iterative (MICE) persists only learned coefficients — the same privacy class as model weights.
Defaults reproduce the historical platform behaviour (median imputation on, label encoding, z-score feature scaling on), so an experiment that never touches these knobs is unchanged.

Data Augmentation for Image Data

Enhance your dataset with real-time image transformations. All parameters support both TensorFlow and PyTorch unless noted otherwise.

Geometric Transformations

On PyTorch, the platform accepts these ranges: rotation 0 to 180 degrees, width and height shift 0 to 1 (fraction of the image), shear 0 to 45 degrees, zoom 0 to 1, channel shift 0 to 255 and brightness 0 to 1.

Color and Intensity Transformations

*PyTorch: Only supported for RGB images

Normalization

On PyTorch, horizontal_flip, vertical_flip, samplewise_center and samplewise_std_normalization are refused: the call records an error and training_plan.start() refuses to run while it stands.

Other Parameters

LLM Parameters (LoRA for Text Tasks)

LoRA (Low-Rank Adaptation) is available for every text task in PyTorch: text classification, sentence pair classification, token classification, masked language modeling, causal language modeling, seq2seq and embeddings. Other tasks print Operation not allowed and keep standard training. LoRA is off by default. When you enable it, the platform adds low-rank adapters to every nn.Linear (and nn.MultiheadAttention) layer of your model and trains only the adapters instead of every weight. Use it to fine-tune a larger pretrained model, for example a GPT-2 style model in causal language modeling, with less compute.
set_lora_parameters only works after enable_lora(True). It validates the values, then runs a short local training check of your model with the adapters before you start the experiment. If you only call enable_lora(True), the defaults below apply. From SDK 1.2.38 the defaults are rank 8 and alpha 16; earlier versions used 256 and 512. Note: LLM parameters are supported only for PyTorch.

Dataset Parameters (Optional)

Customize your dataset configuration and preprocessing options:

Dataset Customization


Next Steps


Need Help?

For more info about available functions and methods, call the help function in your notebook: