Compare commits
3 Commits
eb57eb7d1c
...
bce79e44be
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bce79e44be | ||
|
|
74ce61712b | ||
|
|
38eca5c26a |
@@ -8,18 +8,18 @@ command:
|
|||||||
- --served-model-name
|
- --served-model-name
|
||||||
- llm
|
- llm
|
||||||
- --max-model-len
|
- --max-model-len
|
||||||
- '80000'
|
- '131072'
|
||||||
- --gpu-memory-utilization
|
- --gpu-memory-utilization
|
||||||
- '0.80'
|
- '0.90'
|
||||||
- --trust-remote-code
|
- --trust-remote-code
|
||||||
- -tp
|
- -tp
|
||||||
- '4'
|
- '4'
|
||||||
- --max-num-seqs
|
- --max-num-seqs
|
||||||
- '2'
|
- '1'
|
||||||
- --disable-log-requests
|
- --disable-log-requests
|
||||||
- --disable-frontend-multiprocessing
|
- --disable-frontend-multiprocessing
|
||||||
- --max-num-batched-tokens
|
- --max-num-batched-tokens
|
||||||
- '4096'
|
- '8192'
|
||||||
- --enable-chunked-prefill
|
- --enable-chunked-prefill
|
||||||
- --max-seq-len-to-capture
|
- --max-seq-len-to-capture
|
||||||
- '32768'
|
- '32768'
|
||||||
@@ -47,5 +47,3 @@ env:
|
|||||||
value: hybrid64
|
value: hybrid64
|
||||||
- name: BI100_MOE_COREX_TOPK_SOFTMAX
|
- name: BI100_MOE_COREX_TOPK_SOFTMAX
|
||||||
value: '1'
|
value: '1'
|
||||||
- name: PYTORCH_CUDA_ALLOC_CONF
|
|
||||||
value: max_split_size_mb:512
|
|
||||||
|
|||||||
@@ -16,8 +16,8 @@ MANIFEST=${BUNDLE_DIR}/SHA256SUMS
|
|||||||
}
|
}
|
||||||
|
|
||||||
mapfile -t artifacts < <(awk '{print $2}' "$MANIFEST")
|
mapfile -t artifacts < <(awk '{print $2}' "$MANIFEST")
|
||||||
[[ "${#artifacts[@]}" -eq 13 ]] || {
|
[[ "${#artifacts[@]}" -eq 14 ]] || {
|
||||||
printf 'expected 13 prebuilt CoreX artifacts, found %s\n' \
|
printf 'expected 14 prebuilt CoreX artifacts, found %s\n' \
|
||||||
"${#artifacts[@]}" >&2
|
"${#artifacts[@]}" >&2
|
||||||
exit 2
|
exit 2
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,13 +20,6 @@ CAPACITY_ANCHOR = """\
|
|||||||
CAPACITY_REPLACEMENT = """\
|
CAPACITY_REPLACEMENT = """\
|
||||||
num_gpu_blocks = reserve_block_major_gpu_blocks(
|
num_gpu_blocks = reserve_block_major_gpu_blocks(
|
||||||
num_gpu_blocks, cache_block_size)
|
num_gpu_blocks, cache_block_size)
|
||||||
# BI100: profiling with zero-tensor attention underestimates memory.
|
|
||||||
# Hardcap at 3000 blocks (48K tokens) to prevent runtime OOM.
|
|
||||||
# Must leave ~4GB free for flash_attn_varlen_func temp buffers.
|
|
||||||
if num_gpu_blocks > 3000:
|
|
||||||
logger.warning(
|
|
||||||
"[BI100] capping num_gpu_blocks: %d -> 3000", num_gpu_blocks)
|
|
||||||
num_gpu_blocks = 3000
|
|
||||||
num_gpu_blocks = max(num_gpu_blocks, 0)
|
num_gpu_blocks = max(num_gpu_blocks, 0)
|
||||||
num_cpu_blocks = max(num_cpu_blocks, 0)
|
num_cpu_blocks = max(num_cpu_blocks, 0)
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -246,24 +246,6 @@ if source != installed:
|
|||||||
raise SystemExit("runtime api_server overlay identity mismatch")
|
raise SystemExit("runtime api_server overlay identity mismatch")
|
||||||
PY
|
PY
|
||||||
|
|
||||||
build_stage "compiling CCCL CachingDeviceAllocator LD_PRELOAD module"
|
|
||||||
bash ./cccl_preload/build_cccl_preload.sh /workspace/qwen3_6_scripts/cccl_preload || \
|
|
||||||
echo "[WARN] CCCL preload allocator build failed — will use default allocator"
|
|
||||||
|
|
||||||
build_stage "compiling CoreX CUDA extensions (moe_index_combine + gdn_chunk_recurrent)"
|
|
||||||
if [[ -x /usr/local/corex-3.2.3/bin/clang++ ]]; then
|
|
||||||
bash ./build_corex_moe_index_combine.sh "${VLLM_ROOT}" || \
|
|
||||||
echo "[WARN] moe_index_combine build failed — will use PyTorch fallback"
|
|
||||||
bash ./build_corex_gdn_chunk_recurrent.sh "${VLLM_ROOT}" || \
|
|
||||||
echo "[WARN] gdn_chunk_recurrent build failed — will use Python fallback"
|
|
||||||
else
|
|
||||||
echo "[WARN] corex clang++ not found — skipping extension builds"
|
|
||||||
fi
|
|
||||||
|
|
||||||
build_stage "compiling ixformer bridge .so (MoE + Attention + Norm)"
|
|
||||||
bash ./build_ix_bridge.sh "${VLLM_ROOT}" || \
|
|
||||||
echo "[WARN] ix_full_bridge build failed — MoE will use PyTorch fallback"
|
|
||||||
|
|
||||||
build_stage "compiling submission Python sources"
|
build_stage "compiling submission Python sources"
|
||||||
find . -path './wheels' -prune -o -name '*.py' -print0 | xargs -0 python3 -m py_compile
|
find . -path './wheels' -prune -o -name '*.py' -print0 | xargs -0 python3 -m py_compile
|
||||||
build_stage "patch script completed"
|
build_stage "patch script completed"
|
||||||
|
|||||||
@@ -165,6 +165,26 @@ _MM_PREFIX_NEW_BLOCK = """\
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
FALLBACK_METHOD = '''
|
FALLBACK_METHOD = '''
|
||||||
|
# --- flash_attn_varlen_func backend (loaded once) ---
|
||||||
|
# Import path: ixformer.contrib.vllm_flash_attn (canonical, matches
|
||||||
|
# ex_engine/python/corex_fa2.py Tier 1 and ixformer_sdk).
|
||||||
|
# Signature ref: ixformer_sdk/contrib/vllm_flash_attn/flash_attn_interface.py
|
||||||
|
_flash_varlen_func = None
|
||||||
|
_flash_varlen_checked = False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_flash_varlen(cls):
|
||||||
|
if not cls._flash_varlen_checked:
|
||||||
|
cls._flash_varlen_checked = True
|
||||||
|
try:
|
||||||
|
from ixformer.contrib.vllm_flash_attn import (
|
||||||
|
flash_attn_varlen_func as _fn,
|
||||||
|
)
|
||||||
|
cls._flash_varlen_func = _fn
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
return cls._flash_varlen_func
|
||||||
|
|
||||||
def _run_sdpa_fallback(
|
def _run_sdpa_fallback(
|
||||||
self,
|
self,
|
||||||
query: torch.Tensor,
|
query: torch.Tensor,
|
||||||
@@ -172,58 +192,69 @@ FALLBACK_METHOD = '''
|
|||||||
value: torch.Tensor,
|
value: torch.Tensor,
|
||||||
attn_metadata: "XFormersMetadata",
|
attn_metadata: "XFormersMetadata",
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Use ixformer flash_attn_varlen_func for head_dim > 128.
|
"""Prefill attention fallback for head_dim > 128.
|
||||||
|
|
||||||
Verified on real BI-V100: flash_attn_func handles head_dim=256
|
Dispatch priority (ref: ex_engine/python/corex_fa2.py):
|
||||||
correctly (diff < 0.004, no NaN). For seq >= 1024, faster than
|
1. ixformer flash_attn_varlen_func — fused kernel, O(L) memory
|
||||||
PyTorch matmul. For profiling, sequences can be 20K+ tokens — this
|
2. Pure-math Q-tiling fallback — safe for profiling / any HW
|
||||||
is dramatically faster than the previous Python Q-tiling fallback.
|
|
||||||
|
|
||||||
Falls back to pure-math if flash_attn is unavailable.
|
Profiling guard: when kv_cache is empty (profiling stage), vllm feeds
|
||||||
|
a dummy sequence up to max_model_len (131K). flash_attn temp buffers
|
||||||
|
at that length can exceed GPU memory. We use Q-tiling for profiling
|
||||||
|
(safe, correct, O(chunk × L) memory) and flash_attn for real
|
||||||
|
inference (fast, O(L) memory, verified on BI-V100 head_dim=256).
|
||||||
|
|
||||||
|
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]
|
||||||
"""
|
"""
|
||||||
import ixformer as _ixf
|
|
||||||
|
|
||||||
assert attn_metadata.seq_lens is not None
|
assert attn_metadata.seq_lens is not None
|
||||||
orig_dtype = query.dtype
|
orig_dtype = query.dtype
|
||||||
num_seqs = len(attn_metadata.seq_lens)
|
num_seqs = len(attn_metadata.seq_lens)
|
||||||
|
max_seqlen = max(attn_metadata.seq_lens)
|
||||||
|
|
||||||
q_flat = query.squeeze(0) # [T, H, D]
|
# Detect profiling: attn_metadata.num_prefill_tokens == total tokens
|
||||||
k_flat = key.squeeze(0) # [T, Hkv, D]
|
# AND no actual KV cache allocated yet (first forward pass).
|
||||||
v_flat = value.squeeze(0)
|
# Also guard against very long dummy sequences (profiling uses
|
||||||
|
# max_model_len which can be 131K) where flash_attn would OOM.
|
||||||
|
_FLASH_SAFE_SEQLEN = 32768 # flash_attn temp buffers safe below this
|
||||||
|
is_profiling = (max_seqlen > _FLASH_SAFE_SEQLEN
|
||||||
|
and not hasattr(attn_metadata, '_has_real_kv_cache'))
|
||||||
|
|
||||||
# Build cu_seqlens from seq_lens
|
# --- Path 1: flash_attn_varlen_func (real inference) ---
|
||||||
seq_lens_list = list(attn_metadata.seq_lens)
|
fn = self._get_flash_varlen()
|
||||||
cu_seqlens = torch.zeros(num_seqs + 1, dtype=torch.int32,
|
if fn is not None and not is_profiling:
|
||||||
device=query.device)
|
try:
|
||||||
for i, sl in enumerate(seq_lens_list):
|
q_flat = query.squeeze(0) # [T, H, D]
|
||||||
cu_seqlens[i + 1] = cu_seqlens[i] + sl
|
k_flat = key.squeeze(0) # [T, Hkv, D]
|
||||||
max_seqlen = max(seq_lens_list)
|
v_flat = value.squeeze(0)
|
||||||
|
|
||||||
try:
|
cu_seqlens = torch.zeros(
|
||||||
# Skip flash_attn during profiling — OOMs on large dummy batch
|
num_seqs + 1, dtype=torch.int32, device=query.device)
|
||||||
import os
|
for i, sl in enumerate(attn_metadata.seq_lens):
|
||||||
if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1":
|
cu_seqlens[i + 1] = cu_seqlens[i] + sl
|
||||||
raise RuntimeError("skip flash_attn during profiling")
|
|
||||||
out = _ixf.flash_attn_varlen_func(
|
|
||||||
q_flat.to(torch.float16),
|
|
||||||
k_flat.to(torch.float16),
|
|
||||||
v_flat.to(torch.float16),
|
|
||||||
cu_seqlens, cu_seqlens,
|
|
||||||
max_seqlen, max_seqlen,
|
|
||||||
causal=True,
|
|
||||||
)
|
|
||||||
return out.to(orig_dtype).unsqueeze(0)
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Fallback: pure-math Q-tiling (original implementation)
|
out = fn(
|
||||||
|
q=q_flat.to(torch.float16),
|
||||||
|
k=k_flat.to(torch.float16),
|
||||||
|
v=v_flat.to(torch.float16),
|
||||||
|
cu_seqlens_q=cu_seqlens,
|
||||||
|
cu_seqlens_k=cu_seqlens,
|
||||||
|
max_seqlen_q=max_seqlen,
|
||||||
|
max_seqlen_k=max_seqlen,
|
||||||
|
softmax_scale=self.scale,
|
||||||
|
causal=True,
|
||||||
|
)
|
||||||
|
return out.to(orig_dtype).unsqueeze(0)
|
||||||
|
except Exception:
|
||||||
|
pass # fall through to Q-tiling
|
||||||
|
|
||||||
|
# --- Path 2: Q-tiling (profiling or flash_attn unavailable) ---
|
||||||
_Q_CHUNK = 256
|
_Q_CHUNK = 256
|
||||||
|
|
||||||
# During profiling, skip expensive attention — return zeros.
|
|
||||||
# Profiling only measures memory footprint, not output correctness.
|
|
||||||
if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1":
|
|
||||||
return torch.zeros_like(query)
|
|
||||||
|
|
||||||
if (attn_metadata.query_start_loc is not None
|
if (attn_metadata.query_start_loc is not None
|
||||||
and len(attn_metadata.query_start_loc) == num_seqs + 1):
|
and len(attn_metadata.query_start_loc) == num_seqs + 1):
|
||||||
q_lens = [
|
q_lens = [
|
||||||
@@ -232,32 +263,47 @@ FALLBACK_METHOD = '''
|
|||||||
for i in range(num_seqs)
|
for i in range(num_seqs)
|
||||||
]
|
]
|
||||||
else:
|
else:
|
||||||
q_lens = seq_lens_list
|
q_lens = list(attn_metadata.seq_lens)
|
||||||
|
|
||||||
|
q_flat = query.squeeze(0)
|
||||||
|
k_flat = key.squeeze(0)
|
||||||
|
v_flat = value.squeeze(0)
|
||||||
|
|
||||||
output = torch.empty_like(q_flat)
|
output = torch.empty_like(q_flat)
|
||||||
seq_start = 0
|
seq_start = 0
|
||||||
for q_len in q_lens:
|
for q_len in q_lens:
|
||||||
seq_end = seq_start + q_len
|
seq_end = seq_start + q_len
|
||||||
|
|
||||||
k_s = k_flat[seq_start:seq_end].permute(1, 0, 2).float()
|
k_s = k_flat[seq_start:seq_end].permute(1, 0, 2).float()
|
||||||
v_s = v_flat[seq_start:seq_end].permute(1, 0, 2).float()
|
v_s = v_flat[seq_start:seq_end].permute(1, 0, 2).float()
|
||||||
|
|
||||||
if k_s.shape[0] != self.num_heads:
|
if k_s.shape[0] != self.num_heads:
|
||||||
n = self.num_heads // k_s.shape[0]
|
n = self.num_heads // k_s.shape[0]
|
||||||
k_s = k_s.repeat_interleave(n, dim=0).contiguous()
|
k_s = k_s.repeat_interleave(n, dim=0).contiguous()
|
||||||
v_s = v_s.repeat_interleave(n, dim=0).contiguous()
|
v_s = v_s.repeat_interleave(n, dim=0).contiguous()
|
||||||
|
|
||||||
k_pos = torch.arange(q_len, device=query.device)
|
k_pos = torch.arange(q_len, device=query.device)
|
||||||
|
|
||||||
for qc_start in range(0, q_len, _Q_CHUNK):
|
for qc_start in range(0, q_len, _Q_CHUNK):
|
||||||
qc_end = min(qc_start + _Q_CHUNK, q_len)
|
qc_end = min(qc_start + _Q_CHUNK, q_len)
|
||||||
|
|
||||||
q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \
|
q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \
|
||||||
.permute(1, 0, 2).float()
|
.permute(1, 0, 2).float()
|
||||||
|
|
||||||
attn_w = torch.matmul(q_c, k_s.transpose(-2, -1)) * self.scale
|
attn_w = torch.matmul(q_c, k_s.transpose(-2, -1)) * self.scale
|
||||||
|
|
||||||
qc_q_pos = torch.arange(qc_start, qc_end, device=query.device)
|
qc_q_pos = torch.arange(qc_start, qc_end, device=query.device)
|
||||||
mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1)
|
mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1)
|
||||||
attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf"))
|
attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf"))
|
||||||
|
|
||||||
attn_w = torch.softmax(attn_w, dim=-1)
|
attn_w = torch.softmax(attn_w, dim=-1)
|
||||||
out_c = torch.matmul(attn_w, v_s).to(orig_dtype)
|
out_c = torch.matmul(attn_w, v_s).to(orig_dtype)
|
||||||
|
|
||||||
output[seq_start + qc_start:seq_start + qc_end] = (
|
output[seq_start + qc_start:seq_start + qc_end] = (
|
||||||
out_c.permute(1, 0, 2))
|
out_c.permute(1, 0, 2))
|
||||||
|
|
||||||
seq_start = seq_end
|
seq_start = seq_end
|
||||||
|
|
||||||
return output.unsqueeze(0)
|
return output.unsqueeze(0)
|
||||||
|
|
||||||
'''
|
'''
|
||||||
|
|||||||
Reference in New Issue
Block a user