(self, num_class, input_dim, output_dim, bias=True)
| 35 | |
| 36 | class GroupWiseLinear(nn.Module): |
| 37 | def __init__(self, num_class, input_dim, output_dim, bias=True): |
| 38 | super().__init__() |
| 39 | self.num_class = num_class |
| 40 | self.input_dim = input_dim |
| 41 | self.output_dim = output_dim |
| 42 | self.bias = bias |
| 43 | |
| 44 | self.W = nn.Parameter(torch.Tensor(num_class, input_dim, output_dim)) |
| 45 | if bias: |
| 46 | self.b = nn.Parameter(torch.Tensor(num_class, output_dim)) |
| 47 | self.reset_parameters() |
| 48 | |
| 49 | def reset_parameters(self): |
| 50 | stdv = 1. / math.sqrt(self.W.size(2)) |
nothing calls this directly
no test coverage detected