Implementation:NVIDIA TransformerEngine Multi Tensor C API
Appearance
| Field | Value |
|---|---|
| Sources | TransformerEngine |
| Domains | Deep_Learning, Optimization |
| Last Updated | 2026-02-07 14:00 GMT |
Overview
Declares the C API for multi-tensor fused CUDA kernels that operate on lists of tensors simultaneously, providing L2 norm computation and multiple Adam optimizer variants.
Description
multi_tensor.h provides fused operations across all model parameters in single kernel launches:
- L2 norm:
nvte_multi_tensor_l2norm_cudaandnvte_multi_tensor_unscale_l2norm_cudafor computing L2 norms across tensor lists with optional per-tensor breakdown and FP8 unscaling. - Adam optimizer family:
nvte_multi_tensor_adam_cuda: Standard Adam/AdamWnvte_multi_tensor_adam_param_remainder_cuda: Remainder-bit master parametersnvte_multi_tensor_adam_fp8_cuda: FP8 model parametersnvte_multi_tensor_adam_capturable_cuda: CUDA graph compatible with device-side LR/stepnvte_multi_tensor_adam_capturable_master_cuda: Graph-compatible with FP32 master weights
All functions use chunk-based processing and support a noop flag for conditional execution.
Usage
Use for fused optimizer updates and gradient norm computation across all model parameters to reduce kernel launch overhead.
Code Reference
Source Location
- Repository
NVIDIA/TransformerEngine- File
transformer_engine/common/include/transformer_engine/multi_tensor.h- Lines
- 1--303
Signature
void nvte_multi_tensor_l2norm_cuda(
int chunk_size, NVTETensor noop_flag, NVTETensor **tensor_lists,
const size_t num_tensor_lists, const size_t num_tensors_per_list,
NVTETensor output, NVTETensor output_per_tensor,
NVTETensor ret, NVTETensor ret_per_tensor,
int per_tensor, int max_chunks_per_tensor, cudaStream_t stream);
void nvte_multi_tensor_adam_cuda(
int chunk_size, NVTETensor noop_flag, NVTETensor **tensor_lists,
const size_t num_tensor_lists, const size_t num_tensors_per_list,
const float lr, const float beta1, const float beta2,
const float epsilon, const int step, const int mode,
const int bias_correction, const float weight_decay,
cudaStream_t stream);
Import
#include "transformer_engine/multi_tensor.h"
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
tensor_lists |
NVTETensor** |
Yes | 2D array of input tensors |
chunk_size |
int |
Yes | Number of elements processed per CUDA block |
noop_flag |
NVTETensor |
Yes | Single-element tensor; non-zero skips the kernel |
lr |
float |
Yes (Adam) | Learning rate |
Outputs
| Name | Type | Description |
|---|---|---|
ret |
NVTETensor |
L2 norm result (for norm functions) |
| updated tensors | in-place | Updated parameters, moments (for Adam functions) |
Usage Examples
#include "transformer_engine/multi_tensor.h"
// Compute L2 norm across all parameter tensors
nvte_multi_tensor_l2norm_cuda(2048, noop_flag, tensor_lists,
num_lists, num_tensors,
output, output_per_tensor,
ret, ret_per_tensor,
/*per_tensor=*/1, max_chunks, stream);
// Fused Adam update across all parameters
nvte_multi_tensor_adam_cuda(2048, noop_flag, tensor_lists,
num_lists, num_tensors,
lr, beta1, beta2, epsilon, step,
/*mode=*/1, /*bias_correction=*/1,
weight_decay, stream);
Related Pages
Page Connections
Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment