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

Class StopOnSubsequence

axlearn/common/decoding.py:847–905  ·  view source on GitHub ↗

Early stopping on suffix-matches.

Source from the content-addressed store, hash-verified

845
846
847class 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 )

Callers 3

sample_decodeMethod · 0.90
sample_decodeMethod · 0.90
sample_decodeMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected