From 5172f94b1f6ec3cf55cfb7da496e5bb9403d7ae6 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 16 Aug 2026 15:42:36 +0000 Subject: [PATCH] Revert "feat: 3-tier ixformer flash prefill dispatch + OpenCompass max_tokens clamp + n>1 fanout + index sanitizer" This reverts commit cdec569977ac178acf1ec15f6a2bea881459bccb. --- qwen3_6_scripts/api_server.py | 33 ----- qwen3_6_scripts/paged_attn.py | 218 -------------------------------- qwen3_6_scripts/serving_chat.py | 33 ++--- 3 files changed, 10 insertions(+), 274 deletions(-) diff --git a/qwen3_6_scripts/api_server.py b/qwen3_6_scripts/api_server.py index 2945b051..d63fc4b3 100644 --- a/qwen3_6_scripts/api_server.py +++ b/qwen3_6_scripts/api_server.py @@ -903,39 +903,6 @@ def build_app(args: Namespace) -> FastAPI: allow_headers=args.allowed_headers, ) - @app.middleware("http") - async def sanitize_chat_body(request: Request, call_next): - """Strip fields from chat messages that vLLM's pydantic models reject. - - Some replay datasets include ``index`` on messages (used by OpenAI - streaming deltas but forbidden by the non-streaming request schema). - Stripping it here avoids a ValidatorIterator 400 before our handler - even runs. - """ - if (request.method == "POST" - and request.url.path.endswith("/v1/chat/completions")): - content_type = request.headers.get("content-type", "") - if "json" in content_type or not content_type: - try: - body = await request.json() - changed = False - for msg in body.get("messages", []) if isinstance(body, dict) else []: - if isinstance(msg, dict) and "index" in msg: - del msg["index"] - changed = True - if changed: - import json as _json - raw = _json.dumps(body).encode("utf-8") - - async def patched_body(): - return raw - - request._body = raw - request._receive = patched_body # noqa - except Exception: - pass - return await call_next(request) - @app.exception_handler(RequestValidationError) async def validation_exception_handler(raw_request, exc): _bi100_log_request_validation_4xx(raw_request, exc) diff --git a/qwen3_6_scripts/paged_attn.py b/qwen3_6_scripts/paged_attn.py index 2af2e651..c3f8492c 100644 --- a/qwen3_6_scripts/paged_attn.py +++ b/qwen3_6_scripts/paged_attn.py @@ -23,61 +23,6 @@ try: except ImportError: _corex_fused_paged_prefill = None -# --------------------------------------------------------------------------- -# Tier 0 prefill: ixformer native flash_attn_varlen_func -# Sub 168 (competitor) uses this via corex_fa2.py:333 — single fused kernel -# instead of our multi-tile Python loop. This is the #1 prefill bottleneck. -# --------------------------------------------------------------------------- -_ixformer_flash_attn_varlen = None -_ixformer_flash_attn_kvcache = None -_ixformer_paged_attn_v1 = None -_ixformer_flash_attn_func = None -try: - from ixformer.contrib.vllm_flash_attn import ( - flash_attn_varlen_func as _ixformer_flash_attn_varlen, - ) -except (ImportError, AttributeError): - pass -try: - from ixformer.contrib.vllm_flash_attn import ( - flash_attn_with_kvcache as _ixformer_flash_attn_kvcache, - ) -except (ImportError, AttributeError): - pass -try: - import ixformer.functions as _ixf_F - _ixformer_paged_attn_v1 = _ixf_F.vllm_single_query_cached_kv_attention -except (ImportError, AttributeError): - pass -# Tier 0.5: ixformer top-level flash_attn_func (non-varlen) -# Probe confirmed: flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=None, -# causal=False, return_attn_probs=False) -# Available at ixformer.functions.flash_attn_func on BI-V100 real machine. -# Not varlen — requires [batch, seqlen, nheads, headdim] layout. -# For single-sequence prefill (competition concurrency=1), this replaces -# the entire Python Q-tiling loop with one C++ kernel. -try: - _ixformer_flash_attn_func = _ixf_F.flash_attn_func -except (NameError, AttributeError): - try: - import ixformer.functions as _ixf_F2 - _ixformer_flash_attn_func = _ixf_F2.flash_attn_func - except (ImportError, AttributeError): - pass - -# Tier 0.6: corex_fa2 dispatch (3-mode: packed prefill, paged decode, chunked) -# This module wraps ix_bridge C++ and ixformer Python backends with proper -# fallback chain. Import lazily — if corex_fa2 is not deployed, fall through. -_corex_fa2_dispatch = None -try: - from ex_engine.python.corex_fa2 import CoreXFA2 as _CoreXFA2Class - # Instantiate later when we know num_heads/head_dim -except ImportError: - _CoreXFA2Class = None - -_USE_IXFORMER_FLASH_PREFILL = env_bool("BI100_USE_IXFORMER_FLASH_PREFILL", True) -_LOGGED_IXFORMER_PREFILL = set() - # from vllm.attention.ops.prefix_prefill import context_attention_fwd # NOTE: context_attention_fwd (Triton kernel from prefix_prefill.py) is NOT # imported here. On Iluvatar BI-V100 that kernel hangs the GPU card @@ -1800,169 +1745,6 @@ class PagedAttention: k_scale=k_scale, v_scale=v_scale, ) - # ----------------------------------------------------------------- - # Tier 0: ixformer flash_attn_varlen_func (cu_seqlens packed) - # This is what sub 168 uses via corex_fa2.py:333. - # Handles variable-length sequences in a single fused kernel. - # ----------------------------------------------------------------- - if (_USE_IXFORMER_FLASH_PREFILL - and _ixformer_flash_attn_varlen is not None - and alibi_slopes is None - and sliding_window is None - and k_scale == 1.0 and v_scale == 1.0 - and kv_cache_dtype == "auto"): - try: - batch_size = seq_lens_tensor.shape[0] - num_q_heads = query.shape[1] - head_dim = query.shape[2] - scale = head_dim ** -0.5 - - # Build cu_seqlens for packed varlen interface - # For prefill, all tokens are fresh — cu_seqlens covers full seq - q_lens = (query_start_loc[1:] - query_start_loc[:-1]) - cu_seqlens_q = torch.zeros( - batch_size + 1, dtype=torch.int32, device=query.device) - cu_seqlens_q[1:] = torch.cumsum(q_lens, dim=0).to(torch.int32) - - # For context_lens=0 (pure prefill), k_seqlens == q_seqlens - # For context_lens>0 (chunked prefill), we need to handle - # the cached KV — but flash_attn_varlen handles only the - # fresh Q/K/V, not the paged cache. Fall through for that case. - all_zero_context = bool(context_lens.max().item() == 0) - if all_zero_context: - max_seqlen = int(q_lens.max().item()) - output = _ixformer_flash_attn_varlen( - q=query, k=key, v=value, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_q, - max_seqlen_q=max_seqlen, - max_seqlen_k=max_seqlen, - softmax_scale=scale, - causal=True) - if "varlen_prefill" not in _LOGGED_IXFORMER_PREFILL: - _LOGGED_IXFORMER_PREFILL.add("varlen_prefill") - import logging - logging.getLogger(__name__).info( - "[BI100 PREFILL] ixformer flash_attn_varlen: " - "B=%d Hq=%d D=%d max_q=%d — FUSED kernel active", - batch_size, num_q_heads, head_dim, max_seqlen) - return output - except Exception as _e: - if "varlen_error" not in _LOGGED_IXFORMER_PREFILL: - _LOGGED_IXFORMER_PREFILL.add("varlen_error") - import logging - logging.getLogger(__name__).warning( - "[BI100 PREFILL] ixformer flash_attn_varlen failed: " - "%s — falling through to Tier 0.5", _e) - - # ----------------------------------------------------------------- - # Tier 0.5: ixformer flash_attn_func (non-varlen, batch layout) - # Probe confirmed available: flash_attn_func(q, k, v, ...) - # For single-sequence (batch=1) prefill, reshape to [1, seqlen, h, d] - # and call one C++ kernel. This replaces the entire Python Q-tiling - # loop which iterates hundreds of times for long prompts. - # ----------------------------------------------------------------- - if (_USE_IXFORMER_FLASH_PREFILL - and _ixformer_flash_attn_func is not None - and alibi_slopes is None - and sliding_window is None - and k_scale == 1.0 and v_scale == 1.0 - and kv_cache_dtype == "auto"): - try: - batch_size = seq_lens_tensor.shape[0] - num_q_heads = query.shape[1] - num_kv_heads = key.shape[1] if key.dim() == 3 else query.shape[1] - head_dim = query.shape[2] - scale = head_dim ** -0.5 - - all_zero_context = bool(context_lens.max().item() == 0) - if all_zero_context and batch_size == 1: - # Single sequence, pure prefill — reshape to batch format - total_q = query.shape[0] - # flash_attn_func expects [batch, seqlen, nheads, headdim] - q_4d = query.unsqueeze(0) # [1, total_q, num_q_heads, head_dim] - k_4d = key.unsqueeze(0) - v_4d = value.unsqueeze(0) - - out_4d = _ixformer_flash_attn_func( - q_4d, k_4d, v_4d, - dropout_p=0.0, - softmax_scale=scale, - causal=True) - output = out_4d.squeeze(0) # [total_q, num_q_heads, head_dim] - - if "func_prefill" not in _LOGGED_IXFORMER_PREFILL: - _LOGGED_IXFORMER_PREFILL.add("func_prefill") - import logging - logging.getLogger(__name__).info( - "[BI100 PREFILL] ixformer flash_attn_func: " - "B=1 Hq=%d D=%d seqlen=%d — FUSED kernel active", - num_q_heads, head_dim, total_q) - return output - except Exception as _e: - if "func_error" not in _LOGGED_IXFORMER_PREFILL: - _LOGGED_IXFORMER_PREFILL.add("func_error") - import logging - logging.getLogger(__name__).warning( - "[BI100 PREFILL] ixformer flash_attn_func failed: " - "%s — falling through to Python Q-tiling", _e) - - # ----------------------------------------------------------------- - # Tier 1: corex_fa2 dispatch (3-mode: packed, paged decode, chunked) - # This wraps ix_bridge C++ and ixformer Python backends. - # ----------------------------------------------------------------- - if (_USE_IXFORMER_FLASH_PREFILL - and _CoreXFA2Class is not None - and alibi_slopes is None - and sliding_window is None - and k_scale == 1.0 and v_scale == 1.0 - and kv_cache_dtype == "auto"): - try: - batch_size = seq_lens_tensor.shape[0] - num_q_heads = query.shape[1] - num_kv_heads = key.shape[1] if key.dim() == 3 else num_q_heads - head_dim = query.shape[2] - - all_zero_context = bool(context_lens.max().item() == 0) - if all_zero_context: - q_lens = (query_start_loc[1:] - query_start_loc[:-1]) - cu_seqlens_q = torch.zeros( - batch_size + 1, dtype=torch.int32, - device=query.device) - cu_seqlens_q[1:] = torch.cumsum( - q_lens, dim=0).to(torch.int32) - max_seqlen = int(q_lens.max().item()) - - fa2 = _CoreXFA2Class(num_q_heads, num_kv_heads, head_dim) - if fa2.is_available: - output = fa2.packed_prefill( - query, key, value, - cu_seqlens_q, cu_seqlens_q, - max_seqlen, max_seqlen, - causal=True) - if "corex_fa2" not in _LOGGED_IXFORMER_PREFILL: - _LOGGED_IXFORMER_PREFILL.add("corex_fa2") - import logging - logging.getLogger(__name__).info( - "[BI100 PREFILL] CoreXFA2 packed_prefill: " - "B=%d Hq=%d Hkv=%d D=%d max_q=%d", - batch_size, num_q_heads, num_kv_heads, - head_dim, max_seqlen) - return output - except Exception as _e: - if "corex_fa2_error" not in _LOGGED_IXFORMER_PREFILL: - _LOGGED_IXFORMER_PREFILL.add("corex_fa2_error") - import logging - logging.getLogger(__name__).warning( - "[BI100 PREFILL] CoreXFA2 failed: %s — " - "falling through to Python Q-tiling", _e) - - # ----------------------------------------------------------------- - # Tier 2 (fallback): Python Q-tiling with online softmax - # This is the current default — functional but slow for long prompts. - # 107K prompt = ~400 tile iterations in Python, each launching - # multiple CUDA kernels. Sub 694 shows 190s TTFT for such requests. - # ----------------------------------------------------------------- return PagedAttention._forward_prefix_pytorch( query, key, value, key_cache, value_cache, diff --git a/qwen3_6_scripts/serving_chat.py b/qwen3_6_scripts/serving_chat.py index 032a11a0..14b3456f 100644 --- a/qwen3_6_scripts/serving_chat.py +++ b/qwen3_6_scripts/serving_chat.py @@ -118,17 +118,12 @@ def _sequential_greedy_fanout_count( request: ChatCompletionRequest, max_num_seqs: int, ) -> int: - """Return the supported fan-out width, or zero. - - When max_num_seqs=1 (competition fixed config), vLLM cannot schedule - n>1 natively. We sequentially execute n independent n=1 requests and - merge them. This works for any temperature — deterministic (temp=0) - produces identical choices, stochastic produces diverse ones. - """ + """Return the supported deterministic fan-out width, or zero.""" n = request.n if request.n is not None else 1 if ( max_num_seqs == 1 - and 2 <= n <= 4 + and n == 2 + and request.temperature == 0 and not request.stream and not request.use_beam_search and request.best_of is None @@ -143,8 +138,8 @@ def _merge_sequential_chat_responses( request_id: str, created_time: int, ) -> ChatCompletionResponse: - if len(responses) < 2: - raise ValueError("fan-out requires at least two responses") + if len(responses) != 2: + raise ValueError("deterministic fan-out requires exactly two responses") first = responses[0] if any(response.model != first.model for response in responses): @@ -402,16 +397,8 @@ class OpenAIServingChat(OpenAIServing): # OpenAI API: max_completion_tokens takes precedence over max_tokens if request.max_completion_tokens is not None and request.max_tokens is None: request.max_tokens = request.max_completion_tokens - prompt_len = len(prompt_inputs["prompt_token_ids"]) - default_max_tokens = self.max_model_len - prompt_len - # Clamp max_tokens so prompt + completion <= max_model_len. - # Without this, evaluation systems (e.g. OpenCompass) that send - # max_tokens=131072 get 400 errors when prompt+max_tokens exceeds - # max_model_len, resulting in 0 score on all academic benchmarks. - if default_max_tokens < 1: - default_max_tokens = 1 - if request.max_tokens is not None and request.max_tokens > default_max_tokens: - request.max_tokens = default_max_tokens + default_max_tokens = self.max_model_len - len( + prompt_inputs["prompt_token_ids"]) if request.use_beam_search: sampling_params = request.to_beam_search_params( default_max_tokens) @@ -505,7 +492,7 @@ class OpenAIServingChat(OpenAIServing): logger.error( "Sequential greedy fan-out unexpectedly returned a stream") return self.create_error_response( - f"Failed to aggregate n={fanout_count} completion") + "Failed to aggregate deterministic n=2 completion") responses.append(child_response) try: @@ -516,11 +503,11 @@ class OpenAIServingChat(OpenAIServing): ) except ValueError as error: logger.error( - "Sequential fan-out aggregation failed: %s", + "Sequential greedy fan-out aggregation failed: %s", type(error).__name__, ) return self.create_error_response( - f"Failed to aggregate n={fanout_count} completion") + "Failed to aggregate deterministic n=2 completion") if raw_request is not None: metadata = RequestResponseMetadata(