diff --git a/.dockerignore b/.dockerignore deleted file mode 100644 index 0f46f5e2..00000000 --- a/.dockerignore +++ /dev/null @@ -1,7 +0,0 @@ -# Exclude everything -* - -# Only include what Dockerfile needs -!Dockerfile -!computility-run.yaml -!qwen3_6_scripts/ diff --git a/Dockerfile b/Dockerfile index 0e66cd53..faa0a98a 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,15 +1,10 @@ FROM git.modelhub.org.cn:9443/enginex-iluvatar/bi100-3.2.3-x86-ubuntu20.04-py3.10-poc-llm-infer:v1.2.3 - -ENV PATH=/usr/local/corex/bin:/usr/local/corex-3.2.3/bin:/usr/local/openmpi/bin:${PATH} -ENV PYTHONPATH=/usr/local/corex/lib64/python3/dist-packages:/usr/local/corex/lib/python3/dist-packages -ENV LD_LIBRARY_PATH=/usr/local/corex/lib:/usr/local/corex/lib64:/usr/local/corex-3.2.3/lib:/usr/local/corex-3.2.3/lib64:/usr/local/openmpi/lib -ENV VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 PYTHONUNBUFFERED=1 PYTHONFAULTHANDLER=1 BI100_EXECUTOR_STARTUP_DEBUG=1 ENABLE_CUSTOM_IPC=1 -ENV BI100_PREFIX_MODEL_FINGERPRINT=Qwen3.6-35B-A3B BI100_PREFIX_DTYPE=float16 BI100_PREFIX_TP_SIZE=4 - RUN mkdir -p /workspace WORKDIR /workspace/ +# Copy all our engine patches COPY ./qwen3_6_scripts /workspace/qwen3_6_scripts COPY ./computility-run.yaml /workspace/computility-run.yaml +# Make patch script executable and run it RUN chmod +x /workspace/qwen3_6_scripts/patch_ops.sh && \ bash /workspace/qwen3_6_scripts/patch_ops.sh 2>&1 | tee /workspace/patch_ops.log ; \ echo "[Dockerfile] patch_ops exit code: $?" diff --git a/ex_engine/csrc/ix_full_bridge_v2.cpp b/ex_engine/csrc/ix_full_bridge_v2.cpp new file mode 100644 index 00000000..f5396150 --- /dev/null +++ b/ex_engine/csrc/ix_full_bridge_v2.cpp @@ -0,0 +1,451 @@ +// ix_full_bridge_v2.cpp — Complete bridge to ALL ixformer::infer C++ functions +// +// Base image has ixformer::infer namespace with 14 functions. +// Previous ix_full_bridge.cpp only bridged 4 (silu_and_mul, rms_norm, +// fused_add_rms_norm, linear). This file bridges ALL 14. +// +// The base image's _ixformer_torch.cpython-310.so and libixformer.so +// export these symbols in the ixformer::infer namespace (confirmed by nm -D). +// +// Compile: +// torch.utils.cpp_extension.load( +// name="ix_full_bridge_v2", +// sources=["ix_full_bridge_v2.cpp"], +// extra_ldflags=[, "-Wl,-rpath,..."], +// extra_cflags=["-O2", "-std=c++17"], +// ) +// +// Upstream reference: xllm_latest/core/kernels/ilu/ixformer.h + +#include +#include +#include +#include +#include + +// ============================================================================ +// Forward declarations — ixformer::infer namespace from base image .so +// Signatures EXACTLY match upstream_ref/xllm_latest/core/kernels/ilu/ixformer.h +// ============================================================================ +namespace ixformer { namespace infer { + +// --- Attention --- +torch::Tensor ixinfer_flash_attn_unpad_with_block_tables( + torch::Tensor& query, + torch::Tensor& key_cache, + torch::Tensor& value_cache, + torch::Tensor& out, + torch::Tensor& block_tables, + torch::Tensor& cu_seq_q, + torch::Tensor& cu_seq_k, + int64_t max_seq_q, + int64_t max_seq_k, + bool is_causal, + int64_t window_left, + int64_t window_right, + double scale, + double softcap, + bool sqrt_alibi, + const std::optional& alibi_slopes, + const std::optional& sinks, + std::optional& lse); + +torch::Tensor xllm_paged_attention( + torch::Tensor& out, + torch::Tensor& query, + torch::Tensor& key_cache, + torch::Tensor& value_cache, + int64_t num_kv_heads, + double scale, + torch::Tensor& block_tables, + torch::Tensor& context_lens, + int64_t block_size, + int64_t max_context_len, + const std::optional& alibi_slopes, + bool causal, + int32_t window_left, + int32_t window_right, + double softcap, + bool enable_cuda_graph, + bool use_sqrt_alibi, + const std::optional& sinks); + +// --- Activation --- +void silu_and_mul(torch::Tensor& input, torch::Tensor& output); + +// --- Linear --- +torch::Tensor ixformer_linear(torch::Tensor& input, + torch::Tensor& weight, + int64_t act_type, + const std::optional& bias, + const std::optional& out, + const std::optional persistent); + +torch::Tensor ixformer_linear_ex(torch::Tensor& input, + torch::Tensor& weight, + const c10::optional& bias, + const c10::optional& out); + +// --- Cache --- +void xllm_reshape_and_cache(torch::Tensor& key, + torch::Tensor& value, + torch::Tensor& key_cache, + torch::Tensor& value_cache, + torch::Tensor& slot_mapping, + int64_t key_token_stride, + int64_t value_token_stride); + +// --- RoPE --- +void xllm_rotary_embedding(torch::Tensor& positions, + torch::Tensor& query, + torch::Tensor& key, + int64_t head_size, + torch::Tensor& cos_sin_cache, + bool is_neox); + +// --- Norm --- +void residual_rms_norm(torch::Tensor& input, + torch::Tensor& residual, + torch::Tensor& weight, + torch::Tensor& output, + torch::Tensor& residual_output, + const std::optional& fused_bias, + double alpha, + double eps, + bool is_post); + +void rms_norm(torch::Tensor& input, + torch::Tensor& weight, + torch::Tensor& output, + const std::optional& fused_bias, + double eps); + +// --- MoE --- +void topk_softmax(torch::Tensor& topk_weights, + torch::Tensor& topk_indices, + torch::Tensor& token_expert_indices, + torch::Tensor& gating_output, + bool renormalize); + +void moe_compute_token_index_api( + torch::Tensor& topk_ids, + torch::Tensor& src_dst, + torch::Tensor& dst_src, + torch::Tensor& expert_sizes_gpu, + const c10::optional& expert_mask, + const c10::optional& expert_sizes_cpu, + const c10::optional& expand_tokens_gpu, + int64_t start_expert_id, + int64_t end_expert_id, + int64_t num_experts); + +void moe_expand_input(torch::Tensor outputs, + torch::Tensor inputs, + torch::Tensor dst_to_src, + const c10::optional& src_to_dst, + int64_t dst_tokens, + int64_t expand_factor); + +void moe_w16a16_group_gemm(torch::Tensor output, + torch::Tensor inputs, + torch::Tensor weights, + torch::Tensor tokens_per_experts, + const c10::optional& dst_to_src, + const c10::optional& bias, + std::string format, + int64_t persistent, + int64_t output_n); + +void moe_output_reduce_sum(torch::Tensor outputs, + torch::Tensor inputs, + const c10::optional& mul_weight, + const c10::optional& mask, + const c10::optional& extra_residual, + double scaling_factor); + +}} // namespace ixformer::infer + + +// ============================================================================ +// Python wrappers — thin wrappers that match ix_bridge.py's expected API +// ============================================================================ + +// --- silu_and_mul --- +torch::Tensor ix_silu_and_mul(torch::Tensor input) { + int64_t half_dim = input.size(-1) / 2; + auto output = input.new_empty({input.size(0), half_dim}); + ixformer::infer::silu_and_mul(input, output); + return output; +} + +// --- rms_norm --- +void ix_rms_norm(torch::Tensor output, torch::Tensor input, + torch::Tensor weight, double eps) { + ixformer::infer::rms_norm(input, weight, output, + /*fused_bias=*/std::nullopt, eps); +} + +// --- fused_add_rms_norm --- +// residual_rms_norm does: output = rms_norm(input + alpha*residual, weight, eps) +// residual_output = input + alpha*residual +void ix_fused_add_rms_norm(torch::Tensor input, torch::Tensor residual, + torch::Tensor weight, torch::Tensor output, + torch::Tensor residual_output, double eps) { + ixformer::infer::residual_rms_norm(input, residual, weight, + output, residual_output, + /*fused_bias=*/std::nullopt, + /*alpha=*/1.0, eps, + /*is_post=*/false); +} + +// --- linear --- +torch::Tensor ix_linear(torch::Tensor input, torch::Tensor weight, + const c10::optional& bias) { + auto input_2d = input.view({-1, input.size(-1)}); + int64_t m = input_2d.size(0); + if (m <= 1 && !bias.has_value()) { + return ixformer::infer::ixformer_linear_ex( + input, weight, bias, /*out=*/c10::optional()); + } + return ixformer::infer::ixformer_linear( + input, weight, /*act_type=*/0, bias, + /*out=*/std::nullopt, /*persistent=*/std::nullopt); +} + +// --- rotary_embedding --- +void ix_rotary_embedding(torch::Tensor positions, torch::Tensor query, + torch::Tensor key, int64_t head_size, + torch::Tensor cos_sin_cache, bool is_neox) { + ixformer::infer::xllm_rotary_embedding( + positions, query, key, head_size, cos_sin_cache, is_neox); +} + +// --- reshape_and_cache --- +void ix_reshape_and_cache(torch::Tensor key, torch::Tensor value, + torch::Tensor key_cache, torch::Tensor value_cache, + torch::Tensor slot_mapping) { + // token stride = product of dims after dim 0 for key/value + // key shape: [num_tokens, num_heads, head_dim] + int64_t key_token_stride = 1; + for (int i = 1; i < key.dim(); i++) key_token_stride *= key.size(i); + int64_t value_token_stride = 1; + for (int i = 1; i < value.dim(); i++) value_token_stride *= value.size(i); + + ixformer::infer::xllm_reshape_and_cache( + key, value, key_cache, value_cache, slot_mapping, + key_token_stride, value_token_stride); +} + +// --- paged_attention (decode) --- +torch::Tensor ix_paged_attention( + torch::Tensor output, torch::Tensor query, + torch::Tensor key_cache, torch::Tensor value_cache, + int64_t num_kv_heads, double scale, + torch::Tensor block_tables, torch::Tensor context_lens, + int64_t block_size, int64_t max_context_len, + const c10::optional& alibi_slopes) { + return ixformer::infer::xllm_paged_attention( + output, query, key_cache, value_cache, + num_kv_heads, scale, block_tables, context_lens, + block_size, max_context_len, alibi_slopes, + /*causal=*/true, /*window_left=*/-1, /*window_right=*/-1, + /*softcap=*/0.0, /*enable_cuda_graph=*/false, + /*use_sqrt_alibi=*/false, /*sinks=*/std::nullopt); +} + +// --- flash_attn_prefill --- +torch::Tensor ix_flash_attn_prefill( + torch::Tensor query, torch::Tensor key_cache, torch::Tensor value_cache, + torch::Tensor output, torch::Tensor block_tables, + torch::Tensor cu_seq_q, torch::Tensor cu_seq_k, + int64_t max_query_len, int64_t max_seq_len, + double scale, bool is_causal, + int64_t window_left, int64_t window_right) { + std::optional lse = std::nullopt; + return ixformer::infer::ixinfer_flash_attn_unpad_with_block_tables( + query, key_cache, value_cache, output, block_tables, + cu_seq_q, cu_seq_k, max_query_len, max_seq_len, + is_causal, window_left, window_right, scale, + /*softcap=*/0.0, /*sqrt_alibi=*/false, + /*alibi_slopes=*/std::nullopt, /*sinks=*/std::nullopt, lse); +} + +// --- MoE: topk_softmax --- +// Returns (topk_weights, topk_ids, token_expert_indices) +std::tuple +ix_topk_softmax(torch::Tensor gating_output, int64_t topk, bool renormalize) { + int64_t num_tokens = gating_output.size(0); + auto topk_weights = torch::empty({num_tokens, topk}, + torch::dtype(torch::kFloat32).device(gating_output.device())); + auto topk_ids = torch::empty({num_tokens, topk}, + torch::dtype(torch::kInt32).device(gating_output.device())); + auto token_expert_indices = torch::empty({num_tokens, topk}, + torch::dtype(torch::kInt32).device(gating_output.device())); + + auto gating_f32 = gating_output.to(torch::kFloat32); + ixformer::infer::topk_softmax( + topk_weights, topk_ids, token_expert_indices, gating_f32, renormalize); + + return std::make_tuple(topk_weights, topk_ids, token_expert_indices); +} + +// --- MoE: moe_gen_idx --- +// Equivalent to xllm::kernel::ilu::moe_gen_idx +std::vector +ix_moe_gen_idx(torch::Tensor expert_id, int64_t expert_num) { + auto src_dst = expert_id.new_empty({expert_id.numel()}); + auto dst_src = torch::empty_like(src_dst); + auto expert_sizes_gpu = expert_id.new_empty({expert_num}); + + ixformer::infer::moe_compute_token_index_api( + expert_id, src_dst, dst_src, expert_sizes_gpu, + /*expert_mask=*/c10::nullopt, + /*expert_sizes_cpu=*/c10::nullopt, + /*expand_tokens_gpu=*/c10::nullopt, + /*start_expert_id=*/0, + /*end_expert_id=*/expert_num, + /*num_experts=*/expert_num); + + auto expert_sizes_cumsum = expert_sizes_gpu.cumsum(-1); + return {src_dst, dst_src, expert_sizes_gpu, expert_sizes_cumsum}; +} + +// --- MoE: moe_expand_input --- +torch::Tensor ix_moe_expand_input(torch::Tensor input, + torch::Tensor gather_index, + torch::Tensor combine_idx, + int64_t topk) { + int64_t dst_tokens = input.size(0) * topk; + auto output = input.new_empty({dst_tokens, input.size(1)}); + ixformer::infer::moe_expand_input( + output, input, combine_idx, gather_index, dst_tokens, topk); + return output; +} + +// --- MoE: group_gemm --- +torch::Tensor ix_group_gemm(torch::Tensor inputs, torch::Tensor weights, + torch::Tensor tokens_per_experts, + int64_t output_n) { + int64_t total_tokens = inputs.size(0); + auto output = inputs.new_empty({total_tokens, output_n}); + ixformer::infer::moe_w16a16_group_gemm( + output, inputs, weights, tokens_per_experts, + /*dst_to_src=*/c10::nullopt, + /*bias=*/c10::nullopt, + /*format=*/"default", + /*persistent=*/0, + output_n); + return output; +} + +// --- MoE: moe_combine_result --- +torch::Tensor ix_moe_combine_result(torch::Tensor input, torch::Tensor weight) { + // input: [T*topk, H], weight: [T, topk] + auto input_3d = input.view({-1, weight.size(1), input.size(1)}); + auto output = input.new_empty({input_3d.size(0), input_3d.size(2)}); + ixformer::infer::moe_output_reduce_sum( + output, input_3d, weight, + /*mask=*/c10::nullopt, + /*extra_residual=*/c10::nullopt, + /*scaling_factor=*/1.0); + return output; +} + +// --- MoE: fused_moe_forward (7-step pipeline) --- +// This is the full fused MoE forward: topk → gen_idx → expand → gemm(w13) → +// silu_mul → gemm(w2) → combine +torch::Tensor ix_fused_moe_forward( + torch::Tensor hidden_states, + torch::Tensor router_logits, + torch::Tensor w13, // [num_experts, 2*intermediate, hidden] + torch::Tensor w2, // [num_experts, hidden, intermediate] + int64_t topk, + int64_t num_experts, + bool renormalize) { + + // Step 1: topk_softmax + auto [topk_weights, topk_ids, token_expert_indices] = + ix_topk_softmax(router_logits, topk, renormalize); + + if (renormalize) { + auto sum = topk_weights.sum(-1, /*keepdim=*/true); + topk_weights = topk_weights / sum; + } + + // Step 2: moe_gen_idx + auto idx_results = ix_moe_gen_idx(topk_ids.view({-1}), num_experts); + auto& src_dst = idx_results[0]; + auto& dst_src = idx_results[1]; + auto& expert_sizes_gpu = idx_results[2]; + + // Step 3: moe_expand_input + auto expanded = ix_moe_expand_input(hidden_states, src_dst, dst_src, topk); + + // Step 4: group_gemm (w13: gate_up projection) + int64_t intermediate_2x = w13.size(1); + auto gate_up = ix_group_gemm(expanded, w13.view({-1, w13.size(2)}), + expert_sizes_gpu, intermediate_2x); + + // Step 5: silu_and_mul + auto activated = ix_silu_and_mul(gate_up); + + // Step 6: group_gemm (w2: down projection) + int64_t hidden_size = w2.size(1); + auto down = ix_group_gemm(activated, w2.view({-1, w2.size(2)}), + expert_sizes_gpu, hidden_size); + + // Step 7: moe_combine_result + auto output = ix_moe_combine_result(down, topk_weights); + + return output; +} + + +// ============================================================================ +// Module registration — ALL 14 functions + fused pipeline +// ============================================================================ +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + // Activation + m.def("silu_and_mul", &ix_silu_and_mul, + "Fused SiLU+mul activation via ixformer::infer"); + + // Norm + m.def("rms_norm", &ix_rms_norm, + "RMSNorm via ixformer::infer"); + m.def("fused_add_rms_norm", &ix_fused_add_rms_norm, + "Residual + RMSNorm via ixformer::infer"); + + // Linear + m.def("linear", &ix_linear, + "GEMM via ixformer::infer (linear/linear_ex)"); + + // RoPE + m.def("rotary_embedding", &ix_rotary_embedding, + "Rotary position embedding via ixformer::infer"); + + // Cache + m.def("reshape_and_cache", &ix_reshape_and_cache, + "KV cache reshape+store via ixformer::infer"); + + // Attention + m.def("paged_attention", &ix_paged_attention, + "Paged attention decode via ixformer::infer"); + m.def("flash_attn_prefill", &ix_flash_attn_prefill, + "Flash attention prefill via ixformer::infer"); + + // MoE (individual steps) + m.def("topk_softmax", &ix_topk_softmax, + "MoE topk+softmax routing via ixformer::infer"); + m.def("moe_gen_idx", &ix_moe_gen_idx, + "MoE compute token index via ixformer::infer"); + m.def("moe_expand_input", &ix_moe_expand_input, + "MoE expand input for expert dispatch via ixformer::infer"); + m.def("group_gemm", &ix_group_gemm, + "MoE grouped GEMM via ixformer::infer"); + m.def("moe_combine_result", &ix_moe_combine_result, + "MoE output reduce sum via ixformer::infer"); + + // MoE (fused 7-step pipeline) + m.def("fused_moe_forward", &ix_fused_moe_forward, + "Complete fused MoE forward (7-step pipeline) via ixformer::infer"); +} diff --git a/ex_engine/python/ix_bridge_v2.py b/ex_engine/python/ix_bridge_v2.py new file mode 100644 index 00000000..07bf2546 --- /dev/null +++ b/ex_engine/python/ix_bridge_v2.py @@ -0,0 +1,210 @@ +""" +ix_bridge_v2.py — Complete ixformer bridge loader (14 functions). + +Loads ix_full_bridge_v2.so via JIT compilation, linking against ALL +ixformer .so files in the base image. + +Functions exposed: + MoE: topk_softmax, moe_gen_idx, moe_expand_input, group_gemm, + silu_and_mul, moe_combine_result, fused_moe_forward + Attention: paged_attention, flash_attn_prefill + Norm: rms_norm, fused_add_rms_norm + RoPE: rotary_embedding + Cache: reshape_and_cache + Linear: linear +""" + +import os +import logging +import glob +import torch +from typing import Tuple, Optional, List + +logger = logging.getLogger("ex_engine.ix_bridge_v2") + +_bridge = None +_loaded = False +_available = False + + +def _find_cpp(): + """Find ix_full_bridge_v2.cpp in known locations.""" + here = os.path.dirname(os.path.abspath(__file__)) + candidates = [ + os.path.join(here, "..", "csrc", "ix_full_bridge_v2.cpp"), + os.path.join("/workspace/ex_engine/csrc", "ix_full_bridge_v2.cpp"), + # fallback to v1 + os.path.join(here, "..", "csrc", "ix_full_bridge.cpp"), + os.path.join("/workspace/ex_engine/csrc", "ix_full_bridge.cpp"), + ] + for c in candidates: + p = os.path.normpath(c) + if os.path.exists(p): + return p + return None + + +def _collect_ixformer_libs(): + """Collect all ixformer .so files for linking.""" + extra_ldflags = [] + rpath_dirs = set() + + # From ixformer Python package + try: + import ixformer + ixf_dir = os.path.dirname(ixformer.__file__) + for so in glob.glob(os.path.join(ixf_dir, "*.so")): + extra_ldflags.append(so) + rpath_dirs.add(os.path.dirname(so)) + # Also the _ixformer_torch extension + for so in glob.glob(os.path.join(ixf_dir, "_ixformer_torch*.so")): + if so not in extra_ldflags: + extra_ldflags.append(so) + except ImportError: + pass + + # From corex lib64 + corex_lib = "/usr/local/corex/lib64" + if os.path.isdir(corex_lib): + for lib in ["libixattn.so", "libixformer.so", "libcublas.so", + "libcudart.so", "libcudnn.so"]: + p = os.path.join(corex_lib, lib) + if os.path.exists(p) and p not in extra_ldflags: + extra_ldflags.append(p) + rpath_dirs.add(corex_lib) + + # From ixformer subdirectory + ixf_subdir = os.path.join(corex_lib, "python3/dist-packages/ixformer") + if os.path.isdir(ixf_subdir): + for so in glob.glob(os.path.join(ixf_subdir, "*.so")): + if so not in extra_ldflags: + extra_ldflags.append(so) + rpath_dirs.add(ixf_subdir) + + # Add rpath + for d in rpath_dirs: + extra_ldflags.append(f"-Wl,-rpath,{d}") + + return extra_ldflags + + +def _load_bridge(): + """JIT compile and load the bridge.""" + global _bridge, _loaded, _available + if _loaded: + return _available + _loaded = True + + cpp_path = _find_cpp() + if cpp_path is None: + logger.warning("ix_full_bridge_v2.cpp not found") + return False + + extra_ldflags = _collect_ixformer_libs() + logger.info("ix_bridge_v2: compiling %s", cpp_path) + logger.info("ix_bridge_v2: ldflags count=%d", len(extra_ldflags)) + + try: + from torch.utils.cpp_extension import load + mod_name = "ix_full_bridge_v2" if "v2" in cpp_path else "ix_full_bridge" + _bridge = load( + name=mod_name, + sources=[cpp_path], + extra_cflags=["-O2", "-std=c++17"], + extra_ldflags=extra_ldflags, + verbose=False, + ) + _available = True + fns = [x for x in dir(_bridge) if not x.startswith("_")] + logger.info("ix_bridge_v2 loaded: %s", fns) + return True + except Exception as e: + logger.error("ix_bridge_v2 JIT compile failed: %s", e) + return False + + +def is_available() -> bool: + if not _loaded: + _load_bridge() + return _available + + +def _get(): + if not is_available(): + raise RuntimeError("ix_bridge_v2 not available") + return _bridge + + +# ========================================================================= +# MoE +# ========================================================================= +def topk_softmax(gating_output, topk, renormalize=True): + """Returns (topk_weights, topk_ids, token_expert_indices).""" + return _get().topk_softmax(gating_output, topk, renormalize) + +def moe_gen_idx(expert_id, expert_num): + """Returns [src_dst, dst_src, expert_sizes_gpu, expert_sizes_cumsum].""" + return _get().moe_gen_idx(expert_id, expert_num) + +def moe_expand_input(input, gather_index, combine_idx, topk): + return _get().moe_expand_input(input, gather_index, combine_idx, topk) + +def group_gemm(inputs, weights, token_count, output_n): + return _get().group_gemm(inputs, weights, token_count, output_n) + +def silu_and_mul(input): + return _get().silu_and_mul(input) + +def moe_combine_result(input, weight): + return _get().moe_combine_result(input, weight) + +def fused_moe_forward(hidden_states, router_logits, w13, w2, + topk, num_experts, renormalize=True): + return _get().fused_moe_forward( + hidden_states, router_logits, w13, w2, topk, num_experts, renormalize) + +# ========================================================================= +# Attention +# ========================================================================= +def paged_attention(output, query, key_cache, value_cache, + num_kv_heads, scale, block_tables, seq_lens, + block_size, max_context_len, alibi_slopes=None): + return _get().paged_attention( + output, query, key_cache, value_cache, + num_kv_heads, scale, block_tables, seq_lens, + block_size, max_context_len, alibi_slopes) + +def flash_attn_prefill(query, key_cache, value_cache, output, block_tables, + cu_seq_q, cu_seq_k, max_query_len, max_seq_len, + scale, is_causal=True, window_left=-1, window_right=-1): + return _get().flash_attn_prefill( + query, key_cache, value_cache, output, block_tables, + cu_seq_q, cu_seq_k, max_query_len, max_seq_len, + scale, is_causal, window_left, window_right) + +# ========================================================================= +# Norm +# ========================================================================= +def rms_norm(output, input, weight, eps=1e-6): + return _get().rms_norm(output, input, weight, eps) + +def fused_add_rms_norm(input, residual, weight, output, residual_output, eps=1e-6): + return _get().fused_add_rms_norm(input, residual, weight, output, residual_output, eps) + +# ========================================================================= +# RoPE +# ========================================================================= +def rotary_embedding(positions, query, key, head_size, cos_sin_cache, is_neox=True): + return _get().rotary_embedding(positions, query, key, head_size, cos_sin_cache, is_neox) + +# ========================================================================= +# Cache +# ========================================================================= +def reshape_and_cache(key, value, key_cache, value_cache, slot_mapping): + return _get().reshape_and_cache(key, value, key_cache, value_cache, slot_mapping) + +# ========================================================================= +# Linear +# ========================================================================= +def linear(input, weight, bias=None): + return _get().linear(input, weight, bias)