[Spec] Support sampling in the DSPARK graph-folded draft proposal (#33298)
This commit is contained in:
@@ -226,6 +226,16 @@ class InvariantCheckLevel(IntEnum):
|
|||||||
STRICT = 2
|
STRICT = 2
|
||||||
|
|
||||||
|
|
||||||
|
class DsparkFoldedSampling(IntEnum):
|
||||||
|
"""Sampling support in the graph-folded DSpark draft proposal: OFF =
|
||||||
|
greedy-only folding, AUTO = on when its buffers fit in free GPU memory,
|
||||||
|
FORCE = always."""
|
||||||
|
|
||||||
|
OFF = 0
|
||||||
|
AUTO = 1
|
||||||
|
FORCE = 2
|
||||||
|
|
||||||
|
|
||||||
class Envs:
|
class Envs:
|
||||||
|
|
||||||
# Raise on bare server_args field assignments after resolution; mutation
|
# Raise on bare server_args field assignments after resolution; mutation
|
||||||
@@ -326,6 +336,7 @@ class Envs:
|
|||||||
SGLANG_DSPARK_FAST_KERNEL = EnvBool(True)
|
SGLANG_DSPARK_FAST_KERNEL = EnvBool(True)
|
||||||
SGLANG_DSPARK_FP32_LM_HEAD = EnvBool(False)
|
SGLANG_DSPARK_FP32_LM_HEAD = EnvBool(False)
|
||||||
SGLANG_DSPARK_FAST_SAMPLING = EnvBool(True)
|
SGLANG_DSPARK_FAST_SAMPLING = EnvBool(True)
|
||||||
|
SGLANG_DSPARK_FOLDED_SAMPLING = EnvInt(DsparkFoldedSampling.AUTO)
|
||||||
SGLANG_DSPARK_OPT_MARKOV_W2_BF16 = EnvBool(True)
|
SGLANG_DSPARK_OPT_MARKOV_W2_BF16 = EnvBool(True)
|
||||||
SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD = EnvBool(True)
|
SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD = EnvBool(True)
|
||||||
SGLANG_DSPARK_ENABLE_MULTI_STREAM = EnvBool(True)
|
SGLANG_DSPARK_ENABLE_MULTI_STREAM = EnvBool(True)
|
||||||
|
|||||||
@@ -71,90 +71,6 @@ class DraftProposal(msgspec.Struct, frozen=True):
|
|||||||
folded: bool = False
|
folded: bool = False
|
||||||
|
|
||||||
|
|
||||||
def greedy_step_sampler(step_logits: torch.Tensor, step_idx: int) -> torch.Tensor:
|
|
||||||
del step_idx
|
|
||||||
return torch.argmax(step_logits, dim=-1)
|
|
||||||
|
|
||||||
|
|
||||||
class DsparkDraftSampler:
|
|
||||||
|
|
||||||
def __init__(self, *, model, gamma, max_bs, device, confidence_fn=None, out=None):
|
|
||||||
self.model = model
|
|
||||||
self.markov_head = model.markov_head
|
|
||||||
self.gamma = int(gamma)
|
|
||||||
if out is not None:
|
|
||||||
assert out.shape == (int(max_bs) * self.gamma,) and out.dtype == torch.int64
|
|
||||||
self.out = out
|
|
||||||
else:
|
|
||||||
self.out = torch.empty(
|
|
||||||
(int(max_bs) * self.gamma,), dtype=torch.int64, device=device
|
|
||||||
)
|
|
||||||
self.confidence_fn = confidence_fn
|
|
||||||
self.confidence_out = (
|
|
||||||
torch.empty((int(max_bs), self.gamma), dtype=torch.float32, device=device)
|
|
||||||
if confidence_fn is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
def __call__(self, hidden_states, input_ids):
|
|
||||||
bs = hidden_states.shape[0] // self.gamma
|
|
||||||
base_logits, confidence_tap = self.model.compute_base_logits(hidden_states)
|
|
||||||
base_logits = base_logits.view(bs, self.gamma, -1)
|
|
||||||
anchor = input_ids.view(bs, self.gamma)[:, 0]
|
|
||||||
draft_tokens, _ = self.markov_head.sample_block(
|
|
||||||
base_logits,
|
|
||||||
first_prev_tokens=anchor,
|
|
||||||
hidden_states=hidden_states.view(bs, self.gamma, -1),
|
|
||||||
sampler=greedy_step_sampler,
|
|
||||||
)
|
|
||||||
self.out[: draft_tokens.numel()].copy_(draft_tokens.reshape(-1))
|
|
||||||
if self.confidence_out is not None:
|
|
||||||
confidence = self.confidence_fn(
|
|
||||||
draft_hidden=hidden_states.view(bs, self.gamma, -1),
|
|
||||||
anchor_tokens=anchor,
|
|
||||||
draft_tokens=draft_tokens,
|
|
||||||
confidence_tap=confidence_tap,
|
|
||||||
)
|
|
||||||
self.confidence_out[:bs].copy_(confidence)
|
|
||||||
|
|
||||||
|
|
||||||
def maybe_build_draft_sampler(
|
|
||||||
*,
|
|
||||||
draft_model,
|
|
||||||
gamma: int,
|
|
||||||
max_bs: int,
|
|
||||||
device,
|
|
||||||
tp_rank: int,
|
|
||||||
confidence_fn=None,
|
|
||||||
out=None,
|
|
||||||
) -> Optional[DsparkDraftSampler]:
|
|
||||||
"""Build the graph-folded greedy draft sampler, or return None (with the
|
|
||||||
reason logged) when the draft model cannot support folding and the
|
|
||||||
proposal must stay eager."""
|
|
||||||
|
|
||||||
def _eager(reason):
|
|
||||||
if tp_rank == 0:
|
|
||||||
logger.info("DSpark draft greedy proposal kept eager (reason=%s).", reason)
|
|
||||||
return None
|
|
||||||
|
|
||||||
if gamma <= 0:
|
|
||||||
return _eager("gamma<=0")
|
|
||||||
if not hasattr(draft_model, "compute_base_logits"):
|
|
||||||
return _eager("no compute_base_logits")
|
|
||||||
if getattr(draft_model, "markov_head", None) is None:
|
|
||||||
return _eager("no markov head")
|
|
||||||
if tp_rank == 0:
|
|
||||||
logger.info("DSpark draft greedy proposal folded into the draft cuda graph.")
|
|
||||||
return DsparkDraftSampler(
|
|
||||||
model=draft_model,
|
|
||||||
gamma=gamma,
|
|
||||||
max_bs=max_bs,
|
|
||||||
device=device,
|
|
||||||
confidence_fn=confidence_fn,
|
|
||||||
out=out,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def make_next_draft_input(
|
def make_next_draft_input(
|
||||||
*,
|
*,
|
||||||
bonus_tokens: torch.Tensor,
|
bonus_tokens: torch.Tensor,
|
||||||
@@ -277,6 +193,8 @@ class DraftBlockProposer:
|
|||||||
sampling_info,
|
sampling_info,
|
||||||
) -> DraftProposal:
|
) -> DraftProposal:
|
||||||
embed_module = target_model.get_input_embeddings()
|
embed_module = target_model.get_input_embeddings()
|
||||||
|
draft_sampler = self._draft_sampler
|
||||||
|
all_greedy = sampling_info is None or sampling_info.is_all_greedy
|
||||||
fwd = self._run_forward(
|
fwd = self._run_forward(
|
||||||
batch=batch,
|
batch=batch,
|
||||||
draft_input=draft_input,
|
draft_input=draft_input,
|
||||||
@@ -284,30 +202,51 @@ class DraftBlockProposer:
|
|||||||
bs=bs,
|
bs=bs,
|
||||||
device=device,
|
device=device,
|
||||||
embed_module=embed_module,
|
embed_module=embed_module,
|
||||||
|
draft_sampler=draft_sampler,
|
||||||
|
sampling_info=sampling_info,
|
||||||
)
|
)
|
||||||
draft_block_ids = fwd.draft_block_ids
|
draft_block_ids = fwd.draft_block_ids
|
||||||
|
|
||||||
draft_sampler = self._draft_sampler
|
|
||||||
all_greedy = sampling_info is None or sampling_info.is_all_greedy
|
|
||||||
folded_confidence = None
|
folded_confidence = None
|
||||||
confidence_tap = None
|
confidence_tap = None
|
||||||
folded = False
|
folded = False
|
||||||
if draft_sampler is not None and fwd.can_run_graph and all_greedy:
|
if (
|
||||||
|
draft_sampler is not None
|
||||||
|
and fwd.can_run_graph
|
||||||
|
and (all_greedy or draft_sampler.folded_sampling)
|
||||||
|
):
|
||||||
folded = True
|
folded = True
|
||||||
if sampling_info is None:
|
if draft_sampler.folded_sampling:
|
||||||
temperatures = torch.ones(bs, dtype=torch.float32, device=device)
|
greedy_mask = draft_sampler.greedy_mask[:bs]
|
||||||
else:
|
temperatures = draft_sampler.temperatures[:bs]
|
||||||
temperatures = (
|
# The sampling accept path needs the markov-corrected block
|
||||||
sampling_info.temperatures.view(-1)
|
# logits; greedy accept only compares tokens.
|
||||||
.to(torch.float32)
|
corrected_logits = (
|
||||||
.clamp_min(1e-5)
|
None
|
||||||
|
if all_greedy
|
||||||
|
else draft_sampler.corrected_out[: bs * self.gamma].view(
|
||||||
|
bs, self.gamma, -1
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
# Greedy-only folding: the hook argmaxed every row and kept no
|
||||||
|
# sampling buffers, so derive the params on the fly.
|
||||||
|
greedy_mask = resolve_greedy_mask(
|
||||||
|
bs=bs, sampling_info=sampling_info, device=device
|
||||||
|
)
|
||||||
|
if sampling_info is None:
|
||||||
|
temperatures = torch.ones(bs, dtype=torch.float32, device=device)
|
||||||
|
else:
|
||||||
|
temperatures = (
|
||||||
|
sampling_info.temperatures.view(-1)
|
||||||
|
.to(torch.float32)
|
||||||
|
.clamp_min(1e-5)
|
||||||
|
)
|
||||||
|
corrected_logits = None
|
||||||
draft_block = DraftBlockResult(
|
draft_block = DraftBlockResult(
|
||||||
draft_tokens=draft_sampler.out[: bs * self.gamma].view(bs, self.gamma),
|
draft_tokens=draft_sampler.out[: bs * self.gamma].view(bs, self.gamma),
|
||||||
corrected_logits=None,
|
corrected_logits=corrected_logits,
|
||||||
greedy_mask=resolve_greedy_mask(
|
greedy_mask=greedy_mask,
|
||||||
bs=bs, sampling_info=sampling_info, device=device
|
|
||||||
),
|
|
||||||
temperatures=temperatures,
|
temperatures=temperatures,
|
||||||
)
|
)
|
||||||
if draft_sampler.confidence_out is not None:
|
if draft_sampler.confidence_out is not None:
|
||||||
@@ -367,6 +306,8 @@ class DraftBlockProposer:
|
|||||||
bs: int,
|
bs: int,
|
||||||
device: str,
|
device: str,
|
||||||
embed_module,
|
embed_module,
|
||||||
|
draft_sampler=None,
|
||||||
|
sampling_info=None,
|
||||||
) -> DraftForwardResult:
|
) -> DraftForwardResult:
|
||||||
gamma = self.gamma
|
gamma = self.gamma
|
||||||
prefix_lens = batch.seq_lens
|
prefix_lens = batch.seq_lens
|
||||||
@@ -411,6 +352,13 @@ class DraftBlockProposer:
|
|||||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
capture_hidden_mode=CaptureHiddenMode.NULL,
|
||||||
)
|
)
|
||||||
self._fill_dp_moe_sync_metadata(draft_forward_batch, batch)
|
self._fill_dp_moe_sync_metadata(draft_forward_batch, batch)
|
||||||
|
graph_runner = self.draft_model_runner.decode_cuda_graph_runner
|
||||||
|
if (
|
||||||
|
draft_sampler is not None
|
||||||
|
and graph_runner is not None
|
||||||
|
and graph_runner.can_run_graph(draft_forward_batch)
|
||||||
|
):
|
||||||
|
draft_sampler.stage_sampling_params(bs=bs, sampling_info=sampling_info)
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
draft_out = self.draft_model_runner.forward(draft_forward_batch)
|
draft_out = self.draft_model_runner.forward(draft_forward_batch)
|
||||||
logits_output = draft_out.logits_output
|
logits_output = draft_out.logits_output
|
||||||
|
|||||||
@@ -0,0 +1,205 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_draft_model import (
|
||||||
|
SampleStepTokens,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import DsparkFoldedSampling, envs
|
||||||
|
from sglang.srt.utils import get_available_gpu_memory
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Same free-memory floor init_cuda_graphs requires before draft capture.
|
||||||
|
_CAPTURE_HEADROOM_GB = 1.0
|
||||||
|
|
||||||
|
|
||||||
|
def greedy_step_sampler(step_logits: torch.Tensor, step_idx: int) -> torch.Tensor:
|
||||||
|
del step_idx
|
||||||
|
return torch.argmax(step_logits, dim=-1)
|
||||||
|
|
||||||
|
|
||||||
|
class DsparkDraftSampler:
|
||||||
|
"""Draft proposal head folded into the draft graph as a tail hook; with
|
||||||
|
folded_sampling it also Gumbel-samples non-greedy rows in-graph."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
model,
|
||||||
|
gamma,
|
||||||
|
max_bs,
|
||||||
|
device,
|
||||||
|
confidence_fn=None,
|
||||||
|
out=None,
|
||||||
|
folded_sampling: bool = True,
|
||||||
|
):
|
||||||
|
self.model = model
|
||||||
|
self.markov_head = model.markov_head
|
||||||
|
self.gamma = int(gamma)
|
||||||
|
max_bs = int(max_bs)
|
||||||
|
if out is not None:
|
||||||
|
assert out.shape == (max_bs * self.gamma,) and out.dtype == torch.int64
|
||||||
|
self.out = out
|
||||||
|
else:
|
||||||
|
self.out = torch.empty(
|
||||||
|
(max_bs * self.gamma,), dtype=torch.int64, device=device
|
||||||
|
)
|
||||||
|
self.confidence_fn = confidence_fn
|
||||||
|
self.confidence_out = (
|
||||||
|
torch.empty((max_bs, self.gamma), dtype=torch.float32, device=device)
|
||||||
|
if confidence_fn is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self.folded_sampling = folded_sampling
|
||||||
|
self.temperatures = None
|
||||||
|
self.greedy_mask = None
|
||||||
|
self.exp_noise = None
|
||||||
|
self.corrected_out = None
|
||||||
|
if folded_sampling:
|
||||||
|
vocab = int(model.lm_head.org_vocab_size)
|
||||||
|
self.temperatures = torch.ones(
|
||||||
|
(max_bs,), dtype=torch.float32, device=device
|
||||||
|
)
|
||||||
|
self.greedy_mask = torch.ones((max_bs,), dtype=torch.bool, device=device)
|
||||||
|
self.exp_noise = torch.empty(
|
||||||
|
(max_bs, vocab), dtype=torch.float32, device=device
|
||||||
|
)
|
||||||
|
self.corrected_out = torch.empty(
|
||||||
|
(max_bs * self.gamma, vocab),
|
||||||
|
dtype=model.lm_head.weight.dtype,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
def stage_sampling_params(self, *, bs: int, sampling_info) -> None:
|
||||||
|
"""Host-side refresh of the static sampling params; must run before
|
||||||
|
the draft graph replay that consumes them."""
|
||||||
|
if not self.folded_sampling:
|
||||||
|
return
|
||||||
|
if sampling_info is None:
|
||||||
|
self.temperatures[:bs].fill_(1.0)
|
||||||
|
self.greedy_mask[:bs].fill_(True)
|
||||||
|
return
|
||||||
|
torch.clamp(
|
||||||
|
sampling_info.temperatures.view(-1)[:bs].to(torch.float32),
|
||||||
|
min=1e-5,
|
||||||
|
out=self.temperatures[:bs],
|
||||||
|
)
|
||||||
|
self.greedy_mask[:bs].copy_((sampling_info.top_ks <= 1).view(-1)[:bs])
|
||||||
|
|
||||||
|
def __call__(self, hidden_states, input_ids):
|
||||||
|
bs = hidden_states.shape[0] // self.gamma
|
||||||
|
base_logits, confidence_tap = self.model.compute_base_logits(hidden_states)
|
||||||
|
base_logits = base_logits.view(bs, self.gamma, -1)
|
||||||
|
anchor = input_ids.view(bs, self.gamma)[:, 0]
|
||||||
|
|
||||||
|
if self.folded_sampling:
|
||||||
|
|
||||||
|
def sampler(step_logits: torch.Tensor, step_idx: int) -> torch.Tensor:
|
||||||
|
del step_idx
|
||||||
|
# In-graph philox noise: each replay advances the generator
|
||||||
|
# and redraws.
|
||||||
|
noise = self.exp_noise[:bs].exponential_()
|
||||||
|
return SampleStepTokens.execute(
|
||||||
|
step_logits=step_logits,
|
||||||
|
temperatures=self.temperatures[:bs],
|
||||||
|
greedy_mask=self.greedy_mask[:bs],
|
||||||
|
exp_noise=noise,
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
sampler = greedy_step_sampler
|
||||||
|
|
||||||
|
draft_tokens, corrected_logits = self.markov_head.sample_block(
|
||||||
|
base_logits,
|
||||||
|
first_prev_tokens=anchor,
|
||||||
|
hidden_states=hidden_states.view(bs, self.gamma, -1),
|
||||||
|
sampler=sampler,
|
||||||
|
)
|
||||||
|
self.out[: draft_tokens.numel()].copy_(draft_tokens.reshape(-1))
|
||||||
|
if self.folded_sampling:
|
||||||
|
self.corrected_out[: bs * self.gamma].copy_(
|
||||||
|
corrected_logits.reshape(bs * self.gamma, -1)
|
||||||
|
)
|
||||||
|
if self.confidence_out is not None:
|
||||||
|
confidence = self.confidence_fn(
|
||||||
|
draft_hidden=hidden_states.view(bs, self.gamma, -1),
|
||||||
|
anchor_tokens=anchor,
|
||||||
|
draft_tokens=draft_tokens,
|
||||||
|
confidence_tap=confidence_tap,
|
||||||
|
)
|
||||||
|
self.confidence_out[:bs].copy_(confidence)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_folded_sampling(*, model, gamma, max_bs, device, tp_rank) -> bool:
|
||||||
|
"""The sampling buffers are baked into the captured draft graph, so AUTO
|
||||||
|
must decide before capture from a free-memory probe."""
|
||||||
|
mode = envs.SGLANG_DSPARK_FOLDED_SAMPLING.get()
|
||||||
|
if mode == DsparkFoldedSampling.OFF:
|
||||||
|
return False
|
||||||
|
if mode == DsparkFoldedSampling.FORCE:
|
||||||
|
return True
|
||||||
|
vocab = int(model.lm_head.org_vocab_size)
|
||||||
|
noise_bytes = max_bs * vocab * 4
|
||||||
|
logits_bytes = max_bs * gamma * vocab * model.lm_head.weight.dtype.itemsize
|
||||||
|
need_gb = (noise_bytes + logits_bytes) / (1 << 30)
|
||||||
|
available_gb = get_available_gpu_memory(device, torch.cuda.current_device())
|
||||||
|
if available_gb - need_gb >= _CAPTURE_HEADROOM_GB:
|
||||||
|
return True
|
||||||
|
if tp_rank == 0:
|
||||||
|
logger.warning(
|
||||||
|
"DSpark folded sampling disabled: its static buffers need %.2f GB "
|
||||||
|
"but only %.2f GB GPU memory is free; sampling batches will take "
|
||||||
|
"the eager proposal path. Set SGLANG_DSPARK_FOLDED_SAMPLING=%d "
|
||||||
|
"to force.",
|
||||||
|
need_gb,
|
||||||
|
available_gb,
|
||||||
|
int(DsparkFoldedSampling.FORCE),
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def maybe_build_draft_sampler(
|
||||||
|
*,
|
||||||
|
draft_model,
|
||||||
|
gamma: int,
|
||||||
|
max_bs: int,
|
||||||
|
device,
|
||||||
|
tp_rank: int,
|
||||||
|
confidence_fn=None,
|
||||||
|
out=None,
|
||||||
|
) -> Optional[DsparkDraftSampler]:
|
||||||
|
"""Build the graph-folded draft sampler, or None (reason logged) when the
|
||||||
|
proposal must stay eager."""
|
||||||
|
|
||||||
|
def _eager(reason):
|
||||||
|
if tp_rank == 0:
|
||||||
|
logger.info("DSpark draft proposal kept eager (reason=%s).", reason)
|
||||||
|
return None
|
||||||
|
|
||||||
|
if gamma <= 0:
|
||||||
|
return _eager("gamma<=0")
|
||||||
|
if not hasattr(draft_model, "compute_base_logits"):
|
||||||
|
return _eager("no compute_base_logits")
|
||||||
|
if getattr(draft_model, "markov_head", None) is None:
|
||||||
|
return _eager("no markov head")
|
||||||
|
folded_sampling = _resolve_folded_sampling(
|
||||||
|
model=draft_model, gamma=gamma, max_bs=max_bs, device=device, tp_rank=tp_rank
|
||||||
|
)
|
||||||
|
if tp_rank == 0:
|
||||||
|
logger.info(
|
||||||
|
"DSpark draft proposal (%s) folded into the draft cuda graph.",
|
||||||
|
"greedy + sampling" if folded_sampling else "greedy only",
|
||||||
|
)
|
||||||
|
return DsparkDraftSampler(
|
||||||
|
model=draft_model,
|
||||||
|
gamma=gamma,
|
||||||
|
max_bs=max_bs,
|
||||||
|
device=device,
|
||||||
|
confidence_fn=confidence_fn,
|
||||||
|
out=out,
|
||||||
|
folded_sampling=folded_sampling,
|
||||||
|
)
|
||||||
@@ -33,6 +33,8 @@ from sglang.srt.speculative.dspark_components.dspark_config import (
|
|||||||
from sglang.srt.speculative.dspark_components.dspark_draft import (
|
from sglang.srt.speculative.dspark_components.dspark_draft import (
|
||||||
DraftBlockProposer,
|
DraftBlockProposer,
|
||||||
make_next_draft_input,
|
make_next_draft_input,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.dspark_components.dspark_draft_sampler import (
|
||||||
maybe_build_draft_sampler,
|
maybe_build_draft_sampler,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.dspark_components.dspark_kv_inject import (
|
from sglang.srt.speculative.dspark_components.dspark_kv_inject import (
|
||||||
@@ -592,6 +594,10 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
fold_eligible = (
|
fold_eligible = (
|
||||||
self._verify_executor.verify_epilogue is not None
|
self._verify_executor.verify_epilogue is not None
|
||||||
and proposal.folded
|
and proposal.folded
|
||||||
|
# The epilogue's in-graph accept is greedy (accept_greedy_triton);
|
||||||
|
# sampling batches must take the eager accept path even when the
|
||||||
|
# draft proposal itself folded.
|
||||||
|
and (sampling_info is None or sampling_info.is_all_greedy)
|
||||||
and verify_logits_adjustments_are_noop(sampling_info)
|
and verify_logits_adjustments_are_noop(sampling_info)
|
||||||
and self._simulate_acc_len <= 0
|
and self._simulate_acc_len <= 0
|
||||||
and not batch.has_grammar
|
and not batch.has_grammar
|
||||||
|
|||||||
Reference in New Issue
Block a user