From 20cd2d89047dc663cbee758467e1e2e20d9147ce Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 10 Aug 2026 01:08:38 +0000 Subject: [PATCH] build(SM70): precompile GDN CUDA kernel to .so during docker build MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit precompile_gdn.py: calls torch.utils.cpp_extension.load with build_directory to produce .so at build time. If build env has no GPU/compiler, fails gracefully — kernel JIT compiles at runtime instead. fused_fwd.py: _load_ext() now checks build/ dir for precompiled .so first, skips 2-minute JIT compilation if found. --- qwen3_6_scripts/flash_qla_sm70/fused_fwd.py | 18 +++++++ qwen3_6_scripts/patch_ops.sh | 5 +- qwen3_6_scripts/precompile_gdn.py | 53 +++++++++++++++++++++ 3 files changed, 75 insertions(+), 1 deletion(-) create mode 100644 qwen3_6_scripts/precompile_gdn.py diff --git a/qwen3_6_scripts/flash_qla_sm70/fused_fwd.py b/qwen3_6_scripts/flash_qla_sm70/fused_fwd.py index 6e71c2a3..8308b12d 100644 --- a/qwen3_6_scripts/flash_qla_sm70/fused_fwd.py +++ b/qwen3_6_scripts/flash_qla_sm70/fused_fwd.py @@ -20,6 +20,24 @@ def _load_ext(): raise RuntimeError("SM70 FlashQLA backend requires CUDA.") os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "7.0;7.5") + + # Try precompiled .so first (built during docker build) + build_dir = Path(__file__).with_name("build") + if build_dir.is_dir(): + so_files = list(build_dir.glob("*.so")) + if so_files: + try: + _EXT = load( + name="flash_qla_sm70_gdn_strided", + sources=[], # empty — just load from build_directory + build_directory=str(build_dir), + verbose=False, + ) + return _EXT + except Exception: + pass # fall through to JIT + + # JIT compile (slow, ~2min first time) src = Path(__file__).with_name("csrc") / "gdn_forward.cu" _EXT = load( name="flash_qla_sm70_gdn_strided", diff --git a/qwen3_6_scripts/patch_ops.sh b/qwen3_6_scripts/patch_ops.sh index 31535aa2..985b2744 100755 --- a/qwen3_6_scripts/patch_ops.sh +++ b/qwen3_6_scripts/patch_ops.sh @@ -188,9 +188,12 @@ if [ -d "./flash_qla_sm70" ]; then rm -rf "$FLASH_QLA_DST" 2>/dev/null cp -r ./flash_qla_sm70 "$FLASH_QLA_DST" 2>/dev/null && \ echo "[patch_ops] flash_qla_sm70 deployed to $FLASH_QLA_DST" || true + # Pre-compile CUDA kernel → .so (skipped if no GPU/compiler at build time) + python3 ./precompile_gdn.py "$FLASH_QLA_DST" 2>&1 || \ + echo "[patch_ops] WARNING: precompile failed — kernel will JIT at runtime" # Also deploy to VLLM2 if present if [ -n "$VLLM2" ]; then rm -rf "$VLLM2/model_executor/models/flash_qla_sm70" 2>/dev/null - cp -r ./flash_qla_sm70 "$VLLM2/model_executor/models/flash_qla_sm70" 2>/dev/null || true + cp -r "$FLASH_QLA_DST" "$VLLM2/model_executor/models/flash_qla_sm70" 2>/dev/null || true fi fi diff --git a/qwen3_6_scripts/precompile_gdn.py b/qwen3_6_scripts/precompile_gdn.py new file mode 100644 index 00000000..ec5333bb --- /dev/null +++ b/qwen3_6_scripts/precompile_gdn.py @@ -0,0 +1,53 @@ +""" +Pre-compile SM70 GDN CUDA kernel → .so at Docker build time. +Avoids 2-minute JIT delay at runtime. + +Usage: python3 precompile_gdn.py /path/to/flash_qla_sm70/ +""" +import os +import sys + +def main(): + if len(sys.argv) < 2: + print("[precompile] Usage: python3 precompile_gdn.py ") + sys.exit(1) + + flash_dir = sys.argv[1] + cu_src = os.path.join(flash_dir, "csrc", "gdn_forward.cu") + if not os.path.exists(cu_src): + print(f"[precompile] ERROR: {cu_src} not found") + sys.exit(1) + + # Set arch for BI-V100 (SM70 compatible) + os.environ["TORCH_CUDA_ARCH_LIST"] = "7.0;7.5" + + build_dir = os.path.join(flash_dir, "build") + os.makedirs(build_dir, exist_ok=True) + + print(f"[precompile] Compiling {cu_src} → .so in {build_dir}") + print(f"[precompile] TORCH_CUDA_ARCH_LIST = {os.environ['TORCH_CUDA_ARCH_LIST']}") + + try: + from torch.utils.cpp_extension import load + ext = load( + name="flash_qla_sm70_gdn_strided", + sources=[cu_src], + extra_cuda_cflags=["-O3"], + extra_cflags=["-O3"], + build_directory=build_dir, + verbose=True, + ) + print(f"[precompile] SUCCESS — compiled .so in {build_dir}") + # List the built files + for f in os.listdir(build_dir): + if f.endswith(".so"): + full = os.path.join(build_dir, f) + print(f"[precompile] {f} ({os.path.getsize(full)} bytes)") + except Exception as e: + print(f"[precompile] FAILED: {e}") + print("[precompile] Kernel will JIT compile at runtime instead (~2min)") + sys.exit(1) + + +if __name__ == "__main__": + main()