feat: 3-tier ixformer flash prefill dispatch + OpenCompass max_tokens clamp + n>1 fanout + index sanitizer

paged_attn.py (+218 lines):
  - Tier 0: ixformer flash_attn_varlen_func (cu_seqlens packed)
  - Tier 0.5: ixformer flash_attn_func (non-varlen, batch layout)
  - Tier 1: CoreXFA2 3-mode dispatch (packed/paged/chunked)
  - Tier 2 fallback: existing Python Q-tiling (unchanged)
  - Import chain: flash_attn_func + CoreXFA2Class + varlen

serving_chat.py (+23 lines):
  - max_tokens clamp: fixes OpenCompass 0 score (5 benchmarks all 400)
  - n>1 fanout: remove temperature==0 restriction for t2_n_2 FAIL

api_server.py (+33 lines):
  - HTTP middleware: strip index from messages before pydantic validation
  - Fixes ValidatorIterator 0.index Extra inputs are not permitted x6
This commit is contained in:
Claude
2026-08-16 15:37:23 +00:00
parent 522e8376b6
commit cdec569977
3 changed files with 274 additions and 10 deletions

View File

@@ -903,6 +903,39 @@ 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)

View File

@@ -23,6 +23,61 @@ 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
@@ -1745,6 +1800,169 @@ 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,

View File

@@ -118,12 +118,17 @@ def _sequential_greedy_fanout_count(
request: ChatCompletionRequest,
max_num_seqs: int,
) -> int:
"""Return the supported deterministic fan-out width, or zero."""
"""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.
"""
n = request.n if request.n is not None else 1
if (
max_num_seqs == 1
and n == 2
and request.temperature == 0
and 2 <= n <= 4
and not request.stream
and not request.use_beam_search
and request.best_of is None
@@ -138,8 +143,8 @@ def _merge_sequential_chat_responses(
request_id: str,
created_time: int,
) -> ChatCompletionResponse:
if len(responses) != 2:
raise ValueError("deterministic fan-out requires exactly two responses")
if len(responses) < 2:
raise ValueError("fan-out requires at least two responses")
first = responses[0]
if any(response.model != first.model for response in responses):
@@ -397,8 +402,16 @@ 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
default_max_tokens = self.max_model_len - len(
prompt_inputs["prompt_token_ids"])
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
if request.use_beam_search:
sampling_params = request.to_beam_search_params(
default_max_tokens)
@@ -492,7 +505,7 @@ class OpenAIServingChat(OpenAIServing):
logger.error(
"Sequential greedy fan-out unexpectedly returned a stream")
return self.create_error_response(
"Failed to aggregate deterministic n=2 completion")
f"Failed to aggregate n={fanout_count} completion")
responses.append(child_response)
try:
@@ -503,11 +516,11 @@ class OpenAIServingChat(OpenAIServing):
)
except ValueError as error:
logger.error(
"Sequential greedy fan-out aggregation failed: %s",
"Sequential fan-out aggregation failed: %s",
type(error).__name__,
)
return self.create_error_response(
"Failed to aggregate deterministic n=2 completion")
f"Failed to aggregate n={fanout_count} completion")
if raw_request is not None:
metadata = RequestResponseMetadata(