Pytorch 参数初始化以及Xavier初始化

    def _initialize_weights(self):
        # print(self.modules())

        for m in self.modules():
            print(m)
            if isinstance(m, nn.Linear):
                # print(m.weight.data.type())
                # input()
                # m.weight.data.fill_(1.0)
                init.xavier_uniform_(m.weight, gain=1)
                print(m.weight)

可以再网络的类中定义初始化初始化函数

通过net._initialize_weights()进行初始化

Xavier的具体讲解可以参考

http://www.cnblogs.com/hejunlin1992/p/8723816.html

猜你喜欢

转载自blog.csdn.net/qq_24724109/article/details/82050402