revert: restore a3c45d3b yaml + Q-tiling + remove all OOM hacks
Root cause of 10 consecutive OOM failures: - 'return zeros during profiling' hack → vllm overestimates free memory → allocates 7942 blocks → first real request OOMs - blocks cap 5000 → band-aid that masks profiling bug - gpu-memory-utilization 0.80 → unnecessary reduction from working 0.90 - max-num-seqs 2 → doubles peak activation memory - PYTORCH_CUDA_ALLOC_CONF max_split_size_mb:512 → causes fragmentation Restoringa3c45d3bparameters that actually work: - yaml: max-model-len=131072, gpu-mem=0.90, max-num-seqs=1, batched-tokens=8192 - patch_xformers_sdpa_seq.py: Q-tiling (real memory optimization, not zeros hack) - patch_block_major_worker_capacity.py: no blocks cap, just reserve_block_major - patch_ops.sh: remove all docker-build-time compilation (all .so are prebuilt) Only change froma3c45d3b: BI100_MOE_COREX_TOPK_SOFTMAX=1 (enable corex topk) Kept fixes: - protocol.py extra=allow (recover 180 rejected replay requests) - corex_gdn_chunk_recurrent.so pybind kwargs (prebuilt with fixed signature)
This commit is contained in:
@@ -8,18 +8,18 @@ command:
|
||||
- --served-model-name
|
||||
- llm
|
||||
- --max-model-len
|
||||
- '80000'
|
||||
- '131072'
|
||||
- --gpu-memory-utilization
|
||||
- '0.80'
|
||||
- '0.90'
|
||||
- --trust-remote-code
|
||||
- -tp
|
||||
- '4'
|
||||
- --max-num-seqs
|
||||
- '2'
|
||||
- '1'
|
||||
- --disable-log-requests
|
||||
- --disable-frontend-multiprocessing
|
||||
- --max-num-batched-tokens
|
||||
- '4096'
|
||||
- '8192'
|
||||
- --enable-chunked-prefill
|
||||
- --max-seq-len-to-capture
|
||||
- '32768'
|
||||
@@ -47,5 +47,3 @@ env:
|
||||
value: hybrid64
|
||||
- name: BI100_MOE_COREX_TOPK_SOFTMAX
|
||||
value: '1'
|
||||
- name: PYTORCH_CUDA_ALLOC_CONF
|
||||
value: max_split_size_mb:512
|
||||
|
||||
@@ -20,13 +20,6 @@ CAPACITY_ANCHOR = """\
|
||||
CAPACITY_REPLACEMENT = """\
|
||||
num_gpu_blocks = reserve_block_major_gpu_blocks(
|
||||
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_cpu_blocks = max(num_cpu_blocks, 0)
|
||||
"""
|
||||
|
||||
@@ -246,24 +246,6 @@ if source != installed:
|
||||
raise SystemExit("runtime api_server overlay identity mismatch")
|
||||
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"
|
||||
find . -path './wheels' -prune -o -name '*.py' -print0 | xargs -0 python3 -m py_compile
|
||||
build_stage "patch script completed"
|
||||
|
||||
@@ -172,58 +172,35 @@ FALLBACK_METHOD = '''
|
||||
value: torch.Tensor,
|
||||
attn_metadata: "XFormersMetadata",
|
||||
) -> torch.Tensor:
|
||||
"""Use ixformer flash_attn_varlen_func for head_dim > 128.
|
||||
"""纯数学 causal attention fallback,带 Q-tiling 内存优化。
|
||||
|
||||
Verified on real BI-V100: flash_attn_func handles head_dim=256
|
||||
correctly (diff < 0.004, no NaN). For seq >= 1024, faster than
|
||||
PyTorch matmul. For profiling, sequences can be 20K+ tokens — this
|
||||
is dramatically faster than the previous Python Q-tiling fallback.
|
||||
调用时机:kv_cache.numel()==0(profiling 阶段)。
|
||||
此路径无 KV 缓存前缀,KV 长度 == query 长度。
|
||||
|
||||
Falls back to pure-math if flash_attn is unavailable.
|
||||
内存优化(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]
|
||||
"""
|
||||
import ixformer as _ixf
|
||||
_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)
|
||||
|
||||
q_flat = query.squeeze(0) # [T, H, D]
|
||||
k_flat = key.squeeze(0) # [T, Hkv, D]
|
||||
v_flat = value.squeeze(0)
|
||||
|
||||
# Build cu_seqlens from seq_lens
|
||||
seq_lens_list = list(attn_metadata.seq_lens)
|
||||
cu_seqlens = torch.zeros(num_seqs + 1, dtype=torch.int32,
|
||||
device=query.device)
|
||||
for i, sl in enumerate(seq_lens_list):
|
||||
cu_seqlens[i + 1] = cu_seqlens[i] + sl
|
||||
max_seqlen = max(seq_lens_list)
|
||||
|
||||
try:
|
||||
# Skip flash_attn during profiling — OOMs on large dummy batch
|
||||
import os
|
||||
if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1":
|
||||
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)
|
||||
_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)
|
||||
|
||||
# 推导每条序列的实际 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 = [
|
||||
@@ -232,33 +209,55 @@ FALLBACK_METHOD = '''
|
||||
for i in range(num_seqs)
|
||||
]
|
||||
else:
|
||||
q_lens = seq_lens_list
|
||||
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_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()
|
||||
|
||||
# 当前序列的完整 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)
|
||||
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)
|
||||
|
||||
return output.unsqueeze(0) # [1, T, H, D]
|
||||
|
||||
'''
|
||||
|
||||
|
||||
Reference in New Issue
Block a user