diff --git a/ex_engine/csrc/ix_full_bridge.cpp b/ex_engine/csrc/ix_full_bridge.cpp index 4f82deb8..72ddcd8e 100644 --- a/ex_engine/csrc/ix_full_bridge.cpp +++ b/ex_engine/csrc/ix_full_bridge.cpp @@ -1,339 +1,90 @@ -// ix_full_bridge.cpp — Complete ixformer::infer bridge for BI-V100 +// ix_full_bridge.cpp — Bridge to ixformer C++ functions available in base image // -// Exposes ALL 14 ixformer C++ functions to Python via pybind11. -// Header source: upstream_ref/xllm/xllm/core/kernels/ilu/ixformer.h +// Based on symbol probe of the actual BI-V100 base image: +// _ixformer_torch.so has: silu_and_mul_forward, rms_norm_forward, +// fused_add_rms_norm_forward, ixformer_linear, ixformer_linear_ex +// libixformer.so has: ixinfer_flash_attn_unpad_fwd // -// This replaces the partial ix_moe_bridge.cpp with the full set: -// MoE pipeline: topk_softmax, moe_compute_token_index_api, moe_expand_input, -// moe_w16a16_group_gemm, silu_and_mul, moe_output_reduce_sum -// Attention: ixinfer_flash_attn_unpad_with_block_tables, xllm_paged_attention -// Norm: rms_norm, residual_rms_norm -// RoPE: xllm_rotary_embedding -// Linear: ixformer_linear, ixformer_linear_ex -// Cache: xllm_reshape_and_cache +// MoE functions (topk_softmax, group_gemm, etc.) are NOT in base image. +// They exist only in xllm's compiled library. MoE must use Python fallback. #include +#include #include #include // ============================================================================ -// Compatibility: CoreX torch uses c10::optional which may not implicitly -// convert from kNoneTensor to const std::optional&. -// Use typed empty optionals instead. +// Forward declarations — ACTUAL symbols from base image .so files +// Namespace: ixformer_torch_ext (in _ixformer_torch.cpython-310.so) // ============================================================================ -static const std::optional kNoneTensor = {}; -static const std::optional kNoneBool = {}; +namespace ixformer_torch_ext { + +// silu_and_mul: _ZN18ixformer_torch_ext20silu_and_mul_forwardERN2at6TensorES2_ +void silu_and_mul_forward(at::Tensor& input, at::Tensor& output); + +// rms_norm: _ZN18ixformer_torch_ext16rms_norm_forwardERN2at6TensorES2_S2_d +void rms_norm_forward(at::Tensor& input, at::Tensor& weight, at::Tensor& output, double eps); + +// fused_add_rms_norm: _ZN18ixformer_torch_ext26fused_add_rms_norm_forwardERN2at6TensorES2_S2_dd +void fused_add_rms_norm_forward(at::Tensor& input, at::Tensor& residual, + at::Tensor& weight, double eps, double alpha); + +// ixformer_linear: _ZN18ixformer_torch_ext15ixformer_linearERN2at6TensorES2_RKN3c108optionalIS1_EES7_ +at::Tensor ixformer_linear(at::Tensor& input, at::Tensor& weight, + const c10::optional& bias, + const c10::optional& out); + +// ixformer_linear_ex: _ZN18ixformer_torch_ext18ixformer_linear_exERN2at6TensorES2_RKN3c108optionalIS1_EE +at::Tensor ixformer_linear_ex(at::Tensor& input, at::Tensor& weight, + const c10::optional& bias); + +} // namespace ixformer_torch_ext + // ============================================================================ -// Forward-declare ixformer::infer namespace — matches ixformer.h exactly -// We forward-declare instead of #include to avoid build-time dependency -// on internal headers (ixinfer.h etc) that may not be on include path. -// The symbols resolve at link time against the base image's libixattn.so etc. -// ============================================================================ -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); - -// --- 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); - -// --- Activation --- -void silu_and_mul(torch::Tensor& input, torch::Tensor& output); - -// --- 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); - -// --- KV 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); - -// --- 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 std::optional& bias, - const std::optional& out); - -// --- 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 std::optional& expert_mask, - const std::optional& expert_sizes_cpu, - const std::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 std::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 std::optional& dst_to_src, - const std::optional& bias, - std::string format, int64_t persistent, int64_t output_n); - -void moe_output_reduce_sum( - torch::Tensor outputs, torch::Tensor inputs, - const std::optional& mul_weight, - const std::optional& mask, - const std::optional& extra_residual, - double scaling_factor); - -} // namespace infer -} // namespace ixformer - -// ============================================================================ -// Python wrappers — thin wrappers matching upstream xllm ILU kernel layer -// Source: upstream_ref/xllm/xllm/core/kernels/ilu/*.cpp +// Python wrappers // ============================================================================ -// --- MoE: topk_softmax (from ilu/fused_moe.cpp moe_active_topk) --- -std::tuple ix_topk_softmax( - torch::Tensor gating_output, int64_t topk, bool renormalize) { - auto input = gating_output.to(torch::kFloat32).contiguous(); - int64_t num_tokens = input.size(0); - auto topk_weights = torch::empty({num_tokens, topk}, - torch::dtype(torch::kFloat32).device(input.device())); - auto topk_indices = torch::empty({num_tokens, topk}, - torch::dtype(torch::kInt32).device(input.device())); - auto token_expert_indices = torch::empty({num_tokens, topk}, - torch::dtype(torch::kInt32).device(input.device())); - ixformer::infer::topk_softmax( - topk_weights, topk_indices, token_expert_indices, input, false); - if (renormalize) { - topk_weights = topk_weights / topk_weights.sum(-1, /*keepdim=*/true); - } - return std::make_tuple(topk_weights, topk_indices); -} - -// --- MoE: gen_idx (from ilu/fused_moe.cpp 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, - kNoneTensor, kNoneTensor, kNoneTensor, 0, expert_num, expert_num); - auto expert_sizes_gpu_cumsum = expert_sizes_gpu.cumsum(-1); - return {src_dst, dst_src, expert_sizes_gpu, expert_sizes_gpu_cumsum}; -} - -// --- 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 token_count, 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, token_count, - kNoneTensor, kNoneTensor, "NT", 0, output_n); - return output; -} - -// --- MoE: silu_and_mul --- +// --- 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); + ixformer_torch_ext::silu_and_mul_forward(input, output); return output; } -// --- MoE: combine_result --- -torch::Tensor ix_moe_combine_result(torch::Tensor input, torch::Tensor weight) { - input = input.view({-1, weight.size(1), input.size(1)}); - auto output = input.new_empty({input.size(0), input.size(2)}); - ixformer::infer::moe_output_reduce_sum( - output, input, weight, kNoneTensor, kNoneTensor, 1.0); - return output; +// --- rms_norm --- +void ix_rms_norm(torch::Tensor output, torch::Tensor input, + torch::Tensor weight, double eps) { + ixformer_torch_ext::rms_norm_forward(input, weight, output, eps); } -// --- MoE: full fused forward (from ilu/layers/fused_moe.cpp) --- -torch::Tensor ix_fused_moe_forward( - torch::Tensor hidden_states, torch::Tensor router_logits, - torch::Tensor w13, torch::Tensor w2, - int64_t topk, int64_t num_experts, bool renormalize) { - auto [topk_weights, topk_ids] = ix_topk_softmax(router_logits, topk, renormalize); - auto idx = ix_moe_gen_idx(topk_ids.view({-1}), num_experts); - auto expanded = ix_moe_expand_input(hidden_states, idx[0], idx[1], topk); - int64_t gate_up_dim = w13.size(1); - auto gemm1_out = ix_group_gemm(expanded, w13, idx[2], gate_up_dim); - auto act_out = ix_silu_and_mul(gemm1_out); - int64_t hidden_dim = w2.size(1); - auto gemm2_out = ix_group_gemm(act_out, w2, idx[2], hidden_dim); - return ix_moe_combine_result(gemm2_out, topk_weights); +// --- fused_add_rms_norm --- +void ix_fused_add_rms_norm(torch::Tensor input, torch::Tensor residual, + torch::Tensor weight, double eps) { + ixformer_torch_ext::fused_add_rms_norm_forward(input, residual, weight, eps, 1.0); } -// --- Attention: paged decode (from ilu/attention.cpp batch_decode) --- -void 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 seq_lens, - int64_t block_size, int64_t max_context_len, - const std::optional& alibi_slopes) { - if (query.dim() == 4) { - query = query.view({query.size(0)*query.size(1), query.size(2), query.size(3)}).contiguous(); +// --- linear --- +torch::Tensor ix_linear(torch::Tensor input, torch::Tensor weight, + const c10::optional& bias) { + // Use linear_ex for decode (m<=1), linear for prefill + auto input_2d = input.view({-1, input.size(-1)}); + int64_t m = input_2d.size(0); + if (m <= 1 && !bias.has_value()) { + return ixformer_torch_ext::ixformer_linear_ex(input, weight, bias); } - if (output.dim() == 4) { - output = output.view({output.size(0)*output.size(1), output.size(2), output.size(3)}).contiguous(); - } - ixformer::infer::xllm_paged_attention( - output, query, key_cache, value_cache, num_kv_heads, scale, - block_tables, seq_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=*/kNoneTensor); + return ixformer_torch_ext::ixformer_linear(input, weight, bias, + c10::optional()); } -// --- Attention: prefill flash (from ilu/attention.cpp batch_prefill) --- -void ix_flash_attn_prefill( - torch::Tensor query, torch::Tensor key, torch::Tensor value, - 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 = {}; - ixformer::infer::ixinfer_flash_attn_unpad_with_block_tables( - query, key, value, 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=*/kNoneTensor, /*sinks=*/kNoneTensor, lse); -} - -// --- Norm: rms_norm (from ilu/norm.cpp) --- -void ix_rms_norm( - torch::Tensor output, torch::Tensor input, - torch::Tensor weight, double eps) { - ixformer::infer::rms_norm(input, weight, output, kNoneTensor, eps); -} - -// --- Norm: fused residual + rms_norm (from ilu/norm.cpp) --- -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, - kNoneTensor, 1.0, eps, false); -} - -// --- RoPE (from ilu/rope.cpp) --- -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); -} - -// --- KV Cache reshape (from ilu/attention.cpp reshape_paged_cache) --- -void ix_reshape_and_cache( - torch::Tensor key, torch::Tensor value, - torch::Tensor key_cache, torch::Tensor value_cache, - torch::Tensor slot_mapping) { - slot_mapping = slot_mapping.to(torch::kLong); - int64_t key_stride = key.stride(0); - int64_t val_stride = value.stride(0); - ixformer::infer::xllm_reshape_and_cache( - key, value, key_cache, value_cache, slot_mapping, - key_stride, val_stride); -} - -// --- Linear --- -torch::Tensor ix_linear( - torch::Tensor input, torch::Tensor weight, - const std::optional& bias) { - return ixformer::infer::ixformer_linear( - input, weight, /*act_type=*/-1, bias, kNoneTensor, kNoneBool); -} // ============================================================================ // Module registration // ============================================================================ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - // MoE - m.def("topk_softmax", &ix_topk_softmax, "Fused topk+softmax", - py::arg("gating_output"), py::arg("topk"), py::arg("renormalize")=true); - m.def("moe_gen_idx", &ix_moe_gen_idx); - m.def("moe_expand_input", &ix_moe_expand_input); - m.def("group_gemm", &ix_group_gemm); - m.def("silu_and_mul", &ix_silu_and_mul); - m.def("moe_combine_result", &ix_moe_combine_result); - m.def("fused_moe_forward", &ix_fused_moe_forward, - py::arg("hidden_states"), py::arg("router_logits"), - py::arg("w13"), py::arg("w2"), - py::arg("topk"), py::arg("num_experts"), py::arg("renormalize")=true); - // Attention - m.def("paged_attention", &ix_paged_attention); - m.def("flash_attn_prefill", &ix_flash_attn_prefill); - // Norm - m.def("rms_norm", &ix_rms_norm); - m.def("fused_add_rms_norm", &ix_fused_add_rms_norm); - // RoPE - m.def("rotary_embedding", &ix_rotary_embedding); - // Cache - m.def("reshape_and_cache", &ix_reshape_and_cache); - // Linear - m.def("linear", &ix_linear); + m.def("silu_and_mul", &ix_silu_and_mul, "Fused SiLU+mul activation"); + m.def("rms_norm", &ix_rms_norm, "RMSNorm"); + m.def("fused_add_rms_norm", &ix_fused_add_rms_norm, "Fused residual + RMSNorm"); + m.def("linear", &ix_linear, "ixformer GEMM (linear/linear_ex)"); } diff --git a/verify_single_gpu.py b/verify_single_gpu.py index 08fbd504..28c84573 100644 --- a/verify_single_gpu.py +++ b/verify_single_gpu.py @@ -1,39 +1,23 @@ #!/usr/bin/env python3 """ -verify_single_gpu.py — 单卡 BI-V100 验证 ix_full_bridge + MoE dispatch chain +verify_single_gpu.py — Single-card BI-V100 verification -用法: python3 verify_single_gpu.py -要求: 在真机 Docker 里运行,单卡即可 - -测试链条: - Step 0: JIT 编译 ix_full_bridge.cpp → .so - Step 1: ixformer::infer::topk_softmax - Step 2: ixformer::infer::moe_compute_token_index_api (gen_idx) - Step 3: ixformer::infer::moe_expand_input - Step 4: ixformer::infer::moe_w16a16_group_gemm (w13) - Step 5: ixformer::infer::silu_and_mul - Step 6: ixformer::infer::moe_w16a16_group_gemm (w2) - Step 7: ixformer::infer::moe_output_reduce_sum (combine) - Step 8: fused_moe_forward (全链路一次调用) - Step 9: paged_attention - Step 10: flash_attn_prefill - Step 11: rms_norm +Tests: + Step 0: JIT compile ix_full_bridge.cpp + Step 1: silu_and_mul (from _ixformer_torch.so) + Step 2: rms_norm + Step 3: fused_add_rms_norm + Step 4: linear (ixformer GEMM) + Step 5: ixformer.functions Python-level flash_attn + Step 6: ixformer.functions Python-level paged_attention + Step 7: corex_moe.py Python tiered dispatch (MoE full pipeline) """ +import os, sys, time, traceback, glob -import os -import sys -import time -import traceback - -# ============================================================================ -# Step 0: JIT compile ix_full_bridge.cpp -# ============================================================================ def step0_compile_bridge(): print("=" * 60) print("STEP 0: JIT compile ix_full_bridge.cpp") print("=" * 60) - - # Find the source here = os.path.dirname(os.path.abspath(__file__)) candidates = [ os.path.join(here, "ex_engine", "csrc", "ix_full_bridge.cpp"), @@ -45,16 +29,12 @@ def step0_compile_bridge(): cpp_path = c break if cpp_path is None: - print(" ✗ ix_full_bridge.cpp NOT FOUND") - print(f" Searched: {candidates}") + print(f" ✗ ix_full_bridge.cpp NOT FOUND in {candidates}") return None - print(f" Source: {cpp_path}") from torch.utils.cpp_extension import load - import glob - # Find ixformer .so to link against extra_ldflags = [] try: import ixformer @@ -64,23 +44,14 @@ def step0_compile_bridge(): extra_ldflags.append(so) for so in glob.glob(os.path.join(ixf_dir, "_ixformer_torch*.so")): extra_ldflags.append(so) + extra_ldflags.append(f"-Wl,-rpath,{ixf_dir}") except ImportError: pass corex_lib = "/usr/local/corex/lib64" if os.path.isdir(corex_lib): - for lib in ["libixattn.so", "libixformer.so"]: - p = os.path.join(corex_lib, lib) - if os.path.exists(p) and p not in extra_ldflags: - extra_ldflags.append(p) extra_ldflags.append(f"-Wl,-rpath,{corex_lib}") - try: - import ixformer - extra_ldflags.append(f"-Wl,-rpath,{os.path.dirname(ixformer.__file__)}") - except ImportError: - pass - - print(f" Link libs: {[os.path.basename(x) for x in extra_ldflags if not x.startswith('-')]}") + print(f" Link: {[os.path.basename(x) for x in extra_ldflags if not x.startswith('-')]}") t0 = time.time() try: bridge = load( @@ -92,368 +63,161 @@ def step0_compile_bridge(): ) dt = time.time() - t0 fns = [x for x in dir(bridge) if not x.startswith("_")] - print(f" ✓ Compiled in {dt:.1f}s") - print(f" Functions: {fns}") + print(f" ✓ Compiled in {dt:.1f}s — functions: {fns}") return bridge except Exception as e: - dt = time.time() - t0 - print(f" ✗ Compile FAILED after {dt:.1f}s") - print(f" Error: {e}") + print(f" ✗ FAILED after {time.time()-t0:.1f}s: {e}") traceback.print_exc() return None - -# ============================================================================ -# Step 1-7: MoE dispatch chain (individual steps) -# ============================================================================ -def step1_topk_softmax(bridge): +def step1_silu(bridge): import torch - print("\nSTEP 1: topk_softmax") - gating = torch.randn(4, 64, dtype=torch.float32, device="cuda") # 4 tokens, 64 experts - topk = 8 + print("\nSTEP 1: silu_and_mul") + x = torch.randn(4, 256, dtype=torch.float16, device="cuda") # will split into 128+128 try: - weights, ids = bridge.topk_softmax(gating, topk, True) - print(f" ✓ weights: {weights.shape} {weights.dtype}, ids: {ids.shape} {ids.dtype}") - print(f" weights sum per token: {weights.sum(dim=-1).tolist()}") - print(f" ids range: [{ids.min().item()}, {ids.max().item()}]") - assert weights.shape == (4, 8) - assert ids.shape == (4, 8) - assert ids.max().item() < 64 - return weights, ids - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None, None - - -def step2_gen_idx(bridge, ids): - import torch - print("\nSTEP 2: moe_gen_idx") - try: - flat_ids = ids.view(-1) # (32,) - result = bridge.moe_gen_idx(flat_ids, 64) - print(f" ✓ Returns {len(result)} tensors:") - names = ["src_dst", "dst_src", "expert_sizes", "cumsum"] - for i, (name, t) in enumerate(zip(names, result)): - print(f" [{i}] {name}: {t.shape} {t.dtype}") - return result - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None - - -def step3_expand_input(bridge, idx, topk=8): - import torch - print("\nSTEP 3: moe_expand_input") - hidden = torch.randn(4, 1536, dtype=torch.float16, device="cuda") # 4 tokens, hidden=1536 - try: - expanded = bridge.moe_expand_input(hidden, idx[0], idx[1], topk) - print(f" ✓ expanded: {expanded.shape} {expanded.dtype}") - assert expanded.shape[0] == 4 * topk - return expanded, hidden - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None, None - - -def step4_group_gemm_w13(bridge, expanded, idx): - import torch - print("\nSTEP 4: group_gemm (w13)") - # w13: (64 experts, 2*intermediate, hidden) — simulate small - # Real: (64, 4608, 1536) but we use smaller for test - E, inter2, H = 64, 256, 1536 - w13 = torch.randn(E, inter2, H, dtype=torch.float16, device="cuda") - try: - gemm1 = bridge.group_gemm(expanded, w13, idx[2], inter2) - print(f" ✓ gemm1: {gemm1.shape} {gemm1.dtype}") - return gemm1, w13 - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None, None - - -def step5_silu_and_mul(bridge, gemm1): - import torch - print("\nSTEP 5: silu_and_mul") - try: - act = bridge.silu_and_mul(gemm1) - print(f" ✓ act: {act.shape} {act.dtype}") - assert act.shape[-1] == gemm1.shape[-1] // 2 - return act - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None - - -def step6_group_gemm_w2(bridge, act, idx): - import torch - print("\nSTEP 6: group_gemm (w2)") - E, H, I = 64, 1536, act.shape[-1] - w2 = torch.randn(E, H, I, dtype=torch.float16, device="cuda") - try: - gemm2 = bridge.group_gemm(act, w2, idx[2], H) - print(f" ✓ gemm2: {gemm2.shape} {gemm2.dtype}") - return gemm2 - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None - - -def step7_combine_result(bridge, gemm2, weights): - import torch - print("\nSTEP 7: moe_combine_result") - try: - result = bridge.moe_combine_result(gemm2, weights) - print(f" ✓ result: {result.shape} {result.dtype}") - print(f" NaN: {result.isnan().any().item()}, Inf: {result.isinf().any().item()}") - return result - except Exception as e: - print(f" ✗ FAILED: {e}") - traceback.print_exc() - return None - - -# ============================================================================ -# Step 8: Full fused_moe_forward -# ============================================================================ -def step8_fused_moe(bridge): - import torch - print("\n" + "=" * 60) - print("STEP 8: fused_moe_forward (FULL PIPELINE)") - print("=" * 60) - num_tokens = 4 - hidden_size = 1536 - num_experts = 64 - inter_size = 128 # small for test - topk = 8 - - hidden = torch.randn(num_tokens, hidden_size, dtype=torch.float16, device="cuda") - router = torch.randn(num_tokens, num_experts, dtype=torch.float16, device="cuda") - w13 = torch.randn(num_experts, inter_size * 2, hidden_size, dtype=torch.float16, device="cuda") - w2 = torch.randn(num_experts, hidden_size, inter_size, dtype=torch.float16, device="cuda") - - try: - result = bridge.fused_moe_forward(hidden, router, w13, w2, topk, num_experts, True) - print(f" ✓ result: {result.shape} {result.dtype}") - print(f" NaN: {result.isnan().any().item()}, Inf: {result.isinf().any().item()}") - print(f" abs_mean: {result.abs().mean().item():.4f}") + out = bridge.silu_and_mul(x) + print(f" ✓ {x.shape} → {out.shape}, NaN={out.isnan().any().item()}, abs_mean={out.abs().mean().item():.4f}") return True except Exception as e: - print(f" ✗ FAILED: {e}") + print(f" ✗ {e}") traceback.print_exc() return False - -# ============================================================================ -# Step 9: paged_attention -# ============================================================================ -def step9_paged_attention(bridge): +def step2_rms_norm(bridge): import torch - print("\nSTEP 9: paged_attention") - B, Hq, Hkv, D = 1, 4, 1, 128 - block_size = 16 - max_blocks = 4 - seq_len = 48 # fits in 3 blocks - - query = torch.randn(B, Hq, D, dtype=torch.float16, device="cuda") - # KV cache: (num_blocks, Hkv, block_size, D) - total_blocks = max_blocks - key_cache = torch.randn(total_blocks, Hkv, block_size, D, dtype=torch.float16, device="cuda") - value_cache = torch.randn(total_blocks, Hkv, block_size, D, dtype=torch.float16, device="cuda") - block_tables = torch.arange(max_blocks, dtype=torch.int32, device="cuda").unsqueeze(0) - seq_lens = torch.tensor([seq_len], dtype=torch.int32, device="cuda") - output = torch.empty(B, Hq, D, dtype=torch.float16, device="cuda") - + print("\nSTEP 2: rms_norm") + x = torch.randn(4, 128, dtype=torch.float16, device="cuda") + w = torch.ones(128, dtype=torch.float16, device="cuda") + out = torch.empty_like(x) try: - bridge.paged_attention( - output, query, key_cache, value_cache, - Hkv, D ** -0.5, - block_tables, seq_lens, block_size, seq_len, None) - print(f" ✓ output: {output.shape} {output.dtype}") - print(f" NaN: {output.isnan().any().item()}, abs_mean: {output.abs().mean().item():.4f}") + bridge.rms_norm(out, x, w, 1e-6) + print(f" ✓ {out.shape}, NaN={out.isnan().any().item()}, abs_mean={out.abs().mean().item():.4f}") return True except Exception as e: - print(f" ✗ FAILED: {e}") + print(f" ✗ {e}") traceback.print_exc() return False - -# ============================================================================ -# Step 10: flash_attn_prefill -# ============================================================================ -def step10_flash_attn(bridge): +def step3_fused_add_rms_norm(bridge): import torch - print("\nSTEP 10: flash_attn_prefill") - Hq, Hkv, D = 4, 1, 128 - seq_len = 32 - - query = torch.randn(seq_len, Hq, D, dtype=torch.float16, device="cuda") - key = torch.randn(seq_len, Hkv, D, dtype=torch.float16, device="cuda") - value = torch.randn(seq_len, Hkv, D, dtype=torch.float16, device="cuda") - output = torch.empty_like(query) - block_tables = torch.empty(0, dtype=torch.int32, device="cuda") - cu_q = torch.tensor([0, seq_len], dtype=torch.int32, device="cuda") - cu_k = torch.tensor([0, seq_len], dtype=torch.int32, device="cuda") - + print("\nSTEP 3: fused_add_rms_norm") + x = torch.randn(4, 128, dtype=torch.float16, device="cuda") + res = torch.randn(4, 128, dtype=torch.float16, device="cuda") + w = torch.ones(128, dtype=torch.float16, device="cuda") try: - bridge.flash_attn_prefill( - query, key, value, output, block_tables, - cu_q, cu_k, seq_len, seq_len, D ** -0.5, True, -1, -1) - print(f" ✓ output: {output.shape} {output.dtype}") - print(f" NaN: {output.isnan().any().item()}, abs_mean: {output.abs().mean().item():.4f}") + bridge.fused_add_rms_norm(x, res, w, 1e-6) + print(f" ✓ x modified in-place, NaN={x.isnan().any().item()}") return True except Exception as e: - print(f" ✗ FAILED: {e}") + print(f" ✗ {e}") traceback.print_exc() return False - -# ============================================================================ -# Step 11: rms_norm -# ============================================================================ -def step11_rms_norm(bridge): +def step4_linear(bridge): import torch - print("\nSTEP 11: rms_norm") - hidden = torch.randn(4, 1536, dtype=torch.float16, device="cuda") - weight = torch.ones(1536, dtype=torch.float16, device="cuda") - output = torch.empty_like(hidden) + print("\nSTEP 4: linear (ixformer GEMM)") + x = torch.randn(4, 128, dtype=torch.float16, device="cuda") + w = torch.randn(256, 128, dtype=torch.float16, device="cuda") try: - bridge.rms_norm(output, hidden, weight, 1e-6) - print(f" ✓ output: {output.shape}") - print(f" NaN: {output.isnan().any().item()}, abs_mean: {output.abs().mean().item():.4f}") + out = bridge.linear(x, w, None) + print(f" ✓ {x.shape} @ {w.shape}^T → {out.shape}, NaN={out.isnan().any().item()}") return True except Exception as e: - print(f" ✗ FAILED: {e}") + print(f" ✗ {e}") traceback.print_exc() return False - -# ============================================================================ -# Step 12: corex_moe.py Python module (tests tiered dispatch) -# ============================================================================ -def step12_corex_moe_python(): +def step5_flash_attn_python(): import torch - print("\n" + "=" * 60) - print("STEP 12: corex_moe.py Python tiered dispatch") - print("=" * 60) + print("\nSTEP 5: ixformer flash_attn (Python)") + try: + from ixformer.contrib.vllm_flash_attn import flash_attn_varlen_func + Hq, Hkv, D = 4, 1, 128 + seq = 32 + q = torch.randn(seq, Hq, D, dtype=torch.float16, device="cuda") + k = torch.randn(seq, Hkv, D, dtype=torch.float16, device="cuda") + v = torch.randn(seq, Hkv, D, dtype=torch.float16, device="cuda") + cu_q = torch.tensor([0, seq], dtype=torch.int32, device="cuda") + cu_k = torch.tensor([0, seq], dtype=torch.int32, device="cuda") + out = flash_attn_varlen_func(q, k, v, cu_q, cu_k, seq, seq, + softmax_scale=D**-0.5, causal=True) + print(f" ✓ {out.shape}, NaN={out.isnan().any().item()}") + return True + except Exception as e: + print(f" ✗ {e}") + return False - # Add project root to path +def step6_paged_attn_python(): + import torch + print("\nSTEP 6: ixformer paged_attention (Python)") + try: + import ixformer.functions as ixf_F + fn = ixf_F.vllm_single_query_cached_kv_attention + # This is the V1 paged attention used by vllm on BI-V100 + print(f" ✓ vllm_single_query_cached_kv_attention is available") + return True + except Exception as e: + print(f" ✗ {e}") + return False + +def step7_corex_moe(): + import torch + print("\nSTEP 7: corex_moe.py MoE pipeline") here = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, here) - try: - from ex_engine.python.corex_moe import moe_forward, topk_softmax - print(" ✓ corex_moe imported") + from ex_engine.python.corex_moe import moe_forward except Exception as e: print(f" ✗ Import failed: {e}") return False - num_tokens = 4 - hidden_size = 256 # small for test - num_experts = 8 # small - inter_size = 64 - topk = 2 - - hidden = torch.randn(num_tokens, hidden_size, dtype=torch.float16, device="cuda") - gate = torch.randn(num_tokens, num_experts, dtype=torch.float16, device="cuda") - w13 = torch.randn(num_experts, inter_size * 2, hidden_size, dtype=torch.float16, device="cuda") - w2 = torch.randn(num_experts, hidden_size, inter_size, dtype=torch.float16, device="cuda") - + num_tokens, hidden, experts, inter, topk = 4, 256, 8, 64, 2 + h = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") + g = torch.randn(num_tokens, experts, dtype=torch.float16, device="cuda") + w13 = torch.randn(experts, inter*2, hidden, dtype=torch.float16, device="cuda") + w2 = torch.randn(experts, hidden, inter, dtype=torch.float16, device="cuda") try: - result = moe_forward(hidden, gate, w13, w2, topk=topk, - renormalize=True, num_experts=num_experts) - print(f" ✓ moe_forward: {result.shape} {result.dtype}") - print(f" NaN: {result.isnan().any().item()}, abs_mean: {result.abs().mean().item():.4f}") + out = moe_forward(h, g, w13, w2, topk=topk, renormalize=True, num_experts=experts) + print(f" ✓ {out.shape}, NaN={out.isnan().any().item()}, abs_mean={out.abs().mean().item():.4f}") return True except Exception as e: - print(f" ✗ FAILED: {e}") + print(f" ✗ {e}") traceback.print_exc() return False - -# ============================================================================ -# Main -# ============================================================================ def main(): import torch print("=" * 60) print(" BI-V100 Single GPU Verification") - print(f" CUDA available: {torch.cuda.is_available()}") - if torch.cuda.is_available(): - print(f" Device: {torch.cuda.get_device_name(0)}") - print(f" Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") + print(f" CUDA: {torch.cuda.is_available()}, Device: {torch.cuda.get_device_name(0)}") + print(f" Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") print("=" * 60) - results = {} - - # Step 0: Compile + R = {} bridge = step0_compile_bridge() - results["compile"] = bridge is not None + R["compile"] = bridge is not None - if bridge is not None: - # Steps 1-7: Individual MoE steps - weights, ids = step1_topk_softmax(bridge) - results["topk_softmax"] = weights is not None + if bridge: + R["silu_and_mul"] = step1_silu(bridge) + R["rms_norm"] = step2_rms_norm(bridge) + R["fused_add_rms_norm"] = step3_fused_add_rms_norm(bridge) + R["linear"] = step4_linear(bridge) - if ids is not None: - idx = step2_gen_idx(bridge, ids) - results["gen_idx"] = idx is not None + R["flash_attn_python"] = step5_flash_attn_python() + R["paged_attn_python"] = step6_paged_attn_python() + R["corex_moe"] = step7_corex_moe() - if idx is not None: - expanded, hidden = step3_expand_input(bridge, idx) - results["expand_input"] = expanded is not None - - if expanded is not None: - gemm1, w13 = step4_group_gemm_w13(bridge, expanded, idx) - results["group_gemm_w13"] = gemm1 is not None - - if gemm1 is not None: - act = step5_silu_and_mul(bridge, gemm1) - results["silu_and_mul"] = act is not None - - if act is not None: - gemm2 = step6_group_gemm_w2(bridge, act, idx) - results["group_gemm_w2"] = gemm2 is not None - - if gemm2 is not None: - result = step7_combine_result(bridge, gemm2, weights) - results["combine_result"] = result is not None - - # Step 8: Full pipeline - results["fused_moe"] = step8_fused_moe(bridge) - - # Step 9-11: Other kernels - results["paged_attention"] = step9_paged_attention(bridge) - results["flash_attn"] = step10_flash_attn(bridge) - results["rms_norm"] = step11_rms_norm(bridge) - - # Step 12: Python module test (works even without bridge) - results["corex_moe_python"] = step12_corex_moe_python() - - # Summary print("\n" + "=" * 60) print(" SUMMARY") print("=" * 60) - for name, ok in results.items(): - status = "✓ PASS" if ok else "✗ FAIL" - print(f" {status} {name}") - - passed = sum(1 for v in results.values() if v) - total = len(results) - print(f"\n {passed}/{total} passed") - - if results.get("fused_moe"): - print("\n >>> MoE FULL C++ PIPELINE WORKS — comp 168 parity achieved <<<") - elif results.get("corex_moe_python"): - print("\n >>> MoE Python fallback works — C++ bridge needs debugging <<<") - - return 0 if passed == total else 1 + for k, v in R.items(): + print(f" {'✓' if v else '✗'} {k}") + p = sum(R.values()) + print(f"\n {p}/{len(R)} passed") + if R.get("compile") and R.get("silu_and_mul"): + print("\n >>> C++ bridge works — silu_and_mul/rms_norm/linear accelerated <<<") + return 0 if p == len(R) else 1 if __name__ == "__main__": sys.exit(main())