| 578 | class GroupNorm(nn.Module): |
| 579 | |
| 580 | def __init__( |
| 581 | self, |
| 582 | num_groups: int, |
| 583 | hidden_size: int, |
| 584 | elementwise_affine: bool = True, |
| 585 | bias: bool = False, |
| 586 | eps: float = 1e-5 |
| 587 | ) -> GroupNorm: |
| 588 | super().__init__() |
| 589 | |
| 590 | if hidden_size % num_groups != 0: |
| 591 | raise ValueError('num_channels must be divisible by num_groups') |
| 592 | |
| 593 | self.num_groups = num_groups |
| 594 | self.hidden_size = hidden_size |
| 595 | self.elementwise_affine = elementwise_affine |
| 596 | self.eps = eps |
| 597 | |
| 598 | self.register_parameter("weight", None) |
| 599 | self.register_parameter("bias", None) |
| 600 | if elementwise_affine: |
| 601 | self.weight = nn.Parameter(torch.ones(hidden_size)) |
| 602 | if bias: |
| 603 | self.bias = nn.Parameter(torch.zeros(hidden_size)) |
| 604 | |
| 605 | def __repr__(self) -> str: |
| 606 | s = f"{self.__class__.__name__}({self.num_groups}, {self.hidden_size}" |