feat(CRITICAL): import wudixzy/competition complete corex stack — 12 prebuilt .so + 13 CUDA kernels + 2615-line qwen3_5.py

Source: github.com/wudixzy/competition (1527 files, BI-V100 competition reference)

Imported assets:
- 12 prebuilt CoreX .so extensions (corex-3.2.3-ivcore10):
  corex_gdn_{beta_decay,causal_conv,gated_norm,packed_decode,qk_map}.so
  corex_moe_{direct_routed,exact_reduce,weight_gather}.so
  corex_attn_head_rms_norm.so, corex_paged_kv_gather.so
  corex_block_major_kv_transfer.so, corex_fused_paged_prefill.so

- 13 CUDA kernel sources (.cu) for above extensions
- 11 build scripts (build_corex_*.sh)
- install_prebuilt_corex.sh (SHA256-verified .so deployment)
- qwen3_5.py (2615 lines) with FULL corex kernel integration
- 9 vllm vendor override files (block manager, sampler, etc)
- 19 patch scripts (model_runner, xformers, block_major, etc)
- Complete serving layer (serving_chat, protocol, api_server, etc)
- bi100_env.py, bi100_profile.py, gdn_prefix.py, block_major_kv_cache.py
- Dockerfile aligned with reference build chain
- computility-run.yaml with BI100_MOE_COREX_DIRECT_ROUTED=1

Call chain verified:
  Dockerfile COPY → patch_ops.sh → install_prebuilt_corex.sh → 12 .so to $VLLM_ROOT
  qwen3_5.py imports: from vllm import corex_gdn_* / corex_moe_* / corex_attn_*
This commit is contained in:
project6-dev
2026-08-11 03:55:38 +00:00
parent 81875fff52
commit 5862708b32
86 changed files with 24702 additions and 9860 deletions

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,26 @@
import os
def env_bool(name: str, default: bool = False) -> bool:
raw = os.getenv(name)
if raw is None:
return default
if raw in ("1", "true", "True", "yes", "YES", "on", "ON"):
return True
if raw in ("0", "false", "False", "no", "NO", "off", "OFF"):
return False
raise RuntimeError(f"{name} must be boolean, got {raw!r}")
def env_int(name: str, default: int, min_value: int, max_value: int) -> int:
raw = os.getenv(name)
if raw is None:
return default
try:
value = int(raw)
except ValueError as exc:
raise RuntimeError(f"{name} must be int, got {raw!r}") from exc
if not (min_value <= value <= max_value):
raise RuntimeError(
f"{name}={value} outside [{min_value}, {max_value}]")
return value

View File

@@ -0,0 +1,237 @@
import contextlib
import fnmatch
import functools
import json
import os
import re
import threading
import time
from vllm.logger import init_logger
logger = init_logger(__name__)
_EVENT_SCHEMA = "bi100-profile-event-v1"
_EVENT_VERSION = 1
_NAME_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_.-]{0,63}$")
_FILTER_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_.*?-]{0,63}$")
def _strict_bool(name: str, default: str = "0") -> bool:
value = os.getenv(name, default).strip()
if value not in {"0", "1"}:
raise RuntimeError(f"{name} must be exactly 0 or 1, got {value!r}")
return value == "1"
_ENABLED = _strict_bool("BI100_PROFILE")
_INCLUDE_STARTUP = _strict_bool("BI100_PROFILE_INCLUDE_STARTUP")
_MODE = os.getenv("BI100_PROFILE_MODE", "sync").strip().lower()
_FILTERS = tuple(
item.strip()
for item in os.getenv("BI100_PROFILE_FILTER", "").split(",")
if item.strip()
)
if _ENABLED and _MODE not in {"sync", "event"}:
raise RuntimeError(f"unsupported BI100_PROFILE_MODE={_MODE!r}")
if _ENABLED and any(_FILTER_RE.fullmatch(pattern) is None
for pattern in _FILTERS):
raise RuntimeError("BI100_PROFILE_FILTER contains an invalid pattern")
_EVENT_RECORDS = []
_COUNTERS = {}
_LOCK = threading.Lock()
_FORWARD_INDEX = 0
_LAST_FLUSH_NS = None
_ACTIVE_FORWARD_TOKEN = None
_NEXT_FORWARD_TOKEN = 0
def _enabled_for(name: str) -> bool:
return (_ENABLED
and (not _FILTERS
or any(fnmatch.fnmatchcase(name, pattern)
for pattern in _FILTERS)))
def _skip_startup() -> bool:
return (not _INCLUDE_STARTUP
and os.getenv("BI100_IN_STARTUP_PROFILE") == "1")
def bi100_profile_event_enabled() -> bool:
return _ENABLED and _MODE == "event" and not _skip_startup()
def _begin_profile_forward():
global _ACTIVE_FORWARD_TOKEN, _NEXT_FORWARD_TOKEN
if not bi100_profile_event_enabled():
return None
with _LOCK:
_EVENT_RECORDS.clear()
_COUNTERS.clear()
token = _NEXT_FORWARD_TOKEN
_NEXT_FORWARD_TOKEN += 1
_ACTIVE_FORWARD_TOKEN = token
return token
def _abort_profile_forward(token) -> None:
global _ACTIVE_FORWARD_TOKEN
if token is None:
return
with _LOCK:
if _ACTIVE_FORWARD_TOKEN != token:
return
_EVENT_RECORDS.clear()
_COUNTERS.clear()
_ACTIVE_FORWARD_TOKEN = None
def bi100_profile_transaction(function):
"""Keep one top-level model forward isolated from failed forwards."""
@functools.wraps(function)
def wrapped(*args, **kwargs):
token = _begin_profile_forward()
if token is None:
return function(*args, **kwargs)
try:
result = function(*args, **kwargs)
except BaseException:
_abort_profile_forward(token)
raise
with _LOCK:
was_flushed = _ACTIVE_FORWARD_TOKEN != token
if not was_flushed:
_abort_profile_forward(token)
raise RuntimeError(
"BI100 profile transaction completed without a flush")
return result
return wrapped
def _normalize_metadata(metadata):
normalized = {}
for key, value in metadata.items():
if not isinstance(key, str) or _NAME_RE.fullmatch(key) is None:
raise TypeError("profile metadata keys must be bounded names")
if isinstance(value, bool):
normalized[key] = value
elif isinstance(value, int) and not isinstance(value, bool):
normalized[key] = value
elif isinstance(value, str) and len(value) <= 64:
normalized[key] = value
else:
raise TypeError(
"profile metadata values must be bool, int, or short strings")
return normalized
def bi100_profile_count(name: str, **metadata) -> None:
"""Record privacy-safe path metadata for the current model forward."""
if not bi100_profile_event_enabled() or not _enabled_for(name):
return
if not isinstance(name, str) or _NAME_RE.fullmatch(name) is None:
raise TypeError("profile counter name must be a bounded name")
normalized = _normalize_metadata(metadata)
encoded = json.dumps(
{"name": name, **normalized}, sort_keys=True, separators=(",", ":"))
with _LOCK:
_COUNTERS[encoded] = _COUNTERS.get(encoded, 0) + 1
@contextlib.contextmanager
def bi100_timer(name: str):
if not _enabled_for(name) or _skip_startup():
yield
return
import torch
if _MODE == "event":
started = torch.cuda.Event(enable_timing=True)
finished = torch.cuda.Event(enable_timing=True)
host_started_ns = time.monotonic_ns()
started.record()
try:
yield
finally:
finished.record()
with _LOCK:
_EVENT_RECORDS.append(
(name, started, finished, host_started_ns))
return
torch.cuda.synchronize()
t0 = time.perf_counter()
try:
yield
finally:
torch.cuda.synchronize()
logger.info("[BI100_PROFILE] %s %.3f ms", name,
(time.perf_counter() - t0) * 1000)
def bi100_profile_flush(*, tp_rank, **metadata):
"""Synchronize once and emit one aggregate event record per model forward."""
global _ACTIVE_FORWARD_TOKEN, _FORWARD_INDEX, _LAST_FLUSH_NS
if not bi100_profile_event_enabled():
return None
if (not isinstance(tp_rank, int) or isinstance(tp_rank, bool)
or not 0 <= tp_rank < 256):
raise TypeError("profile TP rank must be an integer in [0, 255]")
normalized_metadata = _normalize_metadata(metadata)
with _LOCK:
records = list(_EVENT_RECORDS)
counters = dict(_COUNTERS)
_EVENT_RECORDS.clear()
_COUNTERS.clear()
_ACTIVE_FORWARD_TOKEN = None
if not records:
return None
import torch
torch.cuda.synchronize()
flushed_ns = time.monotonic_ns()
regions = {}
model_started_ns = []
for name, started, finished, host_started_ns in records:
stats = regions.setdefault(name, {"count": 0, "total_ms": 0.0})
stats["count"] += 1
stats["total_ms"] += float(started.elapsed_time(finished))
if name == "model.forward":
model_started_ns.append(host_started_ns)
counter_rows = []
for encoded, count in sorted(counters.items()):
row = json.loads(encoded)
row["count"] = count
counter_rows.append(row)
first_model_started_ns = (
min(model_started_ns) if model_started_ns else None)
payload = {
"schema": _EVENT_SCHEMA,
"version": _EVENT_VERSION,
"tp_rank": tp_rank,
"forward_index": _FORWARD_INDEX,
"metadata": normalized_metadata,
"event_count": len(records),
"model_forward_event_count": len(model_started_ns),
"regions": regions,
"counters": counter_rows,
"host_model_start_to_flush_ms": (
(flushed_ns - first_model_started_ns) / 1_000_000
if first_model_started_ns is not None else None),
"host_gap_since_previous_flush_ms": (
(first_model_started_ns - _LAST_FLUSH_NS) / 1_000_000
if first_model_started_ns is not None
and _LAST_FLUSH_NS is not None
else None),
}
_FORWARD_INDEX += 1
_LAST_FLUSH_NS = flushed_ns
logger.info("[BI100_PROFILE_EVENT] %s",
json.dumps(payload, sort_keys=True, separators=(",", ":")))
return payload

View File

@@ -0,0 +1,398 @@
from __future__ import annotations
import os
import time
from collections.abc import Mapping
import torch
from vllm.logger import init_logger
logger = init_logger(__name__)
ENABLE_ENV = "BI100_BLOCK_MAJOR_CPU_KV"
TRACE_ENV = "BI100_BLOCK_MAJOR_CPU_KV_TRACE"
CPU_OFFLOAD_ENV = "BI100_CPU_KV_OFFLOAD"
HYBRID_ACCOUNTING_ENV = "BI100_HYBRID_KV_ACCOUNTING"
NUM_ATTENTION_LAYERS = 10
KV_PLANES = 2
ELEMENTS_PER_PLANE_BLOCK = 4096
STAGING_BLOCKS = 512
STAGING_BUFFER_COUNT = 2
BYTES_PER_BLOCK = (
NUM_ATTENTION_LAYERS * KV_PLANES * ELEMENTS_PER_PLANE_BLOCK * 2
)
GPU_STAGING_BYTES = STAGING_BLOCKS * STAGING_BUFFER_COUNT * BYTES_PER_BLOCK
def _strict_binary_selector(
name: str,
environ: Mapping[str, str] | None = None,
) -> bool:
source = os.environ if environ is None else environ
raw = source.get(name, "0")
if raw == "0":
return False
if raw == "1":
return True
raise RuntimeError(f"{name} must be exactly '0' or '1', got {raw!r}")
def block_major_cpu_kv_enabled(
environ: Mapping[str, str] | None = None,
) -> bool:
return _strict_binary_selector(ENABLE_ENV, environ)
def block_major_cpu_kv_trace_enabled(
environ: Mapping[str, str] | None = None,
) -> bool:
return _strict_binary_selector(TRACE_ENV, environ)
def _require_block_major_runtime(
environ: Mapping[str, str] | None = None,
) -> None:
source = os.environ if environ is None else environ
if source.get(CPU_OFFLOAD_ENV, "0") != "1":
raise RuntimeError(
f"{ENABLE_ENV}=1 requires {CPU_OFFLOAD_ENV}=1")
if source.get(HYBRID_ACCOUNTING_ENV, "legacy40") != "full_attention":
raise RuntimeError(
f"{ENABLE_ENV}=1 requires "
f"{HYBRID_ACCOUNTING_ENV}=full_attention")
def reserve_block_major_gpu_blocks(
num_gpu_blocks: int,
cache_block_size: int,
environ: Mapping[str, str] | None = None,
) -> int:
if (not isinstance(num_gpu_blocks, int)
or isinstance(num_gpu_blocks, bool)
or num_gpu_blocks < 0):
raise ValueError("num_gpu_blocks must be a non-negative integer")
if not block_major_cpu_kv_enabled(environ):
return num_gpu_blocks
_require_block_major_runtime(environ)
if cache_block_size != BYTES_PER_BLOCK:
raise RuntimeError(
f"{ENABLE_ENV}=1 requires cache block size "
f"{BYTES_PER_BLOCK}, got {cache_block_size}")
reserved_blocks = (
GPU_STAGING_BYTES + cache_block_size - 1
) // cache_block_size
remaining_blocks = num_gpu_blocks - reserved_blocks
if remaining_blocks <= 0:
raise RuntimeError(
"block-major GPU staging leaves no usable GPU KV blocks")
logger.info(
"[BI100 BLOCK KV] capacity reserve blocks=%d bytes=%d "
"profiled_blocks=%d usable_blocks=%d",
reserved_blocks,
GPU_STAGING_BYTES,
num_gpu_blocks,
remaining_blocks,
)
return remaining_blocks
def validate_block_mapping(
mapping: torch.Tensor,
source_limit: int,
destination_limit: int,
) -> tuple[torch.Tensor, torch.Tensor]:
if not isinstance(mapping, torch.Tensor):
raise TypeError("block mapping must be a torch.Tensor")
if mapping.device.type != "cpu":
raise ValueError("block mapping must be on CPU")
if mapping.dtype != torch.int64:
raise ValueError("block mapping must use torch.int64")
if not mapping.is_contiguous():
raise ValueError("block mapping must be contiguous")
if mapping.dim() != 2 or mapping.shape[1] != 2:
raise ValueError("block mapping must have shape [N, 2]")
if source_limit <= 0 or destination_limit <= 0:
raise ValueError("block mapping limits must be positive")
sources: set[int] = set()
destinations: set[int] = set()
for row, pair in enumerate(mapping.tolist()):
source, destination = pair
if not 0 <= source < source_limit:
raise ValueError(
f"source block out of range at row {row}: {source}")
if not 0 <= destination < destination_limit:
raise ValueError(
f"destination block out of range at row {row}: "
f"{destination}")
if source in sources:
raise ValueError(f"duplicate source block: {source}")
if destination in destinations:
raise ValueError(f"duplicate destination block: {destination}")
sources.add(source)
destinations.add(destination)
return mapping[:, 0].contiguous(), mapping[:, 1].contiguous()
class BlockMajorCpuKVCache:
def __init__(
self,
gpu_cache: list[torch.Tensor],
num_cpu_blocks: int,
pin_memory: bool,
) -> None:
self._validate_gpu_cache(gpu_cache)
if block_major_cpu_kv_enabled():
_require_block_major_runtime()
if num_cpu_blocks <= 0:
raise RuntimeError(
f"{ENABLE_ENV}=1 requires a positive CPU block count")
if not pin_memory:
raise RuntimeError(
f"{ENABLE_ENV}=1 requires pinned CPU memory")
try:
from vllm import corex_block_major_kv_transfer as extension
except ImportError as exc:
raise RuntimeError(
"block-major CoreX extension is unavailable") from exc
self.extension = extension
self.gpu_cache = gpu_cache
self.device = gpu_cache[0].device
self.dtype = gpu_cache[0].dtype
self.num_gpu_blocks = gpu_cache[0].shape[1]
self.num_cpu_blocks = num_cpu_blocks
self.trace_enabled = block_major_cpu_kv_trace_enabled()
self.cpu_pool = torch.zeros(
(
num_cpu_blocks,
NUM_ATTENTION_LAYERS,
KV_PLANES,
ELEMENTS_PER_PLANE_BLOCK,
),
dtype=self.dtype,
device="cpu",
pin_memory=True,
)
if not self.cpu_pool.is_pinned():
raise RuntimeError("block-major CPU pool is not pinned")
# Preserve the public CacheEngine shape without allocating a second
# layer-major CPU cache. Transfer methods use cpu_pool directly.
self.layer_views = [
self.cpu_pool[:, layer, :, :].permute(1, 0, 2)
for layer in range(NUM_ATTENTION_LAYERS)
]
self.cpu_staging = [
torch.empty(
(
STAGING_BLOCKS,
NUM_ATTENTION_LAYERS,
KV_PLANES,
ELEMENTS_PER_PLANE_BLOCK,
),
dtype=self.dtype,
device="cpu",
pin_memory=True,
)
for _ in range(STAGING_BUFFER_COUNT)
]
if not all(staging.is_pinned() for staging in self.cpu_staging):
raise RuntimeError("block-major CPU staging is not pinned")
with torch.cuda.device(self.device):
self.gpu_staging = [
torch.empty_like(staging, device=self.device)
for staging in self.cpu_staging
]
self.events = [
torch.cuda.Event(enable_timing=False)
for _ in range(STAGING_BUFFER_COUNT)
]
self.error_flag = torch.zeros(
1, dtype=torch.int32, device=self.device)
logger.info(
"[BI100 BLOCK KV] enabled device=%s gpu_blocks=%d cpu_blocks=%d "
"layers=%d block_bytes=%d staging_blocks=%d staging_buffers=%d",
self.device,
self.num_gpu_blocks,
self.num_cpu_blocks,
NUM_ATTENTION_LAYERS,
BYTES_PER_BLOCK,
STAGING_BLOCKS,
STAGING_BUFFER_COUNT,
)
@staticmethod
def _validate_gpu_cache(gpu_cache: list[torch.Tensor]) -> None:
if len(gpu_cache) != NUM_ATTENTION_LAYERS:
raise RuntimeError(
f"{ENABLE_ENV}=1 requires exactly "
f"{NUM_ATTENTION_LAYERS} GPU attention caches, got "
f"{len(gpu_cache)}")
first = gpu_cache[0]
if first.device.type != "cuda":
raise RuntimeError("block-major GPU cache must be on CUDA")
if first.dtype != torch.float16:
raise RuntimeError("block-major GPU cache must use float16")
if (first.dim() != 3 or first.shape[0] != KV_PLANES
or first.shape[2] != ELEMENTS_PER_PLANE_BLOCK):
raise RuntimeError(
"block-major GPU cache must have shape [2, blocks, 4096]")
if not first.is_contiguous():
raise RuntimeError("block-major GPU cache must be contiguous")
for layer, tensor in enumerate(gpu_cache):
if tensor.device != first.device:
raise RuntimeError(
f"GPU cache layer {layer} is on a different device")
if tensor.dtype != first.dtype or tensor.shape != first.shape:
raise RuntimeError(
f"GPU cache layer {layer} has inconsistent geometry")
if not tensor.is_contiguous():
raise RuntimeError(
f"GPU cache layer {layer} is not contiguous")
def _to_gpu_ids(self, block_ids: torch.Tensor) -> torch.Tensor:
return block_ids.to(
device=self.device,
dtype=torch.int32,
non_blocking=False,
)
@staticmethod
def _chunks(
source: torch.Tensor,
destination: torch.Tensor,
gpu_ids: torch.Tensor,
):
for start in range(0, source.numel(), STAGING_BLOCKS):
end = min(start + STAGING_BLOCKS, source.numel())
yield (
source[start:end],
destination[start:end],
gpu_ids[start:end],
end - start,
)
def _begin(self) -> None:
self.error_flag.zero_()
def _finish(
self,
direction: str,
block_count: int,
started: float | None,
) -> None:
# check_error performs the final stream synchronization. This also
# makes every staging slot safe to reuse in the next CacheEngine call.
self.extension.check_error(self.error_flag)
if started is not None:
elapsed_ms = (time.perf_counter() - started) * 1000.0
logger.info(
"[BI100 BLOCK KV TRACE] direction=%s blocks=%d bytes=%d "
"elapsed_ms=%.3f",
direction,
block_count,
block_count * BYTES_PER_BLOCK,
elapsed_ms,
)
def swap_out(self, mapping: torch.Tensor) -> None:
started = time.perf_counter() if self.trace_enabled else None
source_gpu, destination_cpu = validate_block_mapping(
mapping,
source_limit=self.num_gpu_blocks,
destination_limit=self.num_cpu_blocks,
)
block_count = source_gpu.numel()
if block_count == 0:
return
source_gpu_ids = self._to_gpu_ids(source_gpu)
self._begin()
pending: tuple[int, torch.Tensor, int] | None = None
for index, (_, destination, gpu_ids, count) in enumerate(
self._chunks(
source_gpu, destination_cpu, source_gpu_ids)):
slot = index % STAGING_BUFFER_COUNT
self.extension.pack(
self.gpu_cache,
gpu_ids,
self.gpu_staging[slot],
self.error_flag,
count,
)
self.cpu_staging[slot][:count].copy_(
self.gpu_staging[slot][:count],
non_blocking=True,
)
self.events[slot].record()
if pending is not None:
pending_slot, pending_destination, pending_count = pending
self.events[pending_slot].synchronize()
self.extension.cpu_scatter(
self.cpu_staging[pending_slot],
self.cpu_pool,
pending_destination,
pending_count,
)
pending = (slot, destination, count)
if pending is not None:
pending_slot, pending_destination, pending_count = pending
self.events[pending_slot].synchronize()
self.extension.cpu_scatter(
self.cpu_staging[pending_slot],
self.cpu_pool,
pending_destination,
pending_count,
)
self._finish("d2h", block_count, started)
def swap_in(self, mapping: torch.Tensor) -> None:
started = time.perf_counter() if self.trace_enabled else None
source_cpu, destination_gpu = validate_block_mapping(
mapping,
source_limit=self.num_cpu_blocks,
destination_limit=self.num_gpu_blocks,
)
block_count = source_cpu.numel()
if block_count == 0:
return
destination_gpu_ids = self._to_gpu_ids(destination_gpu)
self._begin()
for index, (source, _, gpu_ids, count) in enumerate(
self._chunks(
source_cpu, destination_gpu, destination_gpu_ids)):
slot = index % STAGING_BUFFER_COUNT
if index >= STAGING_BUFFER_COUNT:
self.events[slot].synchronize()
self.extension.cpu_gather(
self.cpu_pool,
source,
self.cpu_staging[slot],
count,
)
self.gpu_staging[slot][:count].copy_(
self.cpu_staging[slot][:count],
non_blocking=True,
)
self.extension.scatter(
self.gpu_staging[slot],
gpu_ids,
self.gpu_cache,
self.error_flag,
count,
)
self.events[slot].record()
self._finish("h2d", block_count, started)

View File

@@ -0,0 +1,33 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_attn_head_rms_norm.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_attn_head_rms_norm.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" \
--cuda-gpu-arch=ivcore10 \
--no-cuda-version-check \
-D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_attn_head_rms_norm \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" \
-I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_attn_head_rms_norm.cu" \
-L"${TORCH_ROOT}/lib" \
-L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" \
-Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart \
-o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX attention head RMSNorm extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_fused_paged_prefill_split4.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_fused_paged_prefill_split4.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_fused_paged_prefill \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_fused_paged_prefill_split4.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcublas -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX split4 fused paged-prefill extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,33 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_gdn_beta_decay.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_gdn_beta_decay.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" \
--cuda-gpu-arch=ivcore10 \
--no-cuda-version-check \
-D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_gdn_beta_decay \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" \
-I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_gdn_beta_decay.cu" \
-L"${TORCH_ROOT}/lib" \
-L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" \
-Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart \
-o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX GDN beta/decay extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,33 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_gdn_causal_conv.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_gdn_causal_conv.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" \
--cuda-gpu-arch=ivcore10 \
--no-cuda-version-check \
-D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_gdn_causal_conv \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" \
-I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_gdn_causal_conv.cu" \
-L"${TORCH_ROOT}/lib" \
-L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" \
-Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart \
-o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX GDN causal conv extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,33 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_gdn_gated_norm.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_gdn_gated_norm.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" \
--cuda-gpu-arch=ivcore10 \
--no-cuda-version-check \
-D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_gdn_gated_norm \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" \
-I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_gdn_gated_norm.cu" \
-L"${TORCH_ROOT}/lib" \
-L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" \
-Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart \
-o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX GDN gated norm extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_gdn_packed_decode.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_gdn_packed_decode.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_gdn_packed_decode \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_gdn_packed_decode.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX GDN packed decode extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_gdn_qk_map.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_gdn_qk_map.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_gdn_qk_map \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_gdn_qk_map.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX GDN q/k map extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_moe_direct_routed.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_moe_direct_routed.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_moe_direct_routed \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_moe_direct_routed.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX direct routed-expert extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_moe_exact_reduce.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_moe_exact_reduce.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_moe_exact_reduce \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_moe_exact_reduce.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX MoE exact reduce extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_moe_weight_gather.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_moe_weight_gather.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_moe_weight_gather \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_moe_weight_gather.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX MoE selected-weight gather extension %s\n' "${OUTPUT}"

View File

@@ -0,0 +1,27 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: build_corex_paged_kv_gather.sh VLLM_ROOT}
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
OUTPUT=${VLLM_ROOT}/corex_paged_kv_gather.so
"${COREX_ROOT}/bin/clang++" \
-std=c++17 -O3 -shared -fPIC \
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
-DTORCH_EXTENSION_NAME=corex_paged_kv_gather \
-DTORCH_API_INCLUDE_EXTENSION_H \
-I"${TORCH_ROOT}/include" \
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
-I/usr/local/include/python3.10 \
"${SCRIPT_DIR}/corex_paged_kv_gather.cu" \
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
test -s "${OUTPUT}"
printf '[ok] CoreX paged K/V gather extension %s\n' "${OUTPUT}"

File diff suppressed because it is too large Load Diff

View File

@@ -1,261 +1,261 @@
"""
This file contains the command line arguments for the vLLM's
OpenAI-compatible server. It is kept in a separate file for documentation
purposes.
"""
import argparse
import json
import ssl
from typing import List, Optional, Sequence, Union
from vllm.engine.arg_utils import AsyncEngineArgs, nullable_str
from vllm.entrypoints.chat_utils import validate_chat_template
from vllm.entrypoints.openai.serving_engine import (LoRAModulePath,
PromptAdapterPath)
from vllm.entrypoints.openai.tool_parsers import ToolParserManager
from vllm.utils import FlexibleArgumentParser
class LoRAParserAction(argparse.Action):
def __call__(
self,
parser: argparse.ArgumentParser,
namespace: argparse.Namespace,
values: Optional[Union[str, Sequence[str]]],
option_string: Optional[str] = None,
):
if values is None:
values = []
if isinstance(values, str):
raise TypeError("Expected values to be a list")
lora_list: List[LoRAModulePath] = []
for item in values:
if item in [None, '']: # Skip if item is None or empty string
continue
if '=' in item and ',' not in item: # Old format: name=path
name, path = item.split('=')
lora_list.append(LoRAModulePath(name, path))
else: # Assume JSON format
try:
lora_dict = json.loads(item)
lora = LoRAModulePath(**lora_dict)
lora_list.append(lora)
except json.JSONDecodeError:
parser.error(
f"Invalid JSON format for --lora-modules: {item}")
except TypeError as e:
parser.error(
f"Invalid fields for --lora-modules: {item} - {str(e)}"
)
setattr(namespace, self.dest, lora_list)
class PromptAdapterParserAction(argparse.Action):
def __call__(
self,
parser: argparse.ArgumentParser,
namespace: argparse.Namespace,
values: Optional[Union[str, Sequence[str]]],
option_string: Optional[str] = None,
):
if values is None:
values = []
if isinstance(values, str):
raise TypeError("Expected values to be a list")
adapter_list: List[PromptAdapterPath] = []
for item in values:
name, path = item.split('=')
adapter_list.append(PromptAdapterPath(name, path))
setattr(namespace, self.dest, adapter_list)
def make_arg_parser(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
parser.add_argument("--host",
type=nullable_str,
default=None,
help="host name")
parser.add_argument("--port", type=int, default=8000, help="port number")
parser.add_argument(
"--uvicorn-log-level",
type=str,
default="info",
choices=['debug', 'info', 'warning', 'error', 'critical', 'trace'],
help="log level for uvicorn")
parser.add_argument("--allow-credentials",
action="store_true",
help="allow credentials")
parser.add_argument("--allowed-origins",
type=json.loads,
default=["*"],
help="allowed origins")
parser.add_argument("--allowed-methods",
type=json.loads,
default=["*"],
help="allowed methods")
parser.add_argument("--allowed-headers",
type=json.loads,
default=["*"],
help="allowed headers")
parser.add_argument("--api-key",
type=nullable_str,
default=None,
help="If provided, the server will require this key "
"to be presented in the header.")
parser.add_argument(
"--lora-modules",
type=nullable_str,
default=None,
nargs='+',
action=LoRAParserAction,
help="LoRA module configurations in either 'name=path' format"
"or JSON format. "
"Example (old format): 'name=path' "
"Example (new format): "
"'{\"name\": \"name\", \"local_path\": \"path\", "
"\"base_model_name\": \"id\"}'")
parser.add_argument(
"--prompt-adapters",
type=nullable_str,
default=None,
nargs='+',
action=PromptAdapterParserAction,
help="Prompt adapter configurations in the format name=path. "
"Multiple adapters can be specified.")
parser.add_argument("--chat-template",
type=nullable_str,
default=None,
help="The file path to the chat template, "
"or the template in single-line form "
"for the specified model")
parser.add_argument("--response-role",
type=nullable_str,
default="assistant",
help="The role name to return if "
"`request.add_generation_prompt=true`.")
parser.add_argument("--ssl-keyfile",
type=nullable_str,
default=None,
help="The file path to the SSL key file")
parser.add_argument("--ssl-certfile",
type=nullable_str,
default=None,
help="The file path to the SSL cert file")
parser.add_argument("--ssl-ca-certs",
type=nullable_str,
default=None,
help="The CA certificates file")
parser.add_argument(
"--ssl-cert-reqs",
type=int,
default=int(ssl.CERT_NONE),
help="Whether client certificate is required (see stdlib ssl module's)"
)
parser.add_argument(
"--root-path",
type=nullable_str,
default=None,
help="FastAPI root_path when app is behind a path based routing proxy")
parser.add_argument(
"--middleware",
type=nullable_str,
action="append",
default=[],
help="Additional ASGI middleware to apply to the app. "
"We accept multiple --middleware arguments. "
"The value should be an import path. "
"If a function is provided, vLLM will add it to the server "
"using @app.middleware('http'). "
"If a class is provided, vLLM will add it to the server "
"using app.add_middleware(). ")
parser.add_argument(
"--return-tokens-as-token-ids",
action="store_true",
help="When --max-logprobs is specified, represents single tokens as "
"strings of the form 'token_id:{token_id}' so that tokens that "
"are not JSON-encodable can be identified.")
parser.add_argument(
"--disable-frontend-multiprocessing",
action="store_true",
help="If specified, will run the OpenAI frontend server in the same "
"process as the model serving engine.")
parser.add_argument(
"--enable-auto-tool-choice",
action="store_true",
default=False,
help=
"Enable auto tool choice for supported models. Use --tool-call-parser"
"to specify which parser to use")
valid_tool_parsers = ToolParserManager.tool_parsers.keys()
parser.add_argument(
"--tool-call-parser",
type=str,
metavar="{" + ",".join(valid_tool_parsers) + "} or name registered in "
"--tool-parser-plugin",
default=None,
help=
"Select the tool call parser depending on the model that you're using."
" This is used to parse the model-generated tool call into OpenAI API "
"format. Required for --enable-auto-tool-choice.")
parser.add_argument(
"--tool-parser-plugin",
type=str,
default="",
help=
"Special the tool parser plugin write to parse the model-generated tool"
" into OpenAI API format, the name register in this plugin can be used "
"in --tool-call-parser.")
parser.add_argument(
"--reasoning-parser",
type=str,
default=None,
help=
"Select the reasoning parser to split <think>...</think> content into "
"reasoning_content vs content in the response. "
"Supported: qwen3")
parser = AsyncEngineArgs.add_cli_args(parser)
parser.add_argument('--max-log-len',
type=int,
default=None,
help='Max number of prompt characters or prompt '
'ID numbers being printed in log.'
'\n\nDefault: Unlimited')
parser.add_argument(
"--disable-fastapi-docs",
action='store_true',
default=False,
help="Disable FastAPI's OpenAPI schema, Swagger UI, and ReDoc endpoint"
)
return parser
def validate_parsed_serve_args(args: argparse.Namespace):
"""Quick checks for model serve args that raise prior to loading."""
if hasattr(args, "subparser") and args.subparser != "serve":
return
# Ensure that the chat template is valid; raises if it likely isn't
validate_chat_template(args.chat_template)
# Enable auto tool needs a tool call parser to be valid
if args.enable_auto_tool_choice and not args.tool_call_parser:
raise TypeError("Error: --enable-auto-tool-choice requires "
"--tool-call-parser")
def create_parser_for_docs() -> FlexibleArgumentParser:
parser_for_docs = FlexibleArgumentParser(
prog="-m vllm.entrypoints.openai.api_server")
return make_arg_parser(parser_for_docs)
"""
This file contains the command line arguments for the vLLM's
OpenAI-compatible server. It is kept in a separate file for documentation
purposes.
"""
import argparse
import json
import ssl
from typing import List, Optional, Sequence, Union
from vllm.engine.arg_utils import AsyncEngineArgs, nullable_str
from vllm.entrypoints.chat_utils import validate_chat_template
from vllm.entrypoints.openai.serving_engine import (LoRAModulePath,
PromptAdapterPath)
from vllm.entrypoints.openai.tool_parsers import ToolParserManager
from vllm.utils import FlexibleArgumentParser
class LoRAParserAction(argparse.Action):
def __call__(
self,
parser: argparse.ArgumentParser,
namespace: argparse.Namespace,
values: Optional[Union[str, Sequence[str]]],
option_string: Optional[str] = None,
):
if values is None:
values = []
if isinstance(values, str):
raise TypeError("Expected values to be a list")
lora_list: List[LoRAModulePath] = []
for item in values:
if item in [None, '']: # Skip if item is None or empty string
continue
if '=' in item and ',' not in item: # Old format: name=path
name, path = item.split('=')
lora_list.append(LoRAModulePath(name, path))
else: # Assume JSON format
try:
lora_dict = json.loads(item)
lora = LoRAModulePath(**lora_dict)
lora_list.append(lora)
except json.JSONDecodeError:
parser.error(
f"Invalid JSON format for --lora-modules: {item}")
except TypeError as e:
parser.error(
f"Invalid fields for --lora-modules: {item} - {str(e)}"
)
setattr(namespace, self.dest, lora_list)
class PromptAdapterParserAction(argparse.Action):
def __call__(
self,
parser: argparse.ArgumentParser,
namespace: argparse.Namespace,
values: Optional[Union[str, Sequence[str]]],
option_string: Optional[str] = None,
):
if values is None:
values = []
if isinstance(values, str):
raise TypeError("Expected values to be a list")
adapter_list: List[PromptAdapterPath] = []
for item in values:
name, path = item.split('=')
adapter_list.append(PromptAdapterPath(name, path))
setattr(namespace, self.dest, adapter_list)
def make_arg_parser(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
parser.add_argument("--host",
type=nullable_str,
default=None,
help="host name")
parser.add_argument("--port", type=int, default=8000, help="port number")
parser.add_argument(
"--uvicorn-log-level",
type=str,
default="info",
choices=['debug', 'info', 'warning', 'error', 'critical', 'trace'],
help="log level for uvicorn")
parser.add_argument("--allow-credentials",
action="store_true",
help="allow credentials")
parser.add_argument("--allowed-origins",
type=json.loads,
default=["*"],
help="allowed origins")
parser.add_argument("--allowed-methods",
type=json.loads,
default=["*"],
help="allowed methods")
parser.add_argument("--allowed-headers",
type=json.loads,
default=["*"],
help="allowed headers")
parser.add_argument("--api-key",
type=nullable_str,
default=None,
help="If provided, the server will require this key "
"to be presented in the header.")
parser.add_argument(
"--lora-modules",
type=nullable_str,
default=None,
nargs='+',
action=LoRAParserAction,
help="LoRA module configurations in either 'name=path' format"
"or JSON format. "
"Example (old format): 'name=path' "
"Example (new format): "
"'{\"name\": \"name\", \"local_path\": \"path\", "
"\"base_model_name\": \"id\"}'")
parser.add_argument(
"--prompt-adapters",
type=nullable_str,
default=None,
nargs='+',
action=PromptAdapterParserAction,
help="Prompt adapter configurations in the format name=path. "
"Multiple adapters can be specified.")
parser.add_argument("--chat-template",
type=nullable_str,
default=None,
help="The file path to the chat template, "
"or the template in single-line form "
"for the specified model")
parser.add_argument("--response-role",
type=nullable_str,
default="assistant",
help="The role name to return if "
"`request.add_generation_prompt=true`.")
parser.add_argument("--ssl-keyfile",
type=nullable_str,
default=None,
help="The file path to the SSL key file")
parser.add_argument("--ssl-certfile",
type=nullable_str,
default=None,
help="The file path to the SSL cert file")
parser.add_argument("--ssl-ca-certs",
type=nullable_str,
default=None,
help="The CA certificates file")
parser.add_argument(
"--ssl-cert-reqs",
type=int,
default=int(ssl.CERT_NONE),
help="Whether client certificate is required (see stdlib ssl module's)"
)
parser.add_argument(
"--root-path",
type=nullable_str,
default=None,
help="FastAPI root_path when app is behind a path based routing proxy")
parser.add_argument(
"--middleware",
type=nullable_str,
action="append",
default=[],
help="Additional ASGI middleware to apply to the app. "
"We accept multiple --middleware arguments. "
"The value should be an import path. "
"If a function is provided, vLLM will add it to the server "
"using @app.middleware('http'). "
"If a class is provided, vLLM will add it to the server "
"using app.add_middleware(). ")
parser.add_argument(
"--return-tokens-as-token-ids",
action="store_true",
help="When --max-logprobs is specified, represents single tokens as "
"strings of the form 'token_id:{token_id}' so that tokens that "
"are not JSON-encodable can be identified.")
parser.add_argument(
"--disable-frontend-multiprocessing",
action="store_true",
help="If specified, will run the OpenAI frontend server in the same "
"process as the model serving engine.")
parser.add_argument(
"--enable-auto-tool-choice",
action="store_true",
default=False,
help=
"Enable auto tool choice for supported models. Use --tool-call-parser"
"to specify which parser to use")
valid_tool_parsers = ToolParserManager.tool_parsers.keys()
parser.add_argument(
"--tool-call-parser",
type=str,
metavar="{" + ",".join(valid_tool_parsers) + "} or name registered in "
"--tool-parser-plugin",
default=None,
help=
"Select the tool call parser depending on the model that you're using."
" This is used to parse the model-generated tool call into OpenAI API "
"format. Required for --enable-auto-tool-choice.")
parser.add_argument(
"--tool-parser-plugin",
type=str,
default="",
help=
"Special the tool parser plugin write to parse the model-generated tool"
" into OpenAI API format, the name register in this plugin can be used "
"in --tool-call-parser.")
parser.add_argument(
"--reasoning-parser",
type=str,
default=None,
help=
"Select the reasoning parser to split <think>...</think> content into "
"reasoning_content vs content in the response. "
"Supported: qwen3")
parser = AsyncEngineArgs.add_cli_args(parser)
parser.add_argument('--max-log-len',
type=int,
default=None,
help='Max number of prompt characters or prompt '
'ID numbers being printed in log.'
'\n\nDefault: Unlimited')
parser.add_argument(
"--disable-fastapi-docs",
action='store_true',
default=False,
help="Disable FastAPI's OpenAPI schema, Swagger UI, and ReDoc endpoint"
)
return parser
def validate_parsed_serve_args(args: argparse.Namespace):
"""Quick checks for model serve args that raise prior to loading."""
if hasattr(args, "subparser") and args.subparser != "serve":
return
# Ensure that the chat template is valid; raises if it likely isn't
validate_chat_template(args.chat_template)
# Enable auto tool needs a tool call parser to be valid
if args.enable_auto_tool_choice and not args.tool_call_parser:
raise TypeError("Error: --enable-auto-tool-choice requires "
"--tool-call-parser")
def create_parser_for_docs() -> FlexibleArgumentParser:
parser_for_docs = FlexibleArgumentParser(
prog="-m vllm.entrypoints.openai.api_server")
return make_arg_parser(parser_for_docs)

View File

@@ -0,0 +1,102 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
#include <vector>
namespace {
constexpr int kHeadDim = 256;
constexpr int kThreads = 256;
void check_half_matrix(const torch::Tensor& input, const char* name) {
TORCH_CHECK(input.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(input.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(input.dim() == 2 && input.size(1) == kHeadDim,
name, " must have shape (rows, 256)");
}
__global__ void prepare_kernel(const __half* input, float* converted,
float* squares, int rows) {
const int row = blockIdx.x;
const int column = threadIdx.x;
if (row >= rows || column >= kHeadDim) {
return;
}
const int offset = row * kHeadDim + column;
const float value = __half2float(input[offset]);
converted[offset] = value;
squares[offset] = __fmul_rn(value, value);
}
__global__ void apply_inverse_kernel(
const float* input, const __half* weight, const float* inverse,
__half* output, int rows) {
const int row = blockIdx.x;
const int column = threadIdx.x;
if (row >= rows || column >= kHeadDim) {
return;
}
const int offset = row * kHeadDim + column;
const float scaled = __fmul_rn(input[offset], inverse[row]);
const float factor = __fadd_rn(1.0f, __half2float(weight[column]));
output[offset] = __float2half_rn(__fmul_rn(scaled, factor));
}
} // namespace
std::vector<torch::Tensor> prepare(const torch::Tensor& input) {
check_half_matrix(input, "input");
auto float_options = input.options().dtype(torch::kFloat32);
auto converted = torch::empty(input.sizes(), float_options);
auto squares = torch::empty(input.sizes(), float_options);
const int rows = static_cast<int>(input.size(0));
prepare_kernel<<<rows, kThreads, 0, at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const __half*>(input.data_ptr<at::Half>()),
converted.data_ptr<float>(), squares.data_ptr<float>(), rows);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {converted, squares};
}
torch::Tensor apply_inverse(const torch::Tensor& input,
const torch::Tensor& weight,
const torch::Tensor& inverse) {
TORCH_CHECK(input.is_cuda() && weight.is_cuda() && inverse.is_cuda(),
"all tensors must be CUDA tensors");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"input must have dtype float32");
TORCH_CHECK(weight.scalar_type() == torch::kFloat16,
"weight must have dtype float16");
TORCH_CHECK(inverse.scalar_type() == torch::kFloat32,
"inverse must have dtype float32");
TORCH_CHECK(input.is_contiguous() && weight.is_contiguous()
&& inverse.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(input.dim() == 2 && input.size(1) == kHeadDim,
"input must have shape (rows, 256)");
TORCH_CHECK(weight.dim() == 1 && weight.size(0) == kHeadDim,
"weight must have shape (256,)");
TORCH_CHECK(inverse.numel() == input.size(0),
"inverse must contain one value per row");
auto output = torch::empty(
input.sizes(), input.options().dtype(torch::kFloat16));
const int rows = static_cast<int>(input.size(0));
apply_inverse_kernel<<<rows, kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
input.data_ptr<float>(),
reinterpret_cast<const __half*>(weight.data_ptr<at::Half>()),
inverse.data_ptr<float>(),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()), rows);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("prepare", &prepare,
"Convert FP16 attention heads and compute exact squares");
module.def("apply_inverse", &apply_inverse,
"Apply PyTorch-computed attention head RMSNorm inverse");
}

View File

@@ -0,0 +1,402 @@
#include <ATen/ATen.h>
#include <ATen/Parallel.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <algorithm>
#include <cstdint>
#include <cstring>
#include <vector>
namespace {
constexpr int kAttentionLayers = 10;
constexpr int kKvPlanes = 2;
constexpr int kElementsPerPlaneBlock = 4096;
constexpr int kElementsPerVector = 8;
constexpr int kVectorsPerPlaneBlock =
kElementsPerPlaneBlock / kElementsPerVector;
constexpr int kVectorsPerBlockMajorRow =
kAttentionLayers * kKvPlanes * kVectorsPerPlaneBlock;
constexpr int kThreads = 256;
constexpr int kMaxGridBlocks = 65535;
using PackedVector = uint4;
__device__ __forceinline__ const PackedVector* select_const_layer(
int layer, const PackedVector* layer0, const PackedVector* layer1,
const PackedVector* layer2, const PackedVector* layer3,
const PackedVector* layer4, const PackedVector* layer5,
const PackedVector* layer6, const PackedVector* layer7,
const PackedVector* layer8, const PackedVector* layer9) {
switch (layer) {
case 0:
return layer0;
case 1:
return layer1;
case 2:
return layer2;
case 3:
return layer3;
case 4:
return layer4;
case 5:
return layer5;
case 6:
return layer6;
case 7:
return layer7;
case 8:
return layer8;
default:
return layer9;
}
}
__device__ __forceinline__ PackedVector* select_mutable_layer(
int layer, PackedVector* layer0, PackedVector* layer1,
PackedVector* layer2, PackedVector* layer3, PackedVector* layer4,
PackedVector* layer5, PackedVector* layer6, PackedVector* layer7,
PackedVector* layer8, PackedVector* layer9) {
switch (layer) {
case 0:
return layer0;
case 1:
return layer1;
case 2:
return layer2;
case 3:
return layer3;
case 4:
return layer4;
case 5:
return layer5;
case 6:
return layer6;
case 7:
return layer7;
case 8:
return layer8;
default:
return layer9;
}
}
__global__ void pack_block_major_kernel(
const PackedVector* layer0, const PackedVector* layer1,
const PackedVector* layer2, const PackedVector* layer3,
const PackedVector* layer4, const PackedVector* layer5,
const PackedVector* layer6, const PackedVector* layer7,
const PackedVector* layer8, const PackedVector* layer9,
const int* source_blocks, PackedVector* staging, int* error_flag,
int count, int gpu_blocks) {
const int64_t total =
static_cast<int64_t>(count) * kVectorsPerBlockMajorRow;
for (int64_t linear =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
linear < total;
linear += static_cast<int64_t>(blockDim.x) * gridDim.x) {
int64_t cursor = linear;
const int feature_vector = cursor % kVectorsPerPlaneBlock;
cursor /= kVectorsPerPlaneBlock;
const int kv_plane = cursor % kKvPlanes;
cursor /= kKvPlanes;
const int layer = cursor % kAttentionLayers;
const int row = cursor / kAttentionLayers;
const int source_block = source_blocks[row];
if (static_cast<unsigned int>(source_block) >=
static_cast<unsigned int>(gpu_blocks)) {
atomicExch(error_flag, 1);
continue;
}
const PackedVector* source = select_const_layer(
layer, layer0, layer1, layer2, layer3, layer4, layer5, layer6,
layer7, layer8, layer9);
const int64_t source_index =
((static_cast<int64_t>(kv_plane) * gpu_blocks + source_block)
* kVectorsPerPlaneBlock) +
feature_vector;
staging[linear] = source[source_index];
}
}
__global__ void scatter_block_major_kernel(
const PackedVector* staging, const int* destination_blocks,
PackedVector* layer0, PackedVector* layer1, PackedVector* layer2,
PackedVector* layer3, PackedVector* layer4, PackedVector* layer5,
PackedVector* layer6, PackedVector* layer7, PackedVector* layer8,
PackedVector* layer9, int* error_flag, int count, int gpu_blocks) {
const int64_t total =
static_cast<int64_t>(count) * kVectorsPerBlockMajorRow;
for (int64_t linear =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
linear < total;
linear += static_cast<int64_t>(blockDim.x) * gridDim.x) {
int64_t cursor = linear;
const int feature_vector = cursor % kVectorsPerPlaneBlock;
cursor /= kVectorsPerPlaneBlock;
const int kv_plane = cursor % kKvPlanes;
cursor /= kKvPlanes;
const int layer = cursor % kAttentionLayers;
const int row = cursor / kAttentionLayers;
const int destination_block = destination_blocks[row];
if (static_cast<unsigned int>(destination_block) >=
static_cast<unsigned int>(gpu_blocks)) {
atomicExch(error_flag, 1);
continue;
}
PackedVector* destination = select_mutable_layer(
layer, layer0, layer1, layer2, layer3, layer4, layer5, layer6,
layer7, layer8, layer9);
const int64_t destination_index =
((static_cast<int64_t>(kv_plane) * gpu_blocks + destination_block)
* kVectorsPerPlaneBlock) +
feature_vector;
destination[destination_index] = staging[linear];
}
}
void check_gpu_layers(const std::vector<torch::Tensor>& layers) {
TORCH_CHECK(layers.size() == kAttentionLayers, "expected exactly ",
kAttentionLayers, " GPU attention-layer tensors");
const auto device = layers.front().device();
const int64_t blocks = layers.front().size(1);
for (int layer = 0; layer < kAttentionLayers; ++layer) {
const auto& tensor = layers[layer];
TORCH_CHECK(tensor.is_cuda(), "GPU layer ", layer,
" must be a CUDA tensor");
TORCH_CHECK(tensor.device() == device, "GPU layer ", layer,
" is on a different device");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16, "GPU layer ",
layer, " must use float16");
TORCH_CHECK(tensor.is_contiguous(), "GPU layer ", layer,
" must be contiguous");
TORCH_CHECK(tensor.dim() == 3 && tensor.size(0) == kKvPlanes &&
tensor.size(1) == blocks &&
tensor.size(2) == kElementsPerPlaneBlock,
"GPU layer ", layer, " must have shape [2, blocks, 4096]");
TORCH_CHECK(
reinterpret_cast<uintptr_t>(tensor.data_ptr<at::Half>()) %
alignof(PackedVector) ==
0,
"GPU layer ", layer, " is not 16-byte aligned");
}
}
void check_gpu_transfer_args(const std::vector<torch::Tensor>& layers,
const torch::Tensor& block_ids,
const torch::Tensor& staging,
const torch::Tensor& error_flag,
int64_t count) {
check_gpu_layers(layers);
TORCH_CHECK(block_ids.is_cuda(), "block_ids must be a CUDA tensor");
TORCH_CHECK(block_ids.device() == layers.front().device(),
"block_ids must be on the cache device");
TORCH_CHECK(block_ids.scalar_type() == torch::kInt32,
"block_ids must use int32");
TORCH_CHECK(block_ids.dim() == 1 && block_ids.is_contiguous(),
"block_ids must be a contiguous one-dimensional tensor");
TORCH_CHECK(count > 0 && count <= block_ids.numel(),
"count must be in [1, block_ids.numel()]");
TORCH_CHECK(staging.is_cuda(), "staging must be a CUDA tensor");
TORCH_CHECK(staging.device() == layers.front().device(),
"staging must be on the cache device");
TORCH_CHECK(staging.scalar_type() == torch::kFloat16,
"staging must use float16");
TORCH_CHECK(staging.is_contiguous(), "staging must be contiguous");
TORCH_CHECK(
staging.dim() == 4 && staging.size(0) >= count &&
staging.size(1) == kAttentionLayers &&
staging.size(2) == kKvPlanes &&
staging.size(3) == kElementsPerPlaneBlock,
"staging must have shape [capacity>=count, 10, 2, 4096]");
TORCH_CHECK(
reinterpret_cast<uintptr_t>(staging.data_ptr<at::Half>()) %
alignof(PackedVector) ==
0,
"staging is not 16-byte aligned");
TORCH_CHECK(error_flag.is_cuda(),
"error_flag must be a CUDA tensor");
TORCH_CHECK(error_flag.device() == layers.front().device(),
"error_flag must be on the cache device");
TORCH_CHECK(error_flag.scalar_type() == torch::kInt32,
"error_flag must use int32");
TORCH_CHECK(error_flag.is_contiguous() && error_flag.numel() == 1,
"error_flag must be one contiguous int32 value");
}
int launch_blocks(int64_t count) {
const int64_t total = count * kVectorsPerBlockMajorRow;
return static_cast<int>(std::min<int64_t>(
(total + kThreads - 1) / kThreads, kMaxGridBlocks));
}
void pack_block_major(const std::vector<torch::Tensor>& layers,
const torch::Tensor& source_blocks,
torch::Tensor staging, torch::Tensor error_flag,
int64_t count) {
check_gpu_transfer_args(
layers, source_blocks, staging, error_flag, count);
const int blocks = static_cast<int>(layers.front().size(1));
pack_block_major_kernel<<<launch_blocks(count), kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const PackedVector*>(
layers[0].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[1].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[2].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[3].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[4].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[5].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[6].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[7].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[8].data_ptr<at::Half>()),
reinterpret_cast<const PackedVector*>(
layers[9].data_ptr<at::Half>()),
source_blocks.data_ptr<int>(),
reinterpret_cast<PackedVector*>(staging.data_ptr<at::Half>()),
error_flag.data_ptr<int>(), static_cast<int>(count), blocks);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void scatter_block_major(const torch::Tensor& staging,
const torch::Tensor& destination_blocks,
const std::vector<torch::Tensor>& layers,
torch::Tensor error_flag,
int64_t count) {
check_gpu_transfer_args(
layers, destination_blocks, staging, error_flag, count);
const int blocks = static_cast<int>(layers.front().size(1));
scatter_block_major_kernel<<<launch_blocks(count), kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const PackedVector*>(
staging.data_ptr<at::Half>()),
destination_blocks.data_ptr<int>(),
reinterpret_cast<PackedVector*>(layers[0].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[1].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[2].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[3].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[4].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[5].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[6].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[7].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[8].data_ptr<at::Half>()),
reinterpret_cast<PackedVector*>(layers[9].data_ptr<at::Half>()),
error_flag.data_ptr<int>(), static_cast<int>(count), blocks);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void check_transfer_error(const torch::Tensor& error_flag) {
TORCH_CHECK(error_flag.is_cuda(),
"error_flag must be a CUDA tensor");
TORCH_CHECK(error_flag.scalar_type() == torch::kInt32,
"error_flag must use int32");
TORCH_CHECK(error_flag.is_contiguous() && error_flag.numel() == 1,
"error_flag must be one contiguous int32 value");
TORCH_CHECK(error_flag.item<int>() == 0,
"GPU block mapping contains an out-of-range id");
}
void check_cpu_transfer_args(const torch::Tensor& pool,
const torch::Tensor& block_ids,
const torch::Tensor& staging, int64_t count) {
TORCH_CHECK(!pool.is_cuda() && !staging.is_cuda() &&
!block_ids.is_cuda(),
"CPU gather/scatter tensors must be on CPU");
TORCH_CHECK(pool.scalar_type() == torch::kFloat16 &&
staging.scalar_type() == torch::kFloat16,
"CPU pool and staging must use float16");
TORCH_CHECK(pool.is_contiguous() && staging.is_contiguous(),
"CPU pool and staging must be contiguous");
TORCH_CHECK(
pool.dim() == 4 && pool.size(1) == kAttentionLayers &&
pool.size(2) == kKvPlanes &&
pool.size(3) == kElementsPerPlaneBlock,
"CPU pool must have shape [slots, 10, 2, 4096]");
TORCH_CHECK(
staging.dim() == 4 && staging.size(0) >= count &&
staging.size(1) == kAttentionLayers &&
staging.size(2) == kKvPlanes &&
staging.size(3) == kElementsPerPlaneBlock,
"CPU staging must have shape [capacity>=count, 10, 2, 4096]");
TORCH_CHECK(block_ids.scalar_type() == torch::kInt64,
"CPU block_ids must use int64");
TORCH_CHECK(block_ids.dim() == 1 && block_ids.is_contiguous(),
"CPU block_ids must be contiguous and one-dimensional");
TORCH_CHECK(count > 0 && count <= block_ids.numel(),
"count must be in [1, block_ids.numel()]");
const int64_t* ids = block_ids.data_ptr<int64_t>();
for (int64_t row = 0; row < count; ++row) {
TORCH_CHECK(ids[row] >= 0 && ids[row] < pool.size(0),
"CPU block id out of range at row ", row, ": ", ids[row]);
}
}
void cpu_gather_rows(const torch::Tensor& pool,
const torch::Tensor& source_blocks,
torch::Tensor staging, int64_t count) {
check_cpu_transfer_args(pool, source_blocks, staging, count);
const int64_t row_elements =
kAttentionLayers * kKvPlanes * kElementsPerPlaneBlock;
const size_t row_bytes =
static_cast<size_t>(row_elements) * sizeof(at::Half);
const char* source = reinterpret_cast<const char*>(
pool.data_ptr<at::Half>());
char* destination =
reinterpret_cast<char*>(staging.data_ptr<at::Half>());
const int64_t* ids = source_blocks.data_ptr<int64_t>();
at::parallel_for(0, count, 8, [&](int64_t begin, int64_t end) {
for (int64_t row = begin; row < end; ++row) {
std::memcpy(destination + row * row_bytes,
source + ids[row] * row_bytes, row_bytes);
}
});
}
void cpu_scatter_rows(const torch::Tensor& staging,
torch::Tensor pool,
const torch::Tensor& destination_blocks,
int64_t count) {
check_cpu_transfer_args(pool, destination_blocks, staging, count);
const int64_t row_elements =
kAttentionLayers * kKvPlanes * kElementsPerPlaneBlock;
const size_t row_bytes =
static_cast<size_t>(row_elements) * sizeof(at::Half);
const char* source = reinterpret_cast<const char*>(
staging.data_ptr<at::Half>());
char* destination =
reinterpret_cast<char*>(pool.data_ptr<at::Half>());
const int64_t* ids = destination_blocks.data_ptr<int64_t>();
at::parallel_for(0, count, 8, [&](int64_t begin, int64_t end) {
for (int64_t row = begin; row < end; ++row) {
std::memcpy(destination + ids[row] * row_bytes,
source + row * row_bytes, row_bytes);
}
});
}
} // namespace
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("pack", &pack_block_major,
"Pack ten layer-major FP16 KV caches into block-major staging");
module.def("scatter", &scatter_block_major,
"Scatter block-major FP16 staging into ten layer-major caches");
module.def("check_error", &check_transfer_error,
"Fail fast after a bounds-safe asynchronous transfer");
module.def("cpu_gather", &cpu_gather_rows,
"Gather block-major CPU pool rows into bounded staging");
module.def("cpu_scatter", &cpu_scatter_rows,
"Scatter bounded staging rows into the block-major CPU pool");
}

View File

@@ -0,0 +1,494 @@
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cublas_v2.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <limits>
#include <vector>
namespace {
constexpr int kBlockSize = 16;
constexpr int kHeadDim = 256;
constexpr int kKeyPack = 8;
constexpr int kNumQueryHeads = 4;
constexpr int kNumKvHeads = 1;
constexpr int kTileTokens = 512;
constexpr int kSplitCount = 4;
constexpr int kGroupTokens = kSplitCount * kTileTokens;
constexpr int kThreads = 256;
constexpr int kMaxQueryTokens = 8192;
constexpr int kMaxSequenceTokens = 262144;
void check_half_cuda_contiguous(const torch::Tensor& tensor,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
}
__global__ void convert_query_kernel(const __half* query, float* converted,
int query_len, float scale) {
const int64_t total = static_cast<int64_t>(query_len)
* kNumQueryHeads * kHeadDim;
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < total;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
const int dim = index % kHeadDim;
const int query_index =
(index / kHeadDim) % query_len;
const int head =
index / (static_cast<int64_t>(kHeadDim) * query_len);
const int64_t source =
(static_cast<int64_t>(query_index) * kNumQueryHeads + head)
* kHeadDim + dim;
converted[index] = __half2float(query[source]) * scale;
}
}
__global__ void gather_kv_group_kernel(
const __half* key_new, const __half* value_new,
const __half* key_cache, const __half* value_cache,
const int* block_table, float* key_tiles, float* value_tiles,
int context_len, int query_len, int group_start, int group_tokens,
int active_splits) {
constexpr int kElements = kTileTokens * kHeadDim;
const int64_t total = static_cast<int64_t>(active_splits) * kElements;
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < total;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
const int split = index / kElements;
const int element = index - static_cast<int64_t>(split) * kElements;
const int token_offset = element / kHeadDim;
const int dim = element - token_offset * kHeadDim;
const int remaining_tokens = group_tokens - split * kTileTokens;
const int split_tokens =
remaining_tokens < kTileTokens ? remaining_tokens : kTileTokens;
const int logical_token =
group_start + split * kTileTokens + token_offset;
float key_value = 0.0f;
float value_value = 0.0f;
if (token_offset >= split_tokens) {
// The fixed 512-column GEMMs require zero-filled tail columns.
} else if (logical_token < context_len) {
const int logical_block = logical_token / kBlockSize;
const int block_offset = logical_token % kBlockSize;
const int physical_block = block_table[logical_block];
const int64_t key_index =
(((static_cast<int64_t>(physical_block) * kNumKvHeads)
* (kHeadDim / kKeyPack) + dim / kKeyPack)
* kBlockSize + block_offset) * kKeyPack + dim % kKeyPack;
const int64_t value_index =
((static_cast<int64_t>(physical_block) * kNumKvHeads)
* kHeadDim + dim) * kBlockSize + block_offset;
key_value = __half2float(key_cache[key_index]);
value_value = __half2float(value_cache[value_index]);
} else if (logical_token < context_len + query_len) {
const int query_index = logical_token - context_len;
const int64_t source =
static_cast<int64_t>(query_index) * kHeadDim + dim;
key_value = __half2float(key_new[source]);
value_value = __half2float(value_new[source]);
}
key_tiles[index] = key_value;
value_tiles[index] = value_value;
}
}
__global__ void mask_group_scores_kernel(
float* scores, int query_len, int context_len,
int group_start, int group_tokens, int active_splits,
int rows, bool causal) {
const int64_t split_elements =
static_cast<int64_t>(rows) * kTileTokens;
const int64_t elements = active_splits * split_elements;
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < elements;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
const int split = index / split_elements;
const int split_index = index - split * split_elements;
const int column = split_index % kTileTokens;
const int row = split_index / kTileTokens;
const int query_index = row % query_len;
const int remaining_tokens = group_tokens - split * kTileTokens;
const int split_tokens =
remaining_tokens < kTileTokens ? remaining_tokens : kTileTokens;
const int logical_token =
group_start + split * kTileTokens + column;
if (column >= split_tokens
|| (causal && logical_token > context_len + query_index)) {
scores[index] = -std::numeric_limits<float>::infinity();
}
}
}
__global__ void normalize_split_scores_kernel(
float* scores, float* corrections, float* running_max,
float* running_sum, int active_splits, int rows) {
const int row = blockIdx.x;
if (row >= rows) {
return;
}
__shared__ float reduction[kThreads];
__shared__ float state_max;
__shared__ float state_sum;
__shared__ float next_max;
__shared__ float correction;
if (threadIdx.x == 0) {
state_max = running_max[row];
state_sum = running_sum[row];
}
__syncthreads();
for (int split = 0; split < active_splits; ++split) {
float* row_scores =
scores + (static_cast<int64_t>(split) * rows + row) * kTileTokens;
float local_max = -std::numeric_limits<float>::infinity();
for (int column = threadIdx.x; column < kTileTokens;
column += blockDim.x) {
local_max = fmaxf(local_max, row_scores[column]);
}
reduction[threadIdx.x] = local_max;
__syncthreads();
for (int stride = kThreads / 2; stride > 0; stride /= 2) {
if (threadIdx.x < stride) {
reduction[threadIdx.x] = fmaxf(
reduction[threadIdx.x], reduction[threadIdx.x + stride]);
}
__syncthreads();
}
if (threadIdx.x == 0) {
next_max = fmaxf(state_max, reduction[0]);
correction =
(state_max == -std::numeric_limits<float>::infinity()
&& next_max == -std::numeric_limits<float>::infinity())
? 1.0f
: expf(state_max - next_max);
corrections[static_cast<int64_t>(split) * rows + row] = correction;
}
__syncthreads();
float local_sum = 0.0f;
for (int column = threadIdx.x; column < kTileTokens;
column += blockDim.x) {
const float score = row_scores[column];
const float probability =
(score == -std::numeric_limits<float>::infinity()
&& next_max == -std::numeric_limits<float>::infinity())
? 0.0f
: expf(score - next_max);
row_scores[column] = probability;
local_sum = __fadd_rn(local_sum, probability);
}
reduction[threadIdx.x] = local_sum;
__syncthreads();
for (int stride = kThreads / 2; stride > 0; stride /= 2) {
if (threadIdx.x < stride) {
reduction[threadIdx.x] = __fadd_rn(
reduction[threadIdx.x], reduction[threadIdx.x + stride]);
}
__syncthreads();
}
if (threadIdx.x == 0) {
state_sum = __fadd_rn(
__fmul_rn(state_sum, correction), reduction[0]);
state_max = next_max;
}
__syncthreads();
}
if (threadIdx.x == 0) {
running_max[row] = state_max;
running_sum[row] = state_sum;
}
}
__global__ void merge_split_output_kernel(
float* running_output, const float* split_output,
const float* corrections, int active_splits,
int rows, int64_t output_elements) {
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < output_elements;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
const int row = index / kHeadDim;
float value = running_output[index];
for (int split = 0; split < active_splits; ++split) {
const int64_t row_index =
static_cast<int64_t>(split) * rows + row;
const int64_t output_index =
static_cast<int64_t>(split) * output_elements + index;
value = __fadd_rn(
__fmul_rn(value, corrections[row_index]),
split_output[output_index]);
}
running_output[index] = value;
}
}
__global__ void accumulate_output_kernel(
float* running_output, const float* tile_output,
const float* correction, int64_t elements) {
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < elements;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
const int row = index / kHeadDim;
const float scaled =
__fmul_rn(running_output[index], correction[row]);
running_output[index] = __fadd_rn(scaled, tile_output[index]);
}
}
int launch_blocks(int64_t elements) {
const int64_t needed = (elements + kThreads - 1) / kThreads;
return static_cast<int>(std::min<int64_t>(needed, 65535));
}
void check_cublas(cublasStatus_t status, const char* operation) {
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, operation,
" failed with cuBLAS status ", static_cast<int>(status));
}
cublasStatus_t qk_batched(
cublasHandle_t handle, const float* key_tile, const float* query,
float* scores, int query_len) {
const float alpha = 1.0f;
const float beta = 0.0f;
return cublasSgemmStridedBatched(
handle, CUBLAS_OP_T, CUBLAS_OP_N,
kTileTokens, query_len, kHeadDim,
&alpha, key_tile, kHeadDim, 0,
query, kHeadDim, static_cast<long long>(query_len) * kHeadDim,
&beta, scores, kTileTokens,
static_cast<long long>(query_len) * kTileTokens,
kNumQueryHeads);
}
cublasStatus_t pv_batched(
cublasHandle_t handle, const float* value_tile, const float* scores,
float* output, int query_len) {
const float alpha = 1.0f;
const float beta = 0.0f;
return cublasSgemmStridedBatched(
handle, CUBLAS_OP_N, CUBLAS_OP_N,
kHeadDim, query_len, kTileTokens,
&alpha, value_tile, kHeadDim, 0,
scores, kTileTokens,
static_cast<long long>(query_len) * kTileTokens,
&beta, output, kHeadDim,
static_cast<long long>(query_len) * kHeadDim,
kNumQueryHeads);
}
} // namespace
std::vector<torch::Tensor> fused_paged_prefill_forward(
const torch::Tensor& query, const torch::Tensor& key_new,
const torch::Tensor& value_new, const torch::Tensor& key_cache,
const torch::Tensor& value_cache, const torch::Tensor& block_table,
int64_t context_len_arg, double scale_arg) {
check_half_cuda_contiguous(query, "query");
check_half_cuda_contiguous(key_new, "key_new");
check_half_cuda_contiguous(value_new, "value_new");
check_half_cuda_contiguous(key_cache, "key_cache");
check_half_cuda_contiguous(value_cache, "value_cache");
TORCH_CHECK(block_table.is_cuda(),
"block_table must be a CUDA tensor");
TORCH_CHECK(block_table.scalar_type() == torch::kInt32,
"block_table must have dtype int32");
TORCH_CHECK(block_table.is_contiguous(),
"block_table must be contiguous");
TORCH_CHECK(block_table.dim() == 1,
"block_table must be one-dimensional");
TORCH_CHECK(query.dim() == 3 && query.size(1) == kNumQueryHeads
&& query.size(2) == kHeadDim,
"query must have shape (Q, 4, 256)");
TORCH_CHECK(key_new.dim() == 3 && key_new.size(1) == kNumKvHeads
&& key_new.size(2) == kHeadDim,
"key_new must have shape (Q, 1, 256)");
TORCH_CHECK(value_new.sizes() == key_new.sizes(),
"value_new must match key_new");
TORCH_CHECK(key_new.size(0) == query.size(0),
"query, key_new, and value_new lengths must match");
TORCH_CHECK(key_cache.dim() == 5
&& key_cache.size(1) == kNumKvHeads
&& key_cache.size(2) == kHeadDim / kKeyPack
&& key_cache.size(3) == kBlockSize
&& key_cache.size(4) == kKeyPack,
"key_cache must have shape (N, 1, 32, 16, 8)");
TORCH_CHECK(value_cache.dim() == 4
&& value_cache.size(1) == kNumKvHeads
&& value_cache.size(2) == kHeadDim
&& value_cache.size(3) == kBlockSize,
"value_cache must have shape (N, 1, 256, 16)");
TORCH_CHECK(key_cache.size(0) == value_cache.size(0),
"key/value cache block counts must match");
TORCH_CHECK(query.device() == key_new.device()
&& query.device() == value_new.device()
&& query.device() == key_cache.device()
&& query.device() == value_cache.device()
&& query.device() == block_table.device(),
"all tensors must use the same device");
TORCH_CHECK(context_len_arg >= 0
&& context_len_arg <= kMaxSequenceTokens,
"context_len is out of range");
TORCH_CHECK(context_len_arg % kBlockSize == 0,
"context_len must be block aligned");
const int query_len = static_cast<int>(query.size(0));
const int context_len = static_cast<int>(context_len_arg);
TORCH_CHECK(query_len > 0 && query_len <= kMaxQueryTokens,
"query length must be in [1, 8192]");
TORCH_CHECK(context_len + query_len <= kMaxSequenceTokens,
"context_len + query_len exceeds 262144");
const int required_blocks =
(context_len + kBlockSize - 1) / kBlockSize;
TORCH_CHECK(block_table.numel() >= required_blocks,
"block_table is too short for context_len");
if (required_blocks > 0) {
auto active_blocks = block_table.narrow(0, 0, required_blocks);
const int minimum_block = active_blocks.min().item<int>();
const int maximum_block = active_blocks.max().item<int>();
TORCH_CHECK(minimum_block >= 0
&& maximum_block < key_cache.size(0),
"block_table contains an out-of-range physical block ID");
}
TORCH_CHECK(std::isfinite(scale_arg) && scale_arg > 0.0,
"scale must be finite and positive");
TORCH_CHECK(query_len <= std::numeric_limits<int>::max() / kNumQueryHeads,
"query length overflows row count");
const int rows = kNumQueryHeads * query_len;
const int64_t output_elements =
static_cast<int64_t>(rows) * kHeadDim;
auto float_options = query.options().dtype(torch::kFloat32);
auto converted_query = torch::empty(
{kNumQueryHeads, query_len, kHeadDim}, float_options);
auto key_tiles = torch::empty(
{kSplitCount, kTileTokens, kHeadDim}, float_options);
auto value_tiles = torch::empty(
{kSplitCount, kTileTokens, kHeadDim}, float_options);
auto scores = torch::empty(
{kSplitCount, kNumQueryHeads, query_len, kTileTokens},
float_options);
auto split_output = torch::empty(
{kSplitCount, kNumQueryHeads, query_len, kHeadDim},
float_options);
auto running_max = torch::full(
{kNumQueryHeads, query_len},
-std::numeric_limits<float>::infinity(), float_options);
auto running_sum = torch::zeros(
{kNumQueryHeads, query_len}, float_options);
auto running_output = torch::zeros(
{kNumQueryHeads, query_len, kHeadDim}, float_options);
auto corrections = torch::empty(
{kSplitCount, kNumQueryHeads, query_len}, float_options);
auto stream = at::cuda::getCurrentCUDAStream();
convert_query_kernel<<<launch_blocks(output_elements), kThreads, 0, stream>>>(
reinterpret_cast<const __half*>(query.data_ptr<at::Half>()),
converted_query.data_ptr<float>(), query_len,
static_cast<float>(scale_arg));
C10_CUDA_KERNEL_LAUNCH_CHECK();
cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
check_cublas(cublasSetStream(handle, stream), "cublasSetStream");
const int64_t key_split_stride =
static_cast<int64_t>(kTileTokens) * kHeadDim;
const int64_t score_split_stride =
static_cast<int64_t>(rows) * kTileTokens;
const int64_t output_split_stride = output_elements;
const auto run_group = [&](int group_start, int group_tokens,
bool causal) {
const int active_splits =
(group_tokens + kTileTokens - 1) / kTileTokens;
TORCH_CHECK(active_splits > 0 && active_splits <= kSplitCount,
"invalid split count for paged-prefill group");
constexpr int kGatherBlocks = 512;
gather_kv_group_kernel<<<kGatherBlocks, kThreads, 0, stream>>>(
reinterpret_cast<const __half*>(key_new.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(value_new.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(key_cache.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(value_cache.data_ptr<at::Half>()),
block_table.data_ptr<int>(), key_tiles.data_ptr<float>(),
value_tiles.data_ptr<float>(), context_len, query_len, group_start,
group_tokens, active_splits);
C10_CUDA_KERNEL_LAUNCH_CHECK();
for (int split = 0; split < active_splits; ++split) {
check_cublas(qk_batched(
handle,
key_tiles.data_ptr<float>() + split * key_split_stride,
converted_query.data_ptr<float>(),
scores.data_ptr<float>() + split * score_split_stride,
query_len), "split4 paged prefill QK");
}
const bool needs_mask =
causal || group_tokens != active_splits * kTileTokens;
if (needs_mask) {
const int64_t score_elements =
static_cast<int64_t>(active_splits) * score_split_stride;
mask_group_scores_kernel<<<
launch_blocks(score_elements), kThreads, 0, stream>>>(
scores.data_ptr<float>(), query_len, context_len, group_start,
group_tokens, active_splits, rows, causal);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
normalize_split_scores_kernel<<<rows, kThreads, 0, stream>>>(
scores.data_ptr<float>(), corrections.data_ptr<float>(),
running_max.data_ptr<float>(), running_sum.data_ptr<float>(),
active_splits, rows);
C10_CUDA_KERNEL_LAUNCH_CHECK();
for (int split = 0; split < active_splits; ++split) {
check_cublas(pv_batched(
handle,
value_tiles.data_ptr<float>() + split * key_split_stride,
scores.data_ptr<float>() + split * score_split_stride,
split_output.data_ptr<float>() + split * output_split_stride,
query_len), "split4 paged prefill PV");
}
merge_split_output_kernel<<<
launch_blocks(output_elements), kThreads, 0, stream>>>(
running_output.data_ptr<float>(), split_output.data_ptr<float>(),
corrections.data_ptr<float>(), active_splits, rows,
output_elements);
C10_CUDA_KERNEL_LAUNCH_CHECK();
};
for (int group_start = 0; group_start < context_len;
group_start += kGroupTokens) {
run_group(group_start,
std::min(kGroupTokens, context_len - group_start), false);
}
for (int key_start = 0; key_start < query_len;
key_start += kGroupTokens) {
run_group(context_len + key_start,
std::min(kGroupTokens, query_len - key_start), true);
}
running_output.div_(running_sum.unsqueeze(-1));
auto output = running_output.permute({1, 0, 2})
.to(query.scalar_type()).contiguous();
auto lse = (running_max + at::log(running_sum))
.transpose(0, 1).contiguous();
return {output, lse};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("forward", &fused_paged_prefill_forward,
"Fixed-shape FP32 paged-prefill pipeline for cache-only context");
}

View File

@@ -0,0 +1,84 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
__global__ void beta_decay_kernel(const half* beta_input,
const half* decay_input,
const half* a_log,
const half* dt_bias,
float* output, int elements,
int heads) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index >= elements) {
return;
}
const int head = index % heads;
const float beta_value = __half2float(beta_input[index]);
const float beta_fp32 = 1.0f / (1.0f + expf(-beta_value));
output[index] = __half2float(__float2half(beta_fp32));
const float x = (__half2float(decay_input[index])
+ __half2float(dt_bias[head]));
const float softplus = x > 20.0f ? x : log1pf(expf(x));
output[elements + index] = expf(
-expf(__half2float(a_log[head])) * softplus);
}
void check_half(const torch::Tensor& tensor, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == 2, name, " must have shape (batch, heads)");
}
void check_half_vector(const torch::Tensor& tensor, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == 1, name, " must have shape (heads)");
}
} // namespace
torch::Tensor beta_decay(const torch::Tensor& beta_input,
const torch::Tensor& decay_input,
const torch::Tensor& a_log,
const torch::Tensor& dt_bias) {
check_half(beta_input, "beta_input");
check_half(decay_input, "decay_input");
check_half_vector(a_log, "a_log");
check_half_vector(dt_bias, "dt_bias");
TORCH_CHECK(beta_input.sizes() == decay_input.sizes(),
"beta_input and decay_input shapes must match");
TORCH_CHECK(beta_input.size(1) == a_log.size(0) &&
a_log.sizes() == dt_bias.sizes(),
"parameter heads must match input heads");
const int elements = static_cast<int>(beta_input.numel());
const int heads = static_cast<int>(beta_input.size(1));
torch::Tensor output = torch::empty(
{2, beta_input.size(0), beta_input.size(1)},
beta_input.options().dtype(torch::kFloat32));
constexpr int threads = 128;
const int blocks = (elements + threads - 1) / threads;
beta_decay_kernel<<<blocks, threads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const half*>(beta_input.data_ptr<at::Half>()),
reinterpret_cast<const half*>(decay_input.data_ptr<at::Half>()),
reinterpret_cast<const half*>(a_log.data_ptr<at::Half>()),
reinterpret_cast<const half*>(dt_bias.data_ptr<at::Half>()),
output.data_ptr<float>(), elements, heads);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("beta_decay", &beta_decay,
"Fused GDN beta sigmoid and decay factor");
}

View File

@@ -0,0 +1,89 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
constexpr int kStateLen = 3;
constexpr int kKernelSize = kStateLen + 1;
constexpr int kThreads = 256;
__global__ void causal_conv_update_kernel(
float* state, const __half* hidden, const __half* weight,
__half* output, int channels) {
const int channel = blockIdx.x * blockDim.x + threadIdx.x;
const int batch = blockIdx.y;
if (channel >= channels) {
return;
}
const int state_offset = (batch * channels + channel) * kStateLen;
const int vector_offset = batch * channels + channel;
const int weight_offset = channel * kKernelSize;
const __half current = hidden[vector_offset];
const __half state0 = __float2half_rn(state[state_offset]);
const __half state1 = __float2half_rn(state[state_offset + 1]);
const __half state2 = __float2half_rn(state[state_offset + 2]);
float value = __half2float(state0) * __half2float(weight[weight_offset]);
value += __half2float(state1) * __half2float(weight[weight_offset + 1]);
value += __half2float(state2) * __half2float(weight[weight_offset + 2]);
value += __half2float(current) * __half2float(weight[weight_offset + 3]);
state[state_offset] = __half2float(state1);
state[state_offset + 1] = __half2float(state2);
state[state_offset + 2] = __half2float(current);
const __half convolved = __float2half_rn(value);
const float activation_input = __half2float(convolved);
output[vector_offset] = __float2half_rn(
activation_input / (1.0f + expf(-activation_input)));
}
void check_half_cuda_contiguous(const torch::Tensor& tensor,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
}
} // namespace
torch::Tensor causal_conv_update(torch::Tensor state,
const torch::Tensor& hidden,
const torch::Tensor& weight) {
TORCH_CHECK(state.is_cuda(), "state must be a CUDA tensor");
TORCH_CHECK(state.scalar_type() == torch::kFloat32,
"state must have dtype float32");
TORCH_CHECK(state.is_contiguous(), "state must be contiguous");
check_half_cuda_contiguous(hidden, "hidden");
check_half_cuda_contiguous(weight, "weight");
TORCH_CHECK(state.dim() == 3 && state.size(2) == kStateLen,
"state must have shape (batch, channels, 3)");
TORCH_CHECK(hidden.dim() == 3 && hidden.size(2) == 1 &&
hidden.size(0) == state.size(0) &&
hidden.size(1) == state.size(1),
"hidden must have shape (batch, channels, 1)");
TORCH_CHECK(weight.dim() == 2 && weight.size(0) == state.size(1) &&
weight.size(1) == kKernelSize,
"weight must have shape (channels, 4)");
auto output = torch::empty_like(hidden);
const int channels = static_cast<int>(state.size(1));
const dim3 blocks((channels + kThreads - 1) / kThreads,
static_cast<unsigned int>(state.size(0)));
causal_conv_update_kernel<<<blocks, kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
state.data_ptr<float>(),
reinterpret_cast<const __half*>(hidden.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(weight.data_ptr<at::Half>()),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()), channels);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("causal_conv_update", &causal_conv_update,
"Fused CoreX Gated DeltaNet causal convolution update");
}

View File

@@ -0,0 +1,80 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
constexpr int kHeadDim = 128;
__device__ __forceinline__ float silu(float value) {
return value / (1.0f + expf(-value));
}
__global__ void gated_rms_norm_inverse_kernel(
const float* input, const __half* gate, const __half* weight,
const float* inverse, __half* output, int rows) {
const int row = blockIdx.x;
const int column = threadIdx.x;
if (row >= rows || column >= kHeadDim) {
return;
}
const int offset = row * kHeadDim + column;
const float scaled = __fmul_rn(input[offset], inverse[row]);
const float normalized = __fmul_rn(
__half2float(weight[column]), scaled);
const float activated = silu(__half2float(gate[offset]));
output[offset] = __float2half_rn(__fmul_rn(normalized, activated));
}
void check_input(const torch::Tensor& input, const torch::Tensor& gate,
const torch::Tensor& weight,
const torch::Tensor& inverse) {
TORCH_CHECK(input.is_cuda() && gate.is_cuda() && weight.is_cuda()
&& inverse.is_cuda(),
"all tensors must be CUDA tensors");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"input must have dtype float32");
TORCH_CHECK(gate.scalar_type() == torch::kFloat16,
"gate must have dtype float16");
TORCH_CHECK(weight.scalar_type() == torch::kFloat16,
"weight must have dtype float16");
TORCH_CHECK(inverse.scalar_type() == torch::kFloat32,
"inverse must have dtype float32");
TORCH_CHECK(input.is_contiguous() && gate.is_contiguous()
&& weight.is_contiguous() && inverse.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(input.dim() == 2 && input.size(1) == kHeadDim,
"input must have shape (rows, 128)");
TORCH_CHECK(gate.sizes() == input.sizes(),
"gate must match input shape");
TORCH_CHECK(weight.dim() == 1 && weight.size(0) == kHeadDim,
"weight must have shape (128,)");
TORCH_CHECK(inverse.numel() == input.size(0),
"inverse must contain one value per row");
}
} // namespace
torch::Tensor apply_inverse(const torch::Tensor& input,
const torch::Tensor& gate,
const torch::Tensor& weight,
const torch::Tensor& inverse) {
check_input(input, gate, weight, inverse);
auto output = torch::empty_like(gate);
const int rows = static_cast<int>(input.size(0));
gated_rms_norm_inverse_kernel<<<
rows, kHeadDim, 0, at::cuda::getCurrentCUDAStream()>>>(
input.data_ptr<float>(),
reinterpret_cast<const __half*>(gate.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(weight.data_ptr<at::Half>()),
inverse.data_ptr<float>(),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()), rows);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("apply_inverse", &apply_inverse,
"CoreX gated RMSNorm using a PyTorch-computed inverse");
}

View File

@@ -0,0 +1,165 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
constexpr int kKeyHeads = 4;
constexpr int kValueHeads = 8;
constexpr int kHeadDim = 128;
constexpr int kMixedDim =
(2 * kKeyHeads + kValueHeads) * kHeadDim;
constexpr float kQueryScale = 0.08838834764831845f;
__global__ void gdn_packed_decode_kernel(
float* state, const half* mixed_qkv, const half* beta_input,
const half* decay_input, const half* a_log, const half* dt_bias,
float* output) {
const int batch_head = blockIdx.x;
const int column = threadIdx.x;
const int batch = batch_head / kValueHeads;
const int value_head = batch_head % kValueHeads;
const int key_head = value_head / (kValueHeads / kKeyHeads);
const int mixed_offset = batch * kMixedDim;
const int query_offset = mixed_offset + key_head * kHeadDim;
const int key_offset =
mixed_offset + kKeyHeads * kHeadDim + key_head * kHeadDim;
const int value_offset = mixed_offset + 2 * kKeyHeads * kHeadDim
+ value_head * kHeadDim;
const int vector_offset = batch_head * kHeadDim;
const int state_offset = batch_head * kHeadDim * kHeadDim;
__shared__ half norm_squares[kHeadDim * 2];
__shared__ float normalized_query[kHeadDim];
__shared__ float normalized_key[kHeadDim];
const half raw_query = mixed_qkv[query_offset + column];
const half raw_key = mixed_qkv[key_offset + column];
norm_squares[column] = __hmul(raw_query, raw_query);
norm_squares[kHeadDim + column] = __hmul(raw_key, raw_key);
__syncthreads();
for (int stride = kHeadDim / 2; stride > 0; stride >>= 1) {
if (column < stride) {
norm_squares[column] = __hadd(
norm_squares[column], norm_squares[column + stride]);
norm_squares[kHeadDim + column] = __hadd(
norm_squares[kHeadDim + column],
norm_squares[kHeadDim + column + stride]);
}
__syncthreads();
}
const half epsilon = __float2half(1e-6f);
const half query_inverse = __float2half(rsqrtf(__half2float(
__hadd(norm_squares[0], epsilon))));
const half key_inverse = __float2half(rsqrtf(__half2float(
__hadd(norm_squares[kHeadDim], epsilon))));
normalized_query[column] = __half2float(
__hmul(raw_query, query_inverse)) * kQueryScale;
normalized_key[column] = __half2float(__hmul(raw_key, key_inverse));
__syncthreads();
const int coefficient_offset = batch * kValueHeads + value_head;
const float beta_value = __half2float(beta_input[coefficient_offset]);
const float beta = __half2float(__float2half(
1.0f / (1.0f + expf(-beta_value))));
const float decay_x = __half2float(decay_input[coefficient_offset])
+ __half2float(dt_bias[value_head]);
const float softplus =
decay_x > 20.0f ? decay_x : log1pf(expf(decay_x));
const float decay = expf(
-expf(__half2float(a_log[value_head])) * softplus);
float memory = 0.0f;
#pragma unroll
for (int row = 0; row < kHeadDim; ++row) {
const int index = state_offset + row * kHeadDim + column;
const float decayed = state[index] * decay;
memory += normalized_key[row] * decayed;
}
const float value = __half2float(mixed_qkv[value_offset + column]);
const float delta = (value - memory) * beta;
float result = 0.0f;
#pragma unroll
for (int row = 0; row < kHeadDim; ++row) {
const int index = state_offset + row * kHeadDim + column;
const float decayed = state[index] * decay;
const float updated = decayed + normalized_key[row] * delta;
state[index] = updated;
result += normalized_query[row] * updated;
}
output[vector_offset + column] = result;
}
void check_half_matrix(const torch::Tensor& tensor, const char* name,
int64_t width) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == 2 && tensor.size(1) == width,
name, " must have shape (batch, ", width, ")");
}
void check_half_vector(const torch::Tensor& tensor, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == 1 && tensor.size(0) == kValueHeads,
name, " must have shape (", kValueHeads, ")");
}
} // namespace
torch::Tensor packed_decode(torch::Tensor state,
const torch::Tensor& mixed_qkv,
const torch::Tensor& beta_input,
const torch::Tensor& decay_input,
const torch::Tensor& a_log,
const torch::Tensor& dt_bias) {
TORCH_CHECK(state.is_cuda(), "state must be a CUDA tensor");
TORCH_CHECK(state.scalar_type() == torch::kFloat32,
"state must have dtype float32");
TORCH_CHECK(state.is_contiguous(), "state must be contiguous");
TORCH_CHECK(state.dim() == 4 && state.size(0) == 1
&& state.size(1) == kValueHeads
&& state.size(2) == kHeadDim
&& state.size(3) == kHeadDim,
"state must have shape (1, 8, 128, 128)");
check_half_matrix(mixed_qkv, "mixed_qkv", kMixedDim);
check_half_matrix(beta_input, "beta_input", kValueHeads);
check_half_matrix(decay_input, "decay_input", kValueHeads);
check_half_vector(a_log, "a_log");
check_half_vector(dt_bias, "dt_bias");
TORCH_CHECK(mixed_qkv.size(0) == 1 && beta_input.size(0) == 1
&& decay_input.size(0) == 1,
"packed decode only supports one sequence");
TORCH_CHECK(state.device() == mixed_qkv.device()
&& state.device() == beta_input.device()
&& state.device() == decay_input.device()
&& state.device() == a_log.device()
&& state.device() == dt_bias.device(),
"all inputs must be on the same device");
torch::Tensor output = torch::empty(
{1, kValueHeads, kHeadDim}, state.options());
gdn_packed_decode_kernel<<<kValueHeads, kHeadDim, 0,
at::cuda::getCurrentCUDAStream()>>>(
state.data_ptr<float>(),
reinterpret_cast<const half*>(mixed_qkv.data_ptr<at::Half>()),
reinterpret_cast<const half*>(beta_input.data_ptr<at::Half>()),
reinterpret_cast<const half*>(decay_input.data_ptr<at::Half>()),
reinterpret_cast<const half*>(a_log.data_ptr<at::Half>()),
reinterpret_cast<const half*>(dt_bias.data_ptr<at::Half>()),
output.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("packed_decode", &packed_decode,
"Packed Qwen3.6 GDN single-token decode");
}

View File

@@ -0,0 +1,72 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
constexpr int kHeadDim = 128;
constexpr float kQueryScale = 0.08838834764831845f;
__global__ void qk_map_kernel(const half* query, const half* key,
float* output, int batch, int key_heads,
int value_heads, int expand_ratio) {
const int elements = batch * value_heads * kHeadDim;
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index >= elements) {
return;
}
const int dim = index % kHeadDim;
const int value_head_index = index / kHeadDim;
const int value_head = value_head_index % value_heads;
const int batch_index = value_head_index / value_heads;
const int key_head = value_head / expand_ratio;
const int source = ((batch_index * key_heads + key_head) * kHeadDim + dim);
output[index] = __half2float(query[source]) * kQueryScale;
output[elements + index] = __half2float(key[source]);
}
void check_input(const torch::Tensor& tensor, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == 3 && tensor.size(2) == kHeadDim,
name, " must have shape (batch, key_heads, 128)");
}
} // namespace
torch::Tensor qk_map(const torch::Tensor& query,
const torch::Tensor& key,
int64_t value_heads_arg) {
check_input(query, "query");
check_input(key, "key");
TORCH_CHECK(query.sizes() == key.sizes(),
"query and key shapes must match");
const int batch = static_cast<int>(query.size(0));
const int key_heads = static_cast<int>(query.size(1));
const int value_heads = static_cast<int>(value_heads_arg);
TORCH_CHECK(value_heads > 0 && value_heads % key_heads == 0,
"value_heads must be divisible by key_heads");
torch::Tensor output = torch::empty(
{2, batch, value_heads, kHeadDim},
query.options().dtype(torch::kFloat32));
const int elements = batch * value_heads * kHeadDim;
constexpr int threads = 256;
const int blocks = (elements + threads - 1) / threads;
qk_map_kernel<<<blocks, threads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const half*>(query.data_ptr<at::Half>()),
reinterpret_cast<const half*>(key.data_ptr<at::Half>()),
output.data_ptr<float>(), batch, key_heads, value_heads,
value_heads / key_heads);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("qk_map", &qk_map,
"Map normalized FP16 key heads to FP32 value heads");
}

View File

@@ -0,0 +1,181 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
constexpr int kExperts = 256;
constexpr int kTopK = 8;
constexpr int kHidden = 2048;
constexpr int kIntermediate = 128;
constexpr int kW13Rows = 2 * kIntermediate;
constexpr int kThreads = 256;
constexpr int kWarpSize = 32;
__device__ inline float warp_sum(float value) {
#pragma unroll
for (int offset = kWarpSize / 2; offset > 0; offset /= 2) {
value += __shfl_down_sync(0xffffffff, value, offset);
}
return value;
}
__global__ void direct_w13_kernel(
const __half* input, const __half* w13, const int64_t* expert_ids,
__half* gate_up) {
const int warp =
(static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x) / kWarpSize;
const int lane = threadIdx.x & (kWarpSize - 1);
if (warp >= kTopK * kW13Rows) {
return;
}
const int slot = warp / kW13Rows;
const int local_row = warp - slot * kW13Rows;
const int64_t expert = expert_ids[slot];
const int64_t weight_row =
(expert * kW13Rows + local_row) * static_cast<int64_t>(kHidden);
const __half2* input2 = reinterpret_cast<const __half2*>(input);
const __half2* weight2 =
reinterpret_cast<const __half2*>(w13 + weight_row);
float sum = 0.0f;
for (int index = lane; index < kHidden / 2; index += kWarpSize) {
const __half2 x = input2[index];
const __half2 weight = weight2[index];
sum = fmaf(__half2float(weight.x), __half2float(x.x), sum);
sum = fmaf(__half2float(weight.y), __half2float(x.y), sum);
}
sum = warp_sum(sum);
if (lane == 0) {
gate_up[warp] = __float2half_rn(sum);
}
}
__global__ void direct_w2_reduce_kernel(
const __half* activated, const __half* w2, const int64_t* expert_ids,
const __half* weights, __half* output) {
const int warp =
(static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x) / kWarpSize;
const int lane = threadIdx.x & (kWarpSize - 1);
if (warp >= kHidden) {
return;
}
float weighted_sum = 0.0f;
#pragma unroll
for (int slot = 0; slot < kTopK; ++slot) {
const int64_t expert = expert_ids[slot];
const int64_t weight_row =
(expert * kHidden + warp) * static_cast<int64_t>(kIntermediate);
const __half2* activation2 = reinterpret_cast<const __half2*>(
activated + slot * kIntermediate);
const __half2* weight2 =
reinterpret_cast<const __half2*>(w2 + weight_row);
float expert_sum = 0.0f;
for (int index = lane; index < kIntermediate / 2;
index += kWarpSize) {
const __half2 x = activation2[index];
const __half2 weight = weight2[index];
expert_sum = fmaf(
__half2float(weight.x), __half2float(x.x), expert_sum);
expert_sum = fmaf(
__half2float(weight.y), __half2float(x.y), expert_sum);
}
expert_sum = warp_sum(expert_sum);
if (lane == 0) {
const __half expert_half = __float2half_rn(expert_sum);
const __half product = __hmul(expert_half, weights[slot]);
weighted_sum += __half2float(product);
}
}
if (lane == 0) {
output[warp] = __float2half_rn(weighted_sum);
}
}
void check_half_cuda(const torch::Tensor& tensor, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
}
void check_ids(const torch::Tensor& expert_ids) {
TORCH_CHECK(expert_ids.is_cuda() && expert_ids.is_contiguous(),
"expert_ids must be a contiguous CUDA tensor");
TORCH_CHECK(expert_ids.scalar_type() == torch::kInt64,
"expert_ids must have dtype int64");
TORCH_CHECK(expert_ids.dim() == 1 && expert_ids.numel() == kTopK,
"expert_ids must have shape (8,)");
}
} // namespace
torch::Tensor direct_w13(const torch::Tensor& input,
const torch::Tensor& w13,
const torch::Tensor& expert_ids) {
check_half_cuda(input, "input");
check_half_cuda(w13, "w13");
check_ids(expert_ids);
TORCH_CHECK(input.dim() == 2 && input.size(0) == 1
&& input.size(1) == kHidden,
"input must have shape (1, 2048)");
TORCH_CHECK(w13.dim() == 3 && w13.size(0) == kExperts
&& w13.size(1) == kW13Rows
&& w13.size(2) == kHidden,
"w13 must have shape (256, 256, 2048)");
auto output = torch::empty({kTopK, kW13Rows}, input.options());
constexpr int kWarpsPerBlock = kThreads / kWarpSize;
constexpr int kBlocks =
(kTopK * kW13Rows + kWarpsPerBlock - 1) / kWarpsPerBlock;
direct_w13_kernel<<<kBlocks, kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const __half*>(input.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(w13.data_ptr<at::Half>()),
expert_ids.data_ptr<int64_t>(),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()));
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor direct_w2_reduce(const torch::Tensor& activated,
const torch::Tensor& w2,
const torch::Tensor& expert_ids,
const torch::Tensor& weights) {
check_half_cuda(activated, "activated");
check_half_cuda(w2, "w2");
check_half_cuda(weights, "weights");
check_ids(expert_ids);
TORCH_CHECK(activated.dim() == 2 && activated.size(0) == kTopK
&& activated.size(1) == kIntermediate,
"activated must have shape (8, 128)");
TORCH_CHECK(w2.dim() == 3 && w2.size(0) == kExperts
&& w2.size(1) == kHidden
&& w2.size(2) == kIntermediate,
"w2 must have shape (256, 2048, 128)");
TORCH_CHECK(weights.dim() == 1 && weights.numel() == kTopK,
"weights must have shape (8,)");
auto output = torch::empty({1, kHidden}, activated.options());
constexpr int kWarpsPerBlock = kThreads / kWarpSize;
constexpr int kBlocks =
(kHidden + kWarpsPerBlock - 1) / kWarpsPerBlock;
direct_w2_reduce_kernel<<<kBlocks, kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const __half*>(activated.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(w2.data_ptr<at::Half>()),
expert_ids.data_ptr<int64_t>(),
reinterpret_cast<const __half*>(weights.data_ptr<at::Half>()),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()));
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("w13", &direct_w13,
"Direct selected-expert FP16 W13 matvec");
module.def("w2_reduce", &direct_w2_reduce,
"Direct selected-expert W2 matvec and routed reduction");
}

View File

@@ -0,0 +1,107 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
namespace {
constexpr int kTopK = 8;
constexpr int kThreads = 256;
enum class Mode { kSerialFloat, kTreeFloat, kSerialHalf };
__global__ void exact_reduce_kernel(const __half* expert_output,
const __half* weights,
__half* output, int hidden,
Mode mode) {
const int column = blockIdx.x * blockDim.x + threadIdx.x;
if (column >= hidden) {
return;
}
__half products[kTopK];
#pragma unroll
for (int expert = 0; expert < kTopK; ++expert) {
products[expert] = __hmul(
expert_output[expert * hidden + column], weights[expert]);
}
if (mode == Mode::kSerialHalf) {
__half sum = products[0];
#pragma unroll
for (int expert = 1; expert < kTopK; ++expert) {
sum = __hadd(sum, products[expert]);
}
output[column] = sum;
return;
}
float sum;
if (mode == Mode::kSerialFloat) {
sum = __half2float(products[0]);
#pragma unroll
for (int expert = 1; expert < kTopK; ++expert) {
sum += __half2float(products[expert]);
}
} else {
const float sum01 = __half2float(products[0]) + __half2float(products[1]);
const float sum23 = __half2float(products[2]) + __half2float(products[3]);
const float sum45 = __half2float(products[4]) + __half2float(products[5]);
const float sum67 = __half2float(products[6]) + __half2float(products[7]);
sum = (sum01 + sum23) + (sum45 + sum67);
}
output[column] = __float2half_rn(sum);
}
void check_input(const torch::Tensor& expert_output,
const torch::Tensor& weights) {
TORCH_CHECK(expert_output.is_cuda() && weights.is_cuda(),
"inputs must be CUDA tensors");
TORCH_CHECK(expert_output.scalar_type() == torch::kFloat16
&& weights.scalar_type() == torch::kFloat16,
"inputs must have dtype float16");
TORCH_CHECK(expert_output.is_contiguous() && weights.is_contiguous(),
"inputs must be contiguous");
TORCH_CHECK(expert_output.dim() == 2
&& expert_output.size(0) == kTopK,
"expert_output must have shape (8, hidden)");
TORCH_CHECK(weights.dim() == 1 && weights.size(0) == kTopK,
"weights must have shape (8,)");
}
torch::Tensor launch(const torch::Tensor& expert_output,
const torch::Tensor& weights, Mode mode) {
check_input(expert_output, weights);
auto output = torch::empty(
{1, expert_output.size(1)}, expert_output.options());
const int hidden = static_cast<int>(expert_output.size(1));
const int blocks = (hidden + kThreads - 1) / kThreads;
exact_reduce_kernel<<<blocks, kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const __half*>(expert_output.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(weights.data_ptr<at::Half>()),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()), hidden, mode);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
} // namespace
torch::Tensor serial_float(const torch::Tensor& expert_output,
const torch::Tensor& weights) {
return launch(expert_output, weights, Mode::kSerialFloat);
}
torch::Tensor tree_float(const torch::Tensor& expert_output,
const torch::Tensor& weights) {
return launch(expert_output, weights, Mode::kTreeFloat);
}
torch::Tensor serial_half(const torch::Tensor& expert_output,
const torch::Tensor& weights) {
return launch(expert_output, weights, Mode::kSerialHalf);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("serial_float", &serial_float);
module.def("tree_float", &tree_float);
module.def("serial_half", &serial_half);
}

View File

@@ -0,0 +1,92 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <cstdint>
#include <vector>
namespace {
constexpr int kTopK = 8;
constexpr int kThreads = 256;
constexpr int kGridX = 8;
__global__ void selected_weight_gather_vec16_kernel(
const uint4* w13, const uint4* w2, const int64_t* expert_ids,
uint4* selected_w13, uint4* selected_w2,
int64_t w13_vecs_per_expert, int64_t w2_vecs_per_expert) {
const int segment = blockIdx.y;
const int slot = segment & (kTopK - 1);
const bool copy_w2 = segment >= kTopK;
const int64_t count =
copy_w2 ? w2_vecs_per_expert : w13_vecs_per_expert;
const uint4* source = copy_w2 ? w2 : w13;
uint4* output = copy_w2 ? selected_w2 : selected_w13;
const int64_t source_offset = expert_ids[slot] * count;
const int64_t output_offset = static_cast<int64_t>(slot) * count;
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < count;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
output[output_offset + index] = source[source_offset + index];
}
}
void check_weight(const torch::Tensor& tensor, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == 3, name, " must be rank three");
TORCH_CHECK(tensor.size(1) * tensor.size(2) % 8 == 0,
name, " expert slices must be divisible by 16 bytes");
}
} // namespace
std::vector<torch::Tensor> gather_selected_weights(
const torch::Tensor& w13, const torch::Tensor& w2,
const torch::Tensor& expert_ids) {
check_weight(w13, "w13");
check_weight(w2, "w2");
TORCH_CHECK(w13.device() == w2.device(),
"W13/W2 must be on the same device");
TORCH_CHECK(w13.size(0) == w2.size(0),
"W13/W2 expert counts differ");
TORCH_CHECK(w13.size(2) == w2.size(1),
"W13/W2 hidden dimensions differ");
TORCH_CHECK(w13.size(1) == 2 * w2.size(2),
"W13/W2 intermediate dimensions differ");
TORCH_CHECK(expert_ids.is_cuda() && expert_ids.is_contiguous(),
"expert_ids must be a contiguous CUDA tensor");
TORCH_CHECK(expert_ids.device() == w13.device(),
"weights and expert_ids must be on the same device");
TORCH_CHECK(expert_ids.scalar_type() == torch::kInt64,
"expert_ids must have dtype int64");
TORCH_CHECK(expert_ids.dim() == 1 && expert_ids.numel() == kTopK,
"expert_ids must have shape (8,)");
auto selected_w13 = torch::empty(
{kTopK, w13.size(1), w13.size(2)}, w13.options());
auto selected_w2 = torch::empty(
{kTopK, w2.size(1), w2.size(2)}, w2.options());
const int64_t w13_vecs_per_expert = w13.size(1) * w13.size(2) / 8;
const int64_t w2_vecs_per_expert = w2.size(1) * w2.size(2) / 8;
const dim3 grid(kGridX, 2 * kTopK);
selected_weight_gather_vec16_kernel<<<
grid, kThreads, 0, at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const uint4*>(w13.data_ptr<at::Half>()),
reinterpret_cast<const uint4*>(w2.data_ptr<at::Half>()),
expert_ids.data_ptr<int64_t>(),
reinterpret_cast<uint4*>(selected_w13.data_ptr<at::Half>()),
reinterpret_cast<uint4*>(selected_w2.data_ptr<at::Half>()),
w13_vecs_per_expert, w2_vecs_per_expert);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {selected_w13, selected_w2};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("gather", &gather_selected_weights,
"Gather selected FP16 top-8 MoE weights with 16-byte loads");
}

View File

@@ -0,0 +1,118 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
#include <algorithm>
#include <cstdint>
#include <vector>
namespace {
constexpr int kThreads = 256;
constexpr int kSmallGridBlocks = 256;
constexpr int kSmallGridMaxSeqLen = 96 * 1024;
__global__ void paged_kv_gather_kernel(
const __half* key_cache, const __half* value_cache,
const int* block_table, float* key_output, float* value_output,
int seq_len, int num_kv_heads, int head_size, int block_size,
int key_pack) {
const int64_t total =
static_cast<int64_t>(seq_len) * num_kv_heads * head_size;
for (int64_t index =
static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
index < total;
index += static_cast<int64_t>(blockDim.x) * gridDim.x) {
const int dim = index % head_size;
const int token = (index / head_size) % seq_len;
const int kv_head = index / (static_cast<int64_t>(head_size) * seq_len);
const int logical_block = token / block_size;
const int block_offset = token % block_size;
const int physical_block = block_table[logical_block];
const int64_t key_index =
(((static_cast<int64_t>(physical_block) * num_kv_heads + kv_head)
* (head_size / key_pack) + dim / key_pack)
* block_size + block_offset) * key_pack + dim % key_pack;
const int64_t value_index =
((static_cast<int64_t>(physical_block) * num_kv_heads + kv_head)
* head_size + dim) * block_size + block_offset;
const int64_t key_output_index =
(static_cast<int64_t>(kv_head) * head_size + dim) * seq_len + token;
const int64_t value_output_index =
(static_cast<int64_t>(kv_head) * seq_len + token) * head_size + dim;
key_output[key_output_index] = __half2float(key_cache[key_index]);
value_output[value_output_index] = __half2float(value_cache[value_index]);
}
}
void check_half_cuda_contiguous(const torch::Tensor& tensor,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
}
} // namespace
std::vector<torch::Tensor> gather_paged_kv(
const torch::Tensor& key_cache, const torch::Tensor& value_cache,
const torch::Tensor& block_table, int64_t seq_len) {
check_half_cuda_contiguous(key_cache, "key_cache");
check_half_cuda_contiguous(value_cache, "value_cache");
TORCH_CHECK(block_table.is_cuda(), "block_table must be a CUDA tensor");
TORCH_CHECK(block_table.scalar_type() == torch::kInt32,
"block_table must have dtype int32");
TORCH_CHECK(block_table.is_contiguous(), "block_table must be contiguous");
TORCH_CHECK(key_cache.dim() == 5,
"key_cache must have shape (blocks, kv_heads, d/x, block, x)");
TORCH_CHECK(value_cache.dim() == 4,
"value_cache must have shape (blocks, kv_heads, d, block)");
TORCH_CHECK(block_table.dim() == 1,
"block_table must be a one-dimensional row");
TORCH_CHECK(key_cache.size(0) == value_cache.size(0),
"key/value block counts differ");
TORCH_CHECK(key_cache.size(1) == value_cache.size(1),
"key/value KV-head counts differ");
TORCH_CHECK(key_cache.size(3) == value_cache.size(3),
"key/value block sizes differ");
TORCH_CHECK(key_cache.size(2) * key_cache.size(4) == value_cache.size(2),
"key/value head sizes differ");
TORCH_CHECK(seq_len > 0, "seq_len must be positive");
const int block_size = static_cast<int>(value_cache.size(3));
const int64_t required_blocks = (seq_len + block_size - 1) / block_size;
TORCH_CHECK(required_blocks <= block_table.numel(),
"block_table is too short for seq_len");
const int num_kv_heads = static_cast<int>(value_cache.size(1));
const int head_size = static_cast<int>(value_cache.size(2));
const int key_pack = static_cast<int>(key_cache.size(4));
auto output_options = key_cache.options().dtype(torch::kFloat32);
auto key_output = torch::empty(
{num_kv_heads, head_size, seq_len}, output_options);
auto value_output = torch::empty(
{num_kv_heads, seq_len, head_size}, output_options);
const int64_t total = seq_len * num_kv_heads * head_size;
const int grid_cap =
seq_len <= kSmallGridMaxSeqLen ? kSmallGridBlocks : 65535;
const int blocks = static_cast<int>(std::min<int64_t>(
(total + kThreads - 1) / kThreads, grid_cap));
paged_kv_gather_kernel<<<blocks, kThreads, 0,
at::cuda::getCurrentCUDAStream()>>>(
reinterpret_cast<const __half*>(key_cache.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(value_cache.data_ptr<at::Half>()),
block_table.data_ptr<int>(), key_output.data_ptr<float>(),
value_output.data_ptr<float>(), static_cast<int>(seq_len),
num_kv_heads, head_size, block_size, key_pack);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {key_output, value_output};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("gather", &gather_paged_kv,
"Gather paged FP16 K/V directly into FP32 attention layouts");
}

View File

@@ -0,0 +1,503 @@
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>
#include <mma.h>
#include <torch/extension.h>
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <limits>
#include <vector>
namespace {
constexpr int kBlockSize = 16;
constexpr int kHeadDim = 256;
constexpr int kKeyPack = 8;
constexpr int kNumQueryHeads = 4;
constexpr int kNumKvHeads = 1;
constexpr int kQueryTile = 16;
constexpr int kKeyTile = 16;
constexpr int kReductionTokens = 512;
constexpr int kKeyTilesPerReduction = kReductionTokens / kKeyTile;
constexpr int kPvReductionSplits = 4;
constexpr int kKeyTilesPerPvSplit =
kKeyTilesPerReduction / kPvReductionSplits;
constexpr int kMmaK = 16;
constexpr int kDimTiles = kHeadDim / kMmaK;
constexpr int kWarpSize = 64;
constexpr int kMaxQueryTokens = 8192;
constexpr int kMaxSequenceTokens = 262144;
using namespace nvcuda;
struct __align__(128) SharedStorage {
float matrix_tile[kQueryTile * kKeyTile];
float scores[
kKeyTilesPerReduction * kQueryTile * kKeyTile];
float running_output[kQueryTile * kHeadDim];
float partial_output[
kPvReductionSplits * kQueryTile * kMmaK];
float running_max[kQueryTile];
float running_sum[kQueryTile];
float correction[kQueryTile];
};
void check_half_cuda_contiguous(const torch::Tensor& tensor,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
name, " must have dtype float16");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
}
__device__ __forceinline__ float load_key(
const __half* key_new, const __half* key_cache,
const int* block_table, int logical_token, int context_len, int dim) {
if (logical_token < context_len) {
const int logical_block = logical_token / kBlockSize;
const int block_offset = logical_token % kBlockSize;
const int physical_block = block_table[logical_block];
const int64_t index =
((((static_cast<int64_t>(physical_block) * kNumKvHeads)
* (kHeadDim / kKeyPack) + dim / kKeyPack)
* kBlockSize + block_offset) * kKeyPack + dim % kKeyPack);
return __half2float(key_cache[index]);
}
const int query_index = logical_token - context_len;
return __half2float(
key_new[static_cast<int64_t>(query_index) * kHeadDim + dim]);
}
__device__ __forceinline__ float load_value(
const __half* value_new, const __half* value_cache,
const int* block_table, int logical_token, int context_len, int dim) {
if (logical_token < context_len) {
const int logical_block = logical_token / kBlockSize;
const int block_offset = logical_token % kBlockSize;
const int physical_block = block_table[logical_block];
const int64_t index =
((static_cast<int64_t>(physical_block) * kNumKvHeads)
* kHeadDim + dim) * kBlockSize + block_offset;
return __half2float(value_cache[index]);
}
const int query_index = logical_token - context_len;
return __half2float(
value_new[static_cast<int64_t>(query_index) * kHeadDim + dim]);
}
__global__ void query_tiled_paged_prefill_kernel(
const __half* query, const __half* key_new, const __half* value_new,
const __half* key_cache, const __half* value_cache,
const int* block_table, __half* output, float* lse,
int context_len, int query_len, float scale) {
__shared__ SharedStorage shared;
const int lane = threadIdx.x;
const int query_tile_index = blockIdx.x / kNumQueryHeads;
const int query_head = blockIdx.x % kNumQueryHeads;
const int query_start = query_tile_index * kQueryTile;
const int active_rows = min(kQueryTile, query_len - query_start);
if (active_rows <= 0) {
return;
}
wmma::fragment<wmma::matrix_a, 16, 16, 16, float,
wmma::row_major> query_fragments[kDimTiles];
#pragma unroll
for (int dim_tile = 0; dim_tile < kDimTiles; ++dim_tile) {
#pragma unroll
for (int quarter = 0; quarter < 4; ++quarter) {
const int row = lane / 16 + quarter * 4;
const int column = lane % 16;
float value = 0.0f;
if (row < active_rows) {
const int query_index = query_start + row;
const int dim = dim_tile * kMmaK + column;
const int64_t source =
(static_cast<int64_t>(query_index) * kNumQueryHeads
+ query_head) * kHeadDim + dim;
value = __half2float(query[source]) * scale;
}
const int offset =
wmma::CoordToOffset<32, wmma::layout_t::mem_row_major>(
row, column);
shared.matrix_tile[offset] = value;
}
__syncthreads();
wmma::load_matrix_sync(
query_fragments[dim_tile], shared.matrix_tile, 0);
__syncthreads();
}
for (int index = lane; index < kQueryTile * kHeadDim;
index += kWarpSize) {
shared.running_output[index] = 0.0f;
}
if (lane < kQueryTile) {
shared.running_max[lane] = -std::numeric_limits<float>::infinity();
shared.running_sum[lane] = 0.0f;
shared.correction[lane] = 1.0f;
}
__syncthreads();
const int last_query = min(query_start + kQueryTile, query_len);
// Preserve the installed reference's 512-token reduction boundaries:
// paged context and current causal K/V are separate phases.
for (int phase = 0; phase < 2; ++phase) {
const int phase_base = phase == 0 ? 0 : context_len;
const int phase_tokens = phase == 0 ? context_len : last_query;
for (int group_start = 0; group_start < phase_tokens;
group_start += kReductionTokens) {
const int group_tokens =
min(kReductionTokens, phase_tokens - group_start);
const int group_key_tiles =
(group_tokens + kKeyTile - 1) / kKeyTile;
for (int key_tile_in_group = 0;
key_tile_in_group < group_key_tiles;
++key_tile_in_group) {
const int local_key_start =
group_start + key_tile_in_group * kKeyTile;
const int logical_key_start = phase_base + local_key_start;
wmma::fragment<wmma::accumulator, 16, 16, 16, float>
score_fragment;
wmma::fill_fragment(score_fragment, 0.0f);
#pragma unroll
for (int dim_tile = 0; dim_tile < kDimTiles; ++dim_tile) {
#pragma unroll
for (int quarter = 0; quarter < 4; ++quarter) {
const int row = lane / 16 + quarter * 4;
const int column = lane % 16;
const int logical_token = logical_key_start + column;
const int dim = dim_tile * kMmaK + row;
const float value =
local_key_start + column < phase_tokens
? load_key(key_new, key_cache, block_table,
logical_token, context_len, dim)
: 0.0f;
const int offset =
wmma::CoordToOffset<
32, wmma::layout_t::mem_col_major>(
row, column);
shared.matrix_tile[offset] = value;
}
__syncthreads();
wmma::fragment<wmma::matrix_b, 16, 16, 16, float,
wmma::col_major> key_fragment;
wmma::load_matrix_sync(
key_fragment, shared.matrix_tile, 0);
wmma::mma_sync(
score_fragment,
query_fragments[dim_tile],
key_fragment,
score_fragment);
__syncthreads();
}
float* score_tile =
shared.scores
+ key_tile_in_group * kQueryTile * kKeyTile;
wmma::store_matrix_sync(
score_tile, score_fragment, 0, wmma::mem_row_major);
__syncthreads();
}
if (lane < kQueryTile) {
const int row = lane;
if (row >= active_rows) {
shared.correction[row] = 1.0f;
for (int key_offset = 0; key_offset < group_tokens;
++key_offset) {
const int key_tile = key_offset / kKeyTile;
const int column = key_offset % kKeyTile;
shared.scores[
key_tile * kQueryTile * kKeyTile
+ row * kKeyTile + column] = 0.0f;
}
} else {
const int absolute_query =
context_len + query_start + row;
float block_max =
-std::numeric_limits<float>::infinity();
for (int key_offset = 0; key_offset < group_tokens;
++key_offset) {
const int key_tile = key_offset / kKeyTile;
const int column = key_offset % kKeyTile;
const int score_index =
key_tile * kQueryTile * kKeyTile
+ row * kKeyTile + column;
const int logical_key =
phase_base + group_start + key_offset;
if (logical_key <= absolute_query) {
block_max = fmaxf(
block_max, shared.scores[score_index]);
} else {
shared.scores[score_index] =
-std::numeric_limits<float>::infinity();
}
}
const float old_max = shared.running_max[row];
const float new_max = fmaxf(old_max, block_max);
const float correction =
old_max == -std::numeric_limits<float>::infinity()
? 0.0f
: expf(old_max - new_max);
float group_sum = 0.0f;
for (int key_offset = 0; key_offset < group_tokens;
++key_offset) {
const int key_tile = key_offset / kKeyTile;
const int column = key_offset % kKeyTile;
const int score_index =
key_tile * kQueryTile * kKeyTile
+ row * kKeyTile + column;
const float score = shared.scores[score_index];
const float probability =
score == -std::numeric_limits<float>::infinity()
? 0.0f
: expf(score - new_max);
shared.scores[score_index] = probability;
group_sum += probability;
}
shared.running_sum[row] =
shared.running_sum[row] * correction + group_sum;
shared.running_max[row] = new_max;
shared.correction[row] = correction;
}
for (int key_offset = group_tokens;
key_offset < group_key_tiles * kKeyTile;
++key_offset) {
const int key_tile = key_offset / kKeyTile;
const int column = key_offset % kKeyTile;
shared.scores[
key_tile * kQueryTile * kKeyTile
+ row * kKeyTile + column] = 0.0f;
}
}
__syncthreads();
for (int index = lane;
index < active_rows * kHeadDim;
index += kWarpSize) {
const int row = index / kHeadDim;
shared.running_output[index] *= shared.correction[row];
}
__syncthreads();
#pragma unroll
for (int dim_tile = 0; dim_tile < kDimTiles; ++dim_tile) {
// CoreX's reference matmul reduces a 512-token K dimension
// hierarchically. Preserve that numerical shape with four fixed,
// contiguous 128-token partials and a deterministic binary merge.
#pragma unroll
for (int split = 0; split < kPvReductionSplits; ++split) {
wmma::fragment<wmma::accumulator, 16, 16, 16, float>
output_fragment;
wmma::fill_fragment(output_fragment, 0.0f);
const int split_start = split * kKeyTilesPerPvSplit;
const int split_end =
min(group_key_tiles, split_start + kKeyTilesPerPvSplit);
for (int key_tile_in_group = split_start;
key_tile_in_group < split_end;
++key_tile_in_group) {
const int local_key_start =
group_start + key_tile_in_group * kKeyTile;
const int logical_key_start = phase_base + local_key_start;
const float* score_tile =
shared.scores
+ key_tile_in_group * kQueryTile * kKeyTile;
wmma::fragment<wmma::matrix_a, 16, 16, 16, float,
wmma::row_major> probability_fragment;
wmma::load_matrix_sync(
probability_fragment, score_tile, 0);
#pragma unroll
for (int quarter = 0; quarter < 4; ++quarter) {
const int row = lane / 16 + quarter * 4;
const int column = lane % 16;
const int logical_token = logical_key_start + row;
const int dim = dim_tile * kMmaK + column;
const float value =
local_key_start + row < phase_tokens
? load_value(value_new, value_cache, block_table,
logical_token, context_len, dim)
: 0.0f;
const int offset =
wmma::CoordToOffset<
32, wmma::layout_t::mem_row_major>(
row, column);
shared.matrix_tile[offset] = value;
}
__syncthreads();
wmma::fragment<wmma::matrix_b, 16, 16, 16, float,
wmma::row_major> value_fragment;
wmma::load_matrix_sync(
value_fragment, shared.matrix_tile, 0);
wmma::mma_sync(
output_fragment,
probability_fragment,
value_fragment,
output_fragment);
__syncthreads();
}
wmma::store_matrix_sync(
shared.partial_output
+ split * kQueryTile * kMmaK,
output_fragment,
0,
wmma::mem_row_major);
__syncthreads();
}
#pragma unroll
for (int quarter = 0; quarter < 4; ++quarter) {
const int row = lane / 16 + quarter * 4;
const int column = lane % 16;
if (row < active_rows) {
const int output_index =
row * kHeadDim + dim_tile * kMmaK + column;
const int tile_index = row * kMmaK + column;
const int partial_stride = kQueryTile * kMmaK;
const float left = __fadd_rn(
shared.partial_output[tile_index],
shared.partial_output[partial_stride + tile_index]);
const float right = __fadd_rn(
shared.partial_output[2 * partial_stride + tile_index],
shared.partial_output[3 * partial_stride + tile_index]);
shared.running_output[output_index] = __fadd_rn(
shared.running_output[output_index],
__fadd_rn(left, right));
}
}
__syncthreads();
}
}
}
for (int index = lane; index < active_rows * kHeadDim;
index += kWarpSize) {
const int row = index / kHeadDim;
const int dim = index % kHeadDim;
const int query_index = query_start + row;
const int64_t destination =
(static_cast<int64_t>(query_index) * kNumQueryHeads
+ query_head) * kHeadDim + dim;
output[destination] = __float2half_rn(
shared.running_output[index] / shared.running_sum[row]);
}
if (lane < active_rows) {
const int query_index = query_start + lane;
lse[static_cast<int64_t>(query_index) * kNumQueryHeads
+ query_head] =
shared.running_max[lane] + logf(shared.running_sum[lane]);
}
}
} // namespace
std::vector<torch::Tensor> query_tiled_paged_prefill_forward(
const torch::Tensor& query, const torch::Tensor& key_new,
const torch::Tensor& value_new, const torch::Tensor& key_cache,
const torch::Tensor& value_cache, const torch::Tensor& block_table,
int64_t context_len_arg, double scale_arg) {
check_half_cuda_contiguous(query, "query");
check_half_cuda_contiguous(key_new, "key_new");
check_half_cuda_contiguous(value_new, "value_new");
check_half_cuda_contiguous(key_cache, "key_cache");
check_half_cuda_contiguous(value_cache, "value_cache");
TORCH_CHECK(block_table.is_cuda(),
"block_table must be a CUDA tensor");
TORCH_CHECK(block_table.scalar_type() == torch::kInt32,
"block_table must have dtype int32");
TORCH_CHECK(block_table.is_contiguous(),
"block_table must be contiguous");
TORCH_CHECK(block_table.dim() == 1,
"block_table must be one-dimensional");
TORCH_CHECK(query.dim() == 3 && query.size(1) == kNumQueryHeads
&& query.size(2) == kHeadDim,
"query must have shape (Q, 4, 256)");
TORCH_CHECK(key_new.dim() == 3 && key_new.size(1) == kNumKvHeads
&& key_new.size(2) == kHeadDim,
"key_new must have shape (Q, 1, 256)");
TORCH_CHECK(value_new.sizes() == key_new.sizes(),
"value_new must match key_new");
TORCH_CHECK(key_new.size(0) == query.size(0),
"query, key_new, and value_new lengths must match");
TORCH_CHECK(key_cache.dim() == 5
&& key_cache.size(1) == kNumKvHeads
&& key_cache.size(2) == kHeadDim / kKeyPack
&& key_cache.size(3) == kBlockSize
&& key_cache.size(4) == kKeyPack,
"key_cache must have shape (N, 1, 32, 16, 8)");
TORCH_CHECK(value_cache.dim() == 4
&& value_cache.size(1) == kNumKvHeads
&& value_cache.size(2) == kHeadDim
&& value_cache.size(3) == kBlockSize,
"value_cache must have shape (N, 1, 256, 16)");
TORCH_CHECK(key_cache.size(0) == value_cache.size(0),
"key/value cache block counts must match");
TORCH_CHECK(query.device() == key_new.device()
&& query.device() == value_new.device()
&& query.device() == key_cache.device()
&& query.device() == value_cache.device()
&& query.device() == block_table.device(),
"all tensors must use the same device");
TORCH_CHECK(context_len_arg >= 0
&& context_len_arg <= kMaxSequenceTokens,
"context_len is out of range");
TORCH_CHECK(context_len_arg % kBlockSize == 0,
"context_len must be block aligned");
const int query_len = static_cast<int>(query.size(0));
const int context_len = static_cast<int>(context_len_arg);
TORCH_CHECK(query_len > 0 && query_len <= kMaxQueryTokens,
"query length must be in [1, 8192]");
TORCH_CHECK(context_len + query_len <= kMaxSequenceTokens,
"context_len + query_len exceeds 262144");
const int required_blocks = context_len / kBlockSize;
TORCH_CHECK(block_table.numel() >= required_blocks,
"block_table is too short for context_len");
if (required_blocks > 0) {
auto active_blocks = block_table.narrow(0, 0, required_blocks);
const int minimum_block = active_blocks.min().item<int>();
const int maximum_block = active_blocks.max().item<int>();
TORCH_CHECK(minimum_block >= 0
&& maximum_block < key_cache.size(0),
"block_table contains an out-of-range physical block ID");
}
TORCH_CHECK(std::isfinite(scale_arg) && scale_arg > 0.0,
"scale must be finite and positive");
auto output = torch::empty_like(query);
auto lse = torch::empty(
{query_len, kNumQueryHeads},
query.options().dtype(torch::kFloat32));
const int query_tiles =
(query_len + kQueryTile - 1) / kQueryTile;
const int blocks = query_tiles * kNumQueryHeads;
auto stream = at::cuda::getCurrentCUDAStream();
query_tiled_paged_prefill_kernel<<<
blocks, kWarpSize, 0, stream>>>(
reinterpret_cast<const __half*>(query.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(key_new.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(value_new.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(key_cache.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(value_cache.data_ptr<at::Half>()),
block_table.data_ptr<int>(),
reinterpret_cast<__half*>(output.data_ptr<at::Half>()),
lse.data_ptr<float>(), context_len, query_len,
static_cast<float>(scale_arg));
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {output, lse};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("forward", &query_tiled_paged_prefill_forward,
"Fixed BI100 query-tiled paged-prefill forward");
}

View File

@@ -0,0 +1,291 @@
"""Shared GDN prefix-state cache contracts for the BI100 runtime."""
from __future__ import annotations
import os
from collections import OrderedDict
from dataclasses import dataclass
from typing import Iterable, List, Optional, Sequence, Tuple
GdnPrefixKey = Tuple[int, bytes]
GdnCapturePoint = Tuple[int, GdnPrefixKey]
_VALID_POLICIES = {"fine32", "admission64", "off"}
GDN_KERNEL_CHUNK_TOKENS = 64
GDN_DIRECT_MIN_REPLAY_TOKENS = 2
_VALID_RESTORE_MODES = {"direct", "hybrid64", "chunk64", "aligned"}
def _env_choice(name: str, default: str, choices: set[str]) -> str:
value = os.getenv(name, default).strip().lower()
if value not in choices:
allowed = ", ".join(sorted(choices))
raise RuntimeError(f"invalid {name}={value!r}; expected one of: {allowed}")
return value
def gdn_cache_policy_from_env() -> str:
return _env_choice("BI100_GDN_CACHE_POLICY", "fine32", _VALID_POLICIES)
def gdn_restore_mode_from_env() -> str:
return _env_choice(
"BI100_GDN_RESTORE_MODE", "direct", _VALID_RESTORE_MODES)
def gdn_restore_alignment(restore_mode: str, block_size: int,
scheduler_chunk_tokens: int) -> int:
"""Return the content boundary required by a restore mode."""
if block_size <= 0:
raise ValueError("block_size must be positive")
if restore_mode == "direct":
return block_size
if restore_mode in {"hybrid64", "chunk64"}:
alignment = GDN_KERNEL_CHUNK_TOKENS
elif restore_mode == "aligned":
alignment = scheduler_chunk_tokens
else:
raise ValueError(f"unknown GDN restore mode: {restore_mode}")
if alignment <= 0 or alignment % block_size != 0:
raise ValueError(
f"{restore_mode} GDN restore requires a positive alignment "
f"divisible by block_size={block_size}; got {alignment}")
return alignment
def make_prefix_key(block_count: int, digest: bytes) -> GdnPrefixKey:
if block_count <= 0:
raise ValueError("GDN prefix key requires at least one complete block")
if not isinstance(digest, bytes) or len(digest) != 32:
raise ValueError("GDN prefix digest must be exactly 32 bytes")
return block_count, digest
def keys_from_block_hashes(block_hashes: Sequence[bytes]) -> List[GdnPrefixKey]:
return [make_prefix_key(i + 1, digest)
for i, digest in enumerate(block_hashes)]
def strict_prefix_block_count(token_count: int, block_size: int) -> int:
if block_size <= 0:
raise ValueError("block_size must be positive")
if token_count <= 1:
return 0
return (token_count - 1) // block_size
def key_at_strict_boundary(block_hashes: Sequence[bytes], token_count: int,
block_size: int) -> Optional[GdnPrefixKey]:
block_count = min(
len(block_hashes), strict_prefix_block_count(token_count, block_size))
if block_count <= 0:
return None
return make_prefix_key(block_count, block_hashes[block_count - 1])
def final_capture_key(
block_hashes: Sequence[bytes], prompt_tokens: int, block_size: int,
restore_mode: str, replay_alignment: int) -> Optional[GdnPrefixKey]:
if restore_mode in {"direct", "hybrid64"}:
block_count = min(
len(block_hashes), strict_prefix_block_count(
prompt_tokens, block_size))
if (block_count > 0
and prompt_tokens - block_count * block_size
< GDN_DIRECT_MIN_REPLAY_TOKENS):
block_count -= 1
if block_count <= 0:
return None
return make_prefix_key(block_count, block_hashes[block_count - 1])
if restore_mode not in {"chunk64", "aligned"}:
raise ValueError(f"unknown GDN restore mode: {restore_mode}")
if (replay_alignment <= 0 or replay_alignment % block_size != 0
or prompt_tokens <= 1):
return None
boundary_tokens = ((prompt_tokens - 1) // replay_alignment
* replay_alignment)
block_count = min(len(block_hashes), boundary_tokens // block_size)
if block_count <= 0:
return None
return make_prefix_key(block_count, block_hashes[block_count - 1])
def restore_key_is_eligible(
key: GdnPrefixKey, prompt_tokens: int, block_size: int,
restore_mode: str, replay_alignment: int,
direct_final_key: Optional[GdnPrefixKey] = None) -> bool:
"""Return whether restoring ``key`` preserves the execution contract."""
make_prefix_key(*key)
if block_size <= 0:
raise ValueError("block_size must be positive")
boundary_tokens = key[0] * block_size
remaining_tokens = prompt_tokens - boundary_tokens
if remaining_tokens <= 0:
return False
if restore_mode == "direct":
return remaining_tokens >= GDN_DIRECT_MIN_REPLAY_TOKENS
if restore_mode == "hybrid64":
if direct_final_key is not None:
make_prefix_key(*direct_final_key)
return (remaining_tokens >= GDN_DIRECT_MIN_REPLAY_TOKENS
and replay_alignment > 0
and (boundary_tokens % replay_alignment == 0
or key == direct_final_key))
if restore_mode not in {"chunk64", "aligned"}:
raise ValueError(f"unknown GDN restore mode: {restore_mode}")
return (replay_alignment > 0
and boundary_tokens % replay_alignment == 0)
def capture_points_for_step(
targets: Iterable[GdnPrefixKey], physical_context_tokens: int,
logical_end_tokens: int, block_size: int) -> Tuple[GdnCapturePoint, ...]:
if physical_context_tokens < 0 or logical_end_tokens < 0:
raise ValueError("token positions must be non-negative")
if logical_end_tokens <= physical_context_tokens:
return ()
selected = {}
for key in targets:
make_prefix_key(*key)
boundary_tokens = key[0] * block_size
if physical_context_tokens < boundary_tokens <= logical_end_tokens:
selected[boundary_tokens - physical_context_tokens] = key
points = tuple(sorted(selected.items()))
if len(points) > 2:
raise ValueError("at most two GDN capture points are allowed per step")
return points
def cap_prefill_end_at_capture_boundary(
logical_start_tokens: int, logical_end_tokens: int,
targets: Iterable[GdnPrefixKey], block_size: int) -> int:
"""Stop a physical prefill step at its earliest pending capture boundary."""
if logical_start_tokens < 0 or logical_end_tokens < 0:
raise ValueError("token positions must be non-negative")
if logical_end_tokens < logical_start_tokens:
raise ValueError("logical end must not precede logical start")
if block_size <= 0:
raise ValueError("block_size must be positive")
capped_end = logical_end_tokens
for key in targets:
make_prefix_key(*key)
boundary_tokens = key[0] * block_size
if logical_start_tokens < boundary_tokens < capped_end:
capped_end = boundary_tokens
return capped_end
def canonical_direct_segment_offsets(
block_hashes: Sequence[bytes], physical_context_tokens: int,
logical_end_tokens: int, block_size: int,
scheduler_chunk_tokens: int) -> Tuple[int, ...]:
"""Reproduce cold fine32/direct segment boundaries after fast-forward."""
if physical_context_tokens < 0 or logical_end_tokens < 0:
raise ValueError("token positions must be non-negative")
if block_size <= 0 or scheduler_chunk_tokens <= 0:
raise ValueError("block and scheduler chunk sizes must be positive")
if scheduler_chunk_tokens % block_size != 0:
raise ValueError("scheduler chunk size must be divisible by block size")
if logical_end_tokens <= physical_context_tokens:
return ()
boundaries = set()
step_ends = list(range(scheduler_chunk_tokens, logical_end_tokens,
scheduler_chunk_tokens))
for step_end in (*step_ends, logical_end_tokens):
key = final_capture_key(block_hashes, step_end, block_size,
"direct", block_size)
if key is not None:
boundaries.add(key[0] * block_size)
boundaries.update(step_ends)
return tuple(
boundary - physical_context_tokens
for boundary in sorted(boundaries)
if physical_context_tokens < boundary < logical_end_tokens)
@dataclass(frozen=True)
class GdnCachePlan:
restore_key: Optional[GdnPrefixKey] = None
capture_points: Tuple[GdnCapturePoint, ...] = ()
evict_keys: Tuple[GdnPrefixKey, ...] = ()
class GdnPrefixStatePolicy:
"""Scheduler-owned state index with deterministic worker actions."""
def __init__(self, policy: str) -> None:
if policy not in _VALID_POLICIES:
raise ValueError(f"unknown GDN cache policy: {policy}")
self.policy = policy
self.capacity = {"fine32": 32, "admission64": 64, "off": 0}[policy]
self._resident: OrderedDict[GdnPrefixKey, None] = OrderedDict()
def __len__(self) -> int:
return len(self._resident)
def resident_keys(self) -> Tuple[GdnPrefixKey, ...]:
return tuple(self._resident)
def contains(self, key: GdnPrefixKey) -> bool:
return key in self._resident
def should_capture_final(self, key: GdnPrefixKey) -> bool:
"""Return whether a final state must be materialized on this request."""
make_prefix_key(*key)
if self.policy == "off":
return False
if self.policy == "admission64":
return key not in self._resident
return True
def select_restore(
self, live_prefix_keys: Sequence[GdnPrefixKey],
max_blocks: int) -> Optional[GdnPrefixKey]:
if self.capacity == 0 or max_blocks <= 0:
return None
best = None
for key in live_prefix_keys[:max_blocks]:
if key in self._resident:
best = key
if best is not None:
self._resident.move_to_end(best)
return best
def repeated_branch_candidate(
self, live_prefix_keys: Sequence[GdnPrefixKey],
max_blocks: int) -> Optional[GdnPrefixKey]:
"""Return a repeated raw-KV branch that lacks recurrent state.
A live KV hit proves that the content occurred in an earlier request;
the current request is therefore the second or later occurrence.
"""
if (self.policy != "admission64" or max_blocks <= 0
or not live_prefix_keys):
return None
candidate = live_prefix_keys[min(len(live_prefix_keys), max_blocks) - 1]
if candidate in self._resident:
return None
return candidate
def admit(self, keys: Iterable[GdnPrefixKey]) -> Tuple[GdnPrefixKey, ...]:
evicted: List[GdnPrefixKey] = []
if self.capacity == 0:
return ()
for key in keys:
make_prefix_key(*key)
if key in self._resident:
self._resident.move_to_end(key)
else:
self._resident[key] = None
while len(self._resident) > self.capacity:
evicted_key, _ = self._resident.popitem(last=False)
evicted.append(evicted_key)
return tuple(evicted)
def forget(self, keys: Iterable[GdnPrefixKey]) -> None:
for key in keys:
self._resident.pop(key, None)

View File

@@ -0,0 +1,62 @@
#!/usr/bin/env bash
set -euo pipefail
VLLM_ROOT=${1:?usage: install_prebuilt_corex.sh VLLM_ROOT}
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
BUNDLE_DIR=${SCRIPT_DIR}/prebuilt/corex-3.2.3-ivcore10
MANIFEST=${BUNDLE_DIR}/SHA256SUMS
[[ -d "$VLLM_ROOT" ]] || {
printf 'vLLM root does not exist: %s\n' "$VLLM_ROOT" >&2
exit 2
}
[[ -f "$MANIFEST" ]] || {
printf 'prebuilt CoreX manifest is missing: %s\n' "$MANIFEST" >&2
exit 2
}
mapfile -t artifacts < <(awk '{print $2}' "$MANIFEST")
[[ "${#artifacts[@]}" -eq 12 ]] || {
printf 'expected 12 prebuilt CoreX artifacts, found %s\n' \
"${#artifacts[@]}" >&2
exit 2
}
for artifact in "${artifacts[@]}"; do
[[ "$artifact" == corex_*.so && "$artifact" != */* ]] || {
printf 'invalid prebuilt artifact name: %s\n' "$artifact" >&2
exit 2
}
done
(
cd "$BUNDLE_DIR"
sha256sum --strict --check SHA256SUMS
)
for artifact in "${artifacts[@]}"; do
install -m 0755 "$BUNDLE_DIR/$artifact" "$VLLM_ROOT/$artifact"
done
python3 - "$VLLM_ROOT" "${artifacts[@]}" <<'PY'
import pathlib
import struct
import sys
root = pathlib.Path(sys.argv[1])
for name in sys.argv[2:]:
path = root / name
if not path.is_file() or path.stat().st_size == 0:
raise SystemExit(f"installed CoreX extension is empty: {path}")
header = path.read_bytes()[:20]
if len(header) < 20 or header[:4] != b"\x7fELF":
raise SystemExit(f"installed CoreX extension is not ELF: {path}")
if header[4:6] != b"\x02\x01":
raise SystemExit(
f"installed CoreX extension is not 64-bit little-endian ELF: {path}")
machine = struct.unpack_from("<H", header, 18)[0]
if machine != 62:
raise SystemExit(
f"installed CoreX extension is not x86-64 ELF: {path} machine={machine}")
print(f"[ok] installed prebuilt CoreX extension {path}")
PY

View File

@@ -1,224 +1,224 @@
from typing import Dict, List, Optional
import torch
from vllm.attention.backends.abstract import AttentionMetadata
class MambaCacheManager:
def __init__(self, dtype, num_mamba_layers, max_batch_size,
conv_state_shape, temporal_state_shape):
conv_state = torch.empty(size=(num_mamba_layers, max_batch_size) +
conv_state_shape,
dtype=dtype,
device="cuda")
temporal_state = torch.zeros(size=(num_mamba_layers, max_batch_size) +
temporal_state_shape,
dtype=dtype,
device="cuda")
self.mamba_cache = (conv_state, temporal_state)
# Maps between the request id and a dict that maps between the seq_id
# and its index inside the self.mamba_cache
self.mamba_cache_indices_mapping: Dict[str, Dict[int, int]] = {}
def current_run_tensors(self, input_ids: torch.Tensor,
attn_metadata: AttentionMetadata, **kwargs):
"""
Return the tensors for the current run's conv and ssm state.
"""
if "seqlen_agnostic_capture_inputs" not in kwargs:
# We get here only on Prefill/Eager mode runs
request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"]
finished_requests_ids = kwargs["finished_requests_ids"]
self._release_finished_requests(finished_requests_ids)
mamba_cache_tensors = self._prepare_current_run_mamba_cache(
request_ids_to_seq_ids, finished_requests_ids)
else:
# CUDA graph capturing runs
mamba_cache_tensors = kwargs["seqlen_agnostic_capture_inputs"]
return mamba_cache_tensors
def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs):
"""
Copy the relevant Mamba cache into the CUDA graph input buffer
that was provided during the capture runs
(JambaForCausalLM.mamba_gc_cache_buffer).
"""
assert all(
key in kwargs
for key in ["request_ids_to_seq_ids", "finished_requests_ids"])
finished_requests_ids = kwargs["finished_requests_ids"]
request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"]
self._release_finished_requests(finished_requests_ids)
self._prepare_current_run_mamba_cache(request_ids_to_seq_ids,
finished_requests_ids)
def get_seqlen_agnostic_capture_inputs(self, batch_size: int):
"""
Provide the CUDA graph capture runs with a buffer in adjusted size.
The buffer is used to maintain the Mamba Cache during the CUDA graph
replay runs.
"""
return tuple(buffer[:, :batch_size] for buffer in self.mamba_cache)
def _swap_mamba_cache(self, from_index: int, to_index: int):
assert len(self.mamba_cache) > 0
for cache_t in self.mamba_cache:
cache_t[:, [to_index,from_index]] = \
cache_t[:, [from_index,to_index]]
def _copy_mamba_cache(self, from_index: int, to_index: int):
assert len(self.mamba_cache) > 0
for cache_t in self.mamba_cache:
cache_t[:, to_index].copy_(cache_t[:, from_index],
non_blocking=True)
def _move_out_if_already_occupied(self, index: int,
all_occupied_indices: List[int]):
if index in all_occupied_indices:
first_free_index = self._first_free_index_in_mamba_cache()
# In case occupied, move the occupied to a new empty block
self._move_cache_index_and_mappings(from_index=index,
to_index=first_free_index)
def _assign_seq_id_to_mamba_cache_in_specific_dest(self, cur_rid: str,
seq_id: int,
destination_index: int):
"""
Assign (req_id,seq_id) pair to a `destination_index` index, if
already occupied, move the occupying index to a free index.
"""
all_occupied_indices = self._get_all_occupied_indices()
if cur_rid not in self.mamba_cache_indices_mapping:
self._move_out_if_already_occupied(
index=destination_index,
all_occupied_indices=all_occupied_indices)
for cache_t in self.mamba_cache:
cache_t[:, destination_index].zero_()
self.mamba_cache_indices_mapping[cur_rid] = {
seq_id: destination_index
}
elif seq_id not in (seq_ids2indices :=
self.mamba_cache_indices_mapping[cur_rid]):
# parallel sampling , where n > 1, assume prefill have
# already happened now we only need to copy the already
# existing cache into the siblings seq_ids caches
self._move_out_if_already_occupied(
index=destination_index,
all_occupied_indices=all_occupied_indices)
index_exists = list(seq_ids2indices.values())[0]
# case of decoding n>1, copy prefill cache to decoding indices
self._copy_mamba_cache(from_index=index_exists,
to_index=destination_index)
self.mamba_cache_indices_mapping[cur_rid][
seq_id] = destination_index
else:
# already exists
cache_index_already_exists = self.mamba_cache_indices_mapping[
cur_rid][seq_id]
if cache_index_already_exists != destination_index:
# In case the seq id already exists but not in
# the right destination, swap it with what's occupying it
self._swap_pair_indices_and_mappings(
from_index=cache_index_already_exists,
to_index=destination_index)
def _prepare_current_run_mamba_cache(
self, request_ids_to_seq_ids: Dict[str, list[int]],
finished_requests_ids: List[str]):
running_indices = []
request_ids_to_seq_ids_flatten = [
(req_id, seq_id)
for req_id, seq_ids in request_ids_to_seq_ids.items()
for seq_id in seq_ids
]
batch_size = len(request_ids_to_seq_ids_flatten)
for dest_index, (request_id,
seq_id) in enumerate(request_ids_to_seq_ids_flatten):
if request_id in finished_requests_ids:
# Do not allocate cache index for requests that run
# and finish right after
continue
self._assign_seq_id_to_mamba_cache_in_specific_dest(
request_id, seq_id, dest_index)
running_indices.append(dest_index)
self._clean_up_first_bs_blocks(batch_size, running_indices)
conv_state = self.mamba_cache[0][:, :batch_size]
temporal_state = self.mamba_cache[1][:, :batch_size]
return (conv_state, temporal_state)
def _get_all_occupied_indices(self):
return [
cache_idx
for seq_ids2indices in self.mamba_cache_indices_mapping.values()
for cache_idx in seq_ids2indices.values()
]
def _clean_up_first_bs_blocks(self, batch_size: int,
indices_for_current_run: List[int]):
# move out all of the occupied but currently not running blocks
# outside of the first n blocks
destination_indices = range(batch_size)
max_possible_batch_size = self.mamba_cache[0].shape[1]
for destination_index in destination_indices:
if destination_index in self._get_all_occupied_indices() and \
destination_index not in indices_for_current_run:
# move not running indices outside of the batch
all_other_indices = list(
range(batch_size, max_possible_batch_size))
first_avail_index = self._first_free_index_in_mamba_cache(
all_other_indices)
self._swap_indices(from_index=destination_index,
to_index=first_avail_index)
def _move_cache_index_and_mappings(self, from_index: int, to_index: int):
self._copy_mamba_cache(from_index=from_index, to_index=to_index)
self._update_mapping_index(from_index=from_index, to_index=to_index)
def _swap_pair_indices_and_mappings(self, from_index: int, to_index: int):
self._swap_mamba_cache(from_index=from_index, to_index=to_index)
self._swap_mapping_index(from_index=from_index, to_index=to_index)
def _swap_mapping_index(self, from_index: int, to_index: int):
for seq_ids2index in self.mamba_cache_indices_mapping.values():
for seq_id, index in seq_ids2index.items():
if from_index == index:
seq_ids2index.update({seq_id: to_index})
elif to_index == index:
seq_ids2index.update({seq_id: from_index})
def _update_mapping_index(self, from_index: int, to_index: int):
for seq_ids2index in self.mamba_cache_indices_mapping.values():
for seq_id, index in seq_ids2index.items():
if from_index == index:
seq_ids2index.update({seq_id: to_index})
return
def _release_finished_requests(self,
finished_seq_groups_req_ids: List[str]):
for req_id in finished_seq_groups_req_ids:
if req_id in self.mamba_cache_indices_mapping:
self.mamba_cache_indices_mapping.pop(req_id)
def _first_free_index_in_mamba_cache(
self, indices_range: Optional[List[int]] = None) -> int:
assert self.mamba_cache is not None
if indices_range is None:
max_possible_batch_size = self.mamba_cache[0].shape[1]
indices_range = list(range(max_possible_batch_size))
all_occupied_indices = self._get_all_occupied_indices()
for i in indices_range:
if i not in all_occupied_indices:
return i
raise Exception("Couldn't find a free spot in the mamba cache! This"
"should never happen")
from typing import Dict, List, Optional
import torch
from vllm.attention.backends.abstract import AttentionMetadata
class MambaCacheManager:
def __init__(self, dtype, num_mamba_layers, max_batch_size,
conv_state_shape, temporal_state_shape):
conv_state = torch.empty(size=(num_mamba_layers, max_batch_size) +
conv_state_shape,
dtype=dtype,
device="cuda")
temporal_state = torch.zeros(size=(num_mamba_layers, max_batch_size) +
temporal_state_shape,
dtype=dtype,
device="cuda")
self.mamba_cache = (conv_state, temporal_state)
# Maps between the request id and a dict that maps between the seq_id
# and its index inside the self.mamba_cache
self.mamba_cache_indices_mapping: Dict[str, Dict[int, int]] = {}
def current_run_tensors(self, input_ids: torch.Tensor,
attn_metadata: AttentionMetadata, **kwargs):
"""
Return the tensors for the current run's conv and ssm state.
"""
if "seqlen_agnostic_capture_inputs" not in kwargs:
# We get here only on Prefill/Eager mode runs
request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"]
finished_requests_ids = kwargs["finished_requests_ids"]
self._release_finished_requests(finished_requests_ids)
mamba_cache_tensors = self._prepare_current_run_mamba_cache(
request_ids_to_seq_ids, finished_requests_ids)
else:
# CUDA graph capturing runs
mamba_cache_tensors = kwargs["seqlen_agnostic_capture_inputs"]
return mamba_cache_tensors
def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs):
"""
Copy the relevant Mamba cache into the CUDA graph input buffer
that was provided during the capture runs
(JambaForCausalLM.mamba_gc_cache_buffer).
"""
assert all(
key in kwargs
for key in ["request_ids_to_seq_ids", "finished_requests_ids"])
finished_requests_ids = kwargs["finished_requests_ids"]
request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"]
self._release_finished_requests(finished_requests_ids)
self._prepare_current_run_mamba_cache(request_ids_to_seq_ids,
finished_requests_ids)
def get_seqlen_agnostic_capture_inputs(self, batch_size: int):
"""
Provide the CUDA graph capture runs with a buffer in adjusted size.
The buffer is used to maintain the Mamba Cache during the CUDA graph
replay runs.
"""
return tuple(buffer[:, :batch_size] for buffer in self.mamba_cache)
def _swap_mamba_cache(self, from_index: int, to_index: int):
assert len(self.mamba_cache) > 0
for cache_t in self.mamba_cache:
cache_t[:, [to_index,from_index]] = \
cache_t[:, [from_index,to_index]]
def _copy_mamba_cache(self, from_index: int, to_index: int):
assert len(self.mamba_cache) > 0
for cache_t in self.mamba_cache:
cache_t[:, to_index].copy_(cache_t[:, from_index],
non_blocking=True)
def _move_out_if_already_occupied(self, index: int,
all_occupied_indices: List[int]):
if index in all_occupied_indices:
first_free_index = self._first_free_index_in_mamba_cache()
# In case occupied, move the occupied to a new empty block
self._move_cache_index_and_mappings(from_index=index,
to_index=first_free_index)
def _assign_seq_id_to_mamba_cache_in_specific_dest(self, cur_rid: str,
seq_id: int,
destination_index: int):
"""
Assign (req_id,seq_id) pair to a `destination_index` index, if
already occupied, move the occupying index to a free index.
"""
all_occupied_indices = self._get_all_occupied_indices()
if cur_rid not in self.mamba_cache_indices_mapping:
self._move_out_if_already_occupied(
index=destination_index,
all_occupied_indices=all_occupied_indices)
for cache_t in self.mamba_cache:
cache_t[:, destination_index].zero_()
self.mamba_cache_indices_mapping[cur_rid] = {
seq_id: destination_index
}
elif seq_id not in (seq_ids2indices :=
self.mamba_cache_indices_mapping[cur_rid]):
# parallel sampling , where n > 1, assume prefill have
# already happened now we only need to copy the already
# existing cache into the siblings seq_ids caches
self._move_out_if_already_occupied(
index=destination_index,
all_occupied_indices=all_occupied_indices)
index_exists = list(seq_ids2indices.values())[0]
# case of decoding n>1, copy prefill cache to decoding indices
self._copy_mamba_cache(from_index=index_exists,
to_index=destination_index)
self.mamba_cache_indices_mapping[cur_rid][
seq_id] = destination_index
else:
# already exists
cache_index_already_exists = self.mamba_cache_indices_mapping[
cur_rid][seq_id]
if cache_index_already_exists != destination_index:
# In case the seq id already exists but not in
# the right destination, swap it with what's occupying it
self._swap_pair_indices_and_mappings(
from_index=cache_index_already_exists,
to_index=destination_index)
def _prepare_current_run_mamba_cache(
self, request_ids_to_seq_ids: Dict[str, list[int]],
finished_requests_ids: List[str]):
running_indices = []
request_ids_to_seq_ids_flatten = [
(req_id, seq_id)
for req_id, seq_ids in request_ids_to_seq_ids.items()
for seq_id in seq_ids
]
batch_size = len(request_ids_to_seq_ids_flatten)
for dest_index, (request_id,
seq_id) in enumerate(request_ids_to_seq_ids_flatten):
if request_id in finished_requests_ids:
# Do not allocate cache index for requests that run
# and finish right after
continue
self._assign_seq_id_to_mamba_cache_in_specific_dest(
request_id, seq_id, dest_index)
running_indices.append(dest_index)
self._clean_up_first_bs_blocks(batch_size, running_indices)
conv_state = self.mamba_cache[0][:, :batch_size]
temporal_state = self.mamba_cache[1][:, :batch_size]
return (conv_state, temporal_state)
def _get_all_occupied_indices(self):
return [
cache_idx
for seq_ids2indices in self.mamba_cache_indices_mapping.values()
for cache_idx in seq_ids2indices.values()
]
def _clean_up_first_bs_blocks(self, batch_size: int,
indices_for_current_run: List[int]):
# move out all of the occupied but currently not running blocks
# outside of the first n blocks
destination_indices = range(batch_size)
max_possible_batch_size = self.mamba_cache[0].shape[1]
for destination_index in destination_indices:
if destination_index in self._get_all_occupied_indices() and \
destination_index not in indices_for_current_run:
# move not running indices outside of the batch
all_other_indices = list(
range(batch_size, max_possible_batch_size))
first_avail_index = self._first_free_index_in_mamba_cache(
all_other_indices)
self._swap_indices(from_index=destination_index,
to_index=first_avail_index)
def _move_cache_index_and_mappings(self, from_index: int, to_index: int):
self._copy_mamba_cache(from_index=from_index, to_index=to_index)
self._update_mapping_index(from_index=from_index, to_index=to_index)
def _swap_pair_indices_and_mappings(self, from_index: int, to_index: int):
self._swap_mamba_cache(from_index=from_index, to_index=to_index)
self._swap_mapping_index(from_index=from_index, to_index=to_index)
def _swap_mapping_index(self, from_index: int, to_index: int):
for seq_ids2index in self.mamba_cache_indices_mapping.values():
for seq_id, index in seq_ids2index.items():
if from_index == index:
seq_ids2index.update({seq_id: to_index})
elif to_index == index:
seq_ids2index.update({seq_id: from_index})
def _update_mapping_index(self, from_index: int, to_index: int):
for seq_ids2index in self.mamba_cache_indices_mapping.values():
for seq_id, index in seq_ids2index.items():
if from_index == index:
seq_ids2index.update({seq_id: to_index})
return
def _release_finished_requests(self,
finished_seq_groups_req_ids: List[str]):
for req_id in finished_seq_groups_req_ids:
if req_id in self.mamba_cache_indices_mapping:
self.mamba_cache_indices_mapping.pop(req_id)
def _first_free_index_in_mamba_cache(
self, indices_range: Optional[List[int]] = None) -> int:
assert self.mamba_cache is not None
if indices_range is None:
max_possible_batch_size = self.mamba_cache[0].shape[1]
indices_range = list(range(max_possible_batch_size))
all_occupied_indices = self._get_all_occupied_indices()
for i in indices_range:
if i not in all_occupied_indices:
return i
raise Exception("Couldn't find a free spot in the mamba cache! This"
"should never happen")

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,91 @@
from patch_utils import package_root, replace_once
CACHE_ENGINE = package_root("vllm") / "worker" / "cache_engine.py"
IMPORT_ANCHOR = """\
from vllm.logger import init_logger
"""
IMPORT_REPLACEMENT = """\
from vllm.block_major_kv_cache import (
BlockMajorCpuKVCache,
block_major_cpu_kv_enabled,
)
from vllm.logger import init_logger
"""
ALLOCATION_ANCHOR = """\
self.gpu_cache = self._allocate_kv_cache(
self.num_gpu_blocks, self.device_config.device_type)
self.cpu_cache = self._allocate_kv_cache(self.num_cpu_blocks, "cpu")
"""
ALLOCATION_REPLACEMENT = """\
self.gpu_cache = self._allocate_kv_cache(
self.num_gpu_blocks, self.device_config.device_type)
self._bi100_block_major_cpu_kv = None
if block_major_cpu_kv_enabled():
self._bi100_block_major_cpu_kv = BlockMajorCpuKVCache(
self.gpu_cache,
self.num_cpu_blocks,
pin_memory=is_pin_memory_available(),
)
self.cpu_cache = self._bi100_block_major_cpu_kv.layer_views
else:
self.cpu_cache = self._allocate_kv_cache(
self.num_cpu_blocks, "cpu")
"""
SWAP_ANCHOR = """\
def swap_in(self, src_to_dst: torch.Tensor) -> None:
for i in range(self.num_attention_layers):
self.attn_backend.swap_blocks(self.cpu_cache[i], self.gpu_cache[i],
src_to_dst)
def swap_out(self, src_to_dst: torch.Tensor) -> None:
for i in range(self.num_attention_layers):
self.attn_backend.swap_blocks(self.gpu_cache[i], self.cpu_cache[i],
src_to_dst)
"""
SWAP_REPLACEMENT = """\
def swap_in(self, src_to_dst: torch.Tensor) -> None:
if self._bi100_block_major_cpu_kv is not None:
self._bi100_block_major_cpu_kv.swap_in(src_to_dst)
return
for i in range(self.num_attention_layers):
self.attn_backend.swap_blocks(self.cpu_cache[i], self.gpu_cache[i],
src_to_dst)
def swap_out(self, src_to_dst: torch.Tensor) -> None:
if self._bi100_block_major_cpu_kv is not None:
self._bi100_block_major_cpu_kv.swap_out(src_to_dst)
return
for i in range(self.num_attention_layers):
self.attn_backend.swap_blocks(self.gpu_cache[i], self.cpu_cache[i],
src_to_dst)
"""
replace_once(
CACHE_ENGINE,
IMPORT_ANCHOR,
IMPORT_REPLACEMENT,
required=True,
already_contains="from vllm.block_major_kv_cache import",
)
replace_once(
CACHE_ENGINE,
ALLOCATION_ANCHOR,
ALLOCATION_REPLACEMENT,
required=True,
already_contains="self._bi100_block_major_cpu_kv = None",
)
replace_once(
CACHE_ENGINE,
SWAP_ANCHOR,
SWAP_REPLACEMENT,
required=True,
already_contains="self._bi100_block_major_cpu_kv.swap_in",
)

View File

@@ -0,0 +1,46 @@
from patch_utils import package_root, replace_once
WORKER = package_root("vllm") / "worker" / "worker.py"
IMPORT_ANCHOR = """\
from vllm.logger import init_logger
"""
IMPORT_REPLACEMENT = """\
from vllm.block_major_kv_cache import reserve_block_major_gpu_blocks
from vllm.logger import init_logger
"""
CAPACITY_ANCHOR = """\
num_gpu_blocks = max(num_gpu_blocks, 0)
num_cpu_blocks = max(num_cpu_blocks, 0)
"""
CAPACITY_REPLACEMENT = """\
num_gpu_blocks = reserve_block_major_gpu_blocks(
num_gpu_blocks, cache_block_size)
num_gpu_blocks = max(num_gpu_blocks, 0)
num_cpu_blocks = max(num_cpu_blocks, 0)
"""
replace_once(
WORKER,
IMPORT_ANCHOR,
IMPORT_REPLACEMENT,
required=True,
already_contains=(
"from vllm.block_major_kv_cache import "
"reserve_block_major_gpu_blocks"
),
)
replace_once(
WORKER,
CAPACITY_ANCHOR,
CAPACITY_REPLACEMENT,
required=True,
already_contains=(
"num_gpu_blocks = reserve_block_major_gpu_blocks("
),
)

View File

@@ -0,0 +1,210 @@
"""Install the optional BI100 prefix-cache diagnostic trace."""
from patch_utils import package_root, replace_once, replace_one_of
VLLM_ROOT = package_root("vllm")
TARGET = VLLM_ROOT / "core" / "block_manager_v2.py"
OUTPUTS_TARGET = VLLM_ROOT / "outputs.py"
HELPER = '''
def _bi100_capture_cache_trace(self, seq_group, seq, block_table) -> None:
if os.getenv("BI100_CACHE_TRACE", "0") != "1":
return
session = getattr(self, "_bi100_trace_session", None)
if session is None:
session = hashlib.sha256(os.urandom(16)).hexdigest()[:16]
self._bi100_trace_session = session
self._bi100_trace_ordinal = getattr(self, "_bi100_trace_ordinal", 0) + 1
request_id_sha256 = hashlib.sha256(
str(seq_group.request_id).encode("utf-8")).hexdigest()[:16]
prompt_tokens = len(seq.get_token_ids())
requests = getattr(self, "_bi100_trace_requests", None)
if requests is None:
requests = {}
self._bi100_trace_requests = requests
requests[seq.seq_id] = {
"version": 4,
"trace_session_sha256": session,
"ordinal": self._bi100_trace_ordinal,
"request_id_sha256": request_id_sha256,
"prompt_tokens": prompt_tokens,
"prompt_allocated_blocks": (
(prompt_tokens + self.block_size - 1) // self.block_size
),
"block_size": self.block_size,
"capacity_blocks": self.num_total_gpu_blocks,
}
setattr(seq_group, "_bi100_cache_trace_seq_id", seq.seq_id)
setattr(seq_group, "_bi100_cache_trace_emit",
self._bi100_emit_cache_trace)
def _bi100_update_cache_trace(
self, seq, raw_kv_hit_blocks, restore_key, capture_actions,
evict_keys, policy) -> None:
if os.getenv("BI100_CACHE_TRACE", "0") != "1":
return
requests = getattr(self, "_bi100_trace_requests", None)
if not requests or seq.seq_id not in requests:
return
record = requests[seq.seq_id]
record["gdn_policy"] = policy
if "initial_raw_kv_contiguous_hit_blocks" not in record:
record["initial_raw_kv_contiguous_hit_blocks"] = max(
0, int(raw_kv_hit_blocks))
record["gdn_restore_digest_base64"] = (
base64.b64encode(restore_key[1]).decode("ascii")
if restore_key is not None else None)
record["raw_kv_contiguous_hit_blocks"] = max(
int(raw_kv_hit_blocks),
int(record.get("raw_kv_contiguous_hit_blocks", 0)))
effective_blocks = int(restore_key[0]) if restore_key is not None else 0
record["effective_gdn_hit_blocks"] = max(
effective_blocks, int(record.get("effective_gdn_hit_blocks", 0)))
admissions = record.setdefault("gdn_admissions", [])
for key, reason in capture_actions:
admissions.append({
"block_count": int(key[0]),
"digest_base64": base64.b64encode(key[1]).decode("ascii"),
"reason": str(reason),
})
evictions = record.setdefault("gdn_evictions", [])
for key in evict_keys:
evictions.append({
"block_count": int(key[0]),
"digest_base64": base64.b64encode(key[1]).decode("ascii"),
"reason": "capacity_lru",
})
def _bi100_finalize_cache_trace(self, seq, block_table) -> None:
if os.getenv("BI100_CACHE_TRACE", "0") != "1":
return
requests = getattr(self, "_bi100_trace_requests", None)
if not requests:
return
record = requests.get(seq.seq_id)
if record is None:
return
total_tokens = len(seq.get_token_ids())
block_hashes = block_table.get_content_hashes()
for block_hash in block_hashes:
if not isinstance(block_hash, bytes) or len(block_hash) != 32:
raise RuntimeError(
"BI100 cache trace requires 32-byte content hashes")
full_blocks = len(block_hashes)
record.update({
"total_tokens": total_tokens,
"allocated_blocks": (
(total_tokens + self.block_size - 1) // self.block_size
),
"full_blocks": full_blocks,
"hash_encoding": "sha256_base64",
"block_hashes": base64.b64encode(b"".join(block_hashes)).decode("ascii"),
"_finalized": True,
})
generated_tokens = max(0, total_tokens - record["prompt_tokens"])
record["generated_tokens"] = generated_tokens
def _bi100_emit_cache_trace(self, seq_group) -> None:
if os.getenv("BI100_CACHE_TRACE", "0") != "1":
return
seq_id = getattr(seq_group, "_bi100_cache_trace_seq_id", None)
requests = getattr(self, "_bi100_trace_requests", None)
if seq_id is None or not requests:
return
record = requests.pop(seq_id, None)
if record is None:
return
if record.pop("_finalized", False) is not True:
raise RuntimeError(
"BI100 cache trace emitted before block finalization")
metrics = getattr(seq_group, "metrics", None)
arrival = getattr(metrics, "arrival_time", None)
first_token = getattr(metrics, "first_token_time", None)
finished = getattr(metrics, "finished_time", None)
queue = getattr(metrics, "time_in_queue", None)
cached = getattr(metrics, "num_cached_tokens", None)
if any(value is None for value in (
arrival, first_token, finished, queue)):
raise RuntimeError(
"BI100 cache trace requires finalized request metrics")
record["ttft_s"] = max(0.0, float(first_token - arrival))
record["request_latency_s"] = max(
0.0, float(finished - arrival))
record["time_in_queue_s"] = max(0.0, float(queue))
record["observed_effective_cached_tokens"] = max(
0, int(cached or 0))
ttft_s = record["ttft_s"]
if ttft_s > 0:
record["observed_input_tps"] = record["prompt_tokens"] / ttft_s
generated_tokens = record["generated_tokens"]
if generated_tokens > 1:
decode_s = finished - first_token
if decode_s > 0:
record["observed_output_tps"] = (
(generated_tokens - 1) / decode_s)
print("[BI100_CACHE_TRACE] " + json.dumps(record, separators=(",", ":"),
sort_keys=True), flush=True)
'''
def main():
replace_once(TARGET, "from collections.abc import Mapping\n",
"from collections.abc import Mapping\nimport base64\nimport json\nimport os\n",
required=True, already_contains="import base64\n")
replace_once(TARGET, "class BlockSpaceManagerV2(BlockSpaceManager):\n",
"class BlockSpaceManagerV2(BlockSpaceManager):\n" + HELPER,
required=True, already_contains="def _bi100_capture_cache_trace(")
replace_once(TARGET,
" self.block_tables[seq.seq_id] = block_table\n\n # Track seq",
" self.block_tables[seq.seq_id] = block_table\n self._bi100_capture_cache_trace(\n seq_group, seq, block_table)\n\n # Track seq",
required=True,
already_contains="self.block_tables[seq.seq_id] = block_table\n"
" self._bi100_capture_cache_trace(")
replacements = []
for table_key in ("seq_id", "seq.seq_id"):
prefix = (
" self._last_access_blocks_tracker."
"update_seq_blocks_last_access(\n"
f" seq_id, self.block_tables[{table_key}]."
"physical_block_ids)\n")
replacements.append((
prefix + "\n # Untrack seq",
prefix + " self._bi100_finalize_cache_trace(\n"
f" seq, self.block_tables[{table_key}])\n\n"
" # Untrack seq",
))
replace_one_of(
TARGET,
replacements,
required=True,
already_contains=" self._bi100_finalize_cache_trace(\n"
" seq, self.block_tables[")
replace_once(
OUTPUTS_TARGET,
" seq_group.set_finished_time(finished_time)\n\n"
" init_args = (seq_group.request_id, prompt, prompt_token_ids,\n",
" seq_group.set_finished_time(finished_time)\n"
" if finished_time is not None:\n"
" cache_trace_emit = getattr(\n"
" seq_group, \"_bi100_cache_trace_emit\", None)\n"
" if callable(cache_trace_emit):\n"
" cache_trace_emit(seq_group)\n"
" delattr(seq_group, \"_bi100_cache_trace_emit\")\n"
" delattr(seq_group, \"_bi100_cache_trace_seq_id\")\n\n"
" init_args = (seq_group.request_id, prompt, prompt_token_ids,\n",
required=True,
already_contains="if finished_time is not None:\n"
" cache_trace_emit = getattr(\n",
)
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,65 @@
from patch_utils import package_root, replace_once
CUSTOM_OPS = package_root("vllm") / "_custom_ops.py"
CLEAN_BLOCK = """\
def swap_blocks(src: torch.Tensor, dst: torch.Tensor,
block_mapping: torch.Tensor) -> None:
ixf_F.swap_blocks(src, dst, block_mapping)
"""
COMPATIBLE_BLOCK = """\
def swap_blocks(src: torch.Tensor, dst: torch.Tensor,
block_mapping: torch.Tensor) -> None:
# BI100 CoreX 3.2.3 exposes vllm_swap_blocks, while this vLLM build calls
# the newer swap_blocks name. Normalize the worker's CPU int64 [N, 2]
# tensor only for the legacy public API and fail fast on malformed maps.
native_swap_blocks = getattr(ixf_F, "swap_blocks", None)
if native_swap_blocks is not None:
native_swap_blocks(src, dst, block_mapping)
return
vendor_swap_blocks = getattr(ixf_F, "vllm_swap_blocks", None)
if vendor_swap_blocks is None:
raise RuntimeError(
"ixformer exposes neither swap_blocks nor vllm_swap_blocks")
if isinstance(block_mapping, torch.Tensor):
if block_mapping.device.type != "cpu":
raise ValueError("swap block mapping must be a CPU tensor")
if block_mapping.dtype != torch.int64:
raise ValueError("swap block mapping must use torch.int64")
if block_mapping.dim() != 2 or block_mapping.shape[1] != 2:
raise ValueError("swap block mapping must have shape [N, 2]")
pairs = block_mapping.tolist()
elif isinstance(block_mapping, dict):
pairs = list(block_mapping.items())
else:
raise TypeError("swap block mapping must be a tensor or dict")
normalized_mapping = {}
destinations = set()
for source, destination in pairs:
source = int(source)
destination = int(destination)
if source < 0 or destination < 0:
raise ValueError("swap block indices must be non-negative")
if source in normalized_mapping:
raise ValueError(f"duplicate swap source block: {source}")
if destination in destinations:
raise ValueError(
f"duplicate swap destination block: {destination}")
normalized_mapping[source] = destination
destinations.add(destination)
vendor_swap_blocks(src, dst, normalized_mapping)
"""
replace_once(
CUSTOM_OPS,
CLEAN_BLOCK,
COMPATIBLE_BLOCK,
required=True,
already_contains="BI100 CoreX 3.2.3 exposes vllm_swap_blocks",
)

View File

@@ -0,0 +1,61 @@
from patch_utils import package_root, replace_once
VLLM_ROOT = package_root("vllm")
MULTIPROC_GPU_EXECUTOR = VLLM_ROOT / "executor" / "multiproc_gpu_executor.py"
MULTIPROC_WORKER_UTILS = VLLM_ROOT / "executor" / "multiproc_worker_utils.py"
def ensure_import_os(path):
text = path.read_text()
if "import os\n" in text:
print(f"[skip] import os already present: {path}")
return
for anchor in ("import time\n", "import signal\n", "import sys\n"):
if anchor in text:
replace_once(
path,
anchor,
anchor + "import os\n",
required=True,
already_contains="import os\n",
)
return
raise RuntimeError(f"no import anchor found for os in {path}")
ensure_import_os(MULTIPROC_GPU_EXECUTOR)
ensure_import_os(MULTIPROC_WORKER_UTILS)
replace_once(
MULTIPROC_GPU_EXECUTOR,
"""logger = init_logger(__name__)\n""",
"""logger = init_logger(__name__)\n\n\ndef _bi100_startup_debug(message: str, *args) -> None:\n if os.getenv(\"BI100_EXECUTOR_STARTUP_DEBUG\") == \"1\":\n logger.info(\"[BI100 startup] \" + message, *args)\n""",
required=True,
already_contains="def _bi100_startup_debug(",
)
replace_once(
MULTIPROC_GPU_EXECUTOR,
""" self.driver_worker = self._create_worker(\n distributed_init_method=distributed_init_method)\n self._run_workers(\"init_device\")\n self._run_workers(\"load_model\",\n max_concurrent_workers=self.parallel_config.\n max_parallel_loading_workers)\n""",
""" _bi100_startup_debug(\"creating driver worker\")\n self.driver_worker = self._create_worker(\n distributed_init_method=distributed_init_method)\n _bi100_startup_debug(\"created driver worker\")\n _bi100_startup_debug(\"starting init_device\")\n self._run_workers(\"init_device\")\n _bi100_startup_debug(\"finished init_device\")\n _bi100_startup_debug(\"starting load_model\")\n self._run_workers(\"load_model\",\n max_concurrent_workers=self.parallel_config.\n max_parallel_loading_workers)\n _bi100_startup_debug(\"finished load_model\")\n""",
required=True,
already_contains='_bi100_startup_debug("starting init_device")',
)
replace_once(
MULTIPROC_GPU_EXECUTOR,
""" # Start all remote workers first.\n worker_outputs = [\n worker.execute_method(method, *args, **kwargs)\n for worker in self.workers\n ]\n\n driver_worker_method = getattr(self.driver_worker, method)\n driver_worker_output = driver_worker_method(*args, **kwargs)\n\n # Get the results of the workers.\n return [driver_worker_output\n ] + [output.get() for output in worker_outputs]\n""",
""" _bi100_startup_debug(\"enqueue remote method=%s workers=%d\", method,\n len(self.workers))\n # Start all remote workers first.\n worker_outputs = [\n worker.execute_method(method, *args, **kwargs)\n for worker in self.workers\n ]\n _bi100_startup_debug(\"remote enqueued method=%s\", method)\n\n driver_worker_method = getattr(self.driver_worker, method)\n _bi100_startup_debug(\"driver start method=%s\", method)\n driver_worker_output = driver_worker_method(*args, **kwargs)\n _bi100_startup_debug(\"driver done method=%s\", method)\n\n # Get the results of the workers.\n _bi100_startup_debug(\"waiting remote results method=%s\", method)\n remote_outputs = [output.get() for output in worker_outputs]\n _bi100_startup_debug(\"remote done method=%s\", method)\n return [driver_worker_output] + remote_outputs\n""",
required=True,
already_contains='_bi100_startup_debug("enqueue remote method=%s workers=%d"',
)
replace_once(
MULTIPROC_WORKER_UTILS,
""" task_id, method, args, kwargs = items\n try:\n executor = getattr(worker, method)\n output = executor(*args, **kwargs)\n except SystemExit:\n""",
""" task_id, method, args, kwargs = items\n if os.getenv(\"BI100_EXECUTOR_STARTUP_DEBUG\") == \"1\":\n logger.info(\"[BI100 worker] start method=%s\", method)\n try:\n executor = getattr(worker, method)\n output = executor(*args, **kwargs)\n if os.getenv(\"BI100_EXECUTOR_STARTUP_DEBUG\") == \"1\":\n logger.info(\"[BI100 worker] done method=%s\", method)\n except SystemExit:\n""",
required=True,
already_contains='logger.info("[BI100 worker] start method=%s", method)',
)

View File

@@ -1,46 +1,43 @@
"""
Fix: prefix_cache_hit stays True for chunked-prefill chunk 2+ even when past cache.
"""Patch vLLM 0.6.3 prefix-cache and MRoPE chunk alignment bugs."""
Root cause:
model_runner.py _compute_for_prefix_cache_hit has three cases:
Case 1: prefix_cache_len <= context_len → "already past cache, do normal"
Case 2: context_len < prefix_cache_len < seq_len → partial hit, correct
Case 3: seq_len <= prefix_cache_len → full hit, reduce to 1 token
from __future__ import annotations
Case 1 does nothing (leaves prefix_cache_hit = True). Then in utils.py:
if inter_data.prefix_cache_hit:
block_table = computed_block_nums ← ONLY the original prefix blocks!
import pathlib
But context_len > prefix_cache_len means chunk 1 tokens (between prefix_cache_len
and context_len) are ALSO in KV cache and need to be in block_table.
block_table = computed_block_nums misses all chunk-1 blocks.
from patch_utils import package_root, replace_once
In _forward_prefix_pytorch:
num_ctx_blocks = ceil(context_len / block_size) # e.g. 268
block_tables.shape[1] = len(computed_block_nums) # e.g. 12 <-- too small!
At tile_blk >= 12: blk_ids is empty → k_t shape [..., 0] → amax crash.
Fix:
Set prefix_cache_hit = False for Case 1, so utils.py falls through to:
elif chunked_prefill_enabled:
block_table = block_tables[seq_id] ← full block table (prefix + chunk1)
"""
HELPER_ANCHOR = """\
logger = init_logger(__name__)
import re
import sys
LORA_WARMUP_RANK = 8"""
CANDIDATE_PATHS = [
"/usr/local/corex/lib64/python3/dist-packages/vllm/worker/model_runner.py",
"/usr/local/corex/lib/python3/dist-packages/vllm/worker/model_runner.py",
]
HELPER_REPLACEMENT = """\
logger = init_logger(__name__)
OLD_BLOCK = """\
def _slice_mrope_positions(positions, start, stop, expected_len):
if positions is None or len(positions) != 3:
raise RuntimeError("MRoPE positions must contain three axes")
sliced = [axis[start:stop] for axis in positions]
lengths = [len(axis) for axis in sliced]
if lengths != [expected_len] * 3:
raise RuntimeError(
"MRoPE/input token length mismatch after chunk alignment: "
f"positions={lengths}, input_tokens={expected_len}, "
f"slice=({start}, {stop})")
return sliced
LORA_WARMUP_RANK = 8"""
PREFIX_PAST_ANCHOR = """\
if prefix_cache_len <= context_len:
# We already passed the cache hit region,
# so do normal computation.
pass"""
NEW_BLOCK = """\
PREFIX_PAST_REPLACEMENT = """\
if prefix_cache_len <= context_len:
# We already passed the cache hit region,
# so do normal computation.
@@ -51,28 +48,361 @@ NEW_BLOCK = """\
# causing an empty blk_ids slice and a zero-dim amax() crash.
inter_data.prefix_cache_hit = False"""
import os
PARTIAL_HIT_ANCHOR = """\
inter_data.input_positions[seq_idx] = inter_data.input_positions[
seq_idx][uncomputed_start:]
context_len = prefix_cache_len
patched = False
for path in CANDIDATE_PATHS:
if not os.path.exists(path):
continue
with open(path, "r") as f:
src = f.read()
if OLD_BLOCK not in src:
if NEW_BLOCK in src:
print(f"[patch_model_runner] already patched: {path}")
patched = True
break
print(f"[patch_model_runner] WARNING: expected block not found in {path}, skipping")
continue
patched_src = src.replace(OLD_BLOCK, NEW_BLOCK, 1)
with open(path, "w") as f:
f.write(patched_src)
print(f"[patch_model_runner] patched Case-1 prefix_cache_hit fix in: {path}")
patched = True
break
inter_data.context_lens[seq_idx] = context_len
inter_data.query_lens[
seq_idx] = inter_data.seq_lens[seq_idx] - context_len"""
if not patched:
print("[patch_model_runner] ERROR: could not find model_runner.py at any known path", file=sys.stderr)
sys.exit(1)
PARTIAL_HIT_REPLACEMENT = """\
inter_data.input_positions[seq_idx] = inter_data.input_positions[
seq_idx][uncomputed_start:]
context_len = prefix_cache_len
inter_data.context_lens[seq_idx] = context_len
inter_data.query_lens[
seq_idx] = inter_data.seq_lens[seq_idx] - context_len
if inter_data.mrope_input_positions is not None:
positions = inter_data.mrope_input_positions[seq_idx]
if positions is not None:
inter_data.mrope_input_positions[seq_idx] = \\
_slice_mrope_positions(
positions, uncomputed_start, None,
inter_data.query_lens[seq_idx])"""
FULL_HIT_ANCHOR = """\
inter_data.input_positions[seq_idx] = inter_data.input_positions[
seq_idx][-1:]
inter_data.query_lens[seq_idx] = 1
inter_data.context_lens[seq_idx] = inter_data.seq_lens[seq_idx] - 1"""
FULL_HIT_REPLACEMENT = """\
inter_data.input_positions[seq_idx] = inter_data.input_positions[
seq_idx][-1:]
inter_data.query_lens[seq_idx] = 1
inter_data.context_lens[seq_idx] = inter_data.seq_lens[seq_idx] - 1
if inter_data.mrope_input_positions is not None:
positions = inter_data.mrope_input_positions[seq_idx]
if positions is not None:
inter_data.mrope_input_positions[seq_idx] = \\
_slice_mrope_positions(positions, -1, None, 1)"""
MULTIMODAL_MROPE_ANCHOR = """\
mrope_input_positions, mrope_position_delta = \\
MRotaryEmbedding.get_input_positions(
token_ids,
image_grid_thw=image_grid_thw,
video_grid_thw=video_grid_thw,
image_token_id=hf_config.image_token_id,
video_token_id=hf_config.video_token_id,
vision_start_token_id=hf_config.vision_start_token_id,
vision_end_token_id=hf_config.vision_end_token_id,
spatial_merge_size=hf_config.vision_config.
spatial_merge_size,
context_len=inter_data.context_lens[seq_idx],
)
seq_data.mrope_position_delta = mrope_position_delta
inter_data.mrope_input_positions[
seq_idx] = mrope_input_positions"""
MULTIMODAL_MROPE_REPLACEMENT = """\
# vLLM 0.6.3 returns positions through the end of token_ids,
# while chunked prefill sends only [context_len:seq_len].
# Compute the full MRoPE map once so the delta remains tied to
# the complete request, then select exactly the physical query.
mrope_input_positions, mrope_position_delta = \\
MRotaryEmbedding.get_input_positions(
token_ids,
image_grid_thw=image_grid_thw,
video_grid_thw=video_grid_thw,
image_token_id=hf_config.image_token_id,
video_token_id=hf_config.video_token_id,
vision_start_token_id=hf_config.vision_start_token_id,
vision_end_token_id=hf_config.vision_end_token_id,
spatial_merge_size=hf_config.vision_config.
spatial_merge_size,
context_len=0,
)
mrope_input_positions = _slice_mrope_positions(
mrope_input_positions,
inter_data.context_lens[seq_idx],
inter_data.seq_lens[seq_idx],
len(inter_data.input_tokens[seq_idx]))
seq_data.mrope_position_delta = mrope_position_delta
inter_data.mrope_input_positions[
seq_idx] = mrope_input_positions"""
MODEL_INPUT_FIELDS_ANCHOR = """\
multi_modal_kwargs: Optional[BatchedTensorInputs] = None
request_ids_to_seq_ids: Optional[Dict[str, List[int]]] = None"""
MODEL_INPUT_FIELDS_REPLACEMENT = """\
multi_modal_kwargs: Optional[BatchedTensorInputs] = None
# BI100 scheduler-owned GDN prefix-cache actions. These plain Python
# objects are included in the multiprocess model-input broadcast.
gdn_restore_key: Optional[Tuple[int, bytes]] = None
gdn_capture_points: Optional[List[Tuple[int, Tuple[int, bytes]]]] = None
gdn_evict_keys: Optional[List[Tuple[int, bytes]]] = None
gdn_segment_offsets: Optional[List[int]] = None
request_ids_to_seq_ids: Optional[Dict[str, List[int]]] = None"""
BASE_BROADCAST_ANCHOR = """\
\"multi_modal_kwargs\": self.multi_modal_kwargs,
\"prompt_adapter_mapping\": self.prompt_adapter_mapping,
\"prompt_adapter_requests\": self.prompt_adapter_requests,
\"virtual_engine\": self.virtual_engine,
\"request_ids_to_seq_ids\": self.request_ids_to_seq_ids,
\"finished_requests_ids\": self.finished_requests_ids,
}
_add_attn_metadata_broadcastable_dict(tensor_dict, self.attn_metadata)
return tensor_dict
@classmethod"""
BASE_BROADCAST_REPLACEMENT = """\
\"multi_modal_kwargs\": self.multi_modal_kwargs,
\"gdn_restore_key\": self.gdn_restore_key,
\"gdn_capture_points\": self.gdn_capture_points,
\"gdn_evict_keys\": self.gdn_evict_keys,
\"gdn_segment_offsets\": self.gdn_segment_offsets,
\"prompt_adapter_mapping\": self.prompt_adapter_mapping,
\"prompt_adapter_requests\": self.prompt_adapter_requests,
\"virtual_engine\": self.virtual_engine,
\"request_ids_to_seq_ids\": self.request_ids_to_seq_ids,
\"finished_requests_ids\": self.finished_requests_ids,
}
_add_attn_metadata_broadcastable_dict(tensor_dict, self.attn_metadata)
return tensor_dict
@classmethod"""
SAMPLING_BROADCAST_ANCHOR = """\
\"multi_modal_kwargs\": self.multi_modal_kwargs,
\"prompt_adapter_mapping\": self.prompt_adapter_mapping,
\"prompt_adapter_requests\": self.prompt_adapter_requests,
\"virtual_engine\": self.virtual_engine,
\"request_ids_to_seq_ids\": self.request_ids_to_seq_ids,
\"finished_requests_ids\": self.finished_requests_ids,
}
_add_attn_metadata_broadcastable_dict(tensor_dict, self.attn_metadata)
_add_sampling_metadata_broadcastable_dict(tensor_dict,
self.sampling_metadata)"""
SAMPLING_BROADCAST_REPLACEMENT = """\
\"multi_modal_kwargs\": self.multi_modal_kwargs,
\"gdn_restore_key\": self.gdn_restore_key,
\"gdn_capture_points\": self.gdn_capture_points,
\"gdn_evict_keys\": self.gdn_evict_keys,
\"gdn_segment_offsets\": self.gdn_segment_offsets,
\"prompt_adapter_mapping\": self.prompt_adapter_mapping,
\"prompt_adapter_requests\": self.prompt_adapter_requests,
\"virtual_engine\": self.virtual_engine,
\"request_ids_to_seq_ids\": self.request_ids_to_seq_ids,
\"finished_requests_ids\": self.finished_requests_ids,
}
_add_attn_metadata_broadcastable_dict(tensor_dict, self.attn_metadata)
_add_sampling_metadata_broadcastable_dict(tensor_dict,
self.sampling_metadata)"""
BUILDER_INIT_ANCHOR = """\
self.finished_requests_ids = finished_requests_ids
self.decode_only = True
# Intermediate data"""
BUILDER_INIT_REPLACEMENT = """\
self.finished_requests_ids = finished_requests_ids
self.decode_only = True
self.gdn_restore_key = None
self.gdn_capture_points = None
self.gdn_evict_keys = None
self.gdn_segment_offsets = None
# Intermediate data"""
ADD_SEQ_GROUP_ANCHOR = """\
def add_seq_group(self, seq_group_metadata: SequenceGroupMetadata):
\"\"\"Add a sequence group to the builder.\"\"\"
seq_ids = seq_group_metadata.seq_data.keys()"""
ADD_SEQ_GROUP_REPLACEMENT = """\
def add_seq_group(self, seq_group_metadata: SequenceGroupMetadata):
\"\"\"Add a sequence group to the builder.\"\"\"
gdn_actions = (
seq_group_metadata.gdn_restore_key,
seq_group_metadata.gdn_capture_points,
seq_group_metadata.gdn_evict_keys,
seq_group_metadata.gdn_segment_offsets,
)
if any(value is not None for value in gdn_actions):
if not seq_group_metadata.is_prompt:
raise RuntimeError(\"GDN prefix-cache actions require prefill\")
if any(value is not None for value in (
self.gdn_restore_key, self.gdn_capture_points,
self.gdn_evict_keys, self.gdn_segment_offsets)):
raise RuntimeError(
\"only one GDN prefix-cache action group is supported\")
(self.gdn_restore_key, self.gdn_capture_points,
self.gdn_evict_keys, self.gdn_segment_offsets) = gdn_actions
seq_ids = seq_group_metadata.seq_data.keys()"""
BUILD_RESULT_ANCHOR = """\
lora_mapping=lora_mapping,
lora_requests=lora_requests,
multi_modal_kwargs=multi_modal_kwargs,
request_ids_to_seq_ids=request_ids_to_seq_ids,"""
BUILD_RESULT_REPLACEMENT = """\
lora_mapping=lora_mapping,
lora_requests=lora_requests,
multi_modal_kwargs=multi_modal_kwargs,
gdn_restore_key=self.gdn_restore_key,
gdn_capture_points=self.gdn_capture_points,
gdn_evict_keys=self.gdn_evict_keys,
gdn_segment_offsets=self.gdn_segment_offsets,
request_ids_to_seq_ids=request_ids_to_seq_ids,"""
EXECUTE_KWARGS_ANCHOR = """\
seqlen_agnostic_kwargs = {
\"finished_requests_ids\": model_input.finished_requests_ids,
\"request_ids_to_seq_ids\": model_input.request_ids_to_seq_ids,
} if self.has_inner_state else {}
if (self.observability_config is not None"""
EXECUTE_KWARGS_REPLACEMENT = """\
seqlen_agnostic_kwargs = {
\"finished_requests_ids\": model_input.finished_requests_ids,
\"request_ids_to_seq_ids\": model_input.request_ids_to_seq_ids,
} if self.has_inner_state else {}
gdn_prefix_kwargs = {}
if model_input.gdn_restore_key is not None:
gdn_prefix_kwargs[\"gdn_restore_key\"] = model_input.gdn_restore_key
if model_input.gdn_capture_points is not None:
gdn_prefix_kwargs[\"gdn_capture_points\"] = (
model_input.gdn_capture_points)
if model_input.gdn_evict_keys is not None:
gdn_prefix_kwargs[\"gdn_evict_keys\"] = model_input.gdn_evict_keys
if model_input.gdn_segment_offsets is not None:
gdn_prefix_kwargs[\"gdn_segment_offsets\"] = (
model_input.gdn_segment_offsets)
if (self.observability_config is not None"""
MODEL_CALL_ANCHOR = """\
**MultiModalInputs.as_kwargs(multi_modal_kwargs,
device=self.device),
**seqlen_agnostic_kwargs)"""
MODEL_CALL_REPLACEMENT = """\
**MultiModalInputs.as_kwargs(multi_modal_kwargs,
device=self.device),
**seqlen_agnostic_kwargs,
**gdn_prefix_kwargs)"""
PROFILE_KV_LAYERS_ANCHOR = """\
num_layers = self.model_config.get_num_layers(self.parallel_config)"""
PROFILE_KV_LAYERS_REPLACEMENT = """\
num_layers = self.model_config.get_num_attention_layers(
self.parallel_config)"""
def patch_model_runner(model_runner: pathlib.Path) -> None:
replace_once(
model_runner,
HELPER_ANCHOR,
HELPER_REPLACEMENT,
required=True,
already_contains="def _slice_mrope_positions(",
)
replace_once(
model_runner,
PREFIX_PAST_ANCHOR,
PREFIX_PAST_REPLACEMENT,
required=True,
already_contains="Must clear prefix_cache_hit so _add_seq_group",
)
replace_once(
model_runner,
PARTIAL_HIT_ANCHOR,
PARTIAL_HIT_REPLACEMENT,
required=True,
already_contains="positions, uncomputed_start, None,",
)
replace_once(
model_runner,
FULL_HIT_ANCHOR,
FULL_HIT_REPLACEMENT,
required=True,
already_contains="_slice_mrope_positions(positions, -1, None, 1)",
)
replace_once(
model_runner,
MULTIMODAL_MROPE_ANCHOR,
MULTIMODAL_MROPE_REPLACEMENT,
required=True,
already_contains="Compute the full MRoPE map once",
)
replace_once(
model_runner,
MODEL_INPUT_FIELDS_ANCHOR,
MODEL_INPUT_FIELDS_REPLACEMENT,
already_contains="gdn_restore_key: Optional[Tuple[int, bytes]]",
)
replace_once(
model_runner,
BASE_BROADCAST_ANCHOR,
BASE_BROADCAST_REPLACEMENT,
already_contains=BASE_BROADCAST_REPLACEMENT,
)
replace_once(
model_runner,
SAMPLING_BROADCAST_ANCHOR,
SAMPLING_BROADCAST_REPLACEMENT,
already_contains=SAMPLING_BROADCAST_REPLACEMENT,
)
replace_once(
model_runner,
BUILDER_INIT_ANCHOR,
BUILDER_INIT_REPLACEMENT,
already_contains="self.gdn_restore_key = None",
)
replace_once(
model_runner,
ADD_SEQ_GROUP_ANCHOR,
ADD_SEQ_GROUP_REPLACEMENT,
already_contains="gdn_actions = (",
)
replace_once(
model_runner,
BUILD_RESULT_ANCHOR,
BUILD_RESULT_REPLACEMENT,
already_contains="gdn_restore_key=self.gdn_restore_key",
)
replace_once(
model_runner,
EXECUTE_KWARGS_ANCHOR,
EXECUTE_KWARGS_REPLACEMENT,
already_contains="gdn_prefix_kwargs = {}",
)
replace_once(
model_runner,
MODEL_CALL_ANCHOR,
MODEL_CALL_REPLACEMENT,
already_contains="**gdn_prefix_kwargs)",
)
replace_once(
model_runner,
PROFILE_KV_LAYERS_ANCHOR,
PROFILE_KV_LAYERS_REPLACEMENT,
required=True,
already_contains=PROFILE_KV_LAYERS_REPLACEMENT,
)
if __name__ == "__main__":
patch_model_runner(package_root("vllm") / "worker" / "model_runner.py")

View File

@@ -1,191 +1,247 @@
#!/bin/bash
# ==========================================================================
# PATCH_OPS.SH v2 — Align with comp 168 strategy
#!/usr/bin/env bash
# BI-V100 patch script for Qwen3.6-35B-A3B (Qwen3_5 MoE architecture)
#
# COMP 168 PROOF (dockerrizhi.txt 07-23 lines 310-397):
# corex_gdn.py:56 → dlopen libcorex_gdn.so ✅
# corex_gdn.py:228 → GDN prefill fused
# corex_gdn.py:138 → GDN decode fused ✅
# corex_moe.py:339 → MoE prefill: expert-grouped-wmma ✅
# corex_moe.py:249 → MoE decode fused ✅
# corex_fa2.py:333 → FA2 packed prefill ✅
# corex_fa2.py:507 → FA2 paged chunked prefill ✅
# corex_fa2.py:225 → FA2 paged decode ✅
# Triton situation on BI-V100:
# - Standard Triton 2.3.1 is already present in the image.
# - HAS_TRITON = False (hardcoded in vendor vllm), but Triton is still used
# for TP-mode cache management (custom_cache_manager / libentry).
# - The vendor's triton_utils/__init__.py, custom_cache_manager.py, libentry.py
# are already correct for standard Triton 2.3.1 — do NOT overwrite them.
# - DO NOT install BI-V150 corex Triton 2.1.0 (pkgs/triton): that causes
# GPU hang on BI-V100 because the Triton CUDA PTX kernels are incompatible.
# Recommended server start command for TP=4 support 256K, needs chunked prefill
# CUDA_VISIBLE_DEVICES="4,5,6,7" VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 python3 -m vllm.entrypoints.openai.api_server \
# --model /workspace/models/Qwen3.6-35B-A3B --port 1111 --served-model-name llm \
# --max-model-len 262144 --trust-remote-code -tp 4 --gpu-memory-utilization 0.90 \
# --max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
# --max-num-batched-tokens 8192 --enable-chunked-prefill --enable-prefix-caching \
# --max-seq-len-to-capture 32768 --enable-auto-tool-choice \
# --tool-call-parser qwen3_coder --reasoning-parser qwen3
#
# ALL 3 corex modules are IN THE BASE IMAGE and work correctly.
# Our Sub508 failed because we OVERWROTE qwen3_5.py, breaking the call chain.
#
# STRATEGY: DO NOT TOUCH model layer. Only deploy:
# 1. transformers config (Qwen3_5Config)
# 2. serving layer (protocol/serving_chat/api_server/chat_utils/tool_parser/reasoning)
# 3. ix_bridge.so (fills ixf_F.vllm_moe_topk_softmax gap if base _custom_ops hits it)
# 4. _custom_ops.py patch (make topk_softmax use ix_bridge instead of crashing)
# ==========================================================================
# With prefix caching (GDN align-mode, requires chunked prefill):
# CUDA_VISIBLE_DEVICES="4,5,6,7" VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 python3 -m vllm.entrypoints.openai.api_server \
# --model /workspace/models/Qwen3.6-35B-A3B --port 1111 --served-model-name llm \
# --max-model-len 262144 --trust-remote-code -tp 4 --gpu-memory-utilization 0.90 \
# --max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
# --max-num-batched-tokens 8192 --enable-chunked-prefill --enable-prefix-caching \
# --max-seq-len-to-capture 32768 --enable-auto-tool-choice \
# --tool-call-parser qwen3_coder --reasoning-parser qwen3
cd "$(dirname "$0")"
echo "[patch_ops.v2] START — comp 168 aligned strategy"
set -euo pipefail
VLLM=""
for P in /usr/local/corex/lib/python3/dist-packages/vllm \
/usr/local/corex/lib64/python3/dist-packages/vllm; do
[ -d "$P" ] && VLLM="$P" && echo "[patch_ops] Found vllm at: $VLLM" && break
done
[ -z "$VLLM" ] && echo "[patch_ops] ERROR: vllm not found" && exit 1
build_stage() { printf '[BI100 BUILD] %s\n' "$1" >&2; }
require_file() {
local path=$1
[[ -f "$path" ]] || {
printf 'required patch source is missing: %s\n' "$path" >&2
exit 2
}
}
install_patch_file() {
local source=$1
local target=$2
# ---- PROBE ----
echo "[probe] === Base image state ==="
_QW="$VLLM/model_executor/models/qwen3_5.py"
[ -f "$_QW" ] && echo "[probe] qwen3_5.py: $(wc -c < "$_QW") bytes, $(wc -l < "$_QW") lines" || echo "[probe] qwen3_5.py: MISSING"
for m in corex_gdn.py corex_moe.py corex_fa2.py; do
_F="$VLLM/model_executor/models/$m"
[ -f "$_F" ] && echo "[probe] $m: $(wc -c < "$_F") bytes" || echo "[probe] $m: MISSING"
done
ls -la /usr/local/corex/lib64/libcorex_*.so 2>/dev/null || echo "[probe] no libcorex_*.so"
echo "[probe] ==========================="
# Find secondary vllm path for mirroring
VLLM2=""
for P in /usr/local/corex/lib/python3/dist-packages/vllm \
/usr/local/corex/lib64/python3/dist-packages/vllm; do
[ -d "$P" ] && [ "$P" != "$VLLM" ] && VLLM2="$P" && break
done
# Helper: deploy to both vllm paths
deploy_both() {
local src="$1" dst="$2"
cp "$src" "$VLLM/$dst" 2>/dev/null || true
[ -n "$VLLM2" ] && cp "$src" "$VLLM2/$dst" 2>/dev/null || true
require_file "$source"
mkdir -p "$(dirname "$target")"
install -m 0644 "$source" "$target"
}
# ===========================================================
# 1. Transformers config (Qwen3_5Config support)
# ===========================================================
TMODELS=""
for P in /usr/local/lib/python3.10/site-packages/transformers/models \
/usr/local/corex/lib/python3/dist-packages/transformers/models; do
[ -d "$P" ] && TMODELS="$P" && break
done
if [ -n "$TMODELS" ]; then
pip install transformers==4.55.3 -i https://pypi.tuna.tsinghua.edu.cn/simple --timeout 30 2>&1 || true
apt-get update -qq && apt-get install -y -qq ninja-build 2>&1 || true
cp -r ./qwen3_5 "$TMODELS/" 2>/dev/null || true
cp -r ./qwen3_5_moe "$TMODELS/" 2>/dev/null || true
python3 ./patch_transformers_qwen3_5.py 2>&1 || true
echo "[patch_ops] transformers config deployed"
build_stage "patch script entered"
build_stage "checking offline transformers dependency"
# --- transformers: Qwen3_5 tokenizer / model files --------------------------
TRANSFORMERS_REQUIRED_VERSION="4.55.3"
if ! python3 - "$TRANSFORMERS_REQUIRED_VERSION" <<'PY'
import importlib.metadata
import sys
required = sys.argv[1]
try:
installed = importlib.metadata.version("transformers")
except importlib.metadata.PackageNotFoundError:
raise SystemExit(1)
raise SystemExit(0 if installed == required else 1)
PY
then
WHEEL_DIR="./wheels"
if ! ls "${WHEEL_DIR}/transformers-${TRANSFORMERS_REQUIRED_VERSION}"*.whl >/dev/null 2>&1; then
echo "transformers ${TRANSFORMERS_REQUIRED_VERSION} is required, but no offline wheel was found in ${WHEEL_DIR}" >&2
exit 2
fi
python3 -m pip install --no-index --no-deps --find-links="${WHEEL_DIR}" \
"transformers==${TRANSFORMERS_REQUIRED_VERSION}"
fi
# ===========================================================
# 2. MODEL LAYER — CONDITIONAL deployment
# If base has qwen3_5.py > 1000 bytes → DO NOT OVERWRITE
# This is the comp 168 strategy.
# ===========================================================
_QW_SIZE=0
[ -f "$_QW" ] && _QW_SIZE=$(wc -c < "$_QW")
python3 - "$TRANSFORMERS_REQUIRED_VERSION" <<'PY'
import importlib.metadata
import sys
if [ "$_QW_SIZE" -gt 1000 ]; then
echo "[patch_ops] *** BASE IMAGE HAS qwen3_5.py (${_QW_SIZE} bytes) — KEEPING IT ***"
echo "[patch_ops] *** This is the comp 168 strategy: don't break corex_* call chain ***"
# Only add registry entry if missing
if ! grep -q "Qwen3_5ForCausalLM" "$VLLM/model_executor/models/registry.py" 2>/dev/null; then
cp ./registry.py "$VLLM/model_executor/models/registry.py" 2>/dev/null && \
echo "[patch_ops] registry.py deployed (was missing Qwen3_5)"
[ -n "$VLLM2" ] && cp ./registry.py "$VLLM2/model_executor/models/registry.py" 2>/dev/null || true
fi
else
echo "[patch_ops] *** BASE IMAGE MISSING qwen3_5.py — deploying ours ***"
deploy_both ./qwen3_5.py "model_executor/models/qwen3_5.py"
deploy_both ./registry.py "model_executor/models/registry.py"
deploy_both ./mamba_cache.py "model_executor/models/mamba_cache.py"
# Only deploy corex modules if base doesn't have them
for m in corex_gdn.py corex_moe.py corex_fa2.py; do
if [ ! -f "$VLLM/model_executor/models/$m" ]; then
deploy_both "/workspace/ex_engine/python/$m" "model_executor/models/$m"
echo "[patch_ops] deployed $m (was MISSING)"
fi
done
# flash_qla_sm70 (only if we deployed our qwen3_5.py)
_FLASH_SRC="/workspace/qwen3_6_scripts/flash_qla_sm70"
if [ -d "$_FLASH_SRC" ]; then
for _VPATH in "$VLLM" "$VLLM2"; do
[ -z "$_VPATH" ] && continue
cp -r "$_FLASH_SRC" "$_VPATH/model_executor/models/flash_qla_sm70" 2>/dev/null || true
done
echo "[patch_ops] flash_qla_sm70 deployed"
fi
fi
required = sys.argv[1]
installed = importlib.metadata.version("transformers")
if installed != required:
raise SystemExit(
f"transformers version mismatch: expected {required}, got {installed}")
print(f"[ok] transformers {installed}")
PY
# ===========================================================
# 3. SERVING LAYER — always deploy (comp 168 also used custom serving)
# ===========================================================
mkdir -p "$VLLM/entrypoints/openai/tool_parsers" 2>/dev/null || true
[ -n "$VLLM2" ] && mkdir -p "$VLLM2/entrypoints/openai/tool_parsers" 2>/dev/null || true
build_stage "discovering Python package roots"
python3 - <<'PY' > /tmp/qwen36_patch_paths.env
from patch_utils import package_root, shell_env_line
deploy_both ./protocol.py "entrypoints/openai/protocol.py"
deploy_both ./cli_args.py "entrypoints/openai/cli_args.py"
deploy_both ./serving_chat.py "entrypoints/openai/serving_chat.py"
deploy_both ./api_server.py "entrypoints/openai/api_server.py"
deploy_both ./chat_utils.py "entrypoints/chat_utils.py"
deploy_both ./qwen3coder_tool_parser.py "entrypoints/openai/tool_parsers/qwen3coder_tool_parser.py"
deploy_both ./tool_parsers_init.py "entrypoints/openai/tool_parsers/__init__.py"
python3 ./patch_vllm_tool_parser.py 2>&1 || true
cp -r ./reasoning "$VLLM/" 2>/dev/null || true
[ -n "$VLLM2" ] && cp -r ./reasoning "$VLLM2/" 2>/dev/null || true
echo "[patch_ops] serving layer deployed"
print(shell_env_line("VLLM_ROOT", package_root("vllm")))
print(shell_env_line("TRANSFORMERS_ROOT", package_root("transformers")))
PY
source /tmp/qwen36_patch_paths.env
# ===========================================================
# 4. ix_bridge.so — ONLY PURPOSE: fill ixf_F.vllm_moe_topk_softmax gap
# Even comp 168 had this issue — the base _custom_ops.py tries to call
# ixf_F.vllm_moe_topk_softmax which doesn't exist.
# BUT comp 168's corex_moe.py bypasses _custom_ops entirely.
# So ix_bridge is only needed if base qwen3_5.py path hits _custom_ops.
# ===========================================================
_SITE="/usr/local/corex/lib/python3/dist-packages"
if [ -d "$_SITE" ]; then
_EX_DST="$_SITE/ex_engine"
mkdir -p "$_EX_DST/python" "$_EX_DST/build" "$_EX_DST/csrc"
cp /workspace/ex_engine/python/*.py "$_EX_DST/python/" 2>/dev/null || true
touch "$_EX_DST/__init__.py" "$_EX_DST/python/__init__.py"
# Deploy pre-built .so
if [ -d "/workspace/ex_engine/build" ]; then
cp /workspace/ex_engine/build/*.so "$_EX_DST/build/" 2>/dev/null || true
cp /workspace/ex_engine/build/*.so "$_EX_DST/" 2>/dev/null || true
echo "[patch_ops] ex_engine .so deployed: $(ls /workspace/ex_engine/build/*.so 2>/dev/null | wc -l) files"
fi
# C++ sources for JIT
cp /workspace/ex_engine/csrc/ix_full_bridge.cpp "$_EX_DST/csrc/" 2>/dev/null || true
cp /workspace/ex_engine/csrc/ix_moe_bridge.cpp "$_EX_DST/csrc/" 2>/dev/null || true
echo "[patch_ops] ex_engine package deployed to $_SITE"
fi
echo "VLLM_ROOT=${VLLM_ROOT}"
echo "TRANSFORMERS_ROOT=${TRANSFORMERS_ROOT}"
[[ -d "$VLLM_ROOT" ]] || {
printf 'vLLM root does not exist: %s\n' "$VLLM_ROOT" >&2
exit 2
}
# ===========================================================
# 5. XFormers patches — head_dim=256 bypass for BI-V100
# Comp 168 also had xformers patches (base uses xformers for attention)
# ===========================================================
python3 ./patch_xformers_sdpa_seq.py 2>&1 || true
python3 ./patch_xformers_sdpa_batch.py 2>&1 || true
echo "[patch_ops] xformers patches applied"
VLLM_OVERRIDE_ROOT="./vendor_overrides/vllm"
[[ -d "$VLLM_OVERRIDE_ROOT" ]] || {
printf 'vLLM override directory missing: %s\n' "$VLLM_OVERRIDE_ROOT" >&2
exit 2
}
# ===========================================================
# 6. model_runner patch (prefix_cache_hit fix)
# ===========================================================
python3 ./patch_model_runner.py 2>&1 || true
echo "[patch_ops] model_runner patched"
build_stage "installing authoritative vLLM core block overrides"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/core/evictor_v2.py" \
"${VLLM_ROOT}/core/evictor_v2.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/core/block/cpu_kv_content_cache.py" \
"${VLLM_ROOT}/core/block/cpu_kv_content_cache.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/core/block/cpu_gpu_block_allocator.py" \
"${VLLM_ROOT}/core/block/cpu_gpu_block_allocator.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/core/block/prefix_caching_block.py" \
"${VLLM_ROOT}/core/block/prefix_caching_block.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/core/block/block_table.py" \
"${VLLM_ROOT}/core/block/block_table.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/core/block_manager_v2.py" \
"${VLLM_ROOT}/core/block_manager_v2.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/sampling_params.py" \
"${VLLM_ROOT}/sampling_params.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/model_executor/sampling_metadata.py" \
"${VLLM_ROOT}/model_executor/sampling_metadata.py"
install_patch_file \
"${VLLM_OVERRIDE_ROOT}/model_executor/layers/sampler.py" \
"${VLLM_ROOT}/model_executor/layers/sampler.py"
# ===========================================================
# 7. Deploy precompiled .so files
# ===========================================================
for _SO in /workspace/ex_engine/moe_topk_softmax_v3*.so /tmp/torch_extensions/*/moe_topk_softmax_v3*.so; do
[ -f "$_SO" ] && cp "$_SO" "$_SITE/" 2>/dev/null && echo "[patch_ops] MoE topk .so: $(basename $_SO)" && break
done
for _SO in /workspace/ex_engine/moe_ops_v055*.so /tmp/torch_extensions/*/moe_ops_v055*.so; do
[ -f "$_SO" ] && cp "$_SO" "$_SITE/" 2>/dev/null && echo "[patch_ops] MoE v055 .so: $(basename $_SO)" && break
done
build_stage "installing hash-pinned CoreX 3.2.3 extensions"
bash ./install_prebuilt_corex.sh "${VLLM_ROOT}"
echo "[patch_ops.v2] DONE — comp 168 aligned"
echo "[patch_ops.v2] KEY: base qwen3_5.py $([ "$_QW_SIZE" -gt 1000 ] && echo "KEPT" || echo "REPLACED"), serving layer deployed"
build_stage "installing BI100 runtime modules"
cp ./bi100_env.py "${VLLM_ROOT}/bi100_env.py"
cp ./bi100_profile.py "${VLLM_ROOT}/bi100_profile.py"
cp ./block_major_kv_cache.py "${VLLM_ROOT}/block_major_kv_cache.py"
cp ./gdn_prefix.py "${VLLM_ROOT}/gdn_prefix.py"
build_stage "installing CoreX paged-KV swap compatibility"
python3 ./patch_corex_swap_blocks.py
python3 ./patch_block_major_cache_engine.py
python3 ./patch_worker_cache_transfer_order.py
# --- paged_attn.py: replace forward_prefix with pure-PyTorch fallback -------
# The Triton context_attention_fwd kernel hangs BI-V100 GPUs permanently
# (standard Triton 2.3.1 PTX is not supported by the corex runtime either).
# Our paged_attn.py bypasses it entirely via _forward_prefix_pytorch, which
# utilizes K-tiling techniques, and also have _forward_decode_pytorch to bypass kernel
# when context length is high
cp ./paged_attn.py "${VLLM_ROOT}/attention/ops/paged_attn.py"
# --- model_runner.py: fix prefix_cache_hit stays True in chunked-prefill chunk 2+ ---
# Bug: _compute_for_prefix_cache_hit Case 1 (prefix_cache_len <= context_len)
# leaves prefix_cache_hit=True. Then _add_seq_group uses block_table=computed_block_nums
# (only the original prefix blocks), ignoring chunk-1 KV cache blocks.
# _forward_prefix_pytorch then gets an undersized block_tables and crashes with
# "amax(): Expected reduction dim -1 to have non-zero size" on the 2nd tile.
# Fix: set prefix_cache_hit=False for Case 1 so the full block_tables is used.
python3 ./patch_model_runner.py
build_stage "installing executor startup diagnostics"
python3 ./patch_executor_startup_debug.py
python3 ./patch_worker_startup_profile_guard.py
python3 ./patch_block_major_worker_capacity.py
build_stage "installing transformers Qwen3.5 model support"
cp -r ./qwen3_5 "${TRANSFORMERS_ROOT}/models/"
cp -r ./qwen3_5_moe "${TRANSFORMERS_ROOT}/models/"
python3 ./patch_transformers_qwen3_5.py
build_stage "installing vLLM Qwen3.6 model implementation"
# --- vllm model: Qwen3.6-35B-A3B (Qwen3_5 MoE arch) -------------------------
cp ./mamba_cache.py "${VLLM_ROOT}/model_executor/models/"
cp ./qwen3_5.py "${VLLM_ROOT}/model_executor/models/qwen3_5.py"
python3 ./patch_vllm_qwen3_5.py
# --- sequence.py: fix completion_tokens inflation under chunked prefill ------
# Bug: get_output_token_ids_to_return(delta=True) with num_new_tokens=0
# returns _cached_all_token_ids[-0:] == [0:] (the ENTIRE prompt+output list).
# Each prefill chunk step adds prompt_len to previous_num_tokens, so a 10K
# prompt processed in 3 chunks inflates completion_tokens by ~30K.
# Also adds num_cached_tokens field to RequestMetrics for prefix-cache stats.
cp ./sequence.py "${VLLM_ROOT}/sequence.py"
# --- scheduler.py: record num_cached_tokens in RequestMetrics ----------------
# Reports only the longest prefix backed by both live KV blocks and an exact
# GDN restore state. Raw KV-only hits must not inflate cached_tokens.
# serving_chat.py exposes the value in the OpenAI-compatible usage details.
cp ./scheduler.py "${VLLM_ROOT}/core/scheduler.py"
build_stage "installing diagnostic initial allocation trace"
python3 ./patch_block_manager_cache_trace.py
build_stage "installing scheduler and attention patches"
# --- xformers: bypass cudnnFlashAttnForward (head_dim=256 > 128 limit) ------
# Injects _run_sdpa_fallback (pure matmul+softmax) into xformers.py.
# Required because head_dim=256 > 128 and ixformer flash attention either
# crashes (is_causal=True) or produces wrong output (attn_mask path).
# The fallback uses query_start_loc to derive actual query lengths, so it
# works correctly during profiling runs with chunked-prefill-style batches.
# also bypasses auto chunked prefill on
python3 ./patch_xformers_sdpa_seq.py
python3 ./patch_xformers_profile.py
build_stage "installing API parsers and serving modules"
# --- tool parser: Qwen3 XML tool call format ---------------------------------
# Registers "qwen3_coder" parser for Qwen3.6 XML-style tool calls:
# <tool_call><function=name><parameter=key>\nvalue\n</parameter></function></tool_call>
# Use at server start: --tool-call-parser qwen3_coder --enable-auto-tool-choice
cp ./qwen3coder_tool_parser.py "${VLLM_ROOT}/entrypoints/openai/tool_parsers/"
python3 ./patch_vllm_tool_parser.py
# --- reasoning parser: Qwen3 <think>...</think> split ------------------------
# Adds --reasoning-parser qwen3 support.
# Routes thinking tokens to reasoning_content, rest to content in the delta.
# Works together with --tool-call-parser qwen3_coder (think → tool call flow).
cp -r ./reasoning "${VLLM_ROOT}/"
cp ./protocol.py "${VLLM_ROOT}/entrypoints/openai/protocol.py"
cp ./cli_args.py "${VLLM_ROOT}/entrypoints/openai/cli_args.py"
cp ./serving_chat.py "${VLLM_ROOT}/entrypoints/openai/serving_chat.py"
cp ./serving_tokenization.py \
"${VLLM_ROOT}/entrypoints/openai/serving_tokenization.py"
cp ./api_server.py "${VLLM_ROOT}/entrypoints/openai/api_server.py"
cp ./chat_utils.py "${VLLM_ROOT}/entrypoints/chat_utils.py"
python3 - ./api_server.py \
"${VLLM_ROOT}/entrypoints/openai/api_server.py" <<'PY'
from pathlib import Path
import sys
source = Path(sys.argv[1]).read_bytes()
installed = Path(sys.argv[2]).read_bytes()
if source != installed:
raise SystemExit("runtime api_server overlay identity mismatch")
PY
build_stage "compiling submission Python sources"
find . -path './wheels' -prune -o -name '*.py' -print0 | xargs -0 python3 -m py_compile
build_stage "patch script completed"

View File

@@ -2,54 +2,23 @@
Patches transformers 4.55.3 to register qwen3_5 and qwen3_5_moe model types.
Deploy steps on the remote machine:
1. cp -r modified_scripts/qwen3_5 /usr/local/lib/python3.10/site-packages/transformers/models/qwen3_5
2. cp -r modified_scripts/qwen3_5_moe /usr/local/lib/python3.10/site-packages/transformers/models/qwen3_5_moe
1. patch_ops.sh locates transformers with importlib.util.find_spec.
2. cp -r modified_scripts/qwen3_5* into the detected transformers/models.
3. python3 modified_scripts/patch_transformers_qwen3_5.py
Target: pip-installed transformers at /usr/local/lib/python3.10/site-packages/transformers/
(Not the corex pre-installed path at /usr/local/corex/lib64/python3/dist-packages/)
"""
import sys
TRANSFORMERS_ROOT = None
for _p in ["/usr/local/lib/python3.10/site-packages/transformers",
"/usr/local/corex/lib/python3/dist-packages/transformers",
"/usr/local/corex/lib64/python3/dist-packages/transformers"]:
import os
if os.path.isdir(_p):
TRANSFORMERS_ROOT = _p
break
if TRANSFORMERS_ROOT is None:
TRANSFORMERS_ROOT = "/usr/local/lib/python3.10/site-packages/transformers"
AUTO_CONFIG = f"{TRANSFORMERS_ROOT}/models/auto/configuration_auto.py"
MODELS_INIT = f"{TRANSFORMERS_ROOT}/models/__init__.py"
from patch_utils import package_root, replace_once, replace_one_of
def patch_file(path, replacements):
with open(path, "r") as f:
content = f.read()
patched = False
for old, new in replacements:
if new in content:
print(f" [skip] already patched: {repr(new[:60])}")
continue
if old not in content:
print(f" [warn] anchor not found: {repr(old[:60])}")
continue
content = content.replace(old, new, 1)
patched = True
print(f" [ok] inserted after: {repr(old[:60])}")
if patched:
with open(path, "w") as f:
f.write(content)
TRANSFORMERS_ROOT = package_root("transformers")
AUTO_CONFIG = TRANSFORMERS_ROOT / "models" / "auto" / "configuration_auto.py"
MODELS_INIT = TRANSFORMERS_ROOT / "models" / "__init__.py"
def main():
print(f"=== Patching {AUTO_CONFIG} ===")
patch_file(AUTO_CONFIG, [
replace_one_of(AUTO_CONFIG, [
# CONFIG_MAPPING_NAMES: insert qwen3_5 + qwen3_5_moe right after qwen3
(
'("qwen3", "Qwen3Config"),',
@@ -59,6 +28,8 @@ def main():
'("qwen3", "Qwen3Config")\n',
'("qwen3", "Qwen3Config"),\n ("qwen3_5", "Qwen3_5Config"),\n ("qwen3_5_moe", "Qwen3_5MoeConfig"),\n',
),
], required=True, already_contains='("qwen3_5_moe", "Qwen3_5MoeConfig")')
replace_one_of(AUTO_CONFIG, [
# MODEL_NAMES_MAPPING (model_type -> human readable name)
(
'("qwen3", "Qwen3"),',
@@ -68,15 +39,15 @@ def main():
'("qwen3", "Qwen3")\n',
'("qwen3", "Qwen3"),\n ("qwen3_5", "Qwen3_5"),\n ("qwen3_5_moe", "Qwen3_5_MoE"),\n',
),
])
], required=True, already_contains='("qwen3_5_moe", "Qwen3_5_MoE")')
print(f"\n=== Patching {MODELS_INIT} ===")
patch_file(MODELS_INIT, [
(
"from .qwen3 import *\n",
"from .qwen3 import *\n from .qwen3_5 import *\n from .qwen3_5_moe import *\n",
),
])
replace_once(
MODELS_INIT,
"from .qwen3 import *\n",
"from .qwen3 import *\n from .qwen3_5 import *\n from .qwen3_5_moe import *\n",
required=True,
already_contains="from .qwen3_5_moe import *")
# Verification
print("\n=== Verification ===")
@@ -88,28 +59,31 @@ def main():
mod = importlib.util.module_from_spec(spec)
mod.__package__ = ".".join(module_name.split(".")[:-1])
pkg = sys.modules.setdefault("transformers", types.ModuleType("transformers"))
pkg.__path__ = [TRANSFORMERS_ROOT]
pkg.__path__ = [str(TRANSFORMERS_ROOT)]
cu = sys.modules.setdefault(
"transformers.configuration_utils", types.ModuleType("transformers.configuration_utils"))
class _PC:
def __init__(self, **kwargs): pass
def __init__(self, **kwargs):
return None
cu.PretrainedConfig = _PC
for sub in ("transformers.models", f"transformers.models.{module_name.split('.')[-2]}"):
m = sys.modules.setdefault(sub, types.ModuleType(sub))
m.__path__ = [TRANSFORMERS_ROOT]
m.__path__ = [str(TRANSFORMERS_ROOT)]
spec.loader.exec_module(mod)
return mod
mod27 = _load_config_mod(
"transformers.models.qwen3_5.configuration_qwen3_5",
f"{TRANSFORMERS_ROOT}/models/qwen3_5/configuration_qwen3_5.py",
str(TRANSFORMERS_ROOT / "models" / "qwen3_5" /
"configuration_qwen3_5.py"),
)
cfg = mod27.Qwen3_5Config()
print(f" Qwen3_5Config() smoke-test OK (model_type={cfg.model_type})")
mod35 = _load_config_mod(
"transformers.models.qwen3_5_moe.configuration_qwen3_5_moe",
f"{TRANSFORMERS_ROOT}/models/qwen3_5_moe/configuration_qwen3_5_moe.py",
str(TRANSFORMERS_ROOT / "models" / "qwen3_5_moe" /
"configuration_qwen3_5_moe.py"),
)
moe_cfg = mod35.Qwen3_5MoeConfig()
print(f" Qwen3_5MoeConfig() smoke-test OK (model_type={moe_cfg.model_type})")
@@ -117,7 +91,7 @@ def main():
print(f" num_experts={t.num_experts}, top_k={t.num_experts_per_tok}, "
f"shared={t.shared_expert_intermediate_size}, layers={t.num_hidden_layers}")
except Exception as e:
print(f" [warn] smoke-test failed (may be fine at runtime): {e}")
print(f" [optional] smoke-test failed (may be fine at runtime): {e}")
print("\nDone.")

View File

@@ -0,0 +1,81 @@
from __future__ import annotations
import importlib.util
import pathlib
import shlex
from typing import Iterable, Optional, Sequence, Tuple
def package_root(pkg: str) -> pathlib.Path:
spec = importlib.util.find_spec(pkg)
if spec is None:
raise RuntimeError(f"package not found: {pkg}")
if not spec.submodule_search_locations:
raise RuntimeError(f"package has no package root: {pkg}")
return pathlib.Path(next(iter(spec.submodule_search_locations))).resolve()
def ensure_file(path: pathlib.Path) -> pathlib.Path:
if not path.is_file():
raise FileNotFoundError(str(path))
return path
def ensure_dir(path: pathlib.Path) -> pathlib.Path:
if not path.is_dir():
raise FileNotFoundError(str(path))
return path
def replace_once(path: pathlib.Path,
old: str,
new: str,
*,
required: bool = True,
already_contains: Optional[str] = None) -> bool:
path = ensure_file(path)
text = path.read_text()
marker = already_contains if already_contains is not None else new
if marker in text:
print(f"[skip] already patched: {path}")
return False
if old not in text:
msg = f"anchor not found in {path}: {old[:120]!r}"
if required:
raise RuntimeError(msg)
print(f"[warn] {msg}")
return False
path.write_text(text.replace(old, new, 1))
print(f"[ok] patched: {path}")
return True
def replace_one_of(path: pathlib.Path,
replacements: Sequence[Tuple[str, str]],
*,
required: bool = True,
already_contains: Optional[str] = None) -> bool:
path = ensure_file(path)
text = path.read_text()
if already_contains is not None and already_contains in text:
print(f"[skip] already patched: {path}")
return False
for _, new in replacements:
if new in text:
print(f"[skip] already patched: {path}")
return False
for old, new in replacements:
if old in text:
path.write_text(text.replace(old, new, 1))
print(f"[ok] patched: {path}")
return True
anchors = ", ".join(repr(old[:80]) for old, _ in replacements)
msg = f"anchor not found in {path}; tried: {anchors}"
if required:
raise RuntimeError(msg)
print(f"[warn] {msg}")
return False
def shell_env_line(name: str, value: pathlib.Path) -> str:
return f"{name}={shlex.quote(str(value))}"

View File

@@ -0,0 +1,73 @@
"""
Patches the vLLM model registry and deploys the Qwen3_5 model file.
Deploy steps on the remote machine:
1. patch_ops.sh locates vLLM with importlib.util.find_spec.
2. cp modified_scripts/qwen3_5.py into the detected vllm model directory.
2. python3 modified_scripts/patch_vllm_qwen3_5.py
The registry patch installs Qwen3.6 aliases so /model/config.json does not
need to be edited by hand.
"""
import ast
from patch_utils import package_root, replace_once
VLLM_ROOT = package_root("vllm")
REGISTRY = VLLM_ROOT / "model_executor" / "models" / "registry.py"
MODEL = VLLM_ROOT / "model_executor" / "models" / "qwen3_5.py"
EXPECTED_REGISTRY_ENTRIES = (
'"Qwen3ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM")',
'"Qwen3MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM")',
'"Qwen3_5ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM")',
'"Qwen3_5MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM")',
'"Qwen3_6ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM")',
'"Qwen3_6MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM")',
)
def main():
print(f"=== Patching {REGISTRY} ===")
replace_once(
REGISTRY,
' "Qwen3ForCausalLM": ("qwen3", "Qwen3ForCausalLM"),\n'
' "Qwen3MoeForCausalLM": ("qwen3_moe", "Qwen3MoeForCausalLM"),',
' "Qwen3ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM"),\n'
' "Qwen3MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM"),\n'
' "Qwen3_5ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM"),\n'
' "Qwen3_5MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM"),\n'
' "Qwen3_6ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM"),\n'
' "Qwen3_6MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM"),',
required=True,
already_contains='"Qwen3_6MoeForCausalLM"')
print("\n=== Static verification ===")
model_source = MODEL.read_text(encoding="utf-8")
tree = ast.parse(model_source, filename=str(MODEL))
class_names = {
node.name for node in tree.body if isinstance(node, ast.ClassDef)
}
required_classes = {"Qwen3_5ForCausalLM", "Qwen3_5MoeForCausalLM"}
missing_classes = required_classes - class_names
if missing_classes:
raise RuntimeError(
f"Qwen3.5 model classes missing: {sorted(missing_classes)}")
registry_source = REGISTRY.read_text(encoding="utf-8")
missing_entries = [
entry for entry in EXPECTED_REGISTRY_ENTRIES
if entry not in registry_source
]
if missing_entries:
raise RuntimeError(
f"Qwen3.5 registry entries missing: {missing_entries}")
print(" model syntax and class declarations verified without import")
print(f" registry aliases verified: {len(EXPECTED_REGISTRY_ENTRIES)}")
print("\nDone. Registry aliases installed; do not edit /model/config.json.")
if __name__ == "__main__":
main()

View File

@@ -1,79 +1,57 @@
"""
Patches vLLM 0.6.3 to register Qwen3CoderToolParser under the name "qwen3_coder".
"""
Patches vLLM 0.6.3 to register Qwen3CoderToolParser under the name "qwen3_coder".
Deploy steps on the remote machine (already called by patch_ops.sh):
1. cp qwen3coder_tool_parser.py \
/usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/openai/tool_parsers/
1. patch_ops.sh locates vLLM with importlib.util.find_spec.
2. cp qwen3coder_tool_parser.py into the detected vllm tool_parsers.
2. python3 patch_vllm_tool_parser.py
Usage after patching:
--tool-call-parser qwen3_coder --enable-auto-tool-choice
"""
from patch_utils import ensure_dir, package_root, replace_once
Usage after patching:
--tool-call-parser qwen3_coder --enable-auto-tool-choice
"""
import os
VLLM_ROOT = "/usr/local/corex/lib/python3/dist-packages/vllm"
TOOL_PARSERS_DIR = f"{VLLM_ROOT}/entrypoints/openai/tool_parsers"
INIT_FILE = f"{TOOL_PARSERS_DIR}/__init__.py"
def patch_file(path, replacements):
with open(path, "r") as f:
content = f.read()
patched = False
for old, new in replacements:
if new in content:
print(f" [skip] already patched: {repr(new[:70])}")
continue
if old not in content:
print(f" [warn] anchor not found: {repr(old[:70])}")
continue
content = content.replace(old, new, 1)
patched = True
print(f" [ok] patched: {repr(old[:50])} -> {repr(new[:50])}")
if patched:
with open(path, "w") as f:
f.write(content)
def main():
if not os.path.isdir(TOOL_PARSERS_DIR):
raise FileNotFoundError(
f"Tool parsers directory not found: {TOOL_PARSERS_DIR}\n"
"Verify the vLLM installation path.")
VLLM_ROOT = package_root("vllm")
TOOL_PARSERS_DIR = VLLM_ROOT / "entrypoints" / "openai" / "tool_parsers"
INIT_FILE = TOOL_PARSERS_DIR / "__init__.py"
def main():
ensure_dir(TOOL_PARSERS_DIR)
print(f"=== Patching {INIT_FILE} ===")
patch_file(INIT_FILE, [
(
"from .mistral_tool_parser import MistralToolParser",
"from .mistral_tool_parser import MistralToolParser\n"
"from .qwen3coder_tool_parser import Qwen3CoderToolParser",
),
(
'"MistralToolParser", "Internlm2ToolParser", "Llama3JsonToolParser"\n]',
'"MistralToolParser", "Internlm2ToolParser", "Llama3JsonToolParser",\n'
' "Qwen3CoderToolParser"\n]',
),
])
print("\n=== Verification ===")
try:
import importlib.util
spec = importlib.util.spec_from_file_location(
"qwen3coder_tool_parser",
f"{TOOL_PARSERS_DIR}/qwen3coder_tool_parser.py",
)
mod = importlib.util.module_from_spec(spec)
print(f" Module spec loaded: {spec.name}")
print(" (full import requires torch/vllm runtime — skipping exec)")
except Exception as e:
print(f" [warn] spec check failed: {e}")
print("\nDone. Start vLLM server with:")
print(" --tool-call-parser qwen3_coder --enable-auto-tool-choice")
if __name__ == "__main__":
main()
replace_once(
INIT_FILE,
"from .mistral_tool_parser import MistralToolParser",
"from .mistral_tool_parser import MistralToolParser\n"
"from .qwen3coder_tool_parser import Qwen3CoderToolParser",
required=True,
already_contains="from .qwen3coder_tool_parser import Qwen3CoderToolParser")
replace_once(
INIT_FILE,
'"MistralToolParser", "Internlm2ToolParser", "Llama3JsonToolParser"\n]',
'"MistralToolParser", "Internlm2ToolParser", "Llama3JsonToolParser",\n'
' "Qwen3CoderToolParser"\n]',
required=True,
already_contains='"Qwen3CoderToolParser"')
print("\n=== Verification ===")
try:
import importlib.util
spec = importlib.util.spec_from_file_location(
"qwen3coder_tool_parser",
str(TOOL_PARSERS_DIR / "qwen3coder_tool_parser.py"),
)
mod = importlib.util.module_from_spec(spec)
print(f" Module spec loaded: {spec.name}")
print(" (full import requires torch/vllm runtime — skipping exec)")
except Exception as e:
print(f" [optional] spec check failed: {e}")
print("\nDone. Start vLLM server with:")
print(" --tool-call-parser qwen3_coder --enable-auto-tool-choice")
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,37 @@
from patch_utils import package_root, replace_once
WORKER = package_root("vllm") / "worker" / "worker.py"
CLEAN_BLOCK = """\
if (worker_input.blocks_to_swap_in is not None
and worker_input.blocks_to_swap_in.numel() > 0):
self.cache_engine[virtual_engine].swap_in(
worker_input.blocks_to_swap_in)
if (worker_input.blocks_to_swap_out is not None
and worker_input.blocks_to_swap_out.numel() > 0):
self.cache_engine[virtual_engine].swap_out(
worker_input.blocks_to_swap_out)
"""
ORDERED_BLOCK = """\
# BI100 content-addressed CPU KV tier may preserve a victim and reuse
# that same GPU slot in one step. Complete every D2H before any H2D.
if (worker_input.blocks_to_swap_out is not None
and worker_input.blocks_to_swap_out.numel() > 0):
self.cache_engine[virtual_engine].swap_out(
worker_input.blocks_to_swap_out)
if (worker_input.blocks_to_swap_in is not None
and worker_input.blocks_to_swap_in.numel() > 0):
self.cache_engine[virtual_engine].swap_in(
worker_input.blocks_to_swap_in)
"""
replace_once(
WORKER,
CLEAN_BLOCK,
ORDERED_BLOCK,
required=True,
already_contains="Complete every D2H before any H2D",
)

View File

@@ -0,0 +1,81 @@
from patch_utils import package_root, replace_one_of
WORKER = package_root("vllm") / "worker" / "worker.py"
CLEAN_BLOCK = """\
# Profile the memory usage of the model and get the maximum number of
# cache blocks that can be allocated with the remaining free memory.
torch.cuda.empty_cache()
# Execute a forward pass with dummy inputs to profile the memory usage
# of the model.
self.model_runner.profile_run()
"""
GUARDED_BLOCK = """\
# Profile the memory usage of the model and get the maximum number of
# cache blocks that can be allocated with the remaining free memory.
torch.cuda.empty_cache()
# Execute a forward pass with dummy inputs to profile the memory usage
# of the model. Mark this synthetic pass so BI100_PROFILE can skip
# timing it by default; profiling real requests is the useful signal.
_bi100_prev_startup_profile = os.environ.get("BI100_IN_STARTUP_PROFILE")
os.environ["BI100_IN_STARTUP_PROFILE"] = "1"
try:
self.model_runner.profile_run()
finally:
if _bi100_prev_startup_profile is None:
os.environ.pop("BI100_IN_STARTUP_PROFILE", None)
else:
os.environ["BI100_IN_STARTUP_PROFILE"] = _bi100_prev_startup_profile
"""
NEW_BLOCK = """\
# Profile the memory usage of the model and get the maximum number of
# cache blocks that can be allocated with the remaining free memory.
torch.cuda.empty_cache()
# BI100: Qwen3.6 batched dummy profile_run can trip GDN non-finite
# checks before the server starts. If the operator explicitly provides
# --num-gpu-blocks-override, trust that conservative capacity value and
# skip only the synthetic profile pass. Real inference still uses the
# normal GDN fail-fast path.
if self.cache_config.num_gpu_blocks_override is not None:
cache_block_size = self.get_cache_block_size_bytes()
if cache_block_size == 0:
num_cpu_blocks = 0
else:
num_cpu_blocks = int(self.cache_config.swap_space_bytes //
cache_block_size)
logger.warning(
"[BI100] skipping worker.profile_run because "
"num_gpu_blocks_override=%d was explicitly set",
self.cache_config.num_gpu_blocks_override)
gc.collect()
torch.cuda.empty_cache()
return self.cache_config.num_gpu_blocks_override, max(num_cpu_blocks, 0)
# Execute a forward pass with dummy inputs to profile the memory usage
# of the model. Mark this synthetic pass so BI100_PROFILE can skip
# timing it by default; profiling real requests is the useful signal.
_bi100_prev_startup_profile = os.environ.get("BI100_IN_STARTUP_PROFILE")
os.environ["BI100_IN_STARTUP_PROFILE"] = "1"
try:
self.model_runner.profile_run()
finally:
if _bi100_prev_startup_profile is None:
os.environ.pop("BI100_IN_STARTUP_PROFILE", None)
else:
os.environ["BI100_IN_STARTUP_PROFILE"] = _bi100_prev_startup_profile
"""
replace_one_of(
WORKER,
[
(GUARDED_BLOCK, NEW_BLOCK),
(CLEAN_BLOCK, NEW_BLOCK),
],
required=True,
already_contains="[BI100] skipping worker.profile_run",
)

View File

@@ -0,0 +1,34 @@
from patch_utils import package_root, replace_one_of
WORKER = package_root("vllm") / "worker" / "worker.py"
CLEAN_BLOCK = """\
# Execute a forward pass with dummy inputs to profile the memory usage
# of the model.
self.model_runner.profile_run()
"""
GUARDED_BLOCK = """\
# Execute a forward pass with dummy inputs to profile the memory usage
# of the model. Mark this synthetic pass so BI100_PROFILE can exclude
# it without changing vLLM's normal capacity calculation.
_bi100_prev_startup_profile = os.environ.get("BI100_IN_STARTUP_PROFILE")
os.environ["BI100_IN_STARTUP_PROFILE"] = "1"
try:
self.model_runner.profile_run()
finally:
if _bi100_prev_startup_profile is None:
os.environ.pop("BI100_IN_STARTUP_PROFILE", None)
else:
os.environ["BI100_IN_STARTUP_PROFILE"] = _bi100_prev_startup_profile
"""
replace_one_of(
WORKER,
[(CLEAN_BLOCK, GUARDED_BLOCK)],
required=True,
already_contains=(
"Mark this synthetic pass so BI100_PROFILE can exclude"),
)

View File

@@ -0,0 +1,121 @@
"""Install disabled-by-default M1-48 XFormers timing boundaries."""
from __future__ import annotations
from pathlib import Path
try:
from patch_utils import package_root, replace_once
except ModuleNotFoundError:
from .patch_utils import package_root, replace_once
IMPORT_OLD = "from vllm.logger import init_logger"
IMPORT_NEW = """\
from vllm.bi100_profile import bi100_timer
from vllm.logger import init_logger"""
KV_WRITE_OLD = """\
PagedAttention.write_to_paged_cache(key, value, key_cache,
value_cache,
updated_slot_mapping,
self.kv_cache_dtype,
k_scale, v_scale)"""
KV_WRITE_NEW = """\
with bi100_timer("xformers.kv_write"):
PagedAttention.write_to_paged_cache(
key, value, key_cache, value_cache,
updated_slot_mapping, self.kv_cache_dtype,
k_scale, v_scale)"""
DENSE_OLD = """\
out = self._run_memory_efficient_xformers_forward(
query, key, value, prefill_meta, attn_type=attn_type)"""
DENSE_NEW = """\
with bi100_timer("xformers.dense_prefill"):
out = self._run_memory_efficient_xformers_forward(
query, key, value, prefill_meta, attn_type=attn_type)"""
PAGED_OLD = """\
out = PagedAttention.forward_prefix(
query,
key,
value,
self.kv_cache_dtype,
key_cache,
value_cache,
prefill_meta.block_tables,
prefill_meta.query_start_loc,
prefill_meta.seq_lens_tensor,
prefill_meta.context_lens_tensor,
prefill_meta.max_query_len,
self.alibi_slopes,
self.sliding_window,
k_scale,
v_scale,
is_causal_decoder=(attn_type == AttentionType.DECODER),
)"""
PAGED_NEW = """\
with bi100_timer("xformers.paged_prefill"):
out = PagedAttention.forward_prefix(
query,
key,
value,
self.kv_cache_dtype,
key_cache,
value_cache,
prefill_meta.block_tables,
prefill_meta.query_start_loc,
prefill_meta.seq_lens_tensor,
prefill_meta.context_lens_tensor,
prefill_meta.max_query_len,
self.alibi_slopes,
self.sliding_window,
k_scale,
v_scale,
is_causal_decoder=(attn_type == AttentionType.DECODER),
)"""
def patch_file(path: Path) -> None:
replace_once(
path,
IMPORT_OLD,
IMPORT_NEW,
already_contains="from vllm.bi100_profile import bi100_timer",
)
replace_once(
path,
KV_WRITE_OLD,
KV_WRITE_NEW,
already_contains='bi100_timer("xformers.kv_write")',
)
replace_once(
path,
DENSE_OLD,
DENSE_NEW,
already_contains='bi100_timer("xformers.dense_prefill")',
)
replace_once(
path,
PAGED_OLD,
PAGED_NEW,
already_contains='bi100_timer("xformers.paged_prefill")',
)
text = path.read_text(encoding="utf-8")
canonical = "\n".join(line.rstrip(" \t") for line in text.split("\n"))
if not canonical.endswith("\n"):
canonical += "\n"
if canonical != text:
path.write_text(canonical, encoding="utf-8")
def main() -> None:
path = package_root("vllm") / "attention" / "backends" / "xformers.py"
print("=== patch_xformers_profile (M1-48 diagnostic timers) ===")
print(f"Target: {path}")
patch_file(path)
if __name__ == "__main__":
main()

View File

@@ -1,192 +1,177 @@
"""
策略批量block-diagonalfallback — 纯 PyTorch 数学实现
=============================================================
构建块对角 causal mask对整批序列一次 matmul + softmax
完全绕开所有硬件 flash attention kernel。
背景:
ixformer flshattF: head_dim > 128 报错拒绝
cudnnFlashAttnForward: 接受 head_dim=256但数值结果错误输出全"!"
两者大概率是同一硬件单元ixformer 提前拦截了硬件不支持的配置。
纯 matmul 路径完全绕开硬件 flash attention数值正确。
优点:
数值正确。
并发请求 prefill attention 在 GPU 上真正并行(一次大 matmul
缺点:
峰值显存 = total_tokens² × H × dtype_size
total_tokens 受 --max-num-batched-tokens 控制max-model-len 控制不住。
内存参考fp16H_local=6--max-num-batched-tokens=T
T=2048 → 峰值 ~50 MB
T=4096 → 峰值 ~200 MB
T=8192 → 峰值 ~800 MB
T=16384 → 峰值 ~3.2 GB
策略批量block-diagonalfallback — 纯 PyTorch 数学实现
=============================================================
构建块对角 causal mask对整批序列一次 matmul + softmax
完全绕开所有硬件 flash attention kernel。
背景:
ixformer flshattF: head_dim > 128 报错拒绝
cudnnFlashAttnForward: 接受 head_dim=256但数值结果错误输出全"!"
两者大概率是同一硬件单元ixformer 提前拦截了硬件不支持的配置。
纯 matmul 路径完全绕开硬件 flash attention数值正确。
优点:
数值正确。
并发请求 prefill attention 在 GPU 上真正并行(一次大 matmul
缺点:
峰值显存 = total_tokens² × H × dtype_size
total_tokens 受 --max-num-batched-tokens 控制max-model-len 控制不住。
内存参考fp16H_local=6--max-num-batched-tokens=T
T=2048 → 峰值 ~50 MB
T=4096 → 峰值 ~200 MB
T=8192 → 峰值 ~800 MB
T=16384 → 峰值 ~3.2 GB
Deploy:
python3 modified_scripts/patch_xformers_sdpa_batch.py
"""
XFORMERS_PATH = (
"/usr/local/corex/lib64/python3/dist-packages/"
"vllm/attention/backends/xformers.py"
)
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""批量纯数学 attention fallback。
构建块对角 causal mask等价于 ixformer BlockDiagonalCausalMask
对整批序列一次 matmul + softmaxGPU 并行处理所有序列。
块对角 mask 结构seq1 len=3seq2 len=2
s1,0 s1,1 s1,2 s2,0 s2,1
s1,0 [ 0 -inf -inf -inf -inf ]
s1,1 [ 0 0 -inf -inf -inf ]
s1,2 [ 0 0 0 -inf -inf ]
s2,0 [-inf -inf -inf 0 -inf ]
s2,1 [-inf -inf -inf 0 0 ]
softmax 在 float32 下计算防止 float16 溢出,结果转回原始 dtype。
Args:
query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns:
[1, total_prefill_tokens, num_heads, head_dim]
"""
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
total_tokens = query.shape[1]
# ── 构建块对角 causal mask [T, T] ────────────────────────────────
# 全部初始化为 -inf再对每条序列的对角块填入下三角 0
mask = torch.full(
(total_tokens, total_tokens),
float("-inf"),
dtype=torch.float32,
device=query.device,
)
start = 0
for seq_len in attn_metadata.seq_lens:
end = start + seq_len
mask[start:end, start:end] = torch.tril(
torch.zeros(seq_len, seq_len,
dtype=torch.float32, device=query.device)
)
start = end
# ── [1, H, T, D].contiguous() ──────────────────────────────────
q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
# ── GQA展开 KV heads ────────────────────────────────────────────
if k_all.shape[1] != q_all.shape[1]:
n = q_all.shape[1] // k_all.shape[1]
k_all = k_all.repeat_interleave(n, dim=1).contiguous()
v_all = v_all.repeat_interleave(n, dim=1).contiguous()
# ── 纯数学 attentionfloat32 防溢出)────────────────────────────
# [1, H, T, T]
attn_w = torch.matmul(q_all.float(), k_all.float().transpose(-2, -1))
attn_w = attn_w * self.scale
attn_w = attn_w + mask # 加法广播mask [T,T] → [1, H, T, T]
attn_w = torch.softmax(attn_w, dim=-1)
out = torch.matmul(attn_w, v_all.float()).to(orig_dtype)
# [1, H, T, D] → [1, T, H, D]
return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
from patch_utils import package_root, replace_once
XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py"
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""批量纯数学 attention fallback。
构建块对角 causal mask等价于 ixformer BlockDiagonalCausalMask
对整批序列一次 matmul + softmaxGPU 并行处理所有序列。
块对角 mask 结构seq1 len=3seq2 len=2
s1,0 s1,1 s1,2 s2,0 s2,1
s1,0 [ 0 -inf -inf -inf -inf ]
s1,1 [ 0 0 -inf -inf -inf ]
s1,2 [ 0 0 0 -inf -inf ]
s2,0 [-inf -inf -inf 0 -inf ]
s2,1 [-inf -inf -inf 0 0 ]
softmax 在 float32 下计算防止 float16 溢出,结果转回原始 dtype。
Args:
query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns:
[1, total_prefill_tokens, num_heads, head_dim]
"""
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
total_tokens = query.shape[1]
# ── 构建块对角 causal mask [T, T] ────────────────────────────────
# 全部初始化为 -inf再对每条序列的对角块填入下三角 0
mask = torch.full(
(total_tokens, total_tokens),
float("-inf"),
dtype=torch.float32,
device=query.device,
)
start = 0
for seq_len in attn_metadata.seq_lens:
end = start + seq_len
mask[start:end, start:end] = torch.tril(
torch.zeros(seq_len, seq_len,
dtype=torch.float32, device=query.device)
)
start = end
# ── [1, H, T, D].contiguous() ──────────────────────────────────
q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
# ── GQA展开 KV heads ────────────────────────────────────────────
if k_all.shape[1] != q_all.shape[1]:
n = q_all.shape[1] // k_all.shape[1]
k_all = k_all.repeat_interleave(n, dim=1).contiguous()
v_all = v_all.repeat_interleave(n, dim=1).contiguous()
# ── 纯数学 attentionfloat32 防溢出)────────────────────────────
# [1, H, T, T]
attn_w = torch.matmul(q_all.float(), k_all.float().transpose(-2, -1))
attn_w = attn_w * self.scale
attn_w = attn_w + mask # 加法广播mask [T,T] → [1, H, T, T]
attn_w = torch.softmax(attn_w, dim=-1)
out = torch.matmul(attn_w, v_all.float()).to(orig_dtype)
# [1, H, T, D] → [1, T, H, D]
return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
def patch_file(path):
with open(path, "r") as f:
content = f.read()
changed = False
if "_run_sdpa_fallback" in content:
print(" [skip] _run_sdpa_fallback already present")
elif INJECT_ANCHOR not in content:
print(" [warn] inject anchor not found")
else:
content = content.replace(INJECT_ANCHOR, FALLBACK_METHOD + INJECT_ANCHOR, 1)
print(" [ok] injected _run_sdpa_fallback (batch, pure-math)")
changed = True
if NEW_XFORMER_BLOCK in content:
print(" [skip] dispatch block already patched")
elif OLD_XFORMER_BLOCK in content:
content = content.replace(OLD_XFORMER_BLOCK, NEW_XFORMER_BLOCK, 1)
print(" [ok] patched dispatch block")
changed = True
else:
print(" [warn] dispatch block anchor not found")
if changed:
with open(path, "w") as f:
f.write(content)
print(f" Written: {path}")
def main():
print("=== patch_xformers_sdpa_batch (batch, pure-math) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\nDone.")
if __name__ == "__main__":
main()
replace_once(
path,
INJECT_ANCHOR,
FALLBACK_METHOD + INJECT_ANCHOR,
required=True,
already_contains="def _run_sdpa_fallback(")
replace_once(
path,
OLD_XFORMER_BLOCK,
NEW_XFORMER_BLOCK,
required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
def main():
print("=== patch_xformers_sdpa_batch (batch, pure-math) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\nDone.")
if __name__ == "__main__":
main()

View File

@@ -1,191 +1,176 @@
"""
策略批量block-diagonal— F.scaled_dot_product_attention可走硬件 kernel
=============================================================================
构建块对角 causal mask对整批序列一次 F.scaled_dot_product_attention。
与 patch_xformers_sdpa_batch.py纯 matmul的区别
SDPA 会根据 PyTorch/驱动能力分发到最优 kernelFlash Attention /
mem-efficient attention / math fallback而不是固定走 cublas matmul。
历史说明:
该方案最早因输出全"!"而被弃用,后续排查确认"!"由 mamba_cache.py bug
引起,与 attention 实现无关。当前恢复此方案用于性能对比测试。
已知硬件限制BI-V100
cudnnFlashAttnForward 不支持 is_causal=True报错
本实现使用 is_causal=False + 显式块对角 additive mask 规避此限制。
若 SDPA 仍分发到有问题的 kernel回退到 patch_xformers_sdpa_batch.py。
优点vs 纯 matmul
SDPA 可分发到 Flash Attention kernel → O(L) 显存、更快的 CUDA kernel。
缺点:
依赖硬件 kernel 行为,若 kernel 有 bug 则数值错误(需与 matmul 版对比验证)。
"""
策略批量block-diagonal— F.scaled_dot_product_attention可走硬件 kernel
=============================================================================
构建块对角 causal mask对整批序列一次 F.scaled_dot_product_attention。
与 patch_xformers_sdpa_batch.py纯 matmul的区别
SDPA 会根据 PyTorch/驱动能力分发到最优 kernelFlash Attention /
mem-efficient attention / math fallback而不是固定走 cublas matmul。
历史说明:
该方案最早因输出全"!"而被弃用,后续排查确认"!"由 mamba_cache.py bug
引起,与 attention 实现无关。当前恢复此方案用于性能对比测试。
已知硬件限制BI-V100
cudnnFlashAttnForward 不支持 is_causal=True报错
本实现使用 is_causal=False + 显式块对角 additive mask 规避此限制。
若 SDPA 仍分发到有问题的 kernel回退到 patch_xformers_sdpa_batch.py。
优点vs 纯 matmul
SDPA 可分发到 Flash Attention kernel → O(L) 显存、更快的 CUDA kernel。
缺点:
依赖硬件 kernel 行为,若 kernel 有 bug 则数值错误(需与 matmul 版对比验证)。
Deploy:
python3 modified_scripts/patch_xformers_sdpa_batch_kernel.py
"""
XFORMERS_PATH = (
"/usr/local/corex/lib64/python3/dist-packages/"
"vllm/attention/backends/xformers.py"
)
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""批量 F.scaled_dot_product_attention fallback可走硬件 kernel
构建块对角 causal mask对整批序列一次 SDPA 调用。
SDPA 可分发到 Flash Attention / mem-efficient attention kernel。
is_causal=False + 显式 additive mask规避 cudnnFlashAttnForward
不支持 is_causal=True 的限制。
块对角 maskseq1 len=3seq2 len=2
s1,0 s1,1 s1,2 s2,0 s2,1
s1,0 [ 0 -inf -inf -inf -inf ]
s1,1 [ 0 0 -inf -inf -inf ]
s1,2 [ 0 0 0 -inf -inf ]
s2,0 [-inf -inf -inf 0 -inf ]
s2,1 [-inf -inf -inf 0 0 ]
Args:
query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns:
[1, total_prefill_tokens, num_heads, head_dim]
"""
import torch.nn.functional as F
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
total_tokens = query.shape[1]
# ── 块对角 causal mask [T, T] ─────────────────────────────────────
mask = torch.full(
(total_tokens, total_tokens),
float("-inf"),
dtype=orig_dtype,
device=query.device,
)
start = 0
for seq_len in attn_metadata.seq_lens:
end = start + seq_len
mask[start:end, start:end] = torch.tril(
torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=query.device)
)
start = end
# ── [1, H, T, D] ──────────────────────────────────────────────────
q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
# ── GQA展开 KV heads ────────────────────────────────────────────
if k_all.shape[1] != q_all.shape[1]:
n = q_all.shape[1] // k_all.shape[1]
k_all = k_all.repeat_interleave(n, dim=1).contiguous()
v_all = v_all.repeat_interleave(n, dim=1).contiguous()
# ── F.scaled_dot_product_attention可走硬件 kernel─────────────
# is_causal=False避免 cudnnFlashAttnForward "not support causal mode"
# attn_mask 传 additive float mask非 boolSDPA 选择 math/kernel 路径
out = F.scaled_dot_product_attention(
q_all, k_all, v_all,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
scale=self.scale,
)
# [1, H, T, D] → [1, T, H, D]
return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
from patch_utils import package_root, replace_once
XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py"
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""批量 F.scaled_dot_product_attention fallback可走硬件 kernel
构建块对角 causal mask对整批序列一次 SDPA 调用。
SDPA 可分发到 Flash Attention / mem-efficient attention kernel。
is_causal=False + 显式 additive mask规避 cudnnFlashAttnForward
不支持 is_causal=True 的限制。
块对角 maskseq1 len=3seq2 len=2
s1,0 s1,1 s1,2 s2,0 s2,1
s1,0 [ 0 -inf -inf -inf -inf ]
s1,1 [ 0 0 -inf -inf -inf ]
s1,2 [ 0 0 0 -inf -inf ]
s2,0 [-inf -inf -inf 0 -inf ]
s2,1 [-inf -inf -inf 0 0 ]
Args:
query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns:
[1, total_prefill_tokens, num_heads, head_dim]
"""
import torch.nn.functional as F
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
total_tokens = query.shape[1]
# ── 块对角 causal mask [T, T] ─────────────────────────────────────
mask = torch.full(
(total_tokens, total_tokens),
float("-inf"),
dtype=orig_dtype,
device=query.device,
)
start = 0
for seq_len in attn_metadata.seq_lens:
end = start + seq_len
mask[start:end, start:end] = torch.tril(
torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=query.device)
)
start = end
# ── [1, H, T, D] ──────────────────────────────────────────────────
q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
# ── GQA展开 KV heads ────────────────────────────────────────────
if k_all.shape[1] != q_all.shape[1]:
n = q_all.shape[1] // k_all.shape[1]
k_all = k_all.repeat_interleave(n, dim=1).contiguous()
v_all = v_all.repeat_interleave(n, dim=1).contiguous()
# ── F.scaled_dot_product_attention可走硬件 kernel─────────────
# is_causal=False避免 cudnnFlashAttnForward "not support causal mode"
# attn_mask 传 additive float mask非 boolSDPA 选择 math/kernel 路径
out = F.scaled_dot_product_attention(
q_all, k_all, v_all,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
scale=self.scale,
)
# [1, H, T, D] → [1, T, H, D]
return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
def patch_file(path):
with open(path, "r") as f:
content = f.read()
changed = False
if "_run_sdpa_fallback" in content:
print(" [skip] _run_sdpa_fallback already present")
elif INJECT_ANCHOR not in content:
print(" [warn] inject anchor not found")
else:
content = content.replace(INJECT_ANCHOR, FALLBACK_METHOD + INJECT_ANCHOR, 1)
print(" [ok] injected _run_sdpa_fallback (batch, F.sdpa kernel)")
changed = True
if NEW_XFORMER_BLOCK in content:
print(" [skip] dispatch block already patched")
elif OLD_XFORMER_BLOCK in content:
content = content.replace(OLD_XFORMER_BLOCK, NEW_XFORMER_BLOCK, 1)
print(" [ok] patched dispatch block")
changed = True
else:
print(" [warn] dispatch block anchor not found")
if changed:
with open(path, "w") as f:
f.write(content)
print(f" Written: {path}")
def main():
print("=== patch_xformers_sdpa_batch_kernel (batch, F.sdpa + kernel dispatch) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\nDone.")
if __name__ == "__main__":
main()
replace_once(
path,
INJECT_ANCHOR,
FALLBACK_METHOD + INJECT_ANCHOR,
required=True,
already_contains="def _run_sdpa_fallback(")
replace_once(
path,
OLD_XFORMER_BLOCK,
NEW_XFORMER_BLOCK,
required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
def main():
print("=== patch_xformers_sdpa_batch_kernel (batch, F.sdpa + kernel dispatch) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\nDone.")
if __name__ == "__main__":
main()

View File

@@ -1,321 +1,427 @@
"""
策略顺序per-sequencefallback — 纯 PyTorch 数学实现
==========================================================
逐条序列用 matmul + softmax 手写 attention完全绕开所有硬件
flash attention kernelixformer / cudnnFlashAttnForward
背景:
Iluvatar cudnnFlashAttnForward 存在两个已知问题:
1. 不支持 is_causal=True报错
2. 使用 attn_mask 路径时数值结果不正确(静默错误,输出全为"!"
与华为昇腾 910B4 上 llama.cpp --flash-attn off 修复同类问题的原理相同。
纯数学路径matmul + softmax在任何 PyTorch 后端上结果都正确。
优点:
数值正确,不依赖任何硬件特定 attention kernel。
峰值显存 = max(seq_len)² × H × dtype_size由 --max-model-len 控制。
缺点:
并发请求的 prefill attention 串行执行。
O(L²) 显存(无 flash attention 的 O(L) 优化)。
内存参考fp16H_local=6
max-model-len=4096 → 峰值 ~200 MB
max-model-len=8192 → 峰值 ~800 MB
max-model-len=16384 → 峰值 ~3.2 GB
额外 patcharg_utils.py
vllm 0.6.3 在 max_model_len > 32K 时会自动开启 chunked prefill无命令行
关闭选项),原意是防止 profiling OOM。但 _run_sdpa_fallback 已通过 Q-tiling
解决了该问题chunked prefill 反而会把推理路径从 _run_sdpa_fallback 切换到
_forward_prefix_pytorch属于不必要的行为变更因此一并禁用该自动逻辑。
Deploy:
python3 modified_scripts/patch_xformers_sdpa_seq.py
"""
XFORMERS_PATH = (
"/usr/local/corex/lib64/python3/dist-packages/"
"vllm/attention/backends/xformers.py"
)
ARG_UTILS_PATH = (
"/usr/local/corex/lib64/python3/dist-packages/"
"vllm/engine/arg_utils.py"
)
"""
策略顺序per-sequencefallback — 纯 PyTorch 数学实现
==========================================================
逐条序列用 matmul + softmax 手写 attention完全绕开所有硬件
flash attention kernelixformer / cudnnFlashAttnForward
背景:
Iluvatar cudnnFlashAttnForward 存在两个已知问题:
1. 不支持 is_causal=True报错
2. 使用 attn_mask 路径时数值结果不正确(静默错误,输出全为"!"
与华为昇腾 910B4 上 llama.cpp --flash-attn off 修复同类问题的原理相同。
纯数学路径matmul + softmax在任何 PyTorch 后端上结果都正确。
优点:
数值正确,不依赖任何硬件特定 attention kernel。
峰值显存 = max(seq_len)² × H × dtype_size由 --max-model-len 控制。
缺点:
并发请求的 prefill attention 串行执行。
O(L²) 显存(无 flash attention 的 O(L) 优化)。
内存参考fp16H_local=6
max-model-len=4096 → 峰值 ~200 MB
max-model-len=8192 → 峰值 ~800 MB
max-model-len=16384 → 峰值 ~3.2 GB
额外 patcharg_utils.py
vllm 0.6.3 在 max_model_len > 32K 时会自动开启 chunked prefill无命令行
关闭选项),原意是防止 profiling OOM。但 _run_sdpa_fallback 已通过 Q-tiling
解决了该问题chunked prefill 反而会把推理路径从 _run_sdpa_fallback 切换到
_forward_prefix_pytorch属于不必要的行为变更因此一并禁用该自动逻辑。
Deploy:
python3 modified_scripts/patch_xformers_sdpa_seq.py
"""
from patch_utils import package_root, replace_one_of, replace_once
VLLM_ROOT = package_root("vllm")
XFORMERS_PATH = VLLM_ROOT / "attention" / "backends" / "xformers.py"
ARG_UTILS_PATH = VLLM_ROOT / "engine" / "arg_utils.py"
LOGITS_PROC_PATH = (
"/usr/local/corex/lib64/python3/dist-packages/"
"vllm/model_executor/layers/logits_processor.py"
)
# _apply_logits_processors crashes when seq_groups is None (intermediate
# chunked-prefill chunks on the driver rank). Add an early-return guard.
_LP_OLD_BLOCK = """\
def _apply_logits_processors(
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
found_logits_processors = False\
"""
VLLM_ROOT / "model_executor" / "layers" / "logits_processor.py")
OUTLINES_DECODING_PATH = (
VLLM_ROOT / "model_executor" / "guided_decoding" /
"outlines_decoding.py")
# _apply_logits_processors crashes when seq_groups is None (intermediate
# chunked-prefill chunks on the driver rank). Add an early-return guard.
_LP_OLD_BLOCK = """\
def _apply_logits_processors(
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
found_logits_processors = False\
"""
_LP_NEW_BLOCK = """\
def _apply_logits_processors(
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
if sampling_metadata.seq_groups is None: # intermediate chunked-prefill chunk
return logits
def _apply_logits_processors(
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
if sampling_metadata.seq_groups is None: # intermediate chunked-prefill chunk
return logits
found_logits_processors = False\
"""
# vllm 0.6.3 自动开启 chunked prefill 的原始块
_ARG_OLD_BLOCK = """\
if (is_gpu and not use_sliding_window and not use_spec_decode
and not self.enable_lora
and not self.enable_prompt_adapter):
self.enable_chunked_prefill = True
logger.warning(
"Chunked prefill is enabled by default for models with "
"max_model_len > 32K. Currently, chunked prefill might "
"not work with some features or models. If you "
"encounter any issues, please disable chunked prefill "
"by setting --enable-chunked-prefill=False.")\
# Outlines' UNESCAPED_STRING accepts raw JSON control characters, including
# newlines and tabs. The generated text can therefore satisfy the CFG while
# still failing json.loads(). Use the RFC 8259 string character constraints.
_JSON_STRING_OLD_BLOCK = """\
| UNESCAPED_STRING
| SIGNED_NUMBER -> number
| "true" -> true
| "false" -> false
| "null" -> null
array : "[" [value ("," value)*] "]"
object : "{" [pair ("," pair)*] "}"
pair : UNESCAPED_STRING ":" value
%import common.UNESCAPED_STRING
%import common.SIGNED_NUMBER
%import common.WS
%ignore WS\
"""
_JSON_STRING_V1_BLOCK = r'''| JSON_STRING
| SIGNED_NUMBER -> number
| "true" -> true
| "false" -> false
| "null" -> null
array : "[" [value ("," value)*] "]"
object : "{" [pair ("," pair)*] "}"
pair : JSON_STRING ":" value
JSON_STRING: /"(\\["\\\/bfnrt]|\\u[0-9a-fA-F]{4}|[^"\\\x00-\x1f])*"/
%import common.SIGNED_NUMBER
%import common.WS
%ignore WS'''
_JSON_STRING_NEW_BLOCK = r'''| JSON_STRING
| SIGNED_NUMBER -> number
| "true" -> true
| "false" -> false
| "null" -> null
array : "[" _ws [value (_ws "," _ws value)*] _ws "]"
object : "{" _ws [pair (_ws "," _ws pair)*] _ws "}"
pair : JSON_STRING _ws ":" _ws value
_ws : JSON_WS?
JSON_STRING: /"(\\["\\\/bfnrt]|\\u[0-9a-fA-F]{4}|[^"\\\x00-\x1f])*"/
JSON_WS: /[ \t\r\n]{1,4}/
%import common.SIGNED_NUMBER'''
# vllm 0.6.3 自动开启 chunked prefill 的原始块
_ARG_OLD_BLOCK = """\
if (is_gpu and not use_sliding_window and not use_spec_decode
and not self.enable_lora
and not self.enable_prompt_adapter):
self.enable_chunked_prefill = True
logger.warning(
"Chunked prefill is enabled by default for models with "
"max_model_len > 32K. Currently, chunked prefill might "
"not work with some features or models. If you "
"encounter any issues, please disable chunked prefill "
"by setting --enable-chunked-prefill=False.")\
"""
_ARG_NEW_BLOCK = """\
if (is_gpu and not use_sliding_window and not use_spec_decode
and not self.enable_lora
and not self.enable_prompt_adapter):
pass # skip auto-enable: Q-tiling in _run_sdpa_fallback
if (is_gpu and not use_sliding_window and not use_spec_decode
and not self.enable_lora
and not self.enable_prompt_adapter):
pass # skip auto-enable: Q-tiling in _run_sdpa_fallback
# handles long-context memory without chunked prefill\
"""
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""纯数学 causal attention fallback带 Q-tiling 内存优化。
调用时机kv_cache.numel()==0profiling 阶段)。
此路径无 KV 缓存前缀KV 长度 == query 长度。
内存优化Q-tiling与 Flash Attention 同思路):
将 Q 分成 _Q_CHUNK 大小的子块逐块计算,每块峰值内存
O(_Q_CHUNK × q_len) 而非 O(q_len²)。
profiling 阶段序列可能达到 max_model_len如 20K tokens
不加 Q-tiling 会产生 9.6 GB 矩阵直接 OOM。
softmax 在 float32 下计算以防止 float16 溢出,结果转回原始 dtype。
Args:
query : [1, total_query_tokens, num_heads, head_dim]
key : [1, total_query_tokens, num_kv_heads, head_dim]
value : [1, total_query_tokens, num_kv_heads, head_dim]
Returns:
[1, total_query_tokens, num_heads, head_dim]
"""
_Q_CHUNK = 256 # 与 _forward_prefix_pytorch 的 _ATTN_Q_CHUNK 保持一致
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
num_seqs = len(attn_metadata.seq_lens)
# 推导每条序列的实际 query 长度。
# 正常 prefill 时 q_len == seq_len如果将来遇到 chunked 场景,
# query_start_loc 记录的是真实 query token 数(非全序列长度)。
if (attn_metadata.query_start_loc is not None
and len(attn_metadata.query_start_loc) == num_seqs + 1):
q_lens = [
int(attn_metadata.query_start_loc[i + 1].item()) -
int(attn_metadata.query_start_loc[i].item())
for i in range(num_seqs)
]
else:
q_lens = list(attn_metadata.seq_lens)
q_flat = query.squeeze(0) # [T, H, D]
k_flat = key.squeeze(0) # [T, Hkv, D]
v_flat = value.squeeze(0)
output = torch.empty_like(q_flat)
seq_start = 0
for q_len in q_lens:
seq_end = seq_start + q_len
# 当前序列的完整 K/V此路径无前缀KV == Q
k_s = k_flat[seq_start:seq_end].permute(1, 0, 2).float() # [Hkv, q_len, D]
v_s = v_flat[seq_start:seq_end].permute(1, 0, 2).float() # [Hkv, q_len, D]
# GQA展开 KV heads 至与 query heads 一致
if k_s.shape[0] != self.num_heads:
n = self.num_heads // k_s.shape[0]
k_s = k_s.repeat_interleave(n, dim=0).contiguous()
v_s = v_s.repeat_interleave(n, dim=0).contiguous()
# k_pos 用于因果掩码
k_pos = torch.arange(q_len, device=query.device)
# Q-tiling分块处理 query峰值内存 O(_Q_CHUNK × q_len)
for qc_start in range(0, q_len, _Q_CHUNK):
qc_end = min(qc_start + _Q_CHUNK, q_len)
# [H, qc, D]
q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \
.permute(1, 0, 2).float()
# [H, qc, q_len]
attn_w = torch.matmul(q_c, k_s.transpose(-2, -1)) * self.scale
# 因果掩码q_c 里位置 j 只能看 k_pos <= j相对位置
qc_q_pos = torch.arange(qc_start, qc_end, device=query.device)
mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1)
attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf"))
attn_w = torch.softmax(attn_w, dim=-1)
out_c = torch.matmul(attn_w, v_s).to(orig_dtype) # [H, qc, D]
output[seq_start + qc_start:seq_start + qc_end] = (
out_c.permute(1, 0, 2))
seq_start = seq_end
return output.unsqueeze(0) # [1, T, H, D]
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
_MM_PREFIX_OLD_BLOCK = """\
if model_config.is_multimodal_model:
if self.enable_prefix_caching:
logger.warning(
"--enable-prefix-caching is currently not "
"supported for multimodal models and has been disabled.")
self.enable_prefix_caching = False\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
_MM_PREFIX_NEW_BLOCK = """\
if model_config.is_multimodal_model:
architectures = getattr(model_config.hf_config,
"architectures", []) or []
qwen36_native_vision = "Qwen3_5MoeForCausalLM" in architectures
if self.enable_prefix_caching and qwen36_native_vision:
logger.info(
"Keeping prefix caching enabled for the Qwen3.6 native "
"vision path.")
elif self.enable_prefix_caching:
logger.warning(
"--enable-prefix-caching is currently not "
"supported for multimodal models and has been disabled.")
self.enable_prefix_caching = False\
"""
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""纯数学 causal attention fallback带 Q-tiling 内存优化。
调用时机kv_cache.numel()==0profiling 阶段)。
此路径无 KV 缓存前缀KV 长度 == query 长度。
内存优化Q-tiling与 Flash Attention 同思路):
将 Q 分成 _Q_CHUNK 大小的子块逐块计算,每块峰值内存
O(_Q_CHUNK × q_len) 而非 O(q_len²)。
profiling 阶段序列可能达到 max_model_len如 20K tokens
不加 Q-tiling 会产生 9.6 GB 矩阵直接 OOM。
softmax 在 float32 下计算以防止 float16 溢出,结果转回原始 dtype。
Args:
query : [1, total_query_tokens, num_heads, head_dim]
key : [1, total_query_tokens, num_kv_heads, head_dim]
value : [1, total_query_tokens, num_kv_heads, head_dim]
Returns:
[1, total_query_tokens, num_heads, head_dim]
"""
_Q_CHUNK = 256 # 与 _forward_prefix_pytorch 的 _ATTN_Q_CHUNK 保持一致
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
num_seqs = len(attn_metadata.seq_lens)
# 推导每条序列的实际 query 长度。
# 正常 prefill 时 q_len == seq_len如果将来遇到 chunked 场景,
# query_start_loc 记录的是真实 query token 数(非全序列长度)。
if (attn_metadata.query_start_loc is not None
and len(attn_metadata.query_start_loc) == num_seqs + 1):
q_lens = [
int(attn_metadata.query_start_loc[i + 1].item()) -
int(attn_metadata.query_start_loc[i].item())
for i in range(num_seqs)
]
else:
q_lens = list(attn_metadata.seq_lens)
q_flat = query.squeeze(0) # [T, H, D]
k_flat = key.squeeze(0) # [T, Hkv, D]
v_flat = value.squeeze(0)
output = torch.empty_like(q_flat)
seq_start = 0
for q_len in q_lens:
seq_end = seq_start + q_len
# 当前序列的完整 K/V此路径无前缀KV == Q
k_s = k_flat[seq_start:seq_end].permute(1, 0, 2).float() # [Hkv, q_len, D]
v_s = v_flat[seq_start:seq_end].permute(1, 0, 2).float() # [Hkv, q_len, D]
# GQA展开 KV heads 至与 query heads 一致
if k_s.shape[0] != self.num_heads:
n = self.num_heads // k_s.shape[0]
k_s = k_s.repeat_interleave(n, dim=0).contiguous()
v_s = v_s.repeat_interleave(n, dim=0).contiguous()
# k_pos 用于因果掩码
k_pos = torch.arange(q_len, device=query.device)
# Q-tiling分块处理 query峰值内存 O(_Q_CHUNK × q_len)
for qc_start in range(0, q_len, _Q_CHUNK):
qc_end = min(qc_start + _Q_CHUNK, q_len)
# [H, qc, D]
q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \
.permute(1, 0, 2).float()
# [H, qc, q_len]
attn_w = torch.matmul(q_c, k_s.transpose(-2, -1)) * self.scale
# 因果掩码q_c 里位置 j 只能看 k_pos <= j相对位置
qc_q_pos = torch.arange(qc_start, qc_end, device=query.device)
mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1)
attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf"))
attn_w = torch.softmax(attn_w, dim=-1)
out_c = torch.matmul(attn_w, v_s).to(orig_dtype) # [H, qc, D]
output[seq_start + qc_start:seq_start + qc_end] = (
out_c.permute(1, 0, 2))
seq_start = seq_end
return output.unsqueeze(0) # [1, T, H, D]
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
_PREFIX_CALL_OLD_BLOCK = """\
out = PagedAttention.forward_prefix(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
self.kv_cache_dtype,
key_cache,
value_cache,
prefill_meta.block_tables,
prefill_meta.query_start_loc,
prefill_meta.seq_lens_tensor,
prefill_meta.context_lens_tensor,
prefill_meta.max_query_len,
self.alibi_slopes,
self.sliding_window,
k_scale,
v_scale,
)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
_PREFIX_CALL_NEW_BLOCK = """\
out = PagedAttention.forward_prefix(
query,
key,
value,
self.kv_cache_dtype,
key_cache,
value_cache,
prefill_meta.block_tables,
prefill_meta.query_start_loc,
prefill_meta.seq_lens_tensor,
prefill_meta.context_lens_tensor,
prefill_meta.max_query_len,
self.alibi_slopes,
self.sliding_window,
k_scale,
v_scale,
is_causal_decoder=(attn_type == AttentionType.DECODER),
)\
"""
def patch_file(path):
with open(path, "r") as f:
content = f.read()
changed = False
if "_run_sdpa_fallback" in content:
print(" [skip] _run_sdpa_fallback already present")
elif INJECT_ANCHOR not in content:
print(" [warn] inject anchor not found")
else:
content = content.replace(INJECT_ANCHOR, FALLBACK_METHOD + INJECT_ANCHOR, 1)
print(" [ok] injected _run_sdpa_fallback (sequential, pure-math)")
changed = True
if NEW_XFORMER_BLOCK in content:
print(" [skip] dispatch block already patched")
elif OLD_XFORMER_BLOCK in content:
content = content.replace(OLD_XFORMER_BLOCK, NEW_XFORMER_BLOCK, 1)
print(" [ok] patched dispatch block")
changed = True
else:
print(" [warn] dispatch block anchor not found")
if changed:
with open(path, "w") as f:
f.write(content)
print(f" Written: {path}")
replace_once(
path,
INJECT_ANCHOR,
FALLBACK_METHOD + INJECT_ANCHOR,
required=True,
already_contains="def _run_sdpa_fallback(")
replace_once(
path,
OLD_XFORMER_BLOCK,
NEW_XFORMER_BLOCK,
required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
replace_once(
path,
_PREFIX_CALL_OLD_BLOCK,
_PREFIX_CALL_NEW_BLOCK,
required=True,
already_contains=(
"is_causal_decoder=(attn_type == AttentionType.DECODER)"))
def patch_arg_utils(path):
with open(path, "r") as f:
content = f.read()
changed = False
if "skip auto-enable: Q-tiling" in content:
print(" [skip] chunked-prefill auto-enable already disabled")
elif _ARG_OLD_BLOCK in content:
content = content.replace(_ARG_OLD_BLOCK, _ARG_NEW_BLOCK, 1)
print(" [ok] disabled chunked-prefill auto-enable for 32K+")
changed = True
else:
print(" [warn] target block not found — check arg_utils.py version")
if changed:
with open(path, "w") as f:
f.write(content)
print(f" Written: {path}")
replace_once(
path,
_ARG_OLD_BLOCK,
_ARG_NEW_BLOCK,
required=True,
already_contains="skip auto-enable: Q-tiling")
replace_once(
path,
_MM_PREFIX_OLD_BLOCK,
_MM_PREFIX_NEW_BLOCK,
required=True,
already_contains="Keeping prefix caching enabled for the Qwen3.6")
def patch_logits_processor(path):
with open(path, "r") as f:
content = f.read()
changed = False
if "intermediate chunked-prefill chunk" in content:
print(" [skip] seq_groups=None guard already present")
elif _LP_OLD_BLOCK in content:
content = content.replace(_LP_OLD_BLOCK, _LP_NEW_BLOCK, 1)
print(" [ok] added seq_groups=None guard in _apply_logits_processors")
changed = True
else:
print(" [warn] target block not found — check logits_processor.py version")
if changed:
with open(path, "w") as f:
f.write(content)
print(f" Written: {path}")
replace_once(
path,
_LP_OLD_BLOCK,
_LP_NEW_BLOCK,
required=True,
already_contains="intermediate chunked-prefill chunk")
def main():
print("=== patch_xformers_sdpa_seq (sequential, pure-math) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\n=== patch_arg_utils (disable chunked-prefill auto-enable) ===")
print(f"Target: {ARG_UTILS_PATH}")
patch_arg_utils(ARG_UTILS_PATH)
print("\n=== patch_logits_processor (seq_groups=None guard for chunked prefill) ===")
def patch_outlines_json_grammar(path):
replace_one_of(
path,
[
(_JSON_STRING_V1_BLOCK, _JSON_STRING_NEW_BLOCK),
(_JSON_STRING_OLD_BLOCK, _JSON_STRING_NEW_BLOCK),
],
required=True,
already_contains="JSON_WS:")
def main():
print("=== patch_xformers_sdpa_seq (sequential, pure-math) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\n=== patch_arg_utils (disable chunked-prefill auto-enable) ===")
print(f"Target: {ARG_UTILS_PATH}")
patch_arg_utils(ARG_UTILS_PATH)
print("\n=== patch_logits_processor (seq_groups=None guard for chunked prefill) ===")
print(f"Target: {LOGITS_PROC_PATH}")
patch_logits_processor(LOGITS_PROC_PATH)
print("\nDone.")
if __name__ == "__main__":
main()
print("\n=== patch_outlines_json_grammar (reject raw control chars) ===")
print(f"Target: {OUTLINES_DECODING_PATH}")
patch_outlines_json_grammar(OUTLINES_DECODING_PATH)
print("\nDone.")
if __name__ == "__main__":
main()

View File

@@ -1,181 +1,166 @@
"""
策略顺序per-sequence— F.scaled_dot_product_attention可走硬件 kernel
=============================================================================
逐条序列调用 F.scaled_dot_product_attentionis_causal=False + 显式因果 mask。
与 patch_xformers_sdpa_seq.py纯 matmul的区别
SDPA 可分发到 Flash Attention / mem-efficient attention kernel
而纯 matmul 固定走 cublas。
硬件限制BI-V100
cudnnFlashAttnForward 不支持 is_causal=True直接报错
必须使用 is_causal=False + 显式 additive causal mask。
每条序列单独构造上三角 -inf maskpeak 显存 = max(seq_len)² × dtype
比 batch 版的 total_tokens² 小得多。
与 batch_kernel 的对比:
seq_kernel: 显存小peak = max_single_seq²并发 prefill 串行排队
batch_kernel: 显存大peak = total_tokens²并发 prefill 一次并行处理,
通过 --max-num-batched-tokens 控制 total_tokens 上限
"""
策略顺序per-sequence— F.scaled_dot_product_attention可走硬件 kernel
=============================================================================
逐条序列调用 F.scaled_dot_product_attentionis_causal=False + 显式因果 mask。
与 patch_xformers_sdpa_seq.py纯 matmul的区别
SDPA 可分发到 Flash Attention / mem-efficient attention kernel
而纯 matmul 固定走 cublas。
硬件限制BI-V100
cudnnFlashAttnForward 不支持 is_causal=True直接报错
必须使用 is_causal=False + 显式 additive causal mask。
每条序列单独构造上三角 -inf maskpeak 显存 = max(seq_len)² × dtype
比 batch 版的 total_tokens² 小得多。
与 batch_kernel 的对比:
seq_kernel: 显存小peak = max_single_seq²并发 prefill 串行排队
batch_kernel: 显存大peak = total_tokens²并发 prefill 一次并行处理,
通过 --max-num-batched-tokens 控制 total_tokens 上限
Deploy:
python3 modified_scripts/patch_xformers_sdpa_seq_kernel.py
"""
XFORMERS_PATH = (
"/usr/local/corex/lib64/python3/dist-packages/"
"vllm/attention/backends/xformers.py"
)
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""顺序 F.scaled_dot_product_attention fallback可走硬件 kernel
逐条序列调用 SDPAis_causal=False + 显式上三角 additive mask。
cudnnFlashAttnForward 不支持 is_causal=True必须用显式 mask。
逐序列构造 maskpeak 显存 = max(seq_len)² × dtype远小于 batch 版)。
Args:
query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns:
[1, total_prefill_tokens, num_heads, head_dim]
"""
import torch.nn.functional as F
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
q_flat = query.squeeze(0) # [T, H, D]
k_flat = key.squeeze(0) # [T, Hkv, D]
v_flat = value.squeeze(0)
output = torch.empty_like(q_flat)
start = 0
for seq_len in attn_metadata.seq_lens:
end = start + seq_len
# [1, H, L, D]
q_s = q_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
k_s = k_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
v_s = v_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
# GQA展开 KV heads
if k_s.shape[1] != q_s.shape[1]:
n = q_s.shape[1] // k_s.shape[1]
k_s = k_s.repeat_interleave(n, dim=1).contiguous()
v_s = v_s.repeat_interleave(n, dim=1).contiguous()
# 逐序列因果 mask [L, L],上三角 -inf
causal_mask = torch.tril(
torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=q_s.device)
)
causal_mask = causal_mask.masked_fill(
torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool,
device=q_s.device), diagonal=1),
float("-inf"),
)
# is_causal=False + 显式 mask规避 cudnnFlashAttnForward 不支持 is_causal=True
out_s = F.scaled_dot_product_attention(
q_s, k_s, v_s,
attn_mask=causal_mask,
dropout_p=0.0,
is_causal=False,
scale=self.scale,
)
# [1, H, L, D] → [L, H, D]
output[start:end] = out_s.squeeze(0).permute(1, 0, 2).to(orig_dtype)
start = end
return output.unsqueeze(0) # [1, T, H, D]
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
from patch_utils import package_root, replace_once
XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py"
FALLBACK_METHOD = '''
def _run_sdpa_fallback(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: "XFormersMetadata",
) -> torch.Tensor:
"""顺序 F.scaled_dot_product_attention fallback可走硬件 kernel
逐条序列调用 SDPAis_causal=False + 显式上三角 additive mask。
cudnnFlashAttnForward 不支持 is_causal=True必须用显式 mask。
逐序列构造 maskpeak 显存 = max(seq_len)² × dtype远小于 batch 版)。
Args:
query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns:
[1, total_prefill_tokens, num_heads, head_dim]
"""
import torch.nn.functional as F
assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype
q_flat = query.squeeze(0) # [T, H, D]
k_flat = key.squeeze(0) # [T, Hkv, D]
v_flat = value.squeeze(0)
output = torch.empty_like(q_flat)
start = 0
for seq_len in attn_metadata.seq_lens:
end = start + seq_len
# [1, H, L, D]
q_s = q_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
k_s = k_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
v_s = v_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
# GQA展开 KV heads
if k_s.shape[1] != q_s.shape[1]:
n = q_s.shape[1] // k_s.shape[1]
k_s = k_s.repeat_interleave(n, dim=1).contiguous()
v_s = v_s.repeat_interleave(n, dim=1).contiguous()
# 逐序列因果 mask [L, L],上三角 -inf
causal_mask = torch.tril(
torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=q_s.device)
)
causal_mask = causal_mask.masked_fill(
torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool,
device=q_s.device), diagonal=1),
float("-inf"),
)
# is_causal=False + 显式 mask规避 cudnnFlashAttnForward 不支持 is_causal=True
out_s = F.scaled_dot_product_attention(
q_s, k_s, v_s,
attn_mask=causal_mask,
dropout_p=0.0,
is_causal=False,
scale=self.scale,
)
# [1, H, L, D] → [L, H, D]
output[start:end] = out_s.squeeze(0).permute(1, 0, 2).to(orig_dtype)
start = end
return output.unsqueeze(0) # [1, T, H, D]
'''
OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op = self.attn_op
)
return out.view_as(original_query)\
"""
NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None:
# Add the batch dimension.
query = query.unsqueeze(0)
key = key.unsqueeze(0)
value = value.unsqueeze(0)
if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else:
out = xops.memory_efficient_attention_forward(
query,
key,
value,
attn_bias=attn_bias[0],
p=0.0,
scale=self.scale,
op=self.attn_op,
)
return out.view_as(original_query)\
"""
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
def patch_file(path):
with open(path, "r") as f:
content = f.read()
changed = False
if "_run_sdpa_fallback" in content:
print(" [skip] _run_sdpa_fallback already present")
elif INJECT_ANCHOR not in content:
print(" [warn] inject anchor not found")
else:
content = content.replace(INJECT_ANCHOR, FALLBACK_METHOD + INJECT_ANCHOR, 1)
print(" [ok] injected _run_sdpa_fallback (seq, F.sdpa kernel)")
changed = True
if NEW_XFORMER_BLOCK in content:
print(" [skip] dispatch block already patched")
elif OLD_XFORMER_BLOCK in content:
content = content.replace(OLD_XFORMER_BLOCK, NEW_XFORMER_BLOCK, 1)
print(" [ok] patched dispatch block")
changed = True
else:
print(" [warn] dispatch block anchor not found")
if changed:
with open(path, "w") as f:
f.write(content)
print(f" Written: {path}")
def main():
print("=== patch_xformers_sdpa_seq_kernel (seq, F.sdpa + kernel dispatch) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\nDone.")
if __name__ == "__main__":
main()
replace_once(
path,
INJECT_ANCHOR,
FALLBACK_METHOD + INJECT_ANCHOR,
required=True,
already_contains="def _run_sdpa_fallback(")
replace_once(
path,
OLD_XFORMER_BLOCK,
NEW_XFORMER_BLOCK,
required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
def main():
print("=== patch_xformers_sdpa_seq_kernel (seq, F.sdpa + kernel dispatch) ===")
print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH)
print("\nDone.")
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,12 @@
534019b3c2ad2d2c65492b01a975874ee440026eda2e8666bc3c1dc8a0a0a6f6 corex_attn_head_rms_norm.so
7e2aafd8dc755b0ee16c3b9bb812b95548fc042bbaa840dd9db7d2c51a10474c corex_block_major_kv_transfer.so
ad4ea7707bb2f2bfe04e07a7ad5fd58a647232be70a3056937a0d738c8254bff corex_fused_paged_prefill.so
1856c86e3100415061aa698a48bdeff3fe785994b45b4e72a42cd9158552a7d8 corex_gdn_beta_decay.so
957c7518f5831299fc73f19a4ca2aa3c8231afe9ea7c979127b4f426cd9d6906 corex_gdn_causal_conv.so
ec2d11fa82d9d0816a6da53e62605e962786fa20ecd5f62e50f9d43087fc4d67 corex_gdn_gated_norm.so
27b7ae2ce4fe173336355d72a2678d043df4bd1ed85e9231a99bfb81885a6ce3 corex_gdn_packed_decode.so
015b61046ad73d8f12d754f7a87d4f6cba33070af1c079879e15b71a94571670 corex_gdn_qk_map.so
0eb120e89608bb5b64ca4356a5d3d362121806d081ccc1ccf346dac472a819ec corex_moe_direct_routed.so
d26f2fa39c3921a95793786601e90cf6ebadd06f1d752af541bf82c21acbc1c9 corex_moe_exact_reduce.so
50b0b44c1da779bb2c03419ed549aee9bb922d1f9bab8b7f11a3d91cca0d21c3 corex_moe_weight_gather.so
e944ec0528ed9b6cb74518de3c57e3730543a7bdebc872f993bfdc8424f13e6b corex_paged_kv_gather.so

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,159 @@
from typing import List, Optional, Union
from vllm.config import ModelConfig
from vllm.engine.protocol import EngineClient
from vllm.entrypoints.chat_utils import (apply_hf_chat_template,
apply_mistral_chat_template,
load_chat_template,
parse_chat_messages_futures)
from vllm.entrypoints.logger import RequestLogger
# yapf conflicts with isort for this block
# yapf: disable
from vllm.entrypoints.openai.protocol import (DetokenizeRequest,
DetokenizeResponse,
ErrorResponse,
TokenizeChatRequest,
TokenizeRequest,
TokenizeResponse)
# yapf: enable
from vllm.entrypoints.openai.serving_engine import (BaseModelPath,
LoRAModulePath,
OpenAIServing)
from vllm.logger import init_logger
from vllm.transformers_utils.tokenizer import MistralTokenizer
from vllm.utils import random_uuid
logger = init_logger(__name__)
class OpenAIServingTokenization(OpenAIServing):
def __init__(
self,
engine_client: EngineClient,
model_config: ModelConfig,
base_model_paths: List[BaseModelPath],
*,
lora_modules: Optional[List[LoRAModulePath]],
request_logger: Optional[RequestLogger],
chat_template: Optional[str],
):
super().__init__(engine_client=engine_client,
model_config=model_config,
base_model_paths=base_model_paths,
lora_modules=lora_modules,
prompt_adapters=None,
request_logger=request_logger)
# If this is None we use the tokenizer's default chat template
# the list of commonly-used chat template names for HF named templates
hf_chat_templates: List[str] = ['default', 'tool_use']
self.chat_template = chat_template \
if chat_template in hf_chat_templates \
else load_chat_template(chat_template)
async def create_tokenize(
self,
request: TokenizeRequest,
) -> Union[TokenizeResponse, ErrorResponse]:
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
return error_check_ret
request_id = f"tokn-{random_uuid()}"
(
lora_request,
prompt_adapter_request,
) = self._maybe_get_adapters(request)
tokenizer = await self.engine_client.get_tokenizer(lora_request)
prompt: Union[str, List[int]]
if isinstance(request, TokenizeChatRequest):
model_config = self.model_config
conversation, mm_data_future = parse_chat_messages_futures(
request.messages, model_config, tokenizer)
mm_data = await mm_data_future
if mm_data:
logger.warning(
"Multi-modal inputs are ignored during tokenization")
if isinstance(tokenizer, MistralTokenizer):
prompt = apply_mistral_chat_template(
tokenizer,
messages=request.messages,
chat_template=self.chat_template,
add_generation_prompt=request.add_generation_prompt,
continue_final_message=request.continue_final_message,
**(request.chat_template_kwargs or {}),
)
else:
prompt = apply_hf_chat_template(
tokenizer,
conversation=conversation,
chat_template=self.chat_template,
add_generation_prompt=request.add_generation_prompt,
continue_final_message=request.continue_final_message,
**(request.chat_template_kwargs or {}),
)
else:
prompt = request.prompt
self._log_inputs(request_id,
prompt,
params=None,
lora_request=lora_request,
prompt_adapter_request=prompt_adapter_request)
# Silently ignore prompt adapter since it does not affect tokenization
prompt_input = self._tokenize_prompt_input(
request,
tokenizer,
prompt,
add_special_tokens=request.add_special_tokens,
)
input_ids = prompt_input["prompt_token_ids"]
return TokenizeResponse(tokens=input_ids,
count=len(input_ids),
max_model_len=self.max_model_len)
async def create_detokenize(
self,
request: DetokenizeRequest,
) -> Union[DetokenizeResponse, ErrorResponse]:
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
return error_check_ret
request_id = f"tokn-{random_uuid()}"
(
lora_request,
prompt_adapter_request,
) = self._maybe_get_adapters(request)
tokenizer = await self.engine_client.get_tokenizer(lora_request)
self._log_inputs(request_id,
request.tokens,
params=None,
lora_request=lora_request,
prompt_adapter_request=prompt_adapter_request)
if prompt_adapter_request is not None:
raise NotImplementedError("Prompt adapter is not supported "
"for tokenization")
prompt_input = self._tokenize_prompt_input(
request,
tokenizer,
request.tokens,
)
input_text = prompt_input["prompt"]
return DetokenizeResponse(prompt=input_text)