HSTU LayerNorm-Multiply-SiLU-Dropout (LMSD)
HSTU LayerNorm-Multiply-SiLU-Dropout (LMSD)
This is an experimental API and subject to change.
Overview
HSTU LMSD is a fused training operation used by Hierarchical Sequential
Transduction Unit (HSTU) models. For each row of x, it computes LayerNorm,
applies the learned affine transform, multiplies the result by an optional
SiLU(u), and optionally applies inverted dropout. With
the mandatory result is dropout(ell * a). The output optionally prepends
dropout(a) and dropout(x) according to concat_u and concat_x:
When dropout_ratio > 0, the independent decisions are packed into one
int8 mask per input element: bit 0 corresponds to the mandatory LMSD result,
bit 1 to the optional x segment, and bit 2 to the optional activated-u segment.
A set bit means that the element is kept. Disabled auxiliary segments leave
their bits clear. With dropout_ratio=0, Philox and mask traffic are compiled
out and forward returns mask_tensor=None. Forward also returns the row-wise
mean and reciprocal standard deviation needed by the explicit backward
operation.
This API does not register an autograd operator. Call
hstu_lmsd_backward explicitly with the tensors saved by
hstu_lmsd_forward and the matching four forward configuration parameters.
Installation
Install cuDNN Frontend with the CuTe DSL optional dependencies and a supported PyTorch installation:
From a source checkout, the PyTorch dependencies can instead be installed with
pip install --group torch.
The forward and backward functions are available through lazy top-level exports:
Supported configurations
All tensors must be CUDA tensors on the same device and have 16-byte-aligned storage. Output tensors and backward workspaces must not overlap inputs or one another.
Tensor shapes and layouts
For BF16 matrices with a padded row stride, each row must remain 16-byte
aligned. A cached compiled implementation accepts any runtime N in the
supported range. The x, u, dy, dx, and du row strides are runtime
values and do not participate in the compile-cache key; D, dtypes, devices,
and feature flags remain plan-time configuration.
The launch grid adapts to the runtime row count without changing the compiled
plan. Forward caps its persistent row blocks by the device SM count. Backward
uses at most one persistent tile per input row and caps large inputs at the
device SM count times 64 persistent CTAs per SM. The compiled vector width,
CTA shape, and backward row width are selected from D; they are not tied to
one model shape.
Functions
The allocating functions cache compiled API objects for repeated calls with
the same tensor layout and configuration. The cache key excludes N, so one
compiled kernel is reused across supported row counts:
hstu_lmsd_forward returns a TupleDict in the order y_tensor,
mean_tensor, rstd_tensor, and mask_tensor. hstu_lmsd_backward returns
dx_tensor, du_tensor, dweight_tensor, and dbias_tensor. The backward
function accepts optional caller-owned gradient output tensors. Pass
compute_dweight=False for a non-trainable weight; the returned
dweight_tensor is then None, and neither its FP32 workspace nor its
reduction work is created.
Forward and backward must use the same dropout_ratio, apply_u_silu,
concat_u, concat_x, saved statistics, and optional packed mask. Pass all
four configuration values explicitly to hstu_lmsd_backward; output tensors
do not carry hidden Python metadata. The seed is a signed 64-bit integer.
Pass stream= to enqueue wrapper allocation and execution on a specific
torch.cuda.Stream or CUDA stream handle; None uses the current PyTorch CUDA
stream.
For example, this returns only the mandatory LN(x) * u segment, performs no
dropout work, and skips dWeight:
See the focused tests in test/python/fe_api/hstu/hstu_lmsd/
for complete function calls.