(self, x: Tensor, *, segment_ids: Optional[Tensor] = None)
| 685 | def forward(self, x: Tensor, *, segment_ids: Optional[Tensor] = None) -> Tensor: |
| 686 | cfg = self.config |
| 687 | x_dtype = x.dtype |
| 688 | if cfg.forward_dtype is not None: |
| 689 | x = x.astype(cfg.forward_dtype) |
| 690 | reduction_axis = tuple(range(x.ndim - 1)) |
| 691 | if self.is_training: |
| 692 | mean, variance = _compute_moments_with_segment_ids( |
| 693 | x=x, |
| 694 | segment_ids=segment_ids, |
| 695 | reduction_axis=list(reduction_axis), |
| 696 | keepdims=False, |
| 697 | ) |
| 698 | self.add_state_update( |
| 699 | "moving_mean", |
| 700 | cfg.decay * self.parameters["moving_mean"] + (1 - cfg.decay) * mean, |
| 701 | ) |
| 702 | self.add_state_update( |
| 703 | "moving_variance", |
| 704 | cfg.decay * self.parameters["moving_variance"] + (1 - cfg.decay) * variance, |
| 705 | ) |
| 706 | else: |
| 707 | mean = self.parameters["moving_mean"] |
| 708 | variance = self.parameters["moving_variance"] |
| 709 | x = (x - mean) * jax.lax.rsqrt(variance + cfg.eps) |
| 710 | x = x.astype(x_dtype) |
| 711 | x = x * self.parameters["scale"] + self.parameters["bias"] |
| 712 | return x |
| 713 | |
| 714 | |
| 715 | class Linear(DenseGeneralBaseLayer): |
| 716 | """The linear layer.""" |
| 717 |
nothing calls this directly
no test coverage detected