PyTorch Parameter 클래스
IT 위키
더 많은 작업
Parameter(PyTorch의 torch.nn.Parameter)는 학습 가능한 모델 파라미터를 나타내는 특수 텐서 클래스이다. 일반 텐서와 달리 모듈(nn.Module)에 속성으로 할당되면 자동으로 모델의 학습 대상 파라미터로 등록된다.
- 일반 Tensor는 학습 대상 파라미터로 자동 등록되지 않지만, Parameter 객체는 등록된다.
- Parameter.grad 속성은 역전파로 계산된 gradient를 저장한다.
- requires_grad=False로 설정하면 해당 파라미터는 gradient 계산 대상에서 제외된다.
import torch
from torch.nn import Parameter, Module
class MyModule(Module):
def __init__(self):
super().__init__()
self.weight = Parameter(torch.randn(10, 10))
self.bias = Parameter(torch.zeros(10))
def forward(self, x):
return x @ self.weight + self.bias
model = MyModule()
for name, param in model.named_parameters():
print(name, param.shape, param.requires_grad)