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

Function tree_decoding

eagle/testbug/model/utils.py:295–314  ·  view source on GitHub ↗
(
        model,
        tree_candidates,
        past_key_values,
        tree_position_ids,
        input_ids,
        retrieve_indices,
)

Source from the content-addressed store, hash-verified

293
294
295def tree_decoding(
296 model,
297 tree_candidates,
298 past_key_values,
299 tree_position_ids,
300 input_ids,
301 retrieve_indices,
302):
303 position_ids = tree_position_ids + input_ids.shape[1]
304
305 outputs, tree_logits, hidden_state = model(
306 tree_candidates,
307 output_orig=True,
308 past_key_values=past_key_values,
309 position_ids=position_ids,
310 init=False,
311 )
312
313 logits = tree_logits[0, retrieve_indices]
314 return logits, hidden_state, outputs
315
316
317def evaluate_posterior(

Callers 2

eagenerateMethod · 0.70
eagenerateMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected