refactor(bridge): rewrite ix_full_bridge.cpp for actual base image symbols
Symbol probe revealed ixformer::infer namespace does NOT exist in base image.
That namespace is xllm's own compiled wrapper layer.
Actual available symbols in base image:
_ixformer_torch.so: silu_and_mul_forward, rms_norm_forward,
fused_add_rms_norm_forward, ixformer_linear, ixformer_linear_ex
libixformer.so: ixinfer_flash_attn_unpad_fwd (different signature)
MoE functions (topk_softmax, group_gemm, moe_expand_input, etc.)
are NOT in any base image .so — MoE must use Python path.
Bridge now only wraps: silu_and_mul, rms_norm, fused_add_rms_norm, linear
These accelerate the per-layer ops that run 200x per token.
This commit is contained in:
@@ -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 <torch/extension.h>
|
||||
#include <optional>
|
||||
#include <tuple>
|
||||
#include <vector>
|
||||
|
||||
// ============================================================================
|
||||
// Compatibility: CoreX torch uses c10::optional which may not implicitly
|
||||
// convert from kNoneTensor to const std::optional<T>&.
|
||||
// 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<torch::Tensor> kNoneTensor = {};
|
||||
static const std::optional<bool> 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<at::Tensor>& bias,
|
||||
const c10::optional<at::Tensor>& 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<at::Tensor>& 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<torch::Tensor>& alibi_slopes,
|
||||
const std::optional<torch::Tensor>& sinks,
|
||||
std::optional<torch::Tensor>& 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<torch::Tensor>& 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<torch::Tensor>& 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<torch::Tensor>& fused_bias,
|
||||
double alpha, double eps, bool is_post);
|
||||
|
||||
void rms_norm(
|
||||
torch::Tensor& input, torch::Tensor& weight, torch::Tensor& output,
|
||||
const std::optional<torch::Tensor>& 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<torch::Tensor>& bias,
|
||||
const std::optional<torch::Tensor>& out,
|
||||
const std::optional<bool> persistent);
|
||||
|
||||
torch::Tensor ixformer_linear_ex(
|
||||
torch::Tensor& input, torch::Tensor& weight,
|
||||
const std::optional<torch::Tensor>& bias,
|
||||
const std::optional<torch::Tensor>& 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<torch::Tensor>& expert_mask,
|
||||
const std::optional<torch::Tensor>& expert_sizes_cpu,
|
||||
const std::optional<torch::Tensor>& 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<torch::Tensor>& 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<torch::Tensor>& dst_to_src,
|
||||
const std::optional<torch::Tensor>& 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<torch::Tensor>& mul_weight,
|
||||
const std::optional<torch::Tensor>& mask,
|
||||
const std::optional<torch::Tensor>& 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<torch::Tensor, torch::Tensor> 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<torch::Tensor> 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<torch::Tensor>& 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<torch::Tensor>& 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<at::Tensor>());
|
||||
}
|
||||
|
||||
// --- 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<torch::Tensor> 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<torch::Tensor>& 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)");
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user