PyTorch torch.nn.Parameter Function
PyTorch torch.nn Reference Manual
torch.nn.ParameterIt is a learnable parameter tensor in PyTorch.
It is a wrapper that converts an ordinary tensor into a learnable parameter and is automatically added to the module's parameter list.
Function Definition
torch.nn.Parameter(data=None, requires_grad=True)
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super(MyModule, self).__init__()
# Create a learnable parameter
self.weight = nn.Parameter(torch.randn(10, 5))
self.bias = nn.Parameter(torch.zeros(5))
def forward(self, x):
return x @ self.weight.t() + self.bias
model = MyModule()
print("Parameters:", list(model.named_parameters()))
print("Weight shape:", model.weight.shape)
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super(MyModule, self).__init__()
# Create a learnable parameter
self.weight = nn.Parameter(torch.randn(10, 5))
self.bias = nn.Parameter(torch.zeros(5))
def forward(self, x):
return x @ self.weight.t() + self.bias
model = MyModule()
print("Parameters:", list(model.named_parameters()))
print("Weight shape:", model.weight.shape)
Example 2: Alternative to register_parameter
Example
import torch
import torch.nn as nn
# Method 1: Use nn.Parameter (recommended)
class Net1(nn.Module):
def __init__(self):
super().__init__()
self.param = nn.Parameter(torch.ones(5))
# Method 2: Use register_parameter (equivalent)
class Net2(nn.Module):
def __init__(self):
super().__init__()
self.register_parameter('param', nn.Parameter(torch.ones(5)))
print("Method 1 parameters:", list(Net1().named_parameters()))
print("Method 2 parameters:", list(Net2().named_parameters()))
import torch.nn as nn
# Method 1: Use nn.Parameter (recommended)
class Net1(nn.Module):
def __init__(self):
super().__init__()
self.param = nn.Parameter(torch.ones(5))
# Method 2: Use register_parameter (equivalent)
class Net2(nn.Module):
def __init__(self):
super().__init__()
self.register_parameter('param', nn.Parameter(torch.ones(5)))
print("Method 1 parameters:", list(Net1().named_parameters()))
print("Method 2 parameters:", list(Net2().named_parameters()))
Example 3: Non-learnable Parameters
Example
import torch
import torch.nn as nn
# requires_grad=False creates a non-learnable parameter
param_no_grad = nn.Parameter(torch.ones(5), requires_grad=False)
print("Trainable:", param_no_grad.requires_grad)
print("Still a Parameter:", type(param_no_grad))
import torch.nn as nn
# requires_grad=False creates a non-learnable parameter
param_no_grad = nn.Parameter(torch.ones(5), requires_grad=False)
print("Trainable:", param_no_grad.requires_grad)
print("Still a Parameter:", type(param_no_grad))
Use Cases
- Custom layers: Implement learnable weights
- Special modules: Non-standard parameters
- Direct access: Convenient operation
Tip: nn.Parameter is a subclass of tensor and is automatically added to parameters().
Other Extensions