(input_ids, model, past_key_values, logits_processor)
| 208 | |
| 209 | |
| 210 | def 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 | |
| 232 | def initialize_tree(input_ids, model, past_key_values, logits_processor): |
| 233 | outputs, orig, hidden_states = model( |
nothing calls this directly
no outgoing calls
no test coverage detected