(input_ids, model, past_key_values, logits_processor)
| 230 | return draft_tokens, retrieve_indices,tree_mask,tree_position_ids, logits, hidden_state, sample_token |
| 231 | |
| 232 | def 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 | |
| 257 | def reset_tree_mode( |
no test coverage detected