1602参数管理


点击查看代码
import torch
from torch import nn

# 单隐藏层多层感知机
net = nn.Sequential(
    nn.Linear(4, 8),
    nn.ReLU(),
    nn.Linear(8, 1)
)
X = torch.rand(size=(2, 4))
print(net(X))

# 参数访问
print(net[0].state_dict())
print(net[1].state_dict())
print(net[2].state_dict())

# 目标参数
print("目标参数weight")
print(type(net[2].weight))
print(net[2].weight)
print(net[2].weight.data)
print(net[2].weight.grad)
print("目标参数bias")
print(type(net[2].bias))
print(net[2].bias)
print(net[2].bias.data)
print(net[2].bias.grad)
# 一次性访问所有参数
print("一次性访问所有参数")
print(*[(name, param.shape) for name, param in net[0].named_parameters()])
print(*[(name, param.shape) for name, param in net.named_parameters()])
# 通过参数名访问参数
# tensor([[ 0.2281, -0.0748, -0.3346,  0.2404,  0.0822,  0.2333,  0.2916,  0.1938]])
# 不显示?
print(net.state_dict()['2.weight'])
print(net.state_dict()['2.weight'].data)
print(net.state_dict()['2.weight'].grad)
# 从嵌套块收集参数
print("从嵌套块收集参数")
def block1():
    return nn.Sequential(
        nn.Linear(4, 8),
        nn.ReLU(),
        nn.Linear(8, 4),
        nn.ReLU()
    )

def block2():
    net = nn.Sequential()
    for i in range(5):
        net.add_module(f'block {i}', block1())
    return net

rgnet = nn.Sequential(block2(), nn.Linear(4, 1))
print(rgnet(X))
print(rgnet)
# 内置初始化
print("内置初始化")
def init_normal(m):
    if type(m) == nn.Linear:
        nn.init.normal_(m.weight, mean=0, std=0.01)
        nn.init.zeros_(m.bias)
net.apply(init_normal)
print(net[0].weight.data[0], net[0].bias.data[0])
print(net[2].weight.data[0], net[2].bias.data[0])
# 将所有参数初始化为给定的常数
def init_constant(m):
    if type(m) == nn.Linear:
        nn.init.constant_(m.weight, 1)
        nn.init.zeros_(m.bias)
net.apply(init_constant)
print(net[0].weight.data[0], net[0].bias.data[0])
print(net[2].weight.data[0], net[2].bias.data[0])
# 自定义初始化
print("自定义初始化")
def my_init(m):
    if type(m) == nn.Linear:
        print("Init", *[(name, param.shape)
                        for name, param in m.named_parameters()][0])
        nn.init.uniform_(m.weight, -10, 10)
        # 保留绝对值>5的权重
        m.weight.data *= m.weight.data.abs() >= 5

net.apply(my_init)
print(net[0].weight[:2])
# 直接修改
net[0].weight.data[:] += 1
print(net[0].weight[:2])
net[0].weight.data[0, 0] = 42
print(net[0].weight[:2])
# 参数绑定
print("参数绑定")
# 层之间共享
shared = nn.Linear(8, 8)
net = nn.Sequential(nn.Linear(4, 8), nn.ReLU(),
                    shared, nn.ReLU(),
                    shared, nn.ReLU(),
                    nn.Linear(8, 1))
net(X)
print(net[2].weight.data[0] == net[4].weight.data[0])
net[2].weight.data[0, 0] = 100
print(net[2].weight.data[0] == net[4].weight.data[0])