| 153 | |
| 154 | |
| 155 | def 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 | |
| 184 | def finalize(beam_size, output_ids, parent_ids, out_seq_lens, end_id, max_seq_len=None, args=None): |