PyTorch torch.renorm Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.renormIt is a function in PyTorch used to renormalize tensors. It normalizes the tensor according to the specified norm, so that the norm of elements along the specified dimension does not exceed the given value.

Function Definition

torch.renorm(input, p, dim, maxnorm)

Parameter Description:

  • input: Input tensor
  • p: Norm order
  • dim: Dimension for normalization
  • maxnorm: Maximum norm value

Usage Example

Example

import torch

# Create tensor
x = torch.tensor([[2.0, 4.0, 6.0], [3.0, 6.0, 9.0]])

# Perform L2 norm normalization on dim=1, with maximum norm 1
y = torch.renorm(x, p=2, dim=1, maxnorm=1)
print(y)

The output result is:

tensor([[0.2673, 0.5345, 0.8018],
        [0.2673, 0.5345, 0.8018]])

Example

import torch

# Create tensor
x = torch.tensor([[1.0, 2.0, 3.0]])

# Perform L1 norm normalization on dim=1
y = torch.renorm(x, p=1, dim=1, maxnorm=1)
print(y)

The output result is:

tensor([[0.1667, 0.3333, 0.5000]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions