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 BatchSpeculativeDecoder: def __init__(self, target_model, draft_model, tokenizer, k=5, batch_size=32): self.target_model = target_model self.draft_model = draft_model self.tokenizer = tokenizer self.k = k self.batch_size = batch_size def generate_batch(self, prompts, max_new_tokens): """批量推测解码""" input_ids = self.tokenizer( prompts, return_tensors="pt", padding=True, truncation=True ).input_ids batch_size = input_ids.shape[0] generated = input_ids.clone() while generated.shape[1] < max_new_tokens: draft_tokens, draft_probs = self.batch_draft_generate(generated) accepted, new_tokens = self.batch_target_verify( generated, draft_tokens, draft_probs ) generated = self.update_generated(generated, accepted, new_tokens) return generated def batch_draft_generate(self, input_ids): """批量草稿生成""" all_tokens = [] all_probs = [] for _ in range(self.k): with torch.no_grad(): outputs = self.draft_model(input_ids) probs = F.softmax(outputs.logits[:, -1, :], dim=-1) tokens = torch.multinomial(probs, 1) all_tokens.append(tokens) all_probs.append(probs.gather(1, tokens)) input_ids = torch.cat([input_ids, tokens], dim=1) return torch.cat(all_tokens, dim=1), torch.cat(all_probs, dim=1) def batch_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_list = [] new_tokens_list = [] for i in range(draft_tokens.shape[0]): accepted, new_token = self.verify_single( input_ids[i:i+1], draft_tokens[i:i+1], draft_probs[i:i+1], target_logits[i:i+1] ) if accepted is not None: accepted_list.append(accepted) else: new_tokens_list.append(new_token) return accepted_list, new_tokens_list def verify_single(self, input_ids, draft_tokens, draft_probs, target_logits): """验证单个序列""" accepted_tokens = [] 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)) else: corrected_probs = torch.clamp(target_probs - draft_probs, min=0) corrected_probs = corrected_probs / (corrected_probs.sum() + 1e-8) new_token = torch.multinomial(corrected_probs, 1) if accepted_tokens: return torch.cat(accepted_tokens, dim=1), new_token else: return None, new_token if accepted_tokens: return torch.cat(accepted_tokens, dim=1), None else: return None, draft_tokens[:, :1]
|