MCPcopy Create free account
hub / github.com/apple/axlearn / forward

Method forward

axlearn/common/layers.py:687–714  ·  view source on GitHub ↗
(self, x: Tensor, *, segment_ids: Optional[Tensor] = None)

Source from the content-addressed store, hash-verified

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
715class Linear(DenseGeneralBaseLayer):
716 """The linear layer."""
717

Callers

nothing calls this directly

Calls 3

astypeMethod · 0.80
add_state_updateMethod · 0.45

Tested by

no test coverage detected