MCPcopy Create free account
hub / github.com/NVIDIA/FasterTransformer / gather_tree

Function gather_tree

examples/pytorch/decoding/utils/decoding.py:155–181  ·  view source on GitHub ↗
(step_ids, parent_ids, max_sequence_lengths, end_token)

Source from the content-addressed store, hash-verified

153
154
155def gather_tree(step_ids, parent_ids, max_sequence_lengths, end_token):
156 beams = torch.empty_like(step_ids)
157 beams.fill_(end_token)
158 max_len = step_ids.size(0)
159 batch_size = step_ids.size(1)
160 beam_size = step_ids.size(-1)
161 batch_beam = batch_size * beam_size
162 for i in range(batch_beam):
163 batch = i // beam_size
164 beam = i % beam_size
165 max_seq_len_b = min(max_len, max_sequence_lengths[batch])
166 if max_seq_len_b <= 0:
167 continue
168 beams[max_seq_len_b - 1, batch, beam] = step_ids[max_seq_len_b - 1, batch, beam]
169 parent = parent_ids[max_seq_len_b - 1, batch, beam]
170 for level in range(max_seq_len_b - 2, -1, -1):
171 if parent < 0 or parent > beam_size:
172 raise ValueError("wrong parent id")
173 beams[level, batch, beam] = step_ids[level, batch, parent]
174 parent = parent_ids[level, batch, parent]
175 finished = False
176 for time in range(max_seq_len_b):
177 if finished:
178 beams[time, batch, beam] = end_token
179 elif beams[time, batch, beam] == end_token:
180 finished = True
181 return beams
182
183
184def finalize(beam_size, output_ids, parent_ids, out_seq_lens, end_id, max_seq_len=None, args=None):

Callers 1

finalizeFunction · 0.70

Calls 1

sizeMethod · 0.45

Tested by

no test coverage detected