[spec decoding] support inkling dspark (#31847)
This commit is contained in:
@@ -469,6 +469,210 @@ def fused_conv_window_scatter_with_mask(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_CONV_MULTI_MAX_TYPES = 8
|
||||||
|
_CONV_MULTI_META_COLS = 12
|
||||||
|
_conv_multi_meta_cache: dict = {}
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_conv_window_scatter_multi_kernel(
|
||||||
|
meta_ptr, # int64 [num_types, 12]: src_ptr, dst_ptr, elem, s_l, s_r, s_s, s_d, s_w, d_l, d_r, block_start, last_axis
|
||||||
|
idx1_ptr,
|
||||||
|
step1_ptr,
|
||||||
|
idx2_ptr,
|
||||||
|
step2_ptr,
|
||||||
|
n1,
|
||||||
|
src_req_size,
|
||||||
|
src_step_size,
|
||||||
|
dst_req_size,
|
||||||
|
NUM_TYPES: tl.constexpr,
|
||||||
|
META_COLS: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_req = tl.program_id(0)
|
||||||
|
pid_layer = tl.program_id(1).to(tl.int64)
|
||||||
|
pid_block = tl.program_id(2).to(tl.int64)
|
||||||
|
|
||||||
|
is1 = pid_req < n1
|
||||||
|
is2 = pid_req >= n1
|
||||||
|
off1 = pid_req
|
||||||
|
off2 = pid_req - n1
|
||||||
|
s1 = tl.load(step1_ptr + off1, mask=is1, other=-1).to(tl.int64)
|
||||||
|
s2 = tl.load(step2_ptr + off2, mask=is2, other=-1).to(tl.int64)
|
||||||
|
step_idx = tl.where(is2, s2, s1)
|
||||||
|
if step_idx < 0:
|
||||||
|
return
|
||||||
|
d1 = tl.load(idx1_ptr + off1, mask=is1, other=-1).to(tl.int64)
|
||||||
|
d2 = tl.load(idx2_ptr + off2, mask=is2, other=-1).to(tl.int64)
|
||||||
|
dst_idx = tl.where(is2, d2, d1)
|
||||||
|
src_idx = tl.where(is2, off2, off1).to(tl.int64)
|
||||||
|
|
||||||
|
if not (
|
||||||
|
(dst_idx >= 0)
|
||||||
|
& (dst_idx < dst_req_size)
|
||||||
|
& (src_idx < src_req_size)
|
||||||
|
& (step_idx < src_step_size)
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
for t in tl.static_range(NUM_TYPES):
|
||||||
|
block_start = tl.load(meta_ptr + t * META_COLS + 10)
|
||||||
|
block_end = tl.load(
|
||||||
|
meta_ptr + (t + 1) * META_COLS + 10,
|
||||||
|
mask=t + 1 < NUM_TYPES,
|
||||||
|
other=2147483647,
|
||||||
|
)
|
||||||
|
if (pid_block >= block_start) & (pid_block < block_end):
|
||||||
|
src_ptr = tl.load(meta_ptr + t * META_COLS + 0).to(
|
||||||
|
tl.pointer_type(tl.bfloat16)
|
||||||
|
)
|
||||||
|
dst_ptr = tl.load(meta_ptr + t * META_COLS + 1).to(
|
||||||
|
tl.pointer_type(tl.bfloat16)
|
||||||
|
)
|
||||||
|
elem_per_entry = tl.load(meta_ptr + t * META_COLS + 2)
|
||||||
|
src_layer_stride = tl.load(meta_ptr + t * META_COLS + 3)
|
||||||
|
src_req_stride = tl.load(meta_ptr + t * META_COLS + 4)
|
||||||
|
src_step_stride = tl.load(meta_ptr + t * META_COLS + 5)
|
||||||
|
src_dim_stride = tl.load(meta_ptr + t * META_COLS + 6)
|
||||||
|
src_win_stride = tl.load(meta_ptr + t * META_COLS + 7)
|
||||||
|
dst_layer_stride = tl.load(meta_ptr + t * META_COLS + 8)
|
||||||
|
dst_req_stride = tl.load(meta_ptr + t * META_COLS + 9)
|
||||||
|
last_axis = tl.load(meta_ptr + t * META_COLS + 11)
|
||||||
|
|
||||||
|
start = (pid_block - block_start) * BLOCK_SIZE
|
||||||
|
e = start + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = e < elem_per_entry
|
||||||
|
d = e // last_axis
|
||||||
|
w = e % last_axis
|
||||||
|
src_off = (
|
||||||
|
pid_layer * src_layer_stride
|
||||||
|
+ src_idx * src_req_stride
|
||||||
|
+ step_idx * src_step_stride
|
||||||
|
+ d * src_dim_stride
|
||||||
|
+ w * src_win_stride
|
||||||
|
)
|
||||||
|
dst_off = pid_layer * dst_layer_stride + dst_idx * dst_req_stride + e
|
||||||
|
data = tl.load(src_ptr + src_off, mask=mask, other=0.0)
|
||||||
|
tl.store(dst_ptr + dst_off, data, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def _conv_multi_build_meta(pairs, block_size: int):
|
||||||
|
rows = []
|
||||||
|
block_start = 0
|
||||||
|
for dst, src in pairs:
|
||||||
|
elem = dst.shape[2] * dst.shape[3]
|
||||||
|
rows.append(
|
||||||
|
[
|
||||||
|
src.data_ptr(),
|
||||||
|
dst.data_ptr(),
|
||||||
|
elem,
|
||||||
|
src.stride(0),
|
||||||
|
src.stride(1),
|
||||||
|
src.stride(2),
|
||||||
|
src.stride(3),
|
||||||
|
src.stride(4),
|
||||||
|
dst.stride(0),
|
||||||
|
dst.stride(1),
|
||||||
|
block_start,
|
||||||
|
dst.shape[3],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
block_start += triton.cdiv(elem, block_size)
|
||||||
|
meta = torch.tensor(rows, dtype=torch.int64, device=pairs[0][0].device)
|
||||||
|
return meta, block_start
|
||||||
|
|
||||||
|
|
||||||
|
def _conv_multi_eligible(pairs) -> bool:
|
||||||
|
if not (0 < len(pairs) <= _CONV_MULTI_MAX_TYPES):
|
||||||
|
return False
|
||||||
|
layers = pairs[0][0].shape[0]
|
||||||
|
for dst, src in pairs:
|
||||||
|
if dst.dtype != torch.bfloat16 or src.dtype != torch.bfloat16:
|
||||||
|
return False
|
||||||
|
if dst.ndim != 4 or src.ndim != 5:
|
||||||
|
return False
|
||||||
|
if dst.shape[0] != layers:
|
||||||
|
return False
|
||||||
|
if src.shape[0] != layers or src.shape[3:] != dst.shape[2:]:
|
||||||
|
return False
|
||||||
|
if not dst.is_contiguous():
|
||||||
|
return False
|
||||||
|
if src.shape[1:3] != pairs[0][1].shape[1:3]:
|
||||||
|
return False
|
||||||
|
if dst.shape[1] != pairs[0][0].shape[1]:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def fused_conv_window_scatter_multi(
|
||||||
|
pairs,
|
||||||
|
dst_indices_raw: torch.Tensor,
|
||||||
|
step_indices_raw: torch.Tensor,
|
||||||
|
dst_indices2_raw: torch.Tensor | None = None,
|
||||||
|
step_indices2_raw: torch.Tensor | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Single-launch variant of ``fused_conv_window_scatter_with_mask`` over
|
||||||
|
multiple (dst, src) conv-type pairs and up to two request-index sets (the
|
||||||
|
accept commit plus the optional interval-crossing track set)."""
|
||||||
|
n1 = step_indices_raw.shape[0]
|
||||||
|
n2 = 0 if step_indices2_raw is None else step_indices2_raw.shape[0]
|
||||||
|
if n1 + n2 == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
BLOCK_SIZE = 1024
|
||||||
|
key = tuple(
|
||||||
|
(dst.data_ptr(), src.data_ptr()) + tuple(src.stride()) + tuple(dst.shape)
|
||||||
|
for dst, src in pairs
|
||||||
|
)
|
||||||
|
cached = _conv_multi_meta_cache.get(key)
|
||||||
|
if cached is None:
|
||||||
|
cached = _conv_multi_build_meta(pairs, BLOCK_SIZE)
|
||||||
|
_conv_multi_meta_cache.clear()
|
||||||
|
_conv_multi_meta_cache[key] = cached
|
||||||
|
meta, total_blocks = cached
|
||||||
|
|
||||||
|
idx1 = (
|
||||||
|
dst_indices_raw
|
||||||
|
if dst_indices_raw.is_contiguous()
|
||||||
|
else dst_indices_raw.contiguous()
|
||||||
|
)
|
||||||
|
st1 = (
|
||||||
|
step_indices_raw
|
||||||
|
if step_indices_raw.is_contiguous()
|
||||||
|
else step_indices_raw.contiguous()
|
||||||
|
)
|
||||||
|
if n2 > 0:
|
||||||
|
idx2 = (
|
||||||
|
dst_indices2_raw
|
||||||
|
if dst_indices2_raw.is_contiguous()
|
||||||
|
else dst_indices2_raw.contiguous()
|
||||||
|
)
|
||||||
|
st2 = (
|
||||||
|
step_indices2_raw
|
||||||
|
if step_indices2_raw.is_contiguous()
|
||||||
|
else step_indices2_raw.contiguous()
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
idx2, st2 = idx1, st1
|
||||||
|
|
||||||
|
dst0, src0 = pairs[0]
|
||||||
|
grid = (n1 + n2, dst0.shape[0], total_blocks)
|
||||||
|
_fused_conv_window_scatter_multi_kernel[grid](
|
||||||
|
meta,
|
||||||
|
idx1,
|
||||||
|
st1,
|
||||||
|
idx2,
|
||||||
|
st2,
|
||||||
|
n1,
|
||||||
|
src0.shape[1],
|
||||||
|
src0.shape[2],
|
||||||
|
dst0.shape[1],
|
||||||
|
NUM_TYPES=len(pairs),
|
||||||
|
META_COLS=_CONV_MULTI_META_COLS,
|
||||||
|
BLOCK_SIZE=BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def scatter_mamba_states_after_mtp_verify(
|
def scatter_mamba_states_after_mtp_verify(
|
||||||
mamba_caches,
|
mamba_caches,
|
||||||
state_indices_tensor: torch.Tensor,
|
state_indices_tensor: torch.Tensor,
|
||||||
@@ -488,28 +692,38 @@ def scatter_mamba_states_after_mtp_verify(
|
|||||||
state_indices_tensor,
|
state_indices_tensor,
|
||||||
last_correct_step_indices,
|
last_correct_step_indices,
|
||||||
)
|
)
|
||||||
for conv_states, intermediate_conv_window_cache in zip(
|
if mamba_track_indices is not None:
|
||||||
mamba_caches.conv, mamba_caches.intermediate_conv_window
|
assert mamba_steps_to_track is not None
|
||||||
):
|
|
||||||
fused_conv_window_scatter_with_mask(
|
|
||||||
conv_states,
|
|
||||||
intermediate_conv_window_cache,
|
|
||||||
state_indices_tensor,
|
|
||||||
last_correct_step_indices,
|
|
||||||
)
|
|
||||||
|
|
||||||
if mamba_track_indices is not None:
|
|
||||||
assert mamba_steps_to_track is not None
|
|
||||||
if ssm_states.numel() > 0:
|
|
||||||
fused_mamba_state_scatter_with_mask(
|
fused_mamba_state_scatter_with_mask(
|
||||||
ssm_states,
|
ssm_states,
|
||||||
intermediate_state_cache,
|
intermediate_state_cache,
|
||||||
mamba_track_indices,
|
mamba_track_indices,
|
||||||
mamba_steps_to_track,
|
mamba_steps_to_track,
|
||||||
)
|
)
|
||||||
for conv_states, intermediate_conv_window_cache in zip(
|
|
||||||
mamba_caches.conv, mamba_caches.intermediate_conv_window
|
pairs = list(zip(mamba_caches.conv, mamba_caches.intermediate_conv_window))
|
||||||
):
|
if not pairs:
|
||||||
|
return
|
||||||
|
if mamba_track_indices is not None:
|
||||||
|
assert mamba_steps_to_track is not None
|
||||||
|
if _conv_multi_eligible(pairs):
|
||||||
|
fused_conv_window_scatter_multi(
|
||||||
|
pairs,
|
||||||
|
state_indices_tensor,
|
||||||
|
last_correct_step_indices,
|
||||||
|
mamba_track_indices,
|
||||||
|
mamba_steps_to_track,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
for conv_states, intermediate_conv_window_cache in pairs:
|
||||||
|
fused_conv_window_scatter_with_mask(
|
||||||
|
conv_states,
|
||||||
|
intermediate_conv_window_cache,
|
||||||
|
state_indices_tensor,
|
||||||
|
last_correct_step_indices,
|
||||||
|
)
|
||||||
|
if mamba_track_indices is not None:
|
||||||
|
for conv_states, intermediate_conv_window_cache in pairs:
|
||||||
fused_conv_window_scatter_with_mask(
|
fused_conv_window_scatter_with_mask(
|
||||||
conv_states,
|
conv_states,
|
||||||
intermediate_conv_window_cache,
|
intermediate_conv_window_cache,
|
||||||
@@ -518,6 +732,65 @@ def scatter_mamba_states_after_mtp_verify(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_commit_track_indices_kernel(
|
||||||
|
accept_index_ptr,
|
||||||
|
accept_lens_ptr,
|
||||||
|
seq_lens_ptr,
|
||||||
|
last_correct_out_ptr,
|
||||||
|
track_steps_out_ptr,
|
||||||
|
dtn,
|
||||||
|
interval,
|
||||||
|
HAS_TRACK: tl.constexpr,
|
||||||
|
):
|
||||||
|
b = tl.program_id(0).to(tl.int64)
|
||||||
|
al = tl.load(accept_lens_ptr + b).to(tl.int64)
|
||||||
|
base = b * dtn
|
||||||
|
last = tl.load(accept_index_ptr + base + al - 1).to(tl.int64) - base
|
||||||
|
tl.store(last_correct_out_ptr + b, last)
|
||||||
|
if HAS_TRACK:
|
||||||
|
pre = tl.load(seq_lens_ptr + b).to(tl.int64)
|
||||||
|
post = pre + al
|
||||||
|
cross = (pre // interval) != (post // interval)
|
||||||
|
tp = (post // interval) * interval
|
||||||
|
ti = tp - pre - 1
|
||||||
|
ti = tl.where(ti < 0, 0, ti)
|
||||||
|
cand = tl.load(accept_index_ptr + base + ti).to(tl.int64) - base
|
||||||
|
tl.store(track_steps_out_ptr + b, tl.where(cross, cand, -1))
|
||||||
|
|
||||||
|
|
||||||
|
def fused_commit_track_indices(
|
||||||
|
accept_index: torch.Tensor,
|
||||||
|
accept_lens: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor | None,
|
||||||
|
draft_token_num: int,
|
||||||
|
mamba_track_interval: int,
|
||||||
|
):
|
||||||
|
"""Single-launch replacement for the eager index math in
|
||||||
|
``commit_mamba_states_after_verify`` (index ranges, gathers, floordiv chain)."""
|
||||||
|
bs = accept_lens.shape[0]
|
||||||
|
last_correct_step_indices = torch.empty(
|
||||||
|
bs, dtype=torch.int64, device=accept_lens.device
|
||||||
|
)
|
||||||
|
has_track = seq_lens is not None
|
||||||
|
mamba_steps_to_track = (
|
||||||
|
torch.empty(bs, dtype=torch.int64, device=accept_lens.device)
|
||||||
|
if has_track
|
||||||
|
else last_correct_step_indices
|
||||||
|
)
|
||||||
|
_fused_commit_track_indices_kernel[(bs,)](
|
||||||
|
accept_index,
|
||||||
|
accept_lens,
|
||||||
|
seq_lens if has_track else accept_lens,
|
||||||
|
last_correct_step_indices,
|
||||||
|
mamba_steps_to_track,
|
||||||
|
draft_token_num,
|
||||||
|
mamba_track_interval,
|
||||||
|
HAS_TRACK=has_track,
|
||||||
|
)
|
||||||
|
return last_correct_step_indices, (mamba_steps_to_track if has_track else None)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def track_mamba_states_all_layers_kernel(
|
def track_mamba_states_all_layers_kernel(
|
||||||
conv_states_ptr, # [num_layers, pool_size, ...] full conv pool
|
conv_states_ptr, # [num_layers, pool_size, ...] full conv pool
|
||||||
|
|||||||
@@ -0,0 +1,126 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_kv_norm_rope_write_kernel(
|
||||||
|
kv_ptr,
|
||||||
|
meta_ptr,
|
||||||
|
knw_ptr,
|
||||||
|
cos_sin_ptr,
|
||||||
|
pos_ptr,
|
||||||
|
loc_ptr,
|
||||||
|
commit_lens_ptr,
|
||||||
|
locs_row_width,
|
||||||
|
KV: tl.constexpr,
|
||||||
|
D: tl.constexpr,
|
||||||
|
NH: tl.constexpr,
|
||||||
|
L: tl.constexpr,
|
||||||
|
EPS: tl.constexpr,
|
||||||
|
HAS_COMMIT_LENS: tl.constexpr,
|
||||||
|
):
|
||||||
|
t = tl.program_id(0).to(tl.int64)
|
||||||
|
l = tl.program_id(1).to(tl.int64)
|
||||||
|
if HAS_COMMIT_LENS:
|
||||||
|
row_b = t // locs_row_width
|
||||||
|
col = t - row_b * locs_row_width
|
||||||
|
num_commit = tl.load(commit_lens_ptr + row_b).to(tl.int64)
|
||||||
|
if col >= num_commit:
|
||||||
|
return
|
||||||
|
loc = tl.load(loc_ptr + t).to(tl.int64)
|
||||||
|
if loc < 0:
|
||||||
|
return
|
||||||
|
pos = tl.load(pos_ptr + t).to(tl.int64)
|
||||||
|
|
||||||
|
HALF: tl.constexpr = D // 2
|
||||||
|
half_ar = tl.arange(0, HALF)
|
||||||
|
d_ar = tl.arange(0, D)
|
||||||
|
cos = tl.load(cos_sin_ptr + pos * D + half_ar).to(tl.float32)
|
||||||
|
sin = tl.load(cos_sin_ptr + pos * D + HALF + half_ar).to(tl.float32)
|
||||||
|
knw1 = tl.load(knw_ptr + l * D + half_ar).to(tl.float32)
|
||||||
|
knw2 = tl.load(knw_ptr + l * D + HALF + half_ar).to(tl.float32)
|
||||||
|
|
||||||
|
k_buf = tl.load(meta_ptr + l * 4 + 0).to(tl.pointer_type(tl.bfloat16))
|
||||||
|
v_buf = tl.load(meta_ptr + l * 4 + 1).to(tl.pointer_type(tl.bfloat16))
|
||||||
|
ks0 = tl.load(meta_ptr + l * 4 + 2)
|
||||||
|
vs0 = tl.load(meta_ptr + l * 4 + 3)
|
||||||
|
|
||||||
|
row = kv_ptr + t * (L * 2 * KV) + l * (2 * KV)
|
||||||
|
for h in tl.static_range(NH):
|
||||||
|
k = tl.load(row + h * D + d_ar).to(tl.float32)
|
||||||
|
ms = tl.sum(k * k, 0) / D
|
||||||
|
inv = 1.0 / tl.sqrt(ms + EPS)
|
||||||
|
k1 = tl.load(row + h * D + half_ar).to(tl.float32) * inv * knw1
|
||||||
|
k2 = tl.load(row + h * D + HALF + half_ar).to(tl.float32) * inv * knw2
|
||||||
|
k1 = k1.to(tl.bfloat16).to(tl.float32)
|
||||||
|
k2 = k2.to(tl.bfloat16).to(tl.float32)
|
||||||
|
o1 = k1 * cos - k2 * sin
|
||||||
|
o2 = k2 * cos + k1 * sin
|
||||||
|
tl.store(k_buf + loc * ks0 + h * D + half_ar, o1.to(tl.bfloat16))
|
||||||
|
tl.store(k_buf + loc * ks0 + h * D + HALF + half_ar, o2.to(tl.bfloat16))
|
||||||
|
|
||||||
|
v = tl.load(row + KV + h * D + d_ar)
|
||||||
|
tl.store(v_buf + loc * vs0 + h * D + d_ar, v)
|
||||||
|
|
||||||
|
|
||||||
|
def fused_kv_norm_rope_write(
|
||||||
|
kv: torch.Tensor,
|
||||||
|
meta: torch.Tensor,
|
||||||
|
k_norm_weights: torch.Tensor,
|
||||||
|
cos_sin_cache: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
locs: torch.Tensor,
|
||||||
|
num_layers: int,
|
||||||
|
kv_size: int,
|
||||||
|
head_dim: int,
|
||||||
|
eps: float,
|
||||||
|
commit_lens: Optional[torch.Tensor] = None,
|
||||||
|
locs_row_width: Optional[int] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Write per-layer normed+roped K and raw V rows into the KV pools.
|
||||||
|
|
||||||
|
Rows with loc < 0 are skipped. When commit_lens is given, locs is the
|
||||||
|
flattened [bs, locs_row_width] verify window and only the first
|
||||||
|
commit_lens[b] columns of each row are written — the in-kernel
|
||||||
|
replacement for masking the tail columns to -1 on the host.
|
||||||
|
"""
|
||||||
|
T = kv.shape[0]
|
||||||
|
if T == 0:
|
||||||
|
return
|
||||||
|
has_commit_lens = commit_lens is not None
|
||||||
|
if has_commit_lens != (locs_row_width is not None):
|
||||||
|
raise ValueError(
|
||||||
|
"commit_lens and locs_row_width must be passed together, got "
|
||||||
|
f"commit_lens={'set' if has_commit_lens else None}, "
|
||||||
|
f"locs_row_width={locs_row_width}."
|
||||||
|
)
|
||||||
|
if has_commit_lens:
|
||||||
|
if commit_lens.numel() * locs_row_width != locs.numel():
|
||||||
|
raise ValueError(
|
||||||
|
f"locs must be a flattened [{commit_lens.numel()}, "
|
||||||
|
f"{locs_row_width}] window, got numel={locs.numel()}."
|
||||||
|
)
|
||||||
|
commit_lens_arg = commit_lens.contiguous()
|
||||||
|
else:
|
||||||
|
locs_row_width = 1
|
||||||
|
commit_lens_arg = locs
|
||||||
|
grid = (T, num_layers)
|
||||||
|
_fused_kv_norm_rope_write_kernel[grid](
|
||||||
|
kv,
|
||||||
|
meta,
|
||||||
|
k_norm_weights,
|
||||||
|
cos_sin_cache,
|
||||||
|
positions.to(torch.int64).contiguous(),
|
||||||
|
locs.to(torch.int64).contiguous(),
|
||||||
|
commit_lens_arg,
|
||||||
|
locs_row_width,
|
||||||
|
KV=kv_size,
|
||||||
|
D=head_dim,
|
||||||
|
NH=kv_size // head_dim,
|
||||||
|
L=num_layers,
|
||||||
|
EPS=eps,
|
||||||
|
HAS_COMMIT_LENS=has_commit_lens,
|
||||||
|
)
|
||||||
@@ -149,6 +149,13 @@ class DFlashAttention(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.scaling = head_dim**-0.5
|
self.scaling = head_dim**-0.5
|
||||||
|
rotary = self.rotary_emb
|
||||||
|
self.use_table_qk_norm_rope = (
|
||||||
|
not _is_npu
|
||||||
|
and hasattr(rotary, "cos_sin_cache")
|
||||||
|
and getattr(rotary, "rotary_dim", None) == head_dim
|
||||||
|
and getattr(rotary, "is_neox_style", False)
|
||||||
|
)
|
||||||
self.sliding_window_size, self.attn_type = _get_dflash_layer_attention_params(
|
self.sliding_window_size, self.attn_type = _get_dflash_layer_attention_params(
|
||||||
config, layer_id
|
config, layer_id
|
||||||
)
|
)
|
||||||
@@ -191,6 +198,21 @@ class DFlashAttention(nn.Module):
|
|||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
q, k, v = self.forward_prepare_npu(positions, hidden_states)
|
q, k, v = self.forward_prepare_npu(positions, hidden_states)
|
||||||
|
elif self.use_table_qk_norm_rope and qkv.dtype == torch.bfloat16:
|
||||||
|
from sglang.srt.speculative.dflash_utils import table_qk_norm_rope_
|
||||||
|
|
||||||
|
table_qk_norm_rope_(
|
||||||
|
qkv,
|
||||||
|
positions,
|
||||||
|
self.q_norm.weight,
|
||||||
|
self.k_norm.weight,
|
||||||
|
self.rotary_emb.cos_sin_cache,
|
||||||
|
self.num_heads,
|
||||||
|
self.num_kv_heads,
|
||||||
|
self.head_dim,
|
||||||
|
self.q_norm.variance_epsilon,
|
||||||
|
)
|
||||||
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
else:
|
else:
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
q, k = apply_qk_norm(q, k, self.q_norm, self.k_norm, self.head_dim)
|
q, k = apply_qk_norm(q, k, self.q_norm, self.k_norm, self.head_dim)
|
||||||
|
|||||||
@@ -364,6 +364,8 @@ class DSparkDraftMixin:
|
|||||||
|
|
||||||
def __init__(self, config, quant_config=None, prefix: str = "") -> None:
|
def __init__(self, config, quant_config=None, prefix: str = "") -> None:
|
||||||
super().__init__(config=config, quant_config=quant_config, prefix=prefix)
|
super().__init__(config=config, quant_config=quant_config, prefix=prefix)
|
||||||
|
self._fused_kv_write_cache = None
|
||||||
|
self.logits_mup_width_multiplier = None
|
||||||
dspark_config = parse_dspark_draft_config(draft_hf_config=config)
|
dspark_config = parse_dspark_draft_config(draft_hf_config=config)
|
||||||
if not dspark_config.require_markov():
|
if not dspark_config.require_markov():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -390,11 +392,22 @@ class DSparkDraftMixin:
|
|||||||
def compute_base_logits(
|
def compute_base_logits(
|
||||||
self, hidden: torch.Tensor
|
self, hidden: torch.Tensor
|
||||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||||
|
"""Project the draft's raw final hidden through the target lm_head.
|
||||||
|
|
||||||
|
muP targets (Inkling) train the draft against a FOLDED head (weights
|
||||||
|
pre-divided by logits_mup_width_multiplier) while serving attaches the
|
||||||
|
target's unfolded head, so the division happens here — exactly once,
|
||||||
|
keeping base logits in the scale the markov bias and confidence head
|
||||||
|
were trained against. DSparkWorkerV2 wires the multiplier from the
|
||||||
|
target config; it stays None for non-muP targets.
|
||||||
|
"""
|
||||||
if self.lm_head is None:
|
if self.lm_head is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"DSpark dense draft requires the target lm_head "
|
"DSpark dense draft requires the target lm_head "
|
||||||
"(call attach_shared_modules first)."
|
"(call attach_shared_modules first)."
|
||||||
)
|
)
|
||||||
|
if self.logits_mup_width_multiplier:
|
||||||
|
hidden = hidden / self.logits_mup_width_multiplier
|
||||||
weight = self.lm_head.weight
|
weight = self.lm_head.weight
|
||||||
if hidden.dtype != weight.dtype:
|
if hidden.dtype != weight.dtype:
|
||||||
hidden = hidden.to(weight.dtype)
|
hidden = hidden.to(weight.dtype)
|
||||||
@@ -466,6 +479,75 @@ class DSparkDraftMixin:
|
|||||||
f"or disable the confidence head (enable_confidence_head=False)."
|
f"or disable the confidence head (enable_confidence_head=False)."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _fused_kv_write_bundle(self, pool):
|
||||||
|
cached = self._fused_kv_write_cache
|
||||||
|
if cached is not None and cached[0] == id(pool):
|
||||||
|
return cached[1]
|
||||||
|
bundle = self._build_fused_kv_write_bundle(pool)
|
||||||
|
self._fused_kv_write_cache = (id(pool), bundle)
|
||||||
|
return bundle
|
||||||
|
|
||||||
|
def _build_fused_kv_write_bundle(self, pool):
|
||||||
|
layers = list(self.layers)
|
||||||
|
if not layers:
|
||||||
|
return None
|
||||||
|
if not (hasattr(pool, "get_key_buffer") and hasattr(pool, "get_value_buffer")):
|
||||||
|
return None
|
||||||
|
attn0 = layers[0].self_attn
|
||||||
|
head_dim = attn0.head_dim
|
||||||
|
kv_size = attn0.kv_size
|
||||||
|
rotary = attn0.rotary_emb
|
||||||
|
if type(rotary).__name__ != "RotaryEmbedding":
|
||||||
|
return None
|
||||||
|
if not getattr(rotary, "is_neox_style", False):
|
||||||
|
return None
|
||||||
|
if getattr(rotary, "rotary_dim", None) != head_dim:
|
||||||
|
return None
|
||||||
|
eps = attn0.k_norm.variance_epsilon
|
||||||
|
weights, knws, meta_rows = [], [], []
|
||||||
|
for layer in layers:
|
||||||
|
attn = layer.self_attn
|
||||||
|
ok, _ = can_dflash_slice_qkv_weight(attn.qkv_proj)
|
||||||
|
if not ok:
|
||||||
|
return None
|
||||||
|
if attn.qkv_proj.bias is not None:
|
||||||
|
return None
|
||||||
|
if attn.attn.k_scale is not None or attn.attn.v_scale is not None:
|
||||||
|
return None
|
||||||
|
if attn.head_dim != head_dim or attn.kv_size != kv_size:
|
||||||
|
return None
|
||||||
|
if attn.rotary_emb is not rotary and not torch.equal(
|
||||||
|
attn.rotary_emb.cos_sin_cache, rotary.cos_sin_cache
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
if attn.k_norm.variance_epsilon != eps:
|
||||||
|
return None
|
||||||
|
k_buf = pool.get_key_buffer(attn.attn.layer_id)
|
||||||
|
v_buf = pool.get_value_buffer(attn.attn.layer_id)
|
||||||
|
nh = kv_size // head_dim
|
||||||
|
for buf in (k_buf, v_buf):
|
||||||
|
if buf.dtype != torch.bfloat16:
|
||||||
|
return None
|
||||||
|
if buf.shape[1:] != (nh, head_dim):
|
||||||
|
return None
|
||||||
|
if buf.stride(1) != head_dim or buf.stride(2) != 1:
|
||||||
|
return None
|
||||||
|
kv_slice = slice(attn.q_size, attn.q_size + 2 * attn.kv_size)
|
||||||
|
w = attn.qkv_proj.weight[kv_slice]
|
||||||
|
if w.dtype != torch.bfloat16:
|
||||||
|
return None
|
||||||
|
weights.append(w)
|
||||||
|
knws.append(attn.k_norm.weight.data)
|
||||||
|
meta_rows.append(
|
||||||
|
[k_buf.data_ptr(), v_buf.data_ptr(), k_buf.stride(0), v_buf.stride(0)]
|
||||||
|
)
|
||||||
|
device = weights[0].device
|
||||||
|
w_all = torch.cat(weights, dim=0).contiguous()
|
||||||
|
knw = torch.stack(knws).to(device)
|
||||||
|
meta = torch.tensor(meta_rows, dtype=torch.int64, device=device)
|
||||||
|
cos_sin = rotary.cos_sin_cache.to(device)
|
||||||
|
return (w_all, meta, knw, cos_sin, eps, len(layers), kv_size, head_dim)
|
||||||
|
|
||||||
def _stacked_ctx_kv_params(self) -> Optional[dict]:
|
def _stacked_ctx_kv_params(self) -> Optional[dict]:
|
||||||
"""Stack every layer's KV projection into one weight (exact: the input
|
"""Stack every layer's KV projection into one weight (exact: the input
|
||||||
hidden is shared, so concatenating output columns is equivalent).
|
hidden is shared, so concatenating output columns is equivalent).
|
||||||
@@ -513,6 +595,39 @@ class DSparkDraftMixin:
|
|||||||
commit_lens: Optional[torch.Tensor] = None,
|
commit_lens: Optional[torch.Tensor] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
ctx_hidden = self.project_target_hidden(target_hidden)
|
ctx_hidden = self.project_target_hidden(target_hidden)
|
||||||
|
|
||||||
|
bundle = self._fused_kv_write_bundle(pool)
|
||||||
|
if bundle is not None:
|
||||||
|
from sglang.kernels.ops.speculative.dspark.fused_kv_write import (
|
||||||
|
fused_kv_norm_rope_write,
|
||||||
|
)
|
||||||
|
|
||||||
|
w_all, meta, knw, cos_sin, eps, num_layers, kv_size, head_dim = bundle
|
||||||
|
kv_all = F.linear(ctx_hidden, w_all)
|
||||||
|
if cache_loc_2d is not None and commit_lens is not None:
|
||||||
|
locs = cache_loc_2d.reshape(-1)
|
||||||
|
write_commit_lens = commit_lens
|
||||||
|
locs_row_width = cache_loc_2d.shape[1]
|
||||||
|
else:
|
||||||
|
locs = cache_loc
|
||||||
|
write_commit_lens = None
|
||||||
|
locs_row_width = None
|
||||||
|
fused_kv_norm_rope_write(
|
||||||
|
kv_all,
|
||||||
|
meta,
|
||||||
|
knw,
|
||||||
|
cos_sin,
|
||||||
|
positions,
|
||||||
|
locs,
|
||||||
|
num_layers,
|
||||||
|
kv_size,
|
||||||
|
head_dim,
|
||||||
|
eps,
|
||||||
|
commit_lens=write_commit_lens,
|
||||||
|
locs_row_width=locs_row_width,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
stacked = self._stacked_ctx_kv_params()
|
stacked = self._stacked_ctx_kv_params()
|
||||||
if stacked is not None:
|
if stacked is not None:
|
||||||
k_all, v_all = self._project_ctx_kv_stacked(
|
k_all, v_all = self._project_ctx_kv_stacked(
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ from typing import Any, List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
||||||
from sglang.srt.layers.sampler import apply_custom_logit_processor
|
from sglang.srt.layers.sampler import apply_custom_logit_processor
|
||||||
@@ -570,6 +572,29 @@ def can_dflash_use_fused_qkv_proj(qkv_proj: Any) -> Tuple[bool, str]:
|
|||||||
return True, ""
|
return True, ""
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_correct_drafts_and_bonus_kernel(
|
||||||
|
candidates_ptr,
|
||||||
|
target_predict_ptr,
|
||||||
|
num_correct_drafts_ptr,
|
||||||
|
bonus_tokens_ptr,
|
||||||
|
block_size,
|
||||||
|
BLOCK: tl.constexpr,
|
||||||
|
):
|
||||||
|
b = tl.program_id(0).to(tl.int64)
|
||||||
|
offs = tl.arange(0, BLOCK)
|
||||||
|
in_row = offs < block_size - 1
|
||||||
|
drafts = tl.load(candidates_ptr + b * block_size + 1 + offs, mask=in_row, other=-1)
|
||||||
|
targets = tl.load(target_predict_ptr + b * block_size + offs, mask=in_row, other=-2)
|
||||||
|
eq = (drafts == targets) & in_row
|
||||||
|
# Leading-match count = index of the first mismatch lane; lanes past the
|
||||||
|
# row and all-match rows both resolve to block_size - 1 via the min.
|
||||||
|
num_correct = tl.min(tl.where(eq, BLOCK, offs), 0)
|
||||||
|
bonus_token = tl.load(target_predict_ptr + b * block_size + num_correct)
|
||||||
|
tl.store(num_correct_drafts_ptr + b, num_correct.to(tl.int32))
|
||||||
|
tl.store(bonus_tokens_ptr + b, bonus_token.to(tl.int64))
|
||||||
|
|
||||||
|
|
||||||
def compute_dflash_correct_drafts_and_bonus(
|
def compute_dflash_correct_drafts_and_bonus(
|
||||||
*,
|
*,
|
||||||
candidates: torch.Tensor,
|
candidates: torch.Tensor,
|
||||||
@@ -605,10 +630,25 @@ def compute_dflash_correct_drafts_and_bonus(
|
|||||||
if block_size <= 0:
|
if block_size <= 0:
|
||||||
raise ValueError(f"block_size must be positive, got {block_size}.")
|
raise ValueError(f"block_size must be positive, got {block_size}.")
|
||||||
|
|
||||||
|
if candidates.is_cuda:
|
||||||
|
num_correct_drafts = torch.empty(
|
||||||
|
bs, dtype=torch.int32, device=candidates.device
|
||||||
|
)
|
||||||
|
bonus_tokens = torch.empty(bs, dtype=torch.int64, device=candidates.device)
|
||||||
|
_fused_correct_drafts_and_bonus_kernel[(bs,)](
|
||||||
|
candidates.contiguous(),
|
||||||
|
target_predict.contiguous(),
|
||||||
|
num_correct_drafts,
|
||||||
|
bonus_tokens,
|
||||||
|
block_size,
|
||||||
|
BLOCK=triton.next_power_of_2(max(block_size - 1, 1)),
|
||||||
|
)
|
||||||
|
return num_correct_drafts, bonus_tokens
|
||||||
|
|
||||||
matches = candidates[:, 1:] == target_predict[:, :-1]
|
matches = candidates[:, 1:] == target_predict[:, :-1]
|
||||||
correct_len = matches.to(torch.int32).cumprod(dim=1).sum(dim=1)
|
correct_len = matches.to(torch.int32).cumprod(dim=1).sum(dim=1)
|
||||||
bonus = target_predict[torch.arange(bs, device=target_predict.device), correct_len]
|
bonus = target_predict[torch.arange(bs, device=target_predict.device), correct_len]
|
||||||
return correct_len, bonus.to(torch.int64)
|
return correct_len.to(torch.int32), bonus.to(torch.int64)
|
||||||
|
|
||||||
|
|
||||||
def apply_dflash_simulated_acceptance(
|
def apply_dflash_simulated_acceptance(
|
||||||
@@ -861,3 +901,83 @@ def validate_dflash_request(req: Req, enable_overlap: bool) -> Optional[str]:
|
|||||||
return "DFLASH speculative decoding does not support return_hidden_states yet."
|
return "DFLASH speculative decoding does not support return_hidden_states yet."
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _table_qk_norm_rope_kernel(
|
||||||
|
qkv_ptr,
|
||||||
|
q_weight_ptr,
|
||||||
|
k_weight_ptr,
|
||||||
|
cos_sin_ptr,
|
||||||
|
pos_ptr,
|
||||||
|
row_stride,
|
||||||
|
q_size,
|
||||||
|
NHQ: tl.constexpr,
|
||||||
|
D: tl.constexpr,
|
||||||
|
EPS: tl.constexpr,
|
||||||
|
):
|
||||||
|
t = tl.program_id(0).to(tl.int64)
|
||||||
|
h = tl.program_id(1)
|
||||||
|
pos = tl.load(pos_ptr + t).to(tl.int64)
|
||||||
|
|
||||||
|
HALF: tl.constexpr = D // 2
|
||||||
|
half_ar = tl.arange(0, HALF)
|
||||||
|
d_ar = tl.arange(0, D)
|
||||||
|
cos = tl.load(cos_sin_ptr + pos * D + half_ar).to(tl.float32)
|
||||||
|
sin = tl.load(cos_sin_ptr + pos * D + HALF + half_ar).to(tl.float32)
|
||||||
|
|
||||||
|
is_q = h < NHQ
|
||||||
|
col0 = tl.where(is_q, h * D, q_size + (h - NHQ) * D).to(tl.int64)
|
||||||
|
w_ptr = tl.where(is_q, q_weight_ptr.to(tl.int64), k_weight_ptr.to(tl.int64)).to(
|
||||||
|
tl.pointer_type(tl.bfloat16)
|
||||||
|
)
|
||||||
|
|
||||||
|
row = qkv_ptr + t * row_stride + col0
|
||||||
|
x = tl.load(row + d_ar).to(tl.float32)
|
||||||
|
ms = tl.sum(x * x, 0) / D
|
||||||
|
inv = 1.0 / tl.sqrt(ms + EPS)
|
||||||
|
w1 = tl.load(w_ptr + half_ar).to(tl.float32)
|
||||||
|
w2 = tl.load(w_ptr + HALF + half_ar).to(tl.float32)
|
||||||
|
x1 = tl.load(row + half_ar).to(tl.float32) * inv * w1
|
||||||
|
x2 = tl.load(row + HALF + half_ar).to(tl.float32) * inv * w2
|
||||||
|
x1 = x1.to(tl.bfloat16).to(tl.float32)
|
||||||
|
x2 = x2.to(tl.bfloat16).to(tl.float32)
|
||||||
|
o1 = x1 * cos - x2 * sin
|
||||||
|
o2 = x2 * cos + x1 * sin
|
||||||
|
tl.store(row + half_ar, o1.to(tl.bfloat16))
|
||||||
|
tl.store(row + HALF + half_ar, o2.to(tl.bfloat16))
|
||||||
|
|
||||||
|
|
||||||
|
def table_qk_norm_rope_(
|
||||||
|
qkv: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
q_weight: torch.Tensor,
|
||||||
|
k_weight: torch.Tensor,
|
||||||
|
cos_sin_cache: torch.Tensor,
|
||||||
|
num_q_heads: int,
|
||||||
|
num_k_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
eps: float,
|
||||||
|
) -> None:
|
||||||
|
"""In-place QK RMSNorm + table-lookup neox RoPE on the fused QKV tensor.
|
||||||
|
|
||||||
|
Reads cos/sin from the SAME rotary table as the unfused path, so there is
|
||||||
|
no large-position angle drift (unlike theta-recompute kernels). V columns
|
||||||
|
are untouched.
|
||||||
|
"""
|
||||||
|
T = qkv.shape[0]
|
||||||
|
if T == 0:
|
||||||
|
return
|
||||||
|
grid = (T, num_q_heads + num_k_heads)
|
||||||
|
_table_qk_norm_rope_kernel[grid](
|
||||||
|
qkv,
|
||||||
|
q_weight,
|
||||||
|
k_weight,
|
||||||
|
cos_sin_cache,
|
||||||
|
positions,
|
||||||
|
qkv.stride(0),
|
||||||
|
num_q_heads * head_dim,
|
||||||
|
NHQ=num_q_heads,
|
||||||
|
D=head_dim,
|
||||||
|
EPS=eps,
|
||||||
|
)
|
||||||
|
|||||||
@@ -126,19 +126,27 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
self.draft_model_runner = bundle.draft_model_runner
|
self.draft_model_runner = bundle.draft_model_runner
|
||||||
self.draft_model = bundle.draft_model
|
self.draft_model = bundle.draft_model
|
||||||
self._draft_sampler = None
|
self._draft_sampler = None
|
||||||
|
self._linear_accept_index_cache = None
|
||||||
|
|
||||||
# The mask token needs an embedding row, not a tokenizer entry, so bound it
|
# The mask token is input-only (it is embedded, never sampled), so its
|
||||||
# by the embedding width. A padded vocab reserves rows past the real tokens
|
# bound is the embedding-table row count: the PADDED vocab when the
|
||||||
# and drafts place the mask there (Inkling: 200058 real, 201024 padded).
|
# target pads its embedding (e.g. Inkling true vocab 200058, padded
|
||||||
|
# 201024, mask 200064), else the plain vocab size.
|
||||||
target_model_config = self.target_worker.model_runner.model_config
|
target_model_config = self.target_worker.model_runner.model_config
|
||||||
target_vocab_size = (
|
target_embed_rows = (
|
||||||
getattr(target_model_config.hf_text_config, "padded_vocab_size", None)
|
getattr(target_model_config.hf_text_config, "padded_vocab_size", None)
|
||||||
or target_model_config.vocab_size
|
or target_model_config.vocab_size
|
||||||
)
|
)
|
||||||
|
# muP targets declare logits_mup_width_multiplier; the draft was
|
||||||
|
# trained against the folded head, so compute_base_logits divides.
|
||||||
|
self.draft_model.logits_mup_width_multiplier = getattr(
|
||||||
|
target_model_config.hf_text_config, "logits_mup_width_multiplier", None
|
||||||
|
)
|
||||||
|
self._target_is_mambaish = mambaish_config(target_model_config) is not None
|
||||||
runtime_config = resolve_runtime_config(
|
runtime_config = resolve_runtime_config(
|
||||||
draft_hf_config=self.draft_model_runner.model_config.hf_config,
|
draft_hf_config=self.draft_model_runner.model_config.hf_config,
|
||||||
speculative_num_draft_tokens=server_args.speculative_num_draft_tokens,
|
speculative_num_draft_tokens=server_args.speculative_num_draft_tokens,
|
||||||
target_vocab_size=int(target_vocab_size),
|
target_vocab_size=int(target_embed_rows),
|
||||||
)
|
)
|
||||||
self.gamma = runtime_config.gamma
|
self.gamma = runtime_config.gamma
|
||||||
self.verify_num_draft_tokens = runtime_config.verify_num_draft_tokens
|
self.verify_num_draft_tokens = runtime_config.verify_num_draft_tokens
|
||||||
@@ -806,6 +814,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
mamba_track_indices=batch.mamba_track_indices,
|
mamba_track_indices=batch.mamba_track_indices,
|
||||||
mamba_steps_to_track=mamba_steps_to_track,
|
mamba_steps_to_track=mamba_steps_to_track,
|
||||||
model=self.target_worker.model_runner.model,
|
model=self.target_worker.model_runner.model,
|
||||||
|
req_pool_indices=batch.req_pool_indices,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_confidence_budget_prepare(self):
|
def get_confidence_budget_prepare(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user