Early stopping on suffix-matches.
| 845 | |
| 846 | |
| 847 | class StopOnSubsequence(StopDecodingCondition): |
| 848 | """Early stopping on suffix-matches.""" |
| 849 | |
| 850 | def __init__( |
| 851 | self, stopping_seqs: Union[int, Sequence[int], Sequence[Sequence[int]]], pad_value=-1 |
| 852 | ): |
| 853 | """Stops decoding when a sequence suffix-matches one of `stopping_seqs`. |
| 854 | |
| 855 | Note: we prefix pad targets by -1; if this is a meaningful token, override by setting |
| 856 | pad_value. |
| 857 | |
| 858 | Args: |
| 859 | stopping_seqs: List of lists of ids that mean we should stop decoding. |
| 860 | pad_value: Safe value to use for padding. |
| 861 | |
| 862 | Raises: |
| 863 | ValueError: if stopping_seqs is an empty list. |
| 864 | """ |
| 865 | self.pad_value = pad_value |
| 866 | |
| 867 | if isinstance(stopping_seqs, int): |
| 868 | stopping_seqs = [[stopping_seqs]] |
| 869 | if isinstance(stopping_seqs, list) and stopping_seqs: |
| 870 | if isinstance(stopping_seqs[0], int): |
| 871 | stopping_seqs = [stopping_seqs] |
| 872 | |
| 873 | if any(len(seq) == 0 for seq in stopping_seqs): |
| 874 | indices = np.argwhere([len(seq) == 0 for seq in stopping_seqs]).flatten().tolist() |
| 875 | raise ValueError( |
| 876 | "Zero length stopping seqs are not supported. " |
| 877 | f"Zero length seqs at indices {indices}." |
| 878 | ) |
| 879 | self.longest = max(len(el) for el in stopping_seqs) |
| 880 | self.targets = jnp.stack( |
| 881 | [ |
| 882 | jnp.pad(jnp.array(el), (self.longest - len(el), 0), constant_values=pad_value) |
| 883 | for el in stopping_seqs |
| 884 | ] |
| 885 | ) |
| 886 | |
| 887 | def __call__(self, *, index: Tensor, sequences: Tensor, prefix_len: Tensor) -> Tensor: |
| 888 | sequences = jnp.pad( |
| 889 | sequences, [(0, 0), (0, 0), (self.longest - 1, 0)], constant_values=self.pad_value |
| 890 | ) |
| 891 | index = jnp.reshape(index, (-1, 1, 1)) |
| 892 | |
| 893 | # `index` can take different values across the batch. We slice `self.longest`-length |
| 894 | # sequences starting from each index. |
| 895 | # [batch, num_decodes=1, length + longest - 1]. |
| 896 | index = index + jnp.arange(self.longest) |
| 897 | # TODO(markblee): Compare against dispatch matrix + einsum to understand the performance |
| 898 | # tradeoff between number of updates vs dispatch matrix size. |
| 899 | # [batch, num_decodes, longest]. |
| 900 | sequences = jnp.take_along_axis(sequences, index, axis=-1) |
| 901 | |
| 902 | token_match = (self.targets[None, None, :, :] == sequences[:, :, None, :]) | ( |
| 903 | self.targets == self.pad_value |
| 904 | ) |
no outgoing calls
no test coverage detected