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

Function initialize_tree

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

Source from the content-addressed store, hash-verified

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(
234 input_ids, past_key_values=past_key_values, output_orig=True
235 )
236
237 if logits_processor is not None:
238 logits = orig[:, -1]
239 logits = logits_processor(None, logits)
240 probabilities = torch.nn.functional.softmax(logits, dim=1)
241 token = torch.multinomial(probabilities, 1)
242 else:
243 token = torch.argmax(orig[:, -1])
244 token = token[None, None]
245 input_ids = torch.cat((input_ids, token.to(input_ids.device)), dim=1)
246
247 # Clone the output hidden states
248 if model.use_eagle3:
249 ea_device = model.ea_layer.lm_head.weight.device
250 if outputs["hidden_states"][0].device != ea_device:
251 outputs["hidden_states"] = [x.to(ea_device) for x in outputs["hidden_states"]]
252 hidden_states=torch.cat(outputs["hidden_states"],dim=-1)
253 draft_tokens, retrieve_indices,tree_mask,tree_position_ids = model.ea_layer.topK_genrate(hidden_states, input_ids, model.base_model.lm_head,logits_processor)
254 return draft_tokens, retrieve_indices,tree_mask,tree_position_ids, orig, hidden_states, token
255
256
257def reset_tree_mode(

Callers 4

eagenerateMethod · 0.70
ea_generateMethod · 0.70
ea_forwardFunction · 0.50
ea_forwardFunction · 0.50

Calls 2

catMethod · 0.45
topK_genrateMethod · 0.45

Tested by

no test coverage detected