Implementation:Microsoft Onnxruntime TrainingOptimizer
Appearance
| Knowledge Sources | |
|---|---|
| Domains | Training, API, Optimization |
| Last Updated | 2026-02-10 04:00 GMT |
Overview
Defines the Optimizer class, optimizer state structures, and algorithm variants (AdamW, SGDv2) for performing gradient-based parameter updates in the ORT Training API.
Description
This header provides the complete optimizer infrastructure for ORT Training. Key types include:
- `ParameterOptimizerState`: A map of momentum state names to OrtValues for a single parameter (e.g., first/second order moments for Adam).
- `GroupOptimizerState`: Aggregates step count, initial learning rate, adaptive learning rate, and per-parameter optimizer states for one parameter group.
- `OptimizerCheckpointState`: Contains all group optimizer states and a data transfer manager reference for checkpoint serialization.
- `OptimizerAlgorithmBase`: Base class defining momentum keys and optimizer state input names. Concrete implementations include `AdamWOptimizerAlgorithm` (with momentum0/momentum1) and `SGDOptimizerV2Algorithm` (with momentum0 only).
- `OptimizerAlorithmFactory`: Factory that inspects the optimizer ONNX graph to instantiate the correct algorithm.
- `Optimizer`: The main optimizer struct that wraps an `InferenceSession` loaded with the optimizer ONNX model. It executes gradient update steps, manages learning rate get/set, and handles optimizer state construction from checkpoint or zero-initialization. The optimizer does not own parameters but constructs TensorSequence inputs from the checkpoint state.
Usage
Use this header when building a training loop that requires gradient-based parameter updates. The Optimizer is typically created by `TrainingSession` alongside a `Module`.
Code Reference
Source Location
- Repository: Microsoft_Onnxruntime
- File: orttraining/orttraining/training_api/optimizer.h
- Lines: 1-174
Signature
typedef InlinedHashMap<std::string, OrtValue> ParameterOptimizerState;
struct GroupOptimizerState {
int64_t step = 0;
float initial_lr = 0.001f;
float learning_rate{initial_lr};
InlinedHashMap<std::string, ParameterOptimizerState> param_named_optimizer_states;
};
struct OptimizerCheckpointState {
InlinedHashMap<std::string, std::shared_ptr<GroupOptimizerState>> group_named_optimizer_states;
const DataTransferManager* optimizer_session_data_transfer_mgr;
};
struct Optimizer {
Optimizer(const ModelIdentifiers& model_identifiers,
CheckpointState* state,
const onnxruntime::SessionOptions& session_options,
const Environment& env,
const std::vector<std::shared_ptr<IExecutionProvider>>& providers,
gsl::span<OrtCustomOpDomain* const> op_domains = {});
Status Step();
Status SetLearningRate(float lr);
float GetLearningRate() const noexcept;
Status SetInitialLearningRate(float initial_lr);
Status ConstructOptimizerStateAndInputs();
};
Import
#include "orttraining/training_api/optimizer.h"
I/O Contract
| Method | Inputs | Outputs | Description |
|---|---|---|---|
| Optimizer (ctor) | ModelIdentifiers, CheckpointState*, SessionOptions, Environment, providers | Optimizer instance | Initializes optimizer session and loads/creates optimizer states |
| Step | (none) | Status | Executes one optimizer step (gradient update) on all parameters |
| SetLearningRate | float lr | Status | Sets the current adaptive learning rate |
| GetLearningRate | (none) | float | Returns the current adaptive learning rate |
| SetInitialLearningRate | float initial_lr | Status | Sets both initial and current learning rate |
| ConstructOptimizerStateAndInputs | (none) | Status | Constructs optimizer state tensors and model inputs (for deferred initialization) |
Usage Examples
#include "orttraining/training_api/optimizer.h"
using namespace onnxruntime::training::api;
// Create optimizer
Optimizer optimizer(model_ids, &checkpoint_state, session_options, env, providers);
// Set learning rate
optimizer.SetLearningRate(0.001f);
// Optimizer step after computing gradients
ORT_THROW_IF_ERROR(optimizer.Step());
// Query learning rate
float lr = optimizer.GetLearningRate();
Related Pages
Page Connections
Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment