1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111
| class SpeculativeDecoder: def __init__(self, target_model, eagle3_model, tokenizer, k=5): self.target_model = target_model self.eagle3_model = eagle3_model self.tokenizer = tokenizer self.k = k def generate(self, prompt, max_new_tokens): """推测解码生成""" input_ids = self.tokenizer.encode(prompt, return_tensors="pt") generated = input_ids.clone() total_tokens = 0 accepted_tokens = 0 while generated.shape[1] < max_new_tokens: draft_tokens, draft_probs = self.draft_generate(generated) accepted, new_token = self.target_verify( generated, draft_tokens, draft_probs ) if accepted: generated = torch.cat([generated, accepted], dim=1) accepted_tokens += accepted.shape[1] else: generated = torch.cat([generated, new_token], dim=1) total_tokens += self.k + 1 if self.tokenizer.eos_token_id in generated: break acceptance_rate = accepted_tokens / total_tokens return generated, acceptance_rate def draft_generate(self, input_ids): """草稿模型生成""" features = extract_features(self.target_model, input_ids) tokens = [] probs = [] for _ in range(self.k): with torch.no_grad(): logits = self.eagle3_model(input_ids, features) next_probs = F.softmax(logits[:, -1, :], dim=-1) next_token = torch.multinomial(next_probs, 1) tokens.append(next_token) probs.append(next_probs.gather(1, next_token)) input_ids = torch.cat([input_ids, next_token], dim=1) return torch.cat(tokens, dim=1), torch.cat(probs, dim=1) def target_verify(self, input_ids, draft_tokens, draft_probs): """目标模型验证""" full_input = torch.cat([input_ids, draft_tokens], dim=1) with torch.no_grad(): outputs = self.target_model(full_input) target_logits = outputs.logits[:, input_ids.shape[1]-1:-1, :] accepted_tokens = [] last_accepted_pos = -1 for i in range(draft_tokens.shape[1]): target_probs = F.softmax(target_logits[:, i, :], dim=-1) draft_prob = draft_probs[:, i] token = draft_tokens[:, i] target_prob = target_probs.gather(1, token.unsqueeze(1)) accept_ratio = target_prob / (draft_prob + 1e-8) accept_prob = torch.min(torch.ones_like(accept_ratio), accept_ratio) if torch.rand(1).item() < accept_prob.item(): accepted_tokens.append(token.unsqueeze(1)) last_accepted_pos = i else: 修正_probs = torch.clamp(target_probs - draft_probs, min=0) 修正_probs = 修正_probs / (修正_probs.sum() + 1e-8) new_token = torch.multinomial(修正_probs, 1) return torch.cat(accepted_tokens, dim=1) if accepted_tokens else new_token, new_token if accepted_tokens: return torch.cat(accepted_tokens, dim=1), None else: return None, draft_tokens[:, :1]
|