[sgl] two potential spec_v2 bug fixes (#21589)
Co-authored-by: yilian49 <yilian49@users.noreply.github.com>
This commit is contained in:
@@ -475,6 +475,11 @@ class LogitsProcessor(nn.Module):
|
|||||||
input_logprob_indices_pt = 0
|
input_logprob_indices_pt = 0
|
||||||
input_logprob_indices = []
|
input_logprob_indices = []
|
||||||
pt, pruned_states_list, pruned_states_before_norm_list = 0, [], []
|
pt, pruned_states_list, pruned_states_before_norm_list = 0, [], []
|
||||||
|
aux_pruned_states_lists = (
|
||||||
|
[[] for _ in aux_hidden_states]
|
||||||
|
if aux_hidden_states is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
for idx, (extend_logprob_start_len, extend_len) in enumerate(
|
for idx, (extend_logprob_start_len, extend_len) in enumerate(
|
||||||
zip(
|
zip(
|
||||||
@@ -499,6 +504,11 @@ class LogitsProcessor(nn.Module):
|
|||||||
pruned_states_before_norm_list.append(
|
pruned_states_before_norm_list.append(
|
||||||
hidden_states_before_norm[pt + start_len : pt + extend_len]
|
hidden_states_before_norm[pt + start_len : pt + extend_len]
|
||||||
)
|
)
|
||||||
|
if aux_pruned_states_lists is not None:
|
||||||
|
for j, hidden in enumerate(aux_hidden_states):
|
||||||
|
aux_pruned_states_lists[j].append(
|
||||||
|
hidden[pt + start_len : pt + extend_len]
|
||||||
|
)
|
||||||
# Map each token to its sequence index, for chunked computation
|
# Map each token to its sequence index, for chunked computation
|
||||||
# of input logprobs
|
# of input logprobs
|
||||||
token_to_seq_idx.extend([idx] * (extend_len - start_len))
|
token_to_seq_idx.extend([idx] * (extend_len - start_len))
|
||||||
@@ -518,6 +528,8 @@ class LogitsProcessor(nn.Module):
|
|||||||
pruned_states = torch.cat(pruned_states_list)
|
pruned_states = torch.cat(pruned_states_list)
|
||||||
if hidden_states_before_norm is not None:
|
if hidden_states_before_norm is not None:
|
||||||
pruned_states_before_norm = torch.cat(pruned_states_before_norm_list)
|
pruned_states_before_norm = torch.cat(pruned_states_before_norm_list)
|
||||||
|
if aux_pruned_states_lists is not None:
|
||||||
|
aux_pruned_states = [torch.cat(lst) for lst in aux_pruned_states_lists]
|
||||||
sample_indices = torch.tensor(
|
sample_indices = torch.tensor(
|
||||||
sample_indices, device=pruned_states.device, dtype=torch.int64
|
sample_indices, device=pruned_states.device, dtype=torch.int64
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -531,6 +531,7 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
|
|
||||||
# Run forward
|
# Run forward
|
||||||
forward_batch = ForwardBatch.init_new(batch, self.draft_runner)
|
forward_batch = ForwardBatch.init_new(batch, self.draft_runner)
|
||||||
|
forward_batch.return_logprob = False
|
||||||
if mm_input_embeds is not None:
|
if mm_input_embeds is not None:
|
||||||
forward_batch.mm_input_embeds = mm_input_embeds
|
forward_batch.mm_input_embeds = mm_input_embeds
|
||||||
logits_output = self.draft_runner.forward(forward_batch).logits_output
|
logits_output = self.draft_runner.forward(forward_batch).logits_output
|
||||||
|
|||||||
Reference in New Issue
Block a user