qwen3_5 flex: CUDA graph capture for batched decode
Add static per-layer conv/recurrent state buffers for the Gated DeltaNet sub-layers and a batched decode path that gather/scatters through them. The step is captured into a CUDA graph per bucket size; generate() replays the matching graph when the active batch size exactly matches a captured bucket, else falls back to the eager batched path. Throughput (Qwen3.5-4B, 128 new tokens, post-warmup, B200): | bs | pre-capture tok/s | post-capture tok/s | speedup | |---:|---:|---:|---:| | 1 | 20.1 | 191.2 | 9.5x | | 2 | 20.5 | 338.1 | 16x (extrapolated) | | 4 | 20.9 | 548.6 | 26x | bs=8 is currently a regression (143 tok/s vs 549 at bs=4) — likely FLA kernel sub-linear scaling in B. Tracked as follow-up. Parity (flex-capture vs flex-eager-batched, bs=4 x 16 tokens greedy): 3/4 prompts token-exact, 1 prompt matches 13/16 tokens (bf16 drift).
This commit is contained in:
parent
6abb05fbfc
commit
0058a5fd45
1 changed files with 401 additions and 7 deletions
|
|
@ -204,6 +204,148 @@ def _resolve_text_model(model: torch.nn.Module):
|
|||
)
|
||||
|
||||
|
||||
def _torch_causal_conv1d_update(
|
||||
hidden_states: torch.Tensor,
|
||||
conv_state: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""CUDA-graph-safe causal conv1d update step for Qwen3.5 DeltaNet.
|
||||
|
||||
Mirrors ``transformers.models.qwen3_5.torch_causal_conv1d_update``:
|
||||
rolls ``conv_state`` forward by one step in place, convolves the
|
||||
(state || hidden_states) window with the depthwise conv1d weight, then
|
||||
applies SiLU. Used in place of the ``causal_conv1d_update`` kernel which
|
||||
writes via a CUDA memcpy that cudagraph-capture refuses on some inputs.
|
||||
"""
|
||||
_, hidden_size, seq_len = hidden_states.shape
|
||||
state_len = conv_state.shape[-1]
|
||||
hs_ext = torch.cat([conv_state, hidden_states], dim=-1).to(weight.dtype)
|
||||
conv_state.copy_(hs_ext[:, :, -state_len:])
|
||||
out = F.conv1d(
|
||||
hs_ext, weight.unsqueeze(1), bias, padding=0, groups=hidden_size,
|
||||
)
|
||||
out = F.silu(out[:, :, -seq_len:])
|
||||
return out.to(hidden_states.dtype)
|
||||
|
||||
|
||||
def _deltanet_decode_step_static(
|
||||
layer,
|
||||
hidden_states: torch.Tensor,
|
||||
conv_state_gathered: torch.Tensor,
|
||||
recurrent_state_gathered: torch.Tensor,
|
||||
):
|
||||
"""Single-token decode for ``Qwen3_5GatedDeltaNet`` with caller-owned
|
||||
state tensors (no ``DynamicCache``). ``conv_state_gathered`` is updated
|
||||
in place by :func:`_torch_causal_conv1d_update`; the recurrent state
|
||||
new value is freshly allocated and returned for the caller to scatter.
|
||||
|
||||
Shapes:
|
||||
``hidden_states`` [B, 1, H]
|
||||
``conv_state_gathered`` [B, conv_dim, K] (bf16)
|
||||
``recurrent_state_gathered`` [B, HV, Hk, Hv] (fp32)
|
||||
Returns: output [B, 1, H], updated_conv_state, new_recurrent_state.
|
||||
"""
|
||||
B, S, _ = hidden_states.shape
|
||||
mixed_qkv = layer.in_proj_qkv(hidden_states).transpose(1, 2) # [B, conv_dim, 1]
|
||||
z = layer.in_proj_z(hidden_states).reshape(B, S, -1, layer.head_v_dim)
|
||||
b = layer.in_proj_b(hidden_states)
|
||||
a = layer.in_proj_a(hidden_states)
|
||||
|
||||
mixed_qkv = _torch_causal_conv1d_update(
|
||||
mixed_qkv,
|
||||
conv_state_gathered,
|
||||
layer.conv1d.weight.squeeze(1),
|
||||
layer.conv1d.bias,
|
||||
)
|
||||
|
||||
mixed_qkv = mixed_qkv.transpose(1, 2) # [B, 1, conv_dim]
|
||||
query, key, value = torch.split(
|
||||
mixed_qkv,
|
||||
[layer.key_dim, layer.key_dim, layer.value_dim],
|
||||
dim=-1,
|
||||
)
|
||||
query = query.reshape(B, S, -1, layer.head_k_dim)
|
||||
key = key.reshape(B, S, -1, layer.head_k_dim)
|
||||
value = value.reshape(B, S, -1, layer.head_v_dim)
|
||||
|
||||
beta = b.sigmoid()
|
||||
g = -layer.A_log.float().exp() * F.softplus(a.float() + layer.dt_bias)
|
||||
r = layer.num_v_heads // layer.num_k_heads
|
||||
if r > 1:
|
||||
query = query.repeat_interleave(r, dim=2)
|
||||
key = key.repeat_interleave(r, dim=2)
|
||||
|
||||
core_attn_out, last_recurrent_state = layer.recurrent_gated_delta_rule(
|
||||
query, key, value, g=g, beta=beta,
|
||||
initial_state=recurrent_state_gathered,
|
||||
output_final_state=True,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
|
||||
core_attn_out = core_attn_out.reshape(-1, layer.head_v_dim)
|
||||
z = z.reshape(-1, layer.head_v_dim)
|
||||
core_attn_out = layer.norm(core_attn_out, z)
|
||||
core_attn_out = core_attn_out.reshape(B, S, -1)
|
||||
output = layer.out_proj(core_attn_out)
|
||||
return output, conv_state_gathered, last_recurrent_state
|
||||
|
||||
|
||||
def _call_qwen3_5_decode_static(
|
||||
text_model,
|
||||
input_ids: torch.Tensor,
|
||||
position_ids: torch.Tensor,
|
||||
full_attn_flex_kwargs: dict,
|
||||
batch_idx: torch.Tensor,
|
||||
linear_conv_states: dict,
|
||||
linear_recurrent_states: dict,
|
||||
*,
|
||||
lm_head_fn,
|
||||
):
|
||||
"""Batched decode walker using static per-layer state buffers.
|
||||
|
||||
Linear layers ``index_select`` → deltanet-decode → ``index_copy_`` back.
|
||||
Full-attention layers use the pre-patched flex+paged KV forward.
|
||||
"""
|
||||
cfg = text_model.config
|
||||
inputs_embeds = text_model.embed_tokens(input_ids)
|
||||
position_embeddings = text_model.rotary_emb(inputs_embeds, position_ids)
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
layer_types = cfg.layer_types
|
||||
for layer_idx, layer in enumerate(text_model.layers):
|
||||
residual = hidden_states
|
||||
hidden_states = layer.input_layernorm(hidden_states)
|
||||
|
||||
if layer_types[layer_idx] == "linear_attention":
|
||||
conv_buf = linear_conv_states[layer_idx]
|
||||
rec_buf = linear_recurrent_states[layer_idx]
|
||||
conv_gath = conv_buf.index_select(0, batch_idx)
|
||||
rec_gath = rec_buf.index_select(0, batch_idx)
|
||||
out, conv_new, rec_new = _deltanet_decode_step_static(
|
||||
layer.linear_attn, hidden_states, conv_gath, rec_gath,
|
||||
)
|
||||
conv_buf.index_copy_(0, batch_idx, conv_new)
|
||||
rec_buf.index_copy_(0, batch_idx, rec_new)
|
||||
hidden_states = out
|
||||
else:
|
||||
hidden_states, _ = layer.self_attn(
|
||||
hidden_states,
|
||||
position_embeddings=position_embeddings,
|
||||
**full_attn_flex_kwargs,
|
||||
)
|
||||
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = layer.post_attention_layernorm(hidden_states)
|
||||
hidden_states = layer.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
hidden_states = text_model.norm(hidden_states)
|
||||
return lm_head_fn(hidden_states)
|
||||
|
||||
|
||||
def _call_qwen3_5_with_flex_kwargs(
|
||||
text_model,
|
||||
input_ids: torch.Tensor,
|
||||
|
|
@ -383,6 +525,46 @@ class FlexQwen3_5Inference:
|
|||
self._build_linear_cache() for _ in range(max_batch_size)
|
||||
]
|
||||
|
||||
# Static state buffers for CUDA-graph-captured batched decode.
|
||||
# We keep one conv_state tensor and one recurrent_state tensor per
|
||||
# linear_attention layer, sized ``[max_batch_size, ...]``. Prefill
|
||||
# still writes into per-seq ``DynamicCache`` objects; after
|
||||
# prefill, :meth:`_sync_static_from_dynamic_cache` copies the
|
||||
# prefilled state into the right slot of these buffers, and the
|
||||
# batched decode step gather/scatters through them.
|
||||
cfg = self.text_model.config
|
||||
self._linear_layer_indices = [
|
||||
i for i, t in enumerate(cfg.layer_types) if t == "linear_attention"
|
||||
]
|
||||
self._linear_conv_states = {}
|
||||
self._linear_recurrent_states = {}
|
||||
for layer_idx in self._linear_layer_indices:
|
||||
la = self.text_model.layers[layer_idx].linear_attn
|
||||
self._linear_conv_states[layer_idx] = torch.zeros(
|
||||
max_batch_size, la.conv_dim, la.conv_kernel_size,
|
||||
dtype=la.conv1d.weight.dtype, device=self.device,
|
||||
)
|
||||
self._linear_recurrent_states[layer_idx] = torch.zeros(
|
||||
max_batch_size, la.num_v_heads, la.head_k_dim, la.head_v_dim,
|
||||
dtype=torch.float32, device=self.device,
|
||||
)
|
||||
try:
|
||||
torch._dynamo.mark_static_address(
|
||||
self._linear_conv_states[layer_idx],
|
||||
)
|
||||
torch._dynamo.mark_static_address(
|
||||
self._linear_recurrent_states[layer_idx],
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# CUDA-graph state. Populated by :meth:`capture_decode_cudagraph`
|
||||
# on first decode when ``capture_cudagraph=True`` is passed to
|
||||
# :meth:`generate`.
|
||||
self.graphs = {}
|
||||
self.graph_vars = None
|
||||
self._captured = False
|
||||
|
||||
# ----- cache helpers -------------------------------------------------
|
||||
def _build_linear_cache(self):
|
||||
"""Allocate a ``DynamicCache`` seeded from ``model.config`` so
|
||||
|
|
@ -566,14 +748,223 @@ class FlexQwen3_5Inference:
|
|||
"""Fresh linear-attn cache for a newly-scheduled sequence."""
|
||||
seq._linear_cache = self._build_linear_cache()
|
||||
|
||||
# ----- static-state decode ------------------------------------------
|
||||
def _sync_static_from_dynamic_cache(self, seq: Sequence):
|
||||
"""Copy prefilled conv/recurrent state from ``seq._linear_cache``
|
||||
into the engine's static per-layer buffers at ``seq.batch_idx``.
|
||||
|
||||
The HF prefill path writes into each seq's ``DynamicCache``; to
|
||||
hand off to the batched/captured decode we copy those tensors into
|
||||
the fixed slot in the static buffer."""
|
||||
slot = seq.batch_idx
|
||||
cache = seq._linear_cache
|
||||
for layer_idx in self._linear_layer_indices:
|
||||
layer_cache = cache.layers[layer_idx]
|
||||
conv = getattr(layer_cache, "conv_states", None)
|
||||
rec = getattr(layer_cache, "recurrent_states", None)
|
||||
if conv is not None and conv.numel() > 0:
|
||||
self._linear_conv_states[layer_idx][slot].copy_(conv[0])
|
||||
if rec is not None and rec.numel() > 0:
|
||||
self._linear_recurrent_states[layer_idx][slot].copy_(rec[0])
|
||||
|
||||
def _reset_static_state(self, batch_idx: int):
|
||||
"""Zero the static state for a slot (e.g. when a seq finishes)."""
|
||||
for layer_idx in self._linear_layer_indices:
|
||||
self._linear_conv_states[layer_idx][batch_idx].zero_()
|
||||
self._linear_recurrent_states[layer_idx][batch_idx].zero_()
|
||||
|
||||
@torch.inference_mode()
|
||||
def _decode_step_batched_static(
|
||||
self, batch_idx: torch.Tensor, input_ids: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Eager batched decode using the engine's static state buffers.
|
||||
|
||||
Walks every transformer layer once on the batch: full-attn layers
|
||||
go through the paged KV + flex attention forward; linear layers
|
||||
gather/scatter through static buffers. Returns logits ``[B, V]``
|
||||
at the last position.
|
||||
"""
|
||||
B = input_ids.shape[0]
|
||||
mask, input_pos = self._decode_block_mask(batch_idx)
|
||||
mask = self.page_table.convert_logical_block_mask(mask, batch_idx)
|
||||
ids = input_ids.view(B, 1)
|
||||
pos = input_pos.view(B, 1).to(torch.long)
|
||||
full_attn_kwargs = dict(
|
||||
flex_block_mask=mask,
|
||||
flex_input_pos=pos,
|
||||
flex_batch_idx=batch_idx,
|
||||
flex_kernel_options=self.decode_kernel_options,
|
||||
)
|
||||
logits = _call_qwen3_5_decode_static(
|
||||
self.text_model,
|
||||
ids,
|
||||
pos,
|
||||
full_attn_kwargs,
|
||||
batch_idx,
|
||||
self._linear_conv_states,
|
||||
self._linear_recurrent_states,
|
||||
lm_head_fn=self._lm_head,
|
||||
)
|
||||
return logits[:, -1, :]
|
||||
|
||||
def capture_decode_cudagraph(self):
|
||||
"""Capture decode-step CUDA graphs across a bucket ladder.
|
||||
|
||||
Reserves dummy page-table slots so the paged KV cache machinery
|
||||
has somewhere to write during the capture warmups, then captures
|
||||
one graph per bucket size sharing a single memory pool. Buckets
|
||||
are ``[1, 2, 4, 8, 16, 32, ...]`` capped at ``max_batch_size``.
|
||||
|
||||
After capture the engine zeros the static conv/recurrent states
|
||||
and releases the dummy page reservations so generation starts
|
||||
clean.
|
||||
"""
|
||||
max_bs = self.max_batch_size
|
||||
# Temporarily reserve page-table slots 0..N so the decode block
|
||||
# mask builder has state to index into. This only matters at
|
||||
# capture time — we release these after.
|
||||
reserved = []
|
||||
for _ in range(max_bs):
|
||||
try:
|
||||
bi = self.page_table.allocate()
|
||||
self.page_table.reserve(
|
||||
bi,
|
||||
torch.tensor([bi], device=self.device, dtype=torch.long),
|
||||
self.page_size,
|
||||
)
|
||||
reserved.append(bi)
|
||||
except Exception:
|
||||
break
|
||||
if not reserved:
|
||||
raise RuntimeError(
|
||||
"capture_decode_cudagraph: could not allocate any "
|
||||
"page-table slots for capture.",
|
||||
)
|
||||
|
||||
# Static input + output buffers.
|
||||
input_ids_buf = torch.zeros(
|
||||
max_bs, dtype=torch.int64, device=self.device,
|
||||
)
|
||||
# Use reserved slot ids as the default indirection; replay time
|
||||
# overrides with ``.copy_``.
|
||||
bi_init = reserved + list(range(len(reserved), max_bs))
|
||||
batch_idx_buf = torch.tensor(
|
||||
bi_init[:max_bs], dtype=torch.int64, device=self.device,
|
||||
)
|
||||
vocab_size = getattr(
|
||||
self.text_model.config, "vocab_size", None,
|
||||
) or getattr(self.model.config, "vocab_size", None)
|
||||
if vocab_size is None:
|
||||
# Multimodal wrapper: vocab_size lives on text_config.
|
||||
vocab_size = self.model.config.text_config.vocab_size
|
||||
outputs_buf = torch.zeros(
|
||||
(max_bs, vocab_size),
|
||||
dtype=self.model.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
try:
|
||||
torch._dynamo.mark_static_address(input_ids_buf)
|
||||
torch._dynamo.mark_static_address(batch_idx_buf)
|
||||
torch._dynamo.mark_static_address(outputs_buf)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Bucket ladder. Start small, expand to max_bs.
|
||||
ladder = [1, 2, 4, 8] + list(range(16, max_bs + 1, 16))
|
||||
ladder = sorted(set(bs for bs in ladder if bs <= max_bs))
|
||||
self.graph_bs = ladder
|
||||
|
||||
pool = None
|
||||
for bs in reversed(ladder):
|
||||
torch.cuda.synchronize()
|
||||
# Warmup eager run seeds the flex_attention compiled cache
|
||||
# for this bs and populates any autograd-off kernels.
|
||||
_ = self._decode_step_batched_static(
|
||||
batch_idx_buf[:bs], input_ids_buf[:bs],
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
if pool is None:
|
||||
with torch.cuda.graph(graph):
|
||||
out = self._decode_step_batched_static(
|
||||
batch_idx_buf[:bs], input_ids_buf[:bs],
|
||||
)
|
||||
outputs_buf[:bs].copy_(out)
|
||||
pool = graph.pool()
|
||||
else:
|
||||
with torch.cuda.graph(graph, pool=pool):
|
||||
out = self._decode_step_batched_static(
|
||||
batch_idx_buf[:bs], input_ids_buf[:bs],
|
||||
)
|
||||
outputs_buf[:bs].copy_(out)
|
||||
self.graphs[bs] = graph
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Release capture reservations + wipe polluted static state.
|
||||
for bi in reserved:
|
||||
self.page_table.erase(bi)
|
||||
for layer_idx in self._linear_layer_indices:
|
||||
self._linear_conv_states[layer_idx].zero_()
|
||||
self._linear_recurrent_states[layer_idx].zero_()
|
||||
|
||||
self.graph_vars = dict(
|
||||
input_ids=input_ids_buf,
|
||||
batch_idx=batch_idx_buf,
|
||||
outputs=outputs_buf,
|
||||
)
|
||||
self._captured = True
|
||||
|
||||
def _pick_bucket(self, B: int) -> Optional[int]:
|
||||
"""Exact-match bucket lookup. Padding the batch with duplicate
|
||||
``batch_idx`` entries breaks the ``index_copy_`` scatter into the
|
||||
static state buffers (duplicate indices make the last write win,
|
||||
so the primary slot's conv/recurrent state gets overwritten by
|
||||
the padding slot's advance). Requiring exact match keeps the
|
||||
captured replay correct and falls back to the eager batched
|
||||
decode for sizes without a dedicated graph."""
|
||||
if not self._captured:
|
||||
return None
|
||||
if B in self.graphs:
|
||||
return B
|
||||
return None
|
||||
|
||||
@torch.inference_mode()
|
||||
def _decode_step(
|
||||
self, batch_idx: torch.Tensor, input_ids: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Decode dispatch: captured graph replay when the active batch
|
||||
size exactly matches a captured bucket, otherwise eager batched
|
||||
decode. See :meth:`_pick_bucket` for why we require an exact
|
||||
match (no padding)."""
|
||||
B = batch_idx.shape[0]
|
||||
bucket = self._pick_bucket(B)
|
||||
if bucket is None:
|
||||
return self._decode_step_batched_static(batch_idx, input_ids)
|
||||
gv = self.graph_vars
|
||||
gv["input_ids"][:bucket].copy_(input_ids)
|
||||
gv["batch_idx"][:bucket].copy_(batch_idx)
|
||||
self.graphs[bucket].replay()
|
||||
return gv["outputs"][:B].clone()
|
||||
|
||||
# ----- generate loop -------------------------------------------------
|
||||
@torch.inference_mode()
|
||||
def generate(self, sequences: list, capture_cudagraph: bool = False):
|
||||
"""Main entry; mirror of FlexInference.generate but with the
|
||||
per-sequence decode loop described above. CUDA graph capture is
|
||||
deferred — the per-seq DeltaNet call makes graph capture
|
||||
non-trivial and is a follow-up."""
|
||||
del capture_cudagraph # TODO: add once batched DeltaNet exists.
|
||||
"""Main entry.
|
||||
|
||||
Prefill is single-seq through HF's DeltaNet forward (writes into
|
||||
per-seq ``DynamicCache``). After each prefill we sync that state
|
||||
into the engine's static buffers so decode can run batched.
|
||||
Decode is batched across all active sequences. If
|
||||
``capture_cudagraph=True`` and the bucket ladder hasn't been
|
||||
captured yet, we capture on the first decode; subsequent decodes
|
||||
replay the matching graph.
|
||||
"""
|
||||
# ``UNSLOTH_FLEX_QWEN3_5_NO_CAPTURE=1`` forces the eager batched
|
||||
# decode path for debugging (bypasses graph capture entirely).
|
||||
if os.environ.get("UNSLOTH_FLEX_QWEN3_5_NO_CAPTURE", "0") == "1":
|
||||
capture_cudagraph = False
|
||||
if capture_cudagraph and not self._captured:
|
||||
self.capture_decode_cudagraph()
|
||||
self.tokenize(sequences)
|
||||
waiting = deque(sequences)
|
||||
running = deque()
|
||||
|
|
@ -591,7 +982,11 @@ class FlexQwen3_5Inference:
|
|||
)
|
||||
seq.batch_idx = bi
|
||||
self._reset_linear_cache(seq)
|
||||
self._reset_static_state(bi)
|
||||
next_id = self._prefill_single(seq)
|
||||
# Hand prefilled conv/recurrent state over to the static
|
||||
# buffers so the batched decode step can read it.
|
||||
self._sync_static_from_dynamic_cache(seq)
|
||||
seq.last_token_id = next_id
|
||||
seq.output_ids.append(next_id)
|
||||
if (
|
||||
|
|
@ -659,8 +1054,7 @@ class FlexQwen3_5Inference:
|
|||
self.input_pos_buffer.zero_()
|
||||
self.input_pos_buffer[bi_tensor] = cur_pos
|
||||
|
||||
linear_caches = [s._linear_cache for s in decode_batch]
|
||||
logits = self._decode_step_eager(bi_tensor, last_ids, linear_caches)
|
||||
logits = self._decode_step(bi_tensor, last_ids)
|
||||
next_ids = torch.argmax(logits, dim=-1).tolist()
|
||||
for i, seq in enumerate(decode_batch):
|
||||
seq.last_token_id = next_ids[i]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue