MCPcopy Create free account
hub / github.com/SafeAILab/EAGLE / tree_decoding

Function tree_decoding

eagle/modeling_eagle.py:1321–1349  ·  view source on GitHub ↗
(
        model,
        tree_candidates,
        past_key_values,
        tree_position_ids,
        input_ids,
        retrieve_indices,
        attention_mask=None,
        tree_mask=None,
)

Source from the content-addressed store, hash-verified

1319
1320
1321def 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
1352def evaluate_posterior(

Callers 1

generateMethod · 0.70

Calls 2

forward_with_tree_maskFunction · 0.85
catMethod · 0.45

Tested by

no test coverage detected