PyTorch torch.renorm Function
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 tensorp: Norm orderdim: Dimension for normalizationmaxnorm: 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)
# 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)
# 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]])
Other Extensions