Computes logits from decoder forward outputs. Args: forward_outputs: A dict containing: hidden_states: A float Tensor of shape [batch_size, target_len, hidden_dim]. Returns: A float Tensor of shape [batch_size, target_len, num_classes].
(self, forward_outputs: Nested[Tensor])
| 604 | """Computes logits from decoder forward outputs. |
| 605 | |
| 606 | Args: |
| 607 | forward_outputs: A dict containing: |
| 608 | hidden_states: A float Tensor of shape [batch_size, target_len, hidden_dim]. |
| 609 | |
| 610 | Returns: |
| 611 | A float Tensor of shape [batch_size, target_len, num_classes]. |
| 612 | """ |
| 613 | x: Tensor = forward_outputs["hidden_states"] |
| 614 | |
| 615 | if self.config.logits_forward_dtype == jnp.float32: |
| 616 | logits_context = jax.default_matmul_precision("float32") |
| 617 | else: |
| 618 | logits_context = contextlib.nullcontext() |
| 619 | |
| 620 | with logits_context: |
| 621 | logits_x = ( |
| 622 | x.astype(self.config.logits_forward_dtype) |
| 623 | if self.config.logits_forward_dtype |
| 624 | else x |
| 625 | ) |
| 626 | # Shard hidden states before the logit matmul to prevent XLA from |
| 627 | # materializing the full [batch, seq, vocab] tensor. |
| 628 | spec = self.config.hidden_state_partition_spec |
| 629 | logits_x = maybe_shard( |
| 630 | logits_x, spec if spec is not None else self.config.logits_partition_spec |
| 631 | ) |
| 632 | if "lm_head" in self.children: |
| 633 | logits = self.lm_head(logits_x) |
| 634 | else: |
| 635 | # Reuse the token embedding. |
| 636 | with child_context("emb_attend", module=self.emb): |
| 637 | logits = self.emb.attend(logits_x) |
| 638 | |
| 639 | if self._output_logits_modifier is not None: |
| 640 | logits = self._output_logits_modifier(logits) |
| 641 | logits = maybe_shard(logits, self.config.logits_partition_spec) |
| 642 | return logits |
| 643 | |
| 644 | def forward( |
| 645 | self, |
no test coverage detected