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

Function initialize_tree0

eagle/model/utils.py:210–230  ·  view source on GitHub ↗
(input_ids, model, past_key_values, logits_processor)

Source from the content-addressed store, hash-verified

208
209
210def initialize_tree0(input_ids, model, past_key_values, logits_processor):
211 draft_tokens, retrieve_indices,tree_mask,tree_position_ids, outputs, logits, hidden_state, sample_token = model(
212 input_ids, past_key_values=past_key_values, output_orig=True, logits_processor=logits_processor
213 )
214
215 # if logits_processor is not None:
216 # logits = orig[:, -1]
217 # logits = logits_processor(None, logits)
218 # probabilities = torch.nn.functional.softmax(logits, dim=1)
219 # token = torch.multinomial(probabilities, 1)
220 # else:
221 # token = torch.argmax(orig[:, -1])
222 # token = token[None, None]
223 # input_ids = torch.cat((input_ids, token.to(input_ids.device)), dim=1)
224 # # Clone the output hidden states
225 #
226 # draft_tokens, retrieve_indices,tree_mask,tree_position_ids = self.ea_layer.topK_genrate(hidden_states, input_ids, self.base_model.lm_head)
227 # if output_orig:
228 # return draft_tokens, retrieve_indices,tree_mask,tree_position_ids, outputs, orig, hidden_states, token
229 # return draft_tokens, retrieve_indices,tree_mask,tree_position_ids, hidden_states, token
230 return draft_tokens, retrieve_indices,tree_mask,tree_position_ids, logits, hidden_state, sample_token
231
232def initialize_tree(input_ids, model, past_key_values, logits_processor):
233 outputs, orig, hidden_states = model(

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected