PyTorch torch.nn Reference Manual

PyTorch'storch.nnmodule is the core module for building and training neural networks. It provides a rich set of classes and functions for defining and operating neural networks.

The following aretorch.nnsome key components of the module and their functions:

1. The nn.Module class:

  • nn.Moduleis the base class for all custom neural network models. Users typically derive their own model classes from this class, and define the network layer structure and the forward pass function within them.

2. Predefined Layers (Modules):

  • Includes various types of layer components, such as convolutional layers (nn.Conv1d, nn.Conv2d, nn.Conv3d), fully connected layers (nn.Linear), activation functions (nn.ReLU, nn.Sigmoid, nn.Tanh), etc.

3. Container classes:

  • nn.Sequential: Allow multiple layers to be combined sequentially to form a simple linearly stacked network.
  • nn.ModuleListandnn.ModuleDict: Can dynamically store and access submodules, supporting variable-length or named collections of modules.

4. Loss Functions:

  • torch.nnContains a series of loss functions for measuring the difference between model predictions and true labels, such as mean squared error loss (nn.MSELoss), cross-entropy loss (nn.CrossEntropyLoss), etc.

5. Functional Interface:

  • nn.functional(usually abbreviated asF), contains many functions that can directly operate on tensors. They implement the same functionality as layer objects, but do not have the ability to save and update parameters. For example, you can useF.relu()to directly perform ReLU operations, orF.conv2d()to perform convolution operations.

6. Initialization Methods:

  • torch.nn.initProvides some commonly used weight initialization strategies, such as Xavier initialization (nn.init.xavier_uniform_()) and Kaiming initialization (nn.init.kaiming_uniform_()), which are crucial for successfully training neural networks.

7. Transformer Layers:

  • PyTorch provides complete Transformer architecture components, includingnn.Transformer, nn.TransformerEncoder, nn.TransformerDecoderas well as attention mechanismsnn.MultiheadAttentionetc.

8. Normalization Layers:

  • Include batch normalization (BatchNorm), layer normalization (LayerNorm), group normalization (GroupNorm), instance normalization (InstanceNorm), and RMSNorm, etc.

PyTorch torch.nn Module Reference Manual

Neural Network Containers

Class/Function Description
torch.nn.Module Base class for all neural network modules.
torch.nn.Sequential(*args) Sequentially combines multiple modules.
torch.nn.ModuleList(modules) Stores submodules in a list.
torch.nn.ModuleDict(modules) Stores submodules in a dictionary.
torch.nn.ParameterList(parameters) Stores parameters in a list.
torch.nn.ParameterDict(parameters) Stores parameters in a dictionary.
torch.nn.Parameter(data) Creates a learnable parameter tensor.
torch.nn.Buffer(data) Creates a persistent buffer (non-learnable parameter).
torch.nn.Identity(*args, **kwargs) Identity transformation layer, input is directly output.

Global Hooks

Function Description
register_module_forward_pre_hook(hook) Registers a forward pre-hook.
register_module_forward_hook(hook) Registers a forward hook.
register_module_backward_hook(hook) Registers a backward hook.
register_module_full_backward_pre_hook(hook) Registers a full backward pre-hook.
register_module_full_backward_hook(hook) Registers a full backward hook.

Linear Layers

Class/Function Description
torch.nn.Linear(in_features, out_features, bias) Fully connected layer (linear transformation).
torch.nn.Bilinear(in1_features, in2_features, out_features, bias) Bilinear layer.
torch.nn.LazyLinear(out_features, bias) Linear layer with delayed initialization; automatically infers input dimension during the first forward pass.

Convolution Layers

Class/Function Description
torch.nn.Conv1d(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias, padding_mode) 1D convolution layer, commonly used for text and audio.
torch.nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias, padding_mode) 2D convolution layer, commonly used for images.
torch.nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias, padding_mode) 3D convolution layer, commonly used for videos and volumetric data.
torch.nn.ConvTranspose1d(in_channels, out_channels, kernel_size, stride, padding, output_padding, groups, bias, dilation, padding_mode) 1D transposed convolution (deconvolution), used for upsampling.
torch.nn.ConvTranspose2d(in_channels, out_channels, kernel_size, stride, padding, output_padding, groups, bias, dilation, padding_mode) 2D transposed convolution (deconvolution), used for upsampling.
torch.nn.ConvTranspose3d(in_channels, out_channels, kernel_size, stride, padding, output_padding, groups, bias, dilation, padding_mode) 3D transposed convolution (deconvolution), used for upsampling.
torch.nn.Unfold(kernel_size, dilation, padding, stride) Unfolds the input tensor into sliding window blocks.
torch.nn.Fold(output_size, kernel_size, dilation, padding, stride) Recombines the unfolded blocks into a tensor.

Pooling Layers

Class/Function Description
torch.nn.MaxPool1d(kernel_size, stride, padding, dilation, return_indices) 1D max pooling layer.
torch.nn.MaxPool2d(kernel_size, stride, padding, dilation, return_indices, ceil_mode) 2D max pooling layer.
torch.nn.MaxPool3d(kernel_size, stride, padding, dilation, return_indices, ceil_mode) 3D max pooling layer.
torch.nn.MaxUnpool1d(kernel_size, stride, padding) 1D max unpooling layer.
torch.nn.MaxUnpool2d(kernel_size, stride, padding) 2D max unpooling layer.
torch.nn.MaxUnpool3d(kernel_size, stride, padding) 3D max unpooling layer.
torch.nn.AvgPool1d(kernel_size, stride, padding) 1D average pooling layer.
torch.nn.AvgPool2d(kernel_size, stride, padding, ceil_mode, count_include_pad) 2D average pooling layer.
torch.nn.AvgPool3d(kernel_size, stride, padding, ceil_mode, count_include_pad) 3D average pooling layer.
torch.nn.AdaptiveMaxPool1d(output_size, return_indices) 1D adaptive max pooling, fixed output size.
torch.nn.AdaptiveMaxPool2d(output_size, return_indices) 2D adaptive max pooling, fixed output size.
torch.nn.AdaptiveMaxPool3d(output_size, return_indices) 3D adaptive max pooling, fixed output size.
torch.nn.AdaptiveAvgPool1d(output_size) 1D adaptive average pooling, fixed output size.
torch.nn.AdaptiveAvgPool2d(output_size) 2D adaptive average pooling, fixed output size.
torch.nn.AdaptiveAvgPool3d(output_size) 3D adaptive average pooling, fixed output size.
torch.nn.LPPool1d(norm_type, kernel_size, stride, padding) 1D Lp pooling layer.
torch.nn.LPPool2d(norm_type, kernel_size, stride, padding) 2D Lp pooling layer.
torch.nn.FractionalMaxPool2d(kernel_size, output_size, output_ratio, return_indices) 2D fractional max pooling, using random step sizes.
torch.nn.FractionalMaxPool3d(kernel_size, output_size, output_ratio, return_indices) 3D fractional max pooling, using random step sizes.

Padding Layers

Class/Function Description
torch.nn.ReflectionPad1d(padding) 1D reflection padding, copies by reflecting along the boundary.
torch.nn.ReflectionPad2d(padding) 2D reflection padding, copies by reflecting along the boundary.
torch.nn.ReflectionPad3d(padding) 3D reflection padding, copies by reflecting along the boundary.
torch.nn.ReplicationPad1d(padding) 1D replication padding, copies edge values along the boundary.
torch.nn.ReplicationPad2d(padding) 2D replication padding, copies edge values along the boundary.
torch.nn.ReplicationPad3d(padding) 3D replication padding, copies edge values along the boundary.
torch.nn.ZeroPad1d(padding) 1D zero padding.
torch.nn.ZeroPad2d(padding) 2D zero padding.
torch.nn.ZeroPad3d(padding) 3D zero padding.
torch.nn.ConstantPad1d(padding, value) 1D constant padding, fills with a specified value.
torch.nn.ConstantPad2d(padding, value) 2D constant padding, fills with a specified value.
torch.nn.ConstantPad3d(padding, value) 3D constant padding, fills with a specified value.
torch.nn.CircularPad1d(padding) 1D cyclic padding.
torch.nn.CircularPad2d(padding) 2D cyclic padding.
torch.nn.CircularPad3d(padding) 3D cyclic padding.

Activation Functions (Nonlinear Activations - Weighted Sum)

Class/Function Description
torch.nn.ReLU(inplace) ReLU activation function, f(x) = max(0, x).
torch.nn.ReLU6(inplace) ReLU6 activation function, f(x) = min(max(0, x), 6).
torch.nn.Sigmoid() Sigmoid activation function, f(x) = 1 / (1 + exp(-x)).
torch.nn.Tanh() Tanh activation function, f(x) = (exp(x) - exp(-x)) / (exp(x) + exp(-x)).
torch.nn.LeakyReLU(negative_slope, inplace) LeakyReLU, allows small gradients for negative values.
torch.nn.PReLU(num_parameters, init) Parametric ReLU, with learnable negative slope parameter.
torch.nn.ELU(alpha, inplace) Exponential Linear Unit, uses exponential function for negative values.
torch.nn.CELU(alpha, inplace) Continuously Differentiable Exponential Linear Unit.
torch.nn.SELU(inplace) Self-Normalizing Exponential Linear Unit.
torch.nn.GELU() Gaussian Error Linear Unit, commonly used in Transformers.
torch.nn.SiLU(inplace) Sigmoid Linear Unit (Swish), f(x) = x * sigmoid(x).
torch.nn.Mish(inplace) Mish activation function, f(x) = x * tanh(softplus(x)).
torch.nn.Hardtanh(min_value, max_value, inplace) Hard hyperbolic tangent, limits output range.
torch.nn.Hardswish(inplace) Hard Swish, a smooth version of ReLU6.
torch.nn.Hardsigmoid(inplace) Hard Sigmoid, piecewise linear approximation.
torch.nn.RReLU(lower, upper, inplace) Randomized LeakyReLU, randomly selects negative slope during training.
torch.nn.Softplus(beta, threshold) Softplus, a smooth approximation of ReLU.
torch.nn.Softshrink(lambda) Softshrink activation function.
torch.nn.Hardshrink(lambda) Hardshrink activation function.
torch.nn.Softsign() Softsign activation function, f(x) = x / (1 + |x|).
torch.nn.Tanhshrink() Tanhshrink,f(x) = x - tanh(x)。
torch.nn.LogSigmoid() Log Sigmoid,f(x) = log(sigmoid(x))。
torch.nn.Threshold(threshold, value, inplace) Threshold activation function.
torch.nn.GLU(dim) Gated Linear Unit, splits the input into two parts along a specified dimension and multiplies them element-wise.

Activation Functions (Nonlinear Activations - Others)

Class/Function Description
torch.nn.Softmax(dim) Softmax activation function, converts values into a probability distribution.
torch.nn.Softmax2d() Softmax over spatial dimensions, used for images.
torch.nn.LogSoftmax(dim) Log Softmax, numerically stable version of Softmax.
torch.nn.Softmin(dim) Softmin, the opposite of Softmax.
torch.nn.AdaptiveLogSoftmaxWithLoss(in_features, n_classes, cutoffs, div_value, head_bias) Adaptive Log Softmax, used for large-vocabulary classification.

Normalization Layers

Class/Function Description
torch.nn.BatchNorm1d(num_features, eps, momentum, affine, track_running_stats) One-dimensional batch normalization layer, normalizes mini-batch data.
torch.nn.BatchNorm2d(num_features, eps, momentum, affine, track_running_stats) Two-dimensional batch normalization layer, commonly used in convolutional networks.
torch.nn.BatchNorm3d(num_features, eps, momentum, affine, track_running_stats) Three-dimensional batch normalization layer, used for three-dimensional data such as videos.
torch.nn.LazyBatchNorm1d() One-dimensional batch normalization with deferred initialization.
torch.nn.LazyBatchNorm2d() Two-dimensional batch normalization with deferred initialization.
torch.nn.LazyBatchNorm3d() Three-dimensional batch normalization with deferred initialization.
torch.nn.LayerNorm(normalized_shape, eps, elementwise_affine) Layer normalization, commonly used in Transformers.
torch.nn.GroupNorm(num_groups, num_channels, eps, affine) Group normalization, normalizes after grouping channels.
torch.nn.InstanceNorm1d(num_features, eps, momentum, affine, track_running_stats) One-dimensional instance normalization, used for style transfer.
torch.nn.InstanceNorm2d(num_features, eps, momentum, affine, track_running_stats) Two-dimensional instance normalization, used for style transfer.
torch.nn.InstanceNorm3d(num_features, eps, momentum, affine, track_running_stats) Three-dimensional instance normalization, used for style transfer.
torch.nn.SyncBatchNorm(num_features, eps, momentum, affine, track_running_stats, process_group) Synchronized batch normalization, used for multi-GPU distributed training.
torch.nn.LocalResponseNorm(k, alpha, beta, size) Local response normalization, used in convolutional neural networks for lateral inhibition.
torch.nn.RMSNorm(normalized_shape, eps, elementwise_affine) RMS normalization, commonly used in Transformers.

Recurrent Neural Network Layers

Class/Function Description
torch.nn.RNN(input_size, hidden_size, num_layers, nonlinearity, bias, batch_first, dropout, bidirectional) Simple RNN layer.
torch.nn.LSTM(input_size, hidden_size, num_layers, bias, batch_first, dropout, bidirectional, proj_size) LSTM (Long Short-Term Memory) layer.
torch.nn.GRU(input_size, hidden_size, num_layers, bias, batch_first, dropout, bidirectional, proj_size) GRU (Gated Recurrent Unit) layer.
torch.nn.RNNCell(input_size, hidden_size, bias, nonlinearity) RNN cell (single layer).
torch.nn.LSTMCell(input_size, hidden_size, bias) LSTM cell (single layer).
torch.nn.GRUCell(input_size, hidden_size, bias) GRU cell (single layer).

Transformer Layers

Class/Function Description
torch.nn.Transformer(d_model, nhead, num_encoder_layers, num_decoder_layers, dim_feedforward, dropout, activation, custom_encoder, custom_decoder, layer_norm_eps, normalize_before, need_src_mask, need_tgt_mask, need_memory_mask, batch_first, norm_first, bias) Complete Transformer model.
torch.nn.TransformerEncoder(encoder_layer, num_layers, norm) Transformer encoder.
torch.nn.TransformerDecoder(decoder_layer, num_layers, norm) Transformer decoder.
torch.nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, activation, layer_norm_eps, batch_first, norm_first, bias) Transformer encoder layer.
torch.nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout, activation, layer_norm_eps, batch_first, norm_first, bias) Transformer decoder layer.
torch.nn.MultiheadAttention(embed_dim, num_heads, dropout, bias, add_bias_kv, add_zero_attn, kdim, vdim, batch_first) Multi-head attention mechanism.

Attention Mechanisms

Function Description
torch.nn.functional.scaled_dot_product_attention(query, key, value, attn_mask, dropout_p, is_causal) Scaled dot-product attention, PyTorch's optimized attention implementation.
torch.nn.attention.sdpa_kernel(backends) Sets the SDP (Scaled Dot-Product) attention backend.

Embedding Layers (Sparse Layers)

Class/Function Description
torch.nn.Embedding(num_embeddings, embedding_dim, padding_idx, max_norm, norm_type, scale_grad_by_freq, sparse) Embedding layer, maps discrete indices to dense vectors.
torch.nn.EmbeddingBag(num_embeddings, embedding_dim, max_norm, norm_type, scale_grad_by_freq, mode, sparse, per_sample_weights, include_last_offset, padding_idx) Embedding bag, aggregates multiple embeddings.

Dropout Layers

Class/Function Description
torch.nn.Dropout(p, inplace) Dropout layer, randomly zeroes input elements.
torch.nn.Dropout1d(p, inplace) One-dimensional Dropout, used for one-dimensional inputs.
torch.nn.Dropout2d(p, inplace) Two-dimensional Dropout, used for two-dimensional feature maps.
torch.nn.Dropout3d(p, inplace) Three-dimensional Dropout, used for three-dimensional feature volumes.
torch.nn.AlphaDropout(p, inplace) Alpha Dropout, maintains self-normalizing properties.
torch.nn.FeatureAlphaDropout(p, inplace) Feature Alpha Dropout.

Vision Layers

Class/Function Description
torch.nn.PixelShuffle(upscale_factor) Pixel shuffle, converts channel dimensions to spatial dimensions (upsampling).
torch.nn.PixelUnshuffle(downscale_factor) Inverse pixel shuffle, converts spatial dimensions to channel dimensions (downsampling).
torch.nn.Upsample(size, scale_factor, mode, align_corners, recompute_scale_factor) Upsampling layer.
torch.nn.UpsamplingNearest2d(size, scale_factor) Two-dimensional nearest-neighbor upsampling.
torch.nn.UpsamplingBilinear2d(size, scale_factor, align_corners) Two-dimensional bilinear upsampling.
torch.nn.ChannelShuffle(groups) Channel shuffle, used for ChannelShuffle networks.

Loss Functions

Class/Function Description
torch.nn.MSELoss(size_average, reduce, reduction) Mean Squared Error loss.
torch.nn.L1Loss(size_average, reduce, reduction) L1 loss (Mean Absolute Error).
torch.nn.CrossEntropyLoss(weight, size_average, ignore_index, reduce, reduction, label_smoothing) Cross-entropy loss, used for multi-class classification tasks.
torch.nn.NLLLoss(weight, size_average, ignore_index, reduce, reduction) Negative log-likelihood loss.
torch.nn.BCELoss(weight, size_average, reduce, reduction) Binary cross-entropy loss (binary classification).
torch.nn.BCEWithLogitsLoss(weight, pos_weight, size_average, reduce, reduction, label_smoothing) Binary cross-entropy loss with Sigmoid, numerically more stable.
torch.nn.KLDivLoss(size_average, reduce, reduction, log_target) KL divergence loss, used for distribution matching.
torch.nn.HuberLoss(delta, size_average, reduce, reduction) Huber loss, a combination of L1 and L2, more robust to outliers.
torch.nn.SmoothL1Loss(beta, size_average, reduce, reduction) Smooth L1 loss (a variant of Huber loss).
torch.nn.CTCLoss(blank, reduction, zero_infinity) Connectionist Temporal Classification loss, used for sequence tasks such as speech recognition.
torch.nn.PoissonNLLLoss(log_input, full, size_average, reduce, reduction) Poisson negative log-likelihood loss.
torch.nn.GaussianNLLLoss(full, size_average, reduce, reduction) Gaussian negative log-likelihood loss.
torch.nn.MarginRankingLoss(margin, size_average, reduce, reduction) Margin ranking loss, used for learning to rank.
torch.nn.HingeEmbeddingLoss(margin, size_average, reduce, reduction) Hinge embedding loss, used for metric learning.
torch.nn.MultiLabelMarginLoss(size_average, reduce, reduction) Multi-label margin loss.
torch.nn.SoftMarginLoss(size_average, reduce, reduction) Soft margin loss.
torch.nn.MultiLabelSoftMarginLoss(weight, size_average, reduce, reduction) Multi-label soft margin loss.
torch.nn.CosineEmbeddingLoss(margin, size_average, reduce, reduction) Cosine embedding loss, used for metric learning.
torch.nn.MultiMarginLoss(p, margin, weight, size_average, reduce, reduction) Multi-class margin loss.
torch.nn.TripletMarginLoss(margin, p, eps, swap, size_average, reduce, reduction) Triplet loss, used for metric learning and contrastive learning.
torch.nn.TripletMarginWithDistanceLoss(distance_function, margin, swap, size_average, reduce, reduction) Triplet loss with distance function.

Distance Functions

Class/Function Description
torch.nn.PairwiseDistance(p, eps, keepdim) Pairwise distance computation.
torch.nn.CosineSimilarity(dim, eps) Cosine similarity computation.

Parallel Layers (DataParallel)

Class/Function Description
torch.nn.DataParallel(module, device_ids, output_device, dim) Data parallelism, runs the model in parallel on multiple GPUs.
torch.nn.parallel.DistributedDataParallel(module, device_ids, broadcast_buffers, bucket_cap_mb, find_unused_parameters, gradient_as_bucket_view, static_graph) Distributed data parallelism, used for multi-node distributed training.

Utility Functions (nn.functional)

Function Description
torch.nn.functional.relu(input, inplace) Applies the ReLU activation function.
torch.nn.functional.sigmoid(input) Applies the Sigmoid activation function.
torch.nn.functional.tanh(input) Applies the Tanh activation function.
torch.nn.functional.softmax(input, dim, dtype) Applies the Softmax activation function.
torch.nn.functional.log_softmax(input, dim, dtype) Applies the Log Softmax activation function.
torch.nn.functional.gelu(input) Applies the GELU activation function.
torch.nn.functional.silu(input) Applies the SiLU (Swish) activation function.
torch.nn.functional.mish(input) Applies the Mish activation function.
torch.nn.functional.hardswish(input) Applies the Hardswish activation function.
torch.nn.functional.leaky_relu(input, negative_slope, inplace) Applies the LeakyReLU activation function.
torch.nn.functional.elu(input, alpha, inplace) Applies the ELU activation function.
torch.nn.functional.dropout(input, p, training, inplace) Applies Dropout.
torch.nn.functional.conv1d(input, weight, bias, stride, padding, dilation, groups) One-dimensional convolution operation.
torch.nn.functional.conv2d(input, weight, bias, stride, padding, dilation, groups) Two-dimensional convolution operation.
torch.nn.functional.conv3d(input, weight, bias, stride, padding, dilation, groups) Three-dimensional convolution operation.
torch.nn.functional.max_pool1d(input, kernel_size, stride, padding, dilation, return_indices) One-dimensional max pooling.
torch.nn.functional.max_pool2d(input, kernel_size, stride, padding, dilation, return_indices, ceil_mode) Two-dimensional max pooling.
torch.nn.functional.avg_pool1d(input, kernel_size, stride, padding, ceil_mode, count_include_pad) One-dimensional average pooling.
torch.nn.functional.avg_pool2d(input, kernel_size, stride, padding, ceil_mode, count_include_pad) Two-dimensional average pooling.
torch.nn.functional.linear(input, weight, bias) Linear transformation (matrix multiplication).
torch.nn.functional.cross_entropy(input, target, weight, size_average, ignore_index, reduce, reduction, label_smoothing) Computes cross-entropy loss.
torch.nn.functional.mse_loss(input, target, size_average, reduce, reduction) Computes mean squared error loss.
torch.nn.functional.l1_loss(input, target, size_average, reduce, reduction) Computes L1 loss.
torch.nn.functional.binary_cross_entropy(input, target, weight, size_average, reduce, reduction) Computes binary cross-entropy loss.
torch.nn.functional.nll_loss(input, target, weight, size_average, ignore_index, reduce, reduction) Computes negative log-likelihood loss.
torch.nn.functional.huber_loss(input, target, delta, size_average, reduce, reduction) Computes Huber loss.
torch.nn.functional.batch_norm(input, running_mean, running_var, weight, bias, training, momentum, eps, track_running_stats) Batch normalization operation.
torch.nn.functional.layer_norm(input, normalized_shape, weight, bias, eps) Layer normalization operation.
torch.nn.functional.group_norm(input, num_groups, weight, bias, eps) Group normalization operation.
torch.nn.functional.interpolate(input, size, scale_factor, mode, align_corners, recompute_scale_factor) Interpolation (upsampling/downsampling).
torch.nn.functional.grid_sample(input, grid, mode, padding_mode, align_corners) Grid sampling, used for image registration and spatial transformer networks.
torch.nn.functional.affine_grid(theta, size, align_corners) Affine grid generation, used for spatial transformer networks.
torch.nn.functional.pixel_shuffle(input, upscale_factor) Pixel shuffle.
torch.nn.functional.pixel_unshuffle(input, downscale_factor) Inverse pixel shuffle.
torch.nn.functional.pad(input, pad, mode, value) Padding operation.
torch.nn.functional.one_hot(tensor, num_classes) Converts integers to one-hot encoding.
torch.nn.functional.embedding(input, weight, padding_idx, max_norm, norm_type, scale_grad_by_freq, sparse) Embedding operation.
torch.nn.functional.cosine_similarity(x1, x2, dim, eps) Cosine similarity computation.
torch.nn.functional.pairwise_distance(x1, x2, p, eps, keepdim) Pairwise distance computation.

Utility Functions (nn.utils)

Function Description
torch.nn.utils.clip_grad_norm_(parameters, max_norm, norm_type, error_if_nonfinite) Clips gradient norm (in-place operation).
torch.nn.utils.clip_grad_norm(parameters, max_norm, norm_type) Clips gradient norm (non-in-place).
torch.nn.utils.clip_grad_value_(parameters, clip_value) Clips gradient value range.
torch.nn.utils.get_total_norm(parameters, norm_type) Computes the total gradient norm.
torch.nn.utils.weight_norm(module, name, dim) Applies weight normalization to module parameters.
torch.nn.utils.remove_weight_norm(module, name) Removes weight normalization.
torch.nn.utils.spectral_norm(module, name, n_power_iterations, eps, bias) Applies spectral normalization to module parameters.
torch.nn.utils.remove_spectral_norm(module, name) Removes spectral normalization.
torch.nn.utils.fuse_conv_bn_eval(conv, bn) Fuses convolutional layers and batch normalization layers (inference mode).
torch.nn.utils.fuse_linear_bn_eval(linear, bn) Fuses linear layers and batch normalization layers (inference mode).
torch.nn.utils.skip_init(module_class, *args, **kwargs) Skips parameter initialization.
torch.nn.utils.parameters_to_vector(parameters) Flattens a parameter list into a vector.
torch.nn.utils.vector_to_parameters(vector, parameters) Reshapes a vector into a parameter list.

Parameter Initialization (nn.init)

Function Description
torch.nn.init.zeros_(tensor) Initializes a tensor with zeros.
torch.nn.init.ones_(tensor) Initializes a tensor with ones.
torch.nn.init.uniform_(tensor, a, b) Uniform distribution initialization.
torch.nn.init.normal_(tensor, mean, std) Normal distribution initialization.
torch.nn.init.constant_(tensor, val) Constant value initialization.
torch.nn.init.eye_(tensor) Identity matrix initialization (only applicable to 2D square matrices).
torch.nn.init.dirac_(tensor) Dirac delta initialization (preserves the number of input channels).
torch.nn.init.xavier_uniform_(tensor, gain) Xavier uniform distribution initialization.
torch.nn.init.xavier_normal_(tensor, gain) Xavier normal distribution initialization.
torch.nn.init.kaiming_uniform_(tensor, a, mode, nonlinearity) Kaiming uniform distribution initialization, suitable for ReLU activations.
torch.nn.init.kaiming_normal_(tensor, a, mode, nonlinearity) Kaiming normal distribution initialization, suitable for ReLU activations.
torch.nn.init.trunc_normal_(tensor, mean, std, a, b) Truncated normal distribution initialization.
torch.nn.init.orthogonal_(tensor, gain) Orthogonal initialization.
torch.nn.init.sparse_(tensor, sparsity, std) Sparse initialization (mostly zero).
torch.nn.init.calculate_gain(nonlinearity, param) Compute initialization gain value.

RNN utility functions

Function Description
torch.nn.utils.rnn.PackedSequence Pack sequences to handle variable-length sequences.
torch.nn.utils.rnn.pack_padded_sequence(input, lengths, batch_first, enforce_sorted) Pack padded sequences.
torch.nn.utils.rnn.pad_packed_sequence(input, batch_first, padding_value, total_length) Unpack sequences.
torch.nn.utils.rnn.pad_sequence(sequences, batch_first, padding_value) Pad sequences to the same length.
torch.nn.utils.rnn.pack_sequence(sequences, enforce_sorted) Directly pack a list of sequences.

Pruning functions

Function Description
torch.nn.utils.prune.random_unstructured(module, name, amount) Random unstructured pruning.
torch.nn.utils.prune.l1_unstructured(module, name, amount) L1 unstructured pruning.
torch.nn.utils.prune.global_unstructured(parameters, pruning_method, amount) Global unstructured pruning.
torch.nn.utils.prune.remove(module, name) Remove pruning.
torch.nn.utils.prune.is_pruned(module) Check whether a module has been pruned.

Flatten layer

Class/Function Description
torch.nn.Flatten(start_dim, end_dim) Flatten a tensor, flattening a multi-dimensional tensor into two dimensions.
torch.nn.Unflatten(dim, unflattened_size) Unflatten, reshaping a one-dimensional tensor to multi-dimensional.

Example

Example

import torch
import torch.nn as nn

# Define a simple neural network
class SimpleNet(nn.Module):
    def __init__(self):
        super(SimpleNet, self).__init__()
        self.fc1 = nn.Linear(10, 20)
        self.relu = nn.ReLU()
        self.fc2 = nn.Linear(20, 1)

    def forward(self, x):
        x = self.fc1(x)
        x = self.relu(x)
        x = self.fc2(x)
        return x

# Create model and input
model = SimpleNet()
input = torch.randn(5, 10)
output = model(input)
print(output)

Example: Image Classification Using CNN

import torch
import torch.nn as nn

class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        # Convolutional layer
        self.conv1 = nn.Conv2d(3, 16, kernel_size=3, padding=1)
        self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)
        self.conv3 = nn.Conv2d(32, 64, kernel_size=3, padding=1)

        # Pooling layer
        self.pool = nn.MaxPool2d(2, 2)

        # Batch normalization
        self.bn1 = nn.BatchNorm2d(16)
        self.bn2 = nn.BatchNorm2d(32)
        self.bn3 = nn.BatchNorm2d(64)

        # Activation function
        self.relu = nn.ReLU()

        # Fully connected layer
        self.fc1 = nn.Linear(64 * 4 * 4, 256)
        self.fc2 = nn.Linear(256, 10)

        # Dropout
        self.dropout = nn.Dropout(0.5)

    def forward(self, x):
        # Conv -> BN -> ReLU -> Pool
        x = self.pool(self.relu(self.bn1(self.conv1(x))))
        x = self.pool(self.relu(self.bn2(self.conv2(x))))
        x = self.pool(self.relu(self.bn3(self.conv3(x))))

        # Flatten
        x = x.view(x.size(0), -1)

        # FC -> ReLU -> Dropout -> FC
        x = self.dropout(self.relu(self.fc1(x)))
        x = self.fc2(x)
        return x

# Create model
model = CNN()
print(model)

# Test forward propagation
input_tensor = torch.randn(1, 3, 32, 32)
output = model(input_tensor)
print("Output shape:", output.shape)

Example: Text Classification Using LSTM

import torch
import torch.nn as nn

class LSTMClassifier(nn.Module):
    def __init__(self, vocab_size, embedding_dim, hidden_dim, num_layers, output_dim):
        super(LSTMClassifier, self).__init__()

        # Embedding layer
        self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=0)

        # LSTM layer
        self.lstm = nn.LSTM(
            embedding_dim,
            hidden_dim,
            num_layers=num_layers,
            batch_first=True,
            bidirectional=True,
            dropout=0.5
        )

        # Fully connected layer
        self.fc = nn.Linear(hidden_dim * 2, output_dim)

        # Dropout
        self.dropout = nn.Dropout(0.5)

    def forward(self, text, text_lengths):
        # Embedding
        embedded = self.embedding(text)

        # Pack sequences to handle variable-length inputs
        packed = nn.utils.rnn.pack_padded_sequence(
            embedded, text_lengths.cpu(), batch_first=True, enforce_sorted=False
        )

        # LSTM
        packed_output, (hidden, cell) = self.lstm(packed)

        # Unpack
        output, output_lengths = nn.utils.rnn.pad_packed_sequence(packed_output, batch_first=True)

        # Merge bidirectional final hidden states
        hidden = torch.cat((hidden[-2,:,:], hidden[-1,:,:]), dim=1)

        # Dropout and fully connected
        hidden = self.dropout(hidden)
        output = self.fc(hidden)

        return output

# Parameters
vocab_size = 10000
embedding_dim = 128
hidden_dim = 256
num_layers = 2
output_dim = 5  # 5 classes

# Create model
model = LSTMClassifier(vocab_size, embedding_dim, hidden_dim, num_layers, output_dim)
print(model)

Example: Using Transformer Encoder

import torch
import torch.nn as nn

class TransformerClassifier(nn.Module):
    def __init__(self, input_dim, d_model, nhead, num_layers, dim_feedforward, output_dim, dropout):
        super(TransformerClassifier, self).__init__()

        # Embedding layer
        self.embedding = nn.Linear(input_dim, d_model)

        # Positional encoding
        self.positional_encoding = nn.Parameter(torch.randn(1, 1000, d_model) * 0.1)

        # Transformer encoder layer
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=d_model,
            nhead=nhead,
            dim_feedforward=dim_feedforward,
            dropout=dropout,
            batch_first=True
        )

        # Transformer encoder
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)

        # Classification head
        self.fc = nn.Linear(d_model, output_dim)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x):
        # Add positional encoding
        seq_len = x.size(1)
        x = self.embedding(x) + self.positional_encoding[:, :seq_len, :]

        # Transformer encoding
        x = self.transformer_encoder(x)

        # Use the output at the first position for classification (similar to CLS token)
        x = x[:, 0, :]

        x = self.dropout(x)
        x = self.fc(x)

        return x

# Parameters
input_dim = 512
d_model = 512
nhead = 8
num_layers = 6
dim_feedforward = 2048
output_dim = 10
dropout = 0.1

# Create model
model = TransformerClassifier(input_dim, d_model, nhead, num_layers, dim_feedforward, output_dim, dropout)
print(model)

# Test
x = torch.randn(32, 100, input_dim)  # batch_size=32, seq_len=100
output = model(x)
print("Output shape:", output.shape)  # (32, 10)

If you need more detailed information, you can refer toPyTorch official documentation。

Other extensions