Coverage for hopwise/model/init.py: 100%
16 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2020/9/16
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5"""hopwise.model.init
6########################
7"""
9from torch import nn
10from torch.nn.init import constant_, xavier_normal_, xavier_uniform_
13def xavier_normal_initialization(module):
14 r"""Using `xavier_normal_`_ in PyTorch to initialize the parameters in
15 nn.Embedding and nn.Linear layers. For bias in nn.Linear layers,
16 using constant 0 to initialize.
18 .. _`xavier_normal_`:
19 https://pytorch.org/docs/stable/nn.init.html?highlight=xavier_normal_#torch.nn.init.xavier_normal_
21 Examples:
22 >>> self.apply(xavier_normal_initialization)
23 """
24 if isinstance(module, nn.Embedding):
25 xavier_normal_(module.weight.data)
26 elif isinstance(module, nn.Linear):
27 xavier_normal_(module.weight.data)
28 if module.bias is not None:
29 constant_(module.bias.data, 0)
32def xavier_uniform_initialization(module):
33 r"""Using `xavier_uniform_`_ in PyTorch to initialize the parameters in
34 nn.Embedding and nn.Linear layers. For bias in nn.Linear layers,
35 using constant 0 to initialize.
37 .. _`xavier_uniform_`:
38 https://pytorch.org/docs/stable/nn.init.html?highlight=xavier_uniform_#torch.nn.init.xavier_uniform_
40 Examples:
41 >>> self.apply(xavier_uniform_initialization)
42 """
43 if isinstance(module, nn.Embedding):
44 xavier_uniform_(module.weight.data)
45 elif isinstance(module, nn.Linear):
46 xavier_uniform_(module.weight.data)
47 if module.bias is not None:
48 constant_(module.bias.data, 0)