(self)
| 399 | # self.mask_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) if self.mask_input else None |
| 400 | |
| 401 | def fix_init_weight(self): |
| 402 | def rescale(param, layer_id): |
| 403 | param.div_(math.sqrt(2.0 * layer_id)) |
| 404 | |
| 405 | for layer_id, layer in enumerate(self.blocks): |
| 406 | rescale(layer.attn.proj.weight.data, layer_id + 1) |
| 407 | rescale(layer.mlp.fc2.weight.data, layer_id + 1) |
| 408 | |
| 409 | def _init_weights(self, m): |
| 410 | if isinstance(m, nn.Linear): |