Initialize the weights.
(self, module)
| 702 | self.config = config |
| 703 | |
| 704 | def init_bert_weights(self, module): |
| 705 | """ Initialize the weights. |
| 706 | """ |
| 707 | if isinstance(module, (nn.Linear, nn.Embedding)): |
| 708 | # Slightly different from the TF version which uses truncated_normal for initialization |
| 709 | # cf https://github.com/pytorch/pytorch/pull/5617 |
| 710 | module.weight.data.normal_(mean=0.0, std=self.config.initializer_range) |
| 711 | elif isinstance(module, BertLayerNorm): |
| 712 | module.bias.data.zero_() |
| 713 | module.weight.data.fill_(1.0) |
| 714 | if isinstance(module, nn.Linear) and module.bias is not None: |
| 715 | module.bias.data.zero_() |
| 716 | |
| 717 | @classmethod |
| 718 | def from_pretrained(cls, pretrained_model_name, state_dict=None, cache_dir=None, |
nothing calls this directly
no outgoing calls
no test coverage detected