(
model,
tree_candidates,
past_key_values,
tree_position_ids,
input_ids,
retrieve_indices,
attention_mask=None,
tree_mask=None,
)
| 1319 | |
| 1320 | |
| 1321 | def tree_decoding( |
| 1322 | model, |
| 1323 | tree_candidates, |
| 1324 | past_key_values, |
| 1325 | tree_position_ids, |
| 1326 | input_ids, |
| 1327 | retrieve_indices, |
| 1328 | attention_mask=None, |
| 1329 | tree_mask=None, |
| 1330 | ): |
| 1331 | |
| 1332 | zero_num = attention_mask.shape[1]-attention_mask.long().sum(-1) |
| 1333 | zero_num = zero_num[:, None] |
| 1334 | position_ids = tree_position_ids[None,:] + input_ids.shape[1]-zero_num |
| 1335 | |
| 1336 | |
| 1337 | attention_mask = torch.cat( |
| 1338 | (attention_mask, torch.ones_like(tree_candidates, device=attention_mask.device, dtype=attention_mask.dtype)), dim=1) |
| 1339 | |
| 1340 | hidden_states, past_key_value = forward_with_tree_mask(model.base_model.model, input_ids=tree_candidates,past_key_values=past_key_values, |
| 1341 | attention_mask=attention_mask, tree_mask=tree_mask,position_ids=position_ids) |
| 1342 | |
| 1343 | tree_logits = model.base_model.lm_head(hidden_states) |
| 1344 | |
| 1345 | |
| 1346 | |
| 1347 | |
| 1348 | logits = tree_logits[:, retrieve_indices] |
| 1349 | return logits, hidden_states,past_key_value |
| 1350 | |
| 1351 | |
| 1352 | def evaluate_posterior( |
no test coverage detected