MCPcopy Create free account
hub / github.com/google-deepmind/alphageometry / beam_decode

Method beam_decode

lm_inference.py:146–189  ·  view source on GitHub ↗

Beam search.

(
      self,
      inputs: str,
      eos_tokens: np.ndarray = None,
      mask_tokens: np.ndarray = None,
      dstate: dict[str, np.ndarray] = None,
  )

Source from the content-addressed store, hash-verified

144 return metrics_np
145
146 def beam_decode(
147 self,
148 inputs: str,
149 eos_tokens: np.ndarray = None,
150 mask_tokens: np.ndarray = None,
151 dstate: dict[str, np.ndarray] = None,
152 ) -> MetricsOutput:
153 """Beam search."""
154 inputs = jax.numpy.array([self.vocab.encode(inputs)] * self.batch_size)
155
156 eos = self.eos
157 if eos_tokens is not None:
158 eos_ids = self.encode_list(eos_tokens)
159 eos = np.array(
160 [1 if idx in eos_ids else 0 for idx in range(1024)], dtype=np.bfloat16
161 ).reshape((1, 1, 1024))
162
163 mask = self.mask
164 if mask_tokens is not None:
165 mask_ids = self.encode_list(mask_tokens)
166 mask = np.array(
167 [0 if idx in mask_ids else 1 for idx in range(1024)],
168 dtype=np.bfloat16,
169 ).reshape((1, 1, 1024))
170
171 metrics_np = self.call(inputs, dstate=dstate, eos=eos, mask=mask)
172
173 finished_seqs = metrics_np['finished_seqs']
174 finished_scores = metrics_np['finished_scores']
175
176 seqs = []
177 scores = []
178 for seq, score in zip(finished_seqs, finished_scores):
179 seq = self.decode(seq[1:])
180 seqs.append(seq)
181 scores.append(score)
182
183 return {
184 'finished_seqs': finished_seqs,
185 'finished_scores': finished_scores,
186 'seqs_str': seqs,
187 'scores': scores,
188 'dstate': metrics_np['dstate'],
189 }

Calls 4

encode_listMethod · 0.95
callMethod · 0.95
decodeMethod · 0.95
encodeMethod · 0.80