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 thetraining_plan after linking your model with the dataset.
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.
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()withtype: '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.
(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.
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 isfedavg. 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).
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.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
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 printOperation 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
- Submit your best model: Evaluate Model
Need Help?
For more info about available functions and methods, call the help function in your notebook:- Email us at [email protected]