From 42dbfd161e28504bec9482f19aafc1a61a4802bd Mon Sep 17 00:00:00 2001 From: tangchengxiang <2064027004@qq.com> Date: Thu, 17 Sep 2026 08:37:17 +0000 Subject: [PATCH 1/2] feat(mamba2): add indexed scan kernels for NVIDIA and MetaX Add the descriptor, workspace and graph-aware operator interfaces for packed prefill and single-step updates. Share device scan kernels and verify outputs plus the entire state pool against an independent recurrence. --- include/infinicore/ops.hpp | 1 + include/infinicore/ops/mamba2_scan.hpp | 11 + include/infiniop.h | 1 + include/infiniop/ops/mamba2_scan.h | 25 +++ python/infinicore/nn/functional/__init__.py | 2 + .../infinicore/nn/functional/mamba2_scan.py | 43 ++++ src/infinicore/ops/mamba2_scan/mamba2_scan.cc | 21 ++ .../ops/mamba2_scan/mamba2_scan_infiniop.cc | 25 +++ src/infinicore/pybind11/ops.hpp | 2 + src/infinicore/pybind11/ops/mamba2_scan.hpp | 8 + src/infiniop/ops/mamba2_scan/cuda/kernel.cuh | 141 ++++++++++++ src/infiniop/ops/mamba2_scan/cuda/launch.cuh | 36 ++++ src/infiniop/ops/mamba2_scan/info.h | 109 ++++++++++ src/infiniop/ops/mamba2_scan/mamba2_scan.h | 18 ++ .../ops/mamba2_scan/metax/mamba2_scan_metax.h | 3 + .../mamba2_scan/metax/mamba2_scan_metax.maca | 56 +++++ .../mamba2_scan/nvidia/mamba2_scan_nvidia.cu | 59 +++++ .../mamba2_scan/nvidia/mamba2_scan_nvidia.cuh | 3 + src/infiniop/ops/mamba2_scan/operator.cc | 83 +++++++ test/infinicore/ops/mamba2_scan.py | 204 ++++++++++++++++++ 20 files changed, 851 insertions(+) create mode 100644 include/infinicore/ops/mamba2_scan.hpp create mode 100644 include/infiniop/ops/mamba2_scan.h create mode 100644 python/infinicore/nn/functional/mamba2_scan.py create mode 100644 src/infinicore/ops/mamba2_scan/mamba2_scan.cc create mode 100644 src/infinicore/ops/mamba2_scan/mamba2_scan_infiniop.cc create mode 100644 src/infinicore/pybind11/ops/mamba2_scan.hpp create mode 100644 src/infiniop/ops/mamba2_scan/cuda/kernel.cuh create mode 100644 src/infiniop/ops/mamba2_scan/cuda/launch.cuh create mode 100644 src/infiniop/ops/mamba2_scan/info.h create mode 100644 src/infiniop/ops/mamba2_scan/mamba2_scan.h create mode 100644 src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.h create mode 100644 src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.maca create mode 100644 src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cu create mode 100644 src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cuh create mode 100644 src/infiniop/ops/mamba2_scan/operator.cc create mode 100644 test/infinicore/ops/mamba2_scan.py diff --git a/include/infinicore/ops.hpp b/include/infinicore/ops.hpp index 5e93e1457..4ec17fd15 100644 --- a/include/infinicore/ops.hpp +++ b/include/infinicore/ops.hpp @@ -50,6 +50,7 @@ #include "ops/linear.hpp" #include "ops/linear_allreduce.hpp" #include "ops/linear_mxfp4.hpp" +#include "ops/mamba2_scan.hpp" #include "ops/mamba_selective_scan.hpp" #include "ops/matmul.hpp" #include "ops/moe_align.hpp" diff --git a/include/infinicore/ops/mamba2_scan.hpp b/include/infinicore/ops/mamba2_scan.hpp new file mode 100644 index 000000000..e671aaa3c --- /dev/null +++ b/include/infinicore/ops/mamba2_scan.hpp @@ -0,0 +1,11 @@ +#pragma once +#include "../graph/graph.hpp" +#include "common/op.hpp" + +namespace infinicore::op { +INFINICORE_GRAPH_OP_CLASS(Mamba2Scan, Tensor, const Tensor &, const Tensor &, const Tensor &, const Tensor &, const Tensor &, const Tensor &, const Tensor &, Tensor, const Tensor &, const Tensor &, const Tensor &); + +// Packed Mamba-2 scan. State row zero is read-only; final rows must be unique. +__export Tensor mamba2_scan(const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices); +__export void mamba2_scan_(Tensor out, const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices); +} // namespace infinicore::op diff --git a/include/infiniop.h b/include/infiniop.h index 9f632e27f..4df85b31b 100644 --- a/include/infiniop.h +++ b/include/infiniop.h @@ -94,6 +94,7 @@ #include "infiniop/ops/logcumsumexp.h" #include "infiniop/ops/logdet.h" #include "infiniop/ops/lp_norm.h" +#include "infiniop/ops/mamba2_scan.h" #include "infiniop/ops/mamba_selective_scan.h" #include "infiniop/ops/masked_select.h" #include "infiniop/ops/matmul_all_reduce.h" diff --git a/include/infiniop/ops/mamba2_scan.h b/include/infiniop/ops/mamba2_scan.h new file mode 100644 index 000000000..10cfdf6aa --- /dev/null +++ b/include/infiniop/ops/mamba2_scan.h @@ -0,0 +1,25 @@ +#ifndef INFINIOP_MAMBA2_SCAN_API_H_ +#define INFINIOP_MAMBA2_SCAN_API_H_ +#include "../operator_descriptor.h" +#include + +typedef struct InfiniopDescriptor *infiniopMamba2ScanDescriptor_t; + +// All tensors are contiguous. `out`/`x`: [tokens, heads, head_dim], `dt`: +// [tokens, heads], `b`/`c`: [tokens, groups, state_size], sharing F32/F16/BF16. +// `a`/`d`/`dt_bias`: FP32 [heads]. `a` contains the transformed -exp(A_log). +// `state`: FP32 [pool, heads, head_dim, state_size], with state_size <= 256. +// `offsets`: int32 [requests + 1], strictly increasing from zero to tokens. +// Initial/final indices are int32 [requests] and must name valid pool rows. +// Final rows are unique and nonzero. A request may update its own initial row; +// no request may read another request's final row. Row zero is read-only. +// Output, inputs, state, and workspace must not overlap in storage. Metadata +// values are caller-validated preconditions; execution does not synchronize +// the device to inspect them on the host. +__INFINI_C __export infiniStatus_t infiniopCreateMamba2ScanDescriptor( + infiniopHandle_t handle, infiniopMamba2ScanDescriptor_t *desc_ptr, infiniopTensorDescriptor_t out_desc, infiniopTensorDescriptor_t x_desc, infiniopTensorDescriptor_t dt_desc, infiniopTensorDescriptor_t b_desc, infiniopTensorDescriptor_t c_desc, infiniopTensorDescriptor_t a_desc, infiniopTensorDescriptor_t d_desc, infiniopTensorDescriptor_t dt_bias_desc, infiniopTensorDescriptor_t state_desc, infiniopTensorDescriptor_t offsets_desc, infiniopTensorDescriptor_t initial_indices_desc, infiniopTensorDescriptor_t final_indices_desc); +__INFINI_C __export infiniStatus_t infiniopGetMamba2ScanWorkspaceSize(infiniopMamba2ScanDescriptor_t desc, size_t *size); +__INFINI_C __export infiniStatus_t infiniopMamba2Scan( + infiniopMamba2ScanDescriptor_t desc, void *workspace, size_t workspace_size, void *out, const void *x, const void *dt, const void *b, const void *c, const void *a, const void *d, const void *dt_bias, void *state, const void *offsets, const void *initial_indices, const void *final_indices, void *stream); +__INFINI_C __export infiniStatus_t infiniopDestroyMamba2ScanDescriptor(infiniopMamba2ScanDescriptor_t desc); +#endif diff --git a/python/infinicore/nn/functional/__init__.py b/python/infinicore/nn/functional/__init__.py index e84d5a718..b8b035b94 100644 --- a/python/infinicore/nn/functional/__init__.py +++ b/python/infinicore/nn/functional/__init__.py @@ -26,6 +26,7 @@ from .linear_mxfp4 import linear_mxfp4 from .linear_w8a8i8 import linear_w8a8i8 from .log_softmax import log_softmax +from .mamba2_scan import mamba2_scan from .mamba_selective_scan import mamba_selective_scan from .moe_fused_dense import moe_fused_dense from .multi_margin_loss import multi_margin_loss @@ -88,6 +89,7 @@ "interpolate", "log_softmax", "mamba_selective_scan", + "mamba2_scan", "moe_fused_dense", "upsample_nearest", "triplet_margin_with_distance_loss", diff --git a/python/infinicore/nn/functional/mamba2_scan.py b/python/infinicore/nn/functional/mamba2_scan.py new file mode 100644 index 000000000..4a7908f12 --- /dev/null +++ b/python/infinicore/nn/functional/mamba2_scan.py @@ -0,0 +1,43 @@ +from infinicore.lib import _infinicore +from infinicore.tensor import Tensor + + +def mamba2_scan( + x: Tensor, + dt: Tensor, + b: Tensor, + c: Tensor, + a: Tensor, + d: Tensor, + dt_bias: Tensor, + state: Tensor, + offsets: Tensor, + initial_indices: Tensor, + final_indices: Tensor, +) -> Tensor: + """Run a packed Mamba-2 scan with FP32 recurrent state. + + Inputs are contiguous device tensors. ``x`` is [tokens, heads, head_dim], + ``dt`` is [tokens, heads], and ``b``/``c`` are [tokens, groups, state_size]. + ``a``, ``d`` and ``dt_bias`` are FP32 [heads]; ``a`` contains -exp(A_log). + ``state`` is FP32 [pool, heads, head_dim, state_size]. Offsets and indices + are int32. Offsets delimit nonempty sequences and span all input tokens. + Final state indices must be distinct, nonzero valid pool rows. Initial + rows may be zero or the same request's final row; cross-request read/write + aliasing is unsupported. State row zero is never modified. + """ + return Tensor( + _infinicore.mamba2_scan( + x._underlying, + dt._underlying, + b._underlying, + c._underlying, + a._underlying, + d._underlying, + dt_bias._underlying, + state._underlying, + offsets._underlying, + initial_indices._underlying, + final_indices._underlying, + ) + ) diff --git a/src/infinicore/ops/mamba2_scan/mamba2_scan.cc b/src/infinicore/ops/mamba2_scan/mamba2_scan.cc new file mode 100644 index 000000000..6831e7063 --- /dev/null +++ b/src/infinicore/ops/mamba2_scan/mamba2_scan.cc @@ -0,0 +1,21 @@ +#include "infinicore/ops/mamba2_scan.hpp" +#include "../../utils.hpp" + +namespace infinicore::op { +INFINICORE_GRAPH_OP_DISPATCHERS_IMPL(Mamba2Scan); +Mamba2Scan::Mamba2Scan(Tensor out, const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices) { + INFINICORE_ASSERT_TENSORS_SAME_DEVICE(out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices); + INFINICORE_GRAPH_OP_DISPATCH(out->device().getType(), out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices); +} +void Mamba2Scan::execute(Tensor out, const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices) { + INFINICORE_GRAPH_OP_RECORD_OR_RUN(Mamba2Scan, out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices); +} +Tensor mamba2_scan(const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices) { + auto out = Tensor::empty(x->shape(), x->dtype(), x->device()); + mamba2_scan_(out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices); + return out; +} +void mamba2_scan_(Tensor out, const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices) { + Mamba2Scan::execute(out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices); +} +} // namespace infinicore::op diff --git a/src/infinicore/ops/mamba2_scan/mamba2_scan_infiniop.cc b/src/infinicore/ops/mamba2_scan/mamba2_scan_infiniop.cc new file mode 100644 index 000000000..1c19c2e17 --- /dev/null +++ b/src/infinicore/ops/mamba2_scan/mamba2_scan_infiniop.cc @@ -0,0 +1,25 @@ +#include "../infiniop_impl.hpp" +#include "infinicore/ops/mamba2_scan.hpp" + +namespace infinicore::op::mamba2_scan_impl::infiniop { +INFINIOP_CACHABLE_DESCRIPTOR(Descriptor, Mamba2Scan, 100); +struct PlannedMeta { + std::shared_ptr descriptor; + graph::GraphTensor workspace, out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices; +}; +void *plan(Tensor out, const Tensor &x, const Tensor &dt, const Tensor &b, const Tensor &c, const Tensor &a, const Tensor &d, const Tensor &dt_bias, Tensor state, const Tensor &offsets, const Tensor &initial_indices, const Tensor &final_indices) { + size_t seed = hash_combine(out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices); + INFINIOP_CACHABLE_DESCRIPTOR_GET_OR_CREATE(Descriptor, descriptor, Mamba2Scan, seed, out->desc(), x->desc(), dt->desc(), b->desc(), c->desc(), a->desc(), d->desc(), dt_bias->desc(), state->desc(), offsets->desc(), initial_indices->desc(), final_indices->desc()); + INFINIOP_WORKSPACE_TENSOR(workspace, Mamba2Scan, descriptor); + return new PlannedMeta{descriptor, graph::GraphTensor(workspace), graph::GraphTensor(out), graph::GraphTensor(x), graph::GraphTensor(dt), graph::GraphTensor(b), graph::GraphTensor(c), graph::GraphTensor(a), graph::GraphTensor(d), graph::GraphTensor(dt_bias), graph::GraphTensor(state), graph::GraphTensor(offsets), graph::GraphTensor(initial_indices), graph::GraphTensor(final_indices)}; +} +void run(void *planned_meta) { + auto *p = reinterpret_cast(planned_meta); + INFINICORE_CHECK_ERROR(infiniopMamba2Scan(p->descriptor->desc, p->workspace->data(), p->workspace->numel(), p->out->data(), p->x->data(), p->dt->data(), p->b->data(), p->c->data(), p->a->data(), p->d->data(), p->dt_bias->data(), p->state->data(), p->offsets->data(), p->initial_indices->data(), p->final_indices->data(), context::getStream())); +} +void cleanup(void **planned_meta) { + delete *reinterpret_cast(planned_meta); + *planned_meta = nullptr; +} +INFINICORE_GRAPH_OP_REGISTER_ALLDEVICE(Mamba2Scan, &plan, &run, &cleanup); +} // namespace infinicore::op::mamba2_scan_impl::infiniop diff --git a/src/infinicore/pybind11/ops.hpp b/src/infinicore/pybind11/ops.hpp index e78adf275..3a2e6b3f3 100644 --- a/src/infinicore/pybind11/ops.hpp +++ b/src/infinicore/pybind11/ops.hpp @@ -86,6 +86,7 @@ #include "ops/logdet.hpp" #include "ops/logical_and.hpp" #include "ops/logical_not.hpp" +#include "ops/mamba2_scan.hpp" #include "ops/mamba_selective_scan.hpp" #include "ops/masked_select.hpp" #include "ops/matmul.hpp" @@ -252,6 +253,7 @@ inline void bind(py::module &m) { bind_logdet(m); bind_matmul(m); bind_mamba_selective_scan(m); + bind_mamba2_scan(m); bind_kron(m); bind_mul(m); bind_mul_scalar(m); diff --git a/src/infinicore/pybind11/ops/mamba2_scan.hpp b/src/infinicore/pybind11/ops/mamba2_scan.hpp new file mode 100644 index 000000000..82ccedc25 --- /dev/null +++ b/src/infinicore/pybind11/ops/mamba2_scan.hpp @@ -0,0 +1,8 @@ +#pragma once +#include "infinicore/ops/mamba2_scan.hpp" +#include +namespace infinicore::ops { +inline void bind_mamba2_scan(pybind11::module &m) { + m.def("mamba2_scan", &op::mamba2_scan, pybind11::arg("x"), pybind11::arg("dt"), pybind11::arg("b"), pybind11::arg("c"), pybind11::arg("a"), pybind11::arg("d"), pybind11::arg("dt_bias"), pybind11::arg("state"), pybind11::arg("offsets"), pybind11::arg("initial_indices"), pybind11::arg("final_indices")); +} +} // namespace infinicore::ops diff --git a/src/infiniop/ops/mamba2_scan/cuda/kernel.cuh b/src/infiniop/ops/mamba2_scan/cuda/kernel.cuh new file mode 100644 index 000000000..b2174e83d --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/cuda/kernel.cuh @@ -0,0 +1,141 @@ +#pragma once + +#include "../info.h" + +namespace op::mamba2_scan::cuda { + +constexpr int kWarpSize = 32; +constexpr int kWarpsPerBlock = 4; + +__device__ inline float softplus(float value) { + return value > 20.0f ? value : log1pf(expf(value)); +} + +template +static __global__ void coefficients(float *coeff, const T *dt, const float *a, + const float *dt_bias, size_t tokens, size_t heads) { + const size_t i = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (i < tokens * heads) { + const float step = softplus(static_cast(dt[i]) + dt_bias[i % heads]); + coeff[2 * i] = step; + coeff[2 * i + 1] = expf(step * a[i % heads]); + } +} + +// Each logical 32-lane group owns one channel, including on 64-lane devices. +// Summaries encode an affine state transform for each independent chunk. +template +static __global__ void scan(T *out, const T *x, const T *dt, const T *b, const T *c, + const float *a, const float *d, const float *dt_bias, + float *state, const int32_t *offsets, const int32_t *initial, + const int32_t *final, const float *coeff, float *chunks, + float *decays, Mamba2ScanInfo info) { + const size_t request = blockIdx.y; + const size_t p_blocks = (info.head_dim + kWarpsPerBlock - 1) / kWarpsPerBlock; + const size_t h = blockIdx.x / p_blocks; + const size_t p = (blockIdx.x % p_blocks) * kWarpsPerBlock + threadIdx.x / kWarpSize; + const size_t lane = threadIdx.x % kWarpSize; + if (p >= info.head_dim) { + return; + } + const int32_t start = offsets[request], end = offsets[request + 1]; + const int32_t src = initial[request], dst = final[request]; + if (start < 0 || end <= start || static_cast(end) > info.tokens + || src < 0 || static_cast(src) >= info.pool_size + || dst <= 0 || static_cast(dst) >= info.pool_size) { + return; + } + const size_t first = start + (Chunked ? blockIdx.z * info.chunk_size : 0); + if (first >= static_cast(end)) { + return; + } + const size_t last = Chunked ? min(first + info.chunk_size, static_cast(end)) : end; + const size_t slot = start / info.chunk_size + request + blockIdx.z; + const size_t local_state = (h * info.head_dim + p) * info.state_size; + const size_t chunk_base = slot * (info.heads * info.head_dim * info.state_size) + local_state; + const size_t source_base = static_cast(src) * (info.heads * info.head_dim * info.state_size) + local_state; + float values[8]; + for (size_t j = 0; j < 8; ++j) { + const size_t n = lane + j * kWarpSize; + values[j] = n < info.state_size && !Summary + ? (Chunked ? chunks[chunk_base + n] : state[source_base + n]) + : 0.0f; + } + float decay = 1.0f; + const size_t group = h / (info.heads / info.groups); + for (size_t t = first; t < last; ++t) { + const size_t time_head = t * info.heads + h; + const float step = Chunked ? coeff[2 * time_head] + : softplus(static_cast(dt[time_head]) + dt_bias[h]); + const float alpha = Chunked ? coeff[2 * time_head + 1] : expf(step * a[h]); + const size_t x_index = time_head * info.head_dim + p; + const float input = static_cast(x[x_index]); + const size_t bc_base = (t * info.groups + group) * info.state_size; + float y = 0.0f; + for (size_t j = 0; j < 8; ++j) { + const size_t n = lane + j * kWarpSize; + if (n < info.state_size) { + values[j] = fmaf(alpha, values[j], step * input * static_cast(b[bc_base + n])); + if constexpr (!Summary) { + y = fmaf(values[j], static_cast(c[bc_base + n]), y); + } + } + } + if constexpr (Summary) { + decay *= alpha; + } else { + for (int delta = 16; delta > 0; delta /= 2) { + y += __shfl_down_sync(__activemask(), y, delta, kWarpSize); + } + if (lane == 0) { + out[x_index] = static_cast(y + d[h] * input); + } + } + } + if constexpr (Summary || !Chunked) { + const size_t target_base = Summary ? chunk_base : static_cast(dst) * (info.heads * info.head_dim * info.state_size) + local_state; + for (size_t j = 0; j < 8; ++j) { + const size_t n = lane + j * kWarpSize; + if (n < info.state_size) { + (Summary ? chunks : state)[target_base + n] = values[j]; + } + } + if constexpr (Summary) { + if (p == 0 && lane == 0) { + decays[slot * info.heads + h] = decay; + } + } + } +} + +// Convert chunk summaries to initial states, then publish the request's final state. +static __global__ void carry(float *chunks, const float *decays, float *state, + const int32_t *offsets, const int32_t *initial, + const int32_t *final, Mamba2ScanInfo info) { + const size_t element = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const size_t request = blockIdx.y; + if (element >= (info.heads * info.head_dim * info.state_size)) { + return; + } + const int32_t start = offsets[request], end = offsets[request + 1]; + const int32_t src = initial[request], dst = final[request]; + if (start < 0 || end <= start || static_cast(end) > info.tokens + || src < 0 || static_cast(src) >= info.pool_size + || dst <= 0 || static_cast(dst) >= info.pool_size) { + return; + } + const size_t h = element / (info.head_dim * info.state_size); + const size_t count = (end - start + info.chunk_size - 1) / info.chunk_size; + const size_t first_slot = start / info.chunk_size + request; + float value = state[static_cast(src) * (info.heads * info.head_dim * info.state_size) + element]; + for (size_t chunk = 0; chunk < count; ++chunk) { + const size_t slot = first_slot + chunk; + const size_t i = slot * (info.heads * info.head_dim * info.state_size) + element; + const float summary = chunks[i]; + chunks[i] = value; + value = fmaf(decays[slot * info.heads + h], value, summary); + } + state[static_cast(dst) * (info.heads * info.head_dim * info.state_size) + element] = value; +} + +} // namespace op::mamba2_scan::cuda diff --git a/src/infiniop/ops/mamba2_scan/cuda/launch.cuh b/src/infiniop/ops/mamba2_scan/cuda/launch.cuh new file mode 100644 index 000000000..e95d374fd --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/cuda/launch.cuh @@ -0,0 +1,36 @@ +#pragma once + +#include "kernel.cuh" + +namespace op::mamba2_scan::cuda { + +template +void launch(const Mamba2ScanInfo &info, void *workspace, void *out, + const void *x, const void *dt, const void *b, const void *c, + const void *a, const void *d, const void *dt_bias, void *state, + const void *offsets, const void *initial, const void *final, + Stream stream) { + const dim3 blocks(info.heads * ((info.head_dim + kWarpsPerBlock - 1) / kWarpsPerBlock), info.requests); + constexpr int threads = kWarpSize * kWarpsPerBlock; + auto *out_ptr = static_cast(out); + const auto *x_ptr = static_cast(x), *dt_ptr = static_cast(dt); + const auto *b_ptr = static_cast(b), *c_ptr = static_cast(c); + const auto *a_ptr = static_cast(a), *d_ptr = static_cast(d); + const auto *bias_ptr = static_cast(dt_bias); + auto *state_ptr = static_cast(state); + const auto *offset_ptr = static_cast(offsets); + const auto *init_ptr = static_cast(initial), *final_ptr = static_cast(final); + if (info.single_chunk()) { + scan<<>>(out_ptr, x_ptr, dt_ptr, b_ptr, c_ptr, a_ptr, d_ptr, bias_ptr, state_ptr, offset_ptr, init_ptr, final_ptr, nullptr, nullptr, nullptr, info); + } else { + auto *coeff = static_cast(workspace); + auto *chunks = coeff + 2 * info.tokens * info.heads; + auto *decays = chunks + info.chunk_slots() * info.state_elements(); + coefficients<<<(info.tokens * info.heads + 255) / 256, 256, 0, stream>>>(coeff, dt_ptr, a_ptr, bias_ptr, info.tokens, info.heads); + const dim3 chunk_blocks(blocks.x, blocks.y, info.max_chunks()); + scan<<>>(out_ptr, x_ptr, dt_ptr, b_ptr, c_ptr, a_ptr, d_ptr, bias_ptr, state_ptr, offset_ptr, init_ptr, final_ptr, coeff, chunks, decays, info); + carry<<>>(chunks, decays, state_ptr, offset_ptr, init_ptr, final_ptr, info); + scan<<>>(out_ptr, x_ptr, dt_ptr, b_ptr, c_ptr, a_ptr, d_ptr, bias_ptr, state_ptr, offset_ptr, init_ptr, final_ptr, coeff, chunks, decays, info); + } +} +} // namespace op::mamba2_scan::cuda diff --git a/src/infiniop/ops/mamba2_scan/info.h b/src/infiniop/ops/mamba2_scan/info.h new file mode 100644 index 000000000..202b9806b --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/info.h @@ -0,0 +1,109 @@ +#pragma once + +#include "../../../utils.h" +#include "../../tensor.h" +#include + +namespace op::mamba2_scan { + +struct Mamba2ScanInfo { + infiniDtype_t dtype; + size_t tokens, heads, head_dim, groups, state_size, pool_size, requests; + static constexpr size_t chunk_size = 256; + + bool single_chunk() const { return tokens <= chunk_size || tokens == requests; } + size_t max_chunks() const { return (tokens + chunk_size - 1) / chunk_size; } + size_t chunk_slots() const { return max_chunks() + requests; } + size_t state_elements() const { return heads * head_dim * state_size; } + size_t workspace_bytes() const { + if (single_chunk()) { + return 0; + } + return sizeof(float) * (2 * tokens * heads + chunk_slots() * (state_elements() + heads)); + } + + static utils::Result create( + infiniopTensorDescriptor_t out, infiniopTensorDescriptor_t x, + infiniopTensorDescriptor_t dt, infiniopTensorDescriptor_t b, + infiniopTensorDescriptor_t c, infiniopTensorDescriptor_t a, + infiniopTensorDescriptor_t d, infiniopTensorDescriptor_t dt_bias, + infiniopTensorDescriptor_t state, infiniopTensorDescriptor_t offsets, + infiniopTensorDescriptor_t initial_indices, infiniopTensorDescriptor_t final_indices) { + for (auto tensor : {out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices}) { + if (tensor == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (!tensor->isContiguous()) { + return INFINI_STATUS_BAD_TENSOR_STRIDES; + } + } + const auto dtype = x->dtype(); + CHECK_DTYPE(dtype, INFINI_DTYPE_F16, INFINI_DTYPE_BF16, INFINI_DTYPE_F32); + for (auto tensor : {out, dt, b, c}) { + if (tensor->dtype() != dtype) { + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } + } + for (auto tensor : {a, d, dt_bias, state}) { + if (tensor->dtype() != INFINI_DTYPE_F32) { + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } + } + for (auto tensor : {offsets, initial_indices, final_indices}) { + if (tensor->dtype() != INFINI_DTYPE_I32) { + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } + if (tensor->ndim() != 1) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + } + if (x->ndim() != 3 || dt->ndim() != 2 || b->ndim() != 3 || state->ndim() != 4 + || offsets->dim(0) < 2) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + const size_t tokens = x->dim(0), heads = x->dim(1), head_dim = x->dim(2); + const size_t groups = b->dim(1), state_size = b->dim(2); + const size_t requests = offsets->dim(0) - 1; + if (tokens == 0 || heads == 0 || head_dim == 0 || groups == 0 || state_size == 0 + || state_size > 256 || heads % groups != 0 || state->dim(0) < 2 + || requests > tokens || requests >= state->dim(0) + || tokens > static_cast(std::numeric_limits::max()) + || requests > 65535 || (tokens + chunk_size - 1) / chunk_size > 65535) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + if (out->shape() != x->shape() || c->shape() != b->shape() || b->dim(0) != tokens + || dt->shape() != std::vector{tokens, heads} + || a->shape() != std::vector{heads} || d->shape() != a->shape() + || dt_bias->shape() != a->shape() + || state->shape() != std::vector{state->dim(0), heads, head_dim, state_size} + || initial_indices->shape() != std::vector{requests} + || final_indices->shape() != initial_indices->shape()) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + // Bound address arithmetic and launch dimensions before multiplying sizes. + const size_t max_elements = std::numeric_limits::max() / sizeof(float); + if (heads > max_elements / head_dim + || heads * head_dim > max_elements / state_size + || heads * head_dim > max_elements / tokens + || groups > max_elements / state_size / tokens) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + Mamba2ScanInfo info{dtype, tokens, heads, head_dim, groups, state_size, state->dim(0), requests}; + const size_t state_elements = info.state_elements(); + const size_t max_grid_x = std::numeric_limits::max(); + if (state_elements > max_elements / info.pool_size + || heads > max_grid_x / ((head_dim - 1) / 4 + 1) + || (tokens * heads - 1) / 256 + 1 > max_grid_x + || (state_elements - 1) / 256 + 1 > max_grid_x) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + if (!info.single_chunk() + && (tokens > max_elements / heads / 2 + || info.chunk_slots() > (max_elements - 2 * tokens * heads) / (state_elements + heads))) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + return utils::Result(info); + } +}; + +} // namespace op::mamba2_scan diff --git a/src/infiniop/ops/mamba2_scan/mamba2_scan.h b/src/infiniop/ops/mamba2_scan/mamba2_scan.h new file mode 100644 index 000000000..7fffc838b --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/mamba2_scan.h @@ -0,0 +1,18 @@ +#pragma once +#include "../../operator.h" +#include "info.h" + +#define DESCRIPTOR(NAMESPACE) \ + namespace op::mamba2_scan::NAMESPACE { \ + class Descriptor final : public InfiniopDescriptor { \ + Mamba2ScanInfo _info; \ + size_t _workspace_size; \ + Descriptor(Mamba2ScanInfo info, size_t workspace_size, infiniDevice_t device_type, int device_id) \ + : InfiniopDescriptor{device_type, device_id}, _info(info), _workspace_size(workspace_size) {} \ + \ + public: \ + size_t workspaceSize() const { return _workspace_size; } \ + static infiniStatus_t create(infiniopHandle_t handle, Descriptor **desc_ptr, infiniopTensorDescriptor_t out_desc, infiniopTensorDescriptor_t x_desc, infiniopTensorDescriptor_t dt_desc, infiniopTensorDescriptor_t b_desc, infiniopTensorDescriptor_t c_desc, infiniopTensorDescriptor_t a_desc, infiniopTensorDescriptor_t d_desc, infiniopTensorDescriptor_t dt_bias_desc, infiniopTensorDescriptor_t state_desc, infiniopTensorDescriptor_t offsets_desc, infiniopTensorDescriptor_t initial_indices_desc, infiniopTensorDescriptor_t final_indices_desc); \ + infiniStatus_t calculate(void *workspace, size_t workspace_size, void *out, const void *x, const void *dt, const void *b, const void *c, const void *a, const void *d, const void *dt_bias, void *state, const void *offsets, const void *initial_indices, const void *final_indices, void *stream) const; \ + }; \ + } diff --git a/src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.h b/src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.h new file mode 100644 index 000000000..9ff5ce178 --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.h @@ -0,0 +1,3 @@ +#pragma once +#include "../mamba2_scan.h" +DESCRIPTOR(metax) diff --git a/src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.maca b/src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.maca new file mode 100644 index 000000000..8402394be --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/metax/mamba2_scan_metax.maca @@ -0,0 +1,56 @@ +#include "../../../devices/metax/metax_common.h" +#include "../../../devices/metax/metax_kernel_common.h" +#include "../cuda/launch.cuh" +#include "mamba2_scan_metax.h" + +namespace op::mamba2_scan::metax { +infiniStatus_t Descriptor::create( + infiniopHandle_t handle, Descriptor **desc_ptr, infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t x_desc, infiniopTensorDescriptor_t dt_desc, + infiniopTensorDescriptor_t b_desc, infiniopTensorDescriptor_t c_desc, + infiniopTensorDescriptor_t a_desc, infiniopTensorDescriptor_t d_desc, + infiniopTensorDescriptor_t dt_bias_desc, infiniopTensorDescriptor_t state_desc, + infiniopTensorDescriptor_t offsets_desc, infiniopTensorDescriptor_t initial_indices_desc, + infiniopTensorDescriptor_t final_indices_desc) { + auto result = Mamba2ScanInfo::create(out_desc, x_desc, dt_desc, b_desc, c_desc, a_desc, d_desc, dt_bias_desc, state_desc, offsets_desc, initial_indices_desc, final_indices_desc); + CHECK_RESULT(result); + auto info = result.take(); + *desc_ptr = new Descriptor(info, info.workspace_bytes(), handle->device, handle->device_id); + return INFINI_STATUS_SUCCESS; +} + +infiniStatus_t Descriptor::calculate(void *workspace, size_t workspace_size, void *out, + const void *x, const void *dt, const void *b, const void *c, + const void *a, const void *d, const void *dt_bias, void *state, + const void *offsets, const void *initial_indices, + const void *final_indices, void *stream) const { + if (workspace_size < _workspace_size) { + return INFINI_STATUS_INSUFFICIENT_WORKSPACE; + } + if (_workspace_size && workspace == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + for (const void *pointer : {static_cast(out), x, dt, b, c, a, d, dt_bias, static_cast(state), offsets, initial_indices, final_indices}) { + if (pointer == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + } +#define LAUNCH(T) cuda::launch(_info, workspace, out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices, static_cast(stream)) + switch (_info.dtype) { + case INFINI_DTYPE_F32: + LAUNCH(float); + break; + case INFINI_DTYPE_F16: + LAUNCH(half); + break; + case INFINI_DTYPE_BF16: + LAUNCH(__nv_bfloat16); + break; + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +#undef LAUNCH + CHECK_METAX(hcGetLastError()); + return INFINI_STATUS_SUCCESS; +} +} // namespace op::mamba2_scan::metax diff --git a/src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cu b/src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cu new file mode 100644 index 000000000..290021c9d --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cu @@ -0,0 +1,59 @@ +#include "../../../devices/nvidia/nvidia_common.cuh" +#include "../../../devices/nvidia/nvidia_kernel_common.cuh" +#include "../cuda/launch.cuh" +#include "mamba2_scan_nvidia.cuh" +#include +#include +#include + +namespace op::mamba2_scan::nvidia { +infiniStatus_t Descriptor::create( + infiniopHandle_t handle, Descriptor **desc_ptr, infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t x_desc, infiniopTensorDescriptor_t dt_desc, + infiniopTensorDescriptor_t b_desc, infiniopTensorDescriptor_t c_desc, + infiniopTensorDescriptor_t a_desc, infiniopTensorDescriptor_t d_desc, + infiniopTensorDescriptor_t dt_bias_desc, infiniopTensorDescriptor_t state_desc, + infiniopTensorDescriptor_t offsets_desc, infiniopTensorDescriptor_t initial_indices_desc, + infiniopTensorDescriptor_t final_indices_desc) { + auto result = Mamba2ScanInfo::create(out_desc, x_desc, dt_desc, b_desc, c_desc, a_desc, d_desc, dt_bias_desc, state_desc, offsets_desc, initial_indices_desc, final_indices_desc); + CHECK_RESULT(result); + auto info = result.take(); + *desc_ptr = new Descriptor(info, info.workspace_bytes(), handle->device, handle->device_id); + return INFINI_STATUS_SUCCESS; +} + +infiniStatus_t Descriptor::calculate(void *workspace, size_t workspace_size, void *out, + const void *x, const void *dt, const void *b, const void *c, + const void *a, const void *d, const void *dt_bias, void *state, + const void *offsets, const void *initial_indices, + const void *final_indices, void *stream) const { + if (workspace_size < _workspace_size) { + return INFINI_STATUS_INSUFFICIENT_WORKSPACE; + } + if (_workspace_size && workspace == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + for (const void *pointer : {static_cast(out), x, dt, b, c, a, d, dt_bias, static_cast(state), offsets, initial_indices, final_indices}) { + if (pointer == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + } +#define LAUNCH(T) cuda::launch(_info, workspace, out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices, static_cast(stream)) + switch (_info.dtype) { + case INFINI_DTYPE_F32: + LAUNCH(float); + break; + case INFINI_DTYPE_F16: + LAUNCH(half); + break; + case INFINI_DTYPE_BF16: + LAUNCH(__nv_bfloat16); + break; + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +#undef LAUNCH + CHECK_CUDA(cudaGetLastError()); + return INFINI_STATUS_SUCCESS; +} +} // namespace op::mamba2_scan::nvidia diff --git a/src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cuh b/src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cuh new file mode 100644 index 000000000..8047afced --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/nvidia/mamba2_scan_nvidia.cuh @@ -0,0 +1,3 @@ +#pragma once +#include "../mamba2_scan.h" +DESCRIPTOR(nvidia) diff --git a/src/infiniop/ops/mamba2_scan/operator.cc b/src/infiniop/ops/mamba2_scan/operator.cc new file mode 100644 index 000000000..744e0ee05 --- /dev/null +++ b/src/infiniop/ops/mamba2_scan/operator.cc @@ -0,0 +1,83 @@ +#include "../../operator.h" +#include "../../handle.h" +#include "infiniop/ops/mamba2_scan.h" +#ifdef ENABLE_NVIDIA_API +#include "nvidia/mamba2_scan_nvidia.cuh" +#endif +#ifdef ENABLE_METAX_API +#include "metax/mamba2_scan_metax.h" +#endif + +__INFINI_C infiniStatus_t infiniopCreateMamba2ScanDescriptor( + infiniopHandle_t handle, infiniopMamba2ScanDescriptor_t *desc_ptr, infiniopTensorDescriptor_t out_desc, infiniopTensorDescriptor_t x_desc, infiniopTensorDescriptor_t dt_desc, infiniopTensorDescriptor_t b_desc, infiniopTensorDescriptor_t c_desc, infiniopTensorDescriptor_t a_desc, infiniopTensorDescriptor_t d_desc, infiniopTensorDescriptor_t dt_bias_desc, infiniopTensorDescriptor_t state_desc, infiniopTensorDescriptor_t offsets_desc, infiniopTensorDescriptor_t initial_indices_desc, infiniopTensorDescriptor_t final_indices_desc) { + if (handle == nullptr || desc_ptr == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + switch (handle->device) { +#ifdef ENABLE_NVIDIA_API + case INFINI_DEVICE_NVIDIA: + return op::mamba2_scan::nvidia::Descriptor::create(handle, reinterpret_cast(desc_ptr), out_desc, x_desc, dt_desc, b_desc, c_desc, a_desc, d_desc, dt_bias_desc, state_desc, offsets_desc, initial_indices_desc, final_indices_desc); +#endif +#ifdef ENABLE_METAX_API + case INFINI_DEVICE_METAX: + return op::mamba2_scan::metax::Descriptor::create(handle, reinterpret_cast(desc_ptr), out_desc, x_desc, dt_desc, b_desc, c_desc, a_desc, d_desc, dt_bias_desc, state_desc, offsets_desc, initial_indices_desc, final_indices_desc); +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } +} +__INFINI_C infiniStatus_t infiniopGetMamba2ScanWorkspaceSize(infiniopMamba2ScanDescriptor_t desc, size_t *size) { + if (desc == nullptr || size == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + switch (desc->device_type) { +#ifdef ENABLE_NVIDIA_API + case INFINI_DEVICE_NVIDIA: + *size = reinterpret_cast(desc)->workspaceSize(); + return INFINI_STATUS_SUCCESS; +#endif +#ifdef ENABLE_METAX_API + case INFINI_DEVICE_METAX: + *size = reinterpret_cast(desc)->workspaceSize(); + return INFINI_STATUS_SUCCESS; +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } +} +__INFINI_C infiniStatus_t infiniopMamba2Scan(infiniopMamba2ScanDescriptor_t desc, void *workspace, size_t workspace_size, void *out, const void *x, const void *dt, const void *b, const void *c, const void *a, const void *d, const void *dt_bias, void *state, const void *offsets, const void *initial_indices, const void *final_indices, void *stream) { + if (desc == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + switch (desc->device_type) { +#ifdef ENABLE_NVIDIA_API + case INFINI_DEVICE_NVIDIA: + return reinterpret_cast(desc)->calculate(workspace, workspace_size, out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices, stream); +#endif +#ifdef ENABLE_METAX_API + case INFINI_DEVICE_METAX: + return reinterpret_cast(desc)->calculate(workspace, workspace_size, out, x, dt, b, c, a, d, dt_bias, state, offsets, initial_indices, final_indices, stream); +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } +} +__INFINI_C infiniStatus_t infiniopDestroyMamba2ScanDescriptor(infiniopMamba2ScanDescriptor_t desc) { + if (desc == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + switch (desc->device_type) { +#ifdef ENABLE_NVIDIA_API + case INFINI_DEVICE_NVIDIA: + delete reinterpret_cast(desc); + return INFINI_STATUS_SUCCESS; +#endif +#ifdef ENABLE_METAX_API + case INFINI_DEVICE_METAX: + delete reinterpret_cast(desc); + return INFINI_STATUS_SUCCESS; +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } +} diff --git a/test/infinicore/ops/mamba2_scan.py b/test/infinicore/ops/mamba2_scan.py new file mode 100644 index 000000000..1ec282b5e --- /dev/null +++ b/test/infinicore/ops/mamba2_scan.py @@ -0,0 +1,204 @@ +"""Compare packed Mamba-2 outputs and state against an independent recurrence.""" + +import os +import sys + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +import pytest +import torch +from framework import ( + BaseOperatorTest, + GenericTestRunner, + TensorInitializer, + TensorSpec, +) +from framework import ( + TestCase as OperatorTestCase, +) + +import infinicore + + +def torch_mamba2_scan(x, dt, b, c, a, d, dt_bias, state, offsets, initial, final): + """Use token-by-token FP32 recurrence, independent of the chunk algorithm.""" + output = torch.empty_like(x) + boundaries = offsets.cpu().tolist() + source_rows = initial.cpu().tolist() + target_rows = final.cpu().tolist() + heads_per_group = x.shape[1] // b.shape[1] + source = state.clone() + for request, (start, end) in enumerate(zip(boundaries, boundaries[1:])): + current = source[source_rows[request]].clone() + for token in range(start, end): + step = torch.nn.functional.softplus(dt[token].float() + dt_bias) + decay = torch.exp(step * a) + bt = b[token].float().repeat_interleave(heads_per_group, dim=0) + ct = c[token].float().repeat_interleave(heads_per_group, dim=0) + xt = x[token].float() + current = ( + decay[:, None, None] * current + + step[:, None, None] * xt[:, :, None] * bt[:, None, :] + ) + output[token] = (current * ct[:, None, :]).sum(-1) + d[:, None] * xt + state[target_rows[request]] = current + return output, state + + +def _spec(tensor, dtype): + return TensorSpec.from_tensor( + tuple(tensor.shape), + None, + dtype, + init_mode=TensorInitializer.MANUAL, + set_tensor=tensor, + ) + + +def parse_test_cases(): + cases = [] + shapes = [ + ([length], 4, 7, 2, 9) + for length in (1, 2, 3, 4, 5, 255, 256, 257, 511, 512, 513) + ] + shapes += [([1, 3, 257], 4, 7, 1, 33), ([257, 2, 256], 4, 7, 2, 128)] + shapes += [([1], 24, 64, 1, 128), ([257], 24, 64, 1, 128)] + shapes += [([1, 1, 1, 1], 24, 64, 1, 128), ([5], 4, 7, 2, 256)] + shapes += [([1025], 4, 7, 2, 33)] + generator = torch.Generator().manual_seed(20260916) + for lengths, heads, head_dim, groups, state_size in shapes: + tokens, requests = sum(lengths), len(lengths) + pool = 2 * requests + 2 + offsets = torch.tensor([0] + lengths, dtype=torch.int32).cumsum(0).int() + for torch_dtype, infini_dtype, tolerance in ( + (torch.float32, infinicore.float32, 1e-4), + (torch.float16, infinicore.float16, 3e-3), + (torch.bfloat16, infinicore.bfloat16, 2e-2), + ): + + def random(shape): + return (torch.randn(shape, generator=generator) * 0.2).to(torch_dtype) + + state = ( + torch.randn(pool, heads, head_dim, state_size, generator=generator) + * 0.1 + ) + state[0].zero_() + initial = torch.arange(1, requests + 1, dtype=torch.int32) + initial[0] = 0 + # Nonzero requests update their own slots; the first uses a distinct slot. + final = torch.arange(1, requests + 1, dtype=torch.int32) + final[0] = pool - 1 + tensors = [ + _spec(random((tokens, heads, head_dim)), infini_dtype), + _spec(random((tokens, heads)), infini_dtype), + _spec(random((tokens, groups, state_size)), infini_dtype), + _spec(random((tokens, groups, state_size)), infini_dtype), + _spec( + -torch.arange(1, heads + 1, dtype=torch.float32), infinicore.float32 + ), + _spec(torch.ones(heads), infinicore.float32), + _spec( + torch.linspace(-80, 80, heads) + if lengths == [1025] + else torch.full((heads,), -3.0), + infinicore.float32, + ), + _spec(state, infinicore.float32), + _spec(offsets, infinicore.int32), + _spec(initial, infinicore.int32), + _spec(final, infinicore.int32), + ] + cases.append( + OperatorTestCase( + inputs=tensors, + kwargs={}, + output_spec=None, + comparison_target=[7], + tolerance={"atol": tolerance, "rtol": tolerance}, + description=f"Mamba2Scan lengths={lengths}, groups={groups}: output and entire state pool", + output_count=2, + ) + ) + return cases + + +class OpTest(BaseOperatorTest): + def __init__(self): + super().__init__("Mamba2Scan") + + def get_test_cases(self): + return parse_test_cases() + + def torch_operator(self, *args, **kwargs): + return torch_mamba2_scan(*args) + + def infinicore_operator(self, *args, **kwargs): + output = infinicore.nn.functional.mamba2_scan(*args) + return output, args[7] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="A CUDA device is required.") +@pytest.mark.parametrize( + "invalid", + [ + "activation_dtype", + "parameter_dtype", + "state_dtype", + "offset_dtype", + "index_count", + "bc_shape", + "zero_only_pool", + "noncontiguous", + "head_groups", + ], +) +def test_invalid_descriptor_is_rejected(invalid): + args = [ + torch.zeros(4, 4, 8, device="cuda"), + torch.zeros(4, 4, device="cuda"), + torch.zeros(4, 2, 16, device="cuda"), + torch.zeros(4, 2, 16, device="cuda"), + -torch.ones(4, device="cuda"), + torch.ones(4, device="cuda"), + torch.zeros(4, device="cuda"), + torch.zeros(4, 4, 8, 16, device="cuda"), + torch.tensor([0, 4], dtype=torch.int32, device="cuda"), + torch.tensor([0], dtype=torch.int32, device="cuda"), + torch.tensor([1], dtype=torch.int32, device="cuda"), + ] + if invalid == "activation_dtype": + args[1] = args[1].half() + elif invalid == "parameter_dtype": + args[4] = args[4].half() + elif invalid == "state_dtype": + args[7] = args[7].half() + elif invalid == "offset_dtype": + args[8] = args[8].long() + elif invalid == "index_count": + args[9] = args[9].repeat(2) + elif invalid == "bc_shape": + args[3] = args[3][:, :1].contiguous() + elif invalid == "zero_only_pool": + args[7] = args[7][:1] + elif invalid == "noncontiguous": + args[0] = args[0].transpose(0, 1) + elif invalid == "head_groups": + args[2] = args[3] = torch.zeros(4, 3, 16, device="cuda") + torch.cuda.synchronize() + wrapped = [ + infinicore.strided_from_blob( + value.data_ptr(), + list(value.shape), + list(value.stride()), + dtype=infinicore.utils.to_infinicore_dtype(value.dtype), + device=infinicore.device("cuda", value.device.index), + ) + for value in args + ] + with pytest.raises(RuntimeError): + infinicore.nn.functional.mamba2_scan(*wrapped) + + +if __name__ == "__main__": + GenericTestRunner(OpTest).run_and_exit() From f634435bb683895e5b084cf1e43f41f28da460ea Mon Sep 17 00:00:00 2001 From: tangchengxiang <2064027004@qq.com> Date: Sun, 20 Sep 2026 11:05:07 +0000 Subject: [PATCH 2/2] issue/1571 fix(mamba2): include MetaX FP32 precision policy Consolidate #1561 into the Mamba-2 backend contribution. Preserve the default TF32 behavior and the existing process-isolated strict-FP32 regression; no scan algorithm changes. --- README.md | 6 ++ src/infiniop/ops/gemm/metax/gemm_metax.cc | 10 +++- .../ops/test_metax_gemm_precision.py | 59 +++++++++++++++++++ 3 files changed, 73 insertions(+), 2 deletions(-) create mode 100644 test/infinicore/ops/test_metax_gemm_precision.py diff --git a/README.md b/README.md index 4569c0157..b84c41fde 100644 --- a/README.md +++ b/README.md @@ -109,6 +109,12 @@ python scripts/install.py [XMAKE_CONFIG_FLAGS] | `--ccl=[y\|n]` | 是否编译 InfiniCCL 通信库接口实现 | n | `--graph=[y\|n]` | 是否编译 cuda graph 接口实现 | n +MetaX FP32 GEMM permits TF32 by default. Set `INFINIOP_METAX_ALLOW_TF32=0` +before starting the process to request `MCBLAS_COMPUTE_32F` instead of +`MCBLAS_COMPUTE_32F_FAST_TF32`. The policy is captured when each descriptor +is created and remains fixed during graph replay; restart the process when +changing it. FP16 and BF16 GEMM retain their existing FP32 accumulation. + ##### 手动安装底层库 0. 生成九齿算子(可选) diff --git a/src/infiniop/ops/gemm/metax/gemm_metax.cc b/src/infiniop/ops/gemm/metax/gemm_metax.cc index 9d45099dc..0e3f758d8 100644 --- a/src/infiniop/ops/gemm/metax/gemm_metax.cc +++ b/src/infiniop/ops/gemm/metax/gemm_metax.cc @@ -1,11 +1,14 @@ #include "gemm_metax.h" #include "../../../devices/metax/metax_common.h" #include "../../../devices/metax/metax_handle.h" +#include +#include namespace op::gemm::metax { struct Descriptor::Opaque { std::shared_ptr internal; + bool allow_tf32; }; Descriptor::~Descriptor() { @@ -26,9 +29,12 @@ infiniStatus_t Descriptor::create( auto result = MatmulInfo::create(c_desc, a_desc, b_desc, MatrixLayout::COL_MAJOR); CHECK_RESULT(result); + // Capture the precision policy with the descriptor, including graph replay. + const char *allow_tf32 = std::getenv("INFINIOP_METAX_ALLOW_TF32"); + const bool use_tf32 = allow_tf32 == nullptr || std::strcmp(allow_tf32, "0") != 0; *desc_ptr = new Descriptor( dtype, result.take(), 0, - new Opaque{handle->internal()}, + new Opaque{handle->internal(), use_tf32}, handle->device, handle->device_id); return INFINI_STATUS_SUCCESS; } @@ -57,7 +63,7 @@ infiniStatus_t Descriptor::calculate( break; case INFINI_DTYPE_F32: a_type = b_type = c_type = HPCC_R_32F; - compute_type = HCBLAS_COMPUTE_32F_FAST_TF32; + compute_type = _opaque->allow_tf32 ? HCBLAS_COMPUTE_32F_FAST_TF32 : HCBLAS_COMPUTE_32F; break; default: diff --git a/test/infinicore/ops/test_metax_gemm_precision.py b/test/infinicore/ops/test_metax_gemm_precision.py new file mode 100644 index 000000000..6a179debe --- /dev/null +++ b/test/infinicore/ops/test_metax_gemm_precision.py @@ -0,0 +1,59 @@ +"""Check MetaX GEMM precision in processes with independent descriptor caches.""" + +import os +import subprocess +import sys + +import pytest +import torch +from infinicore.lib import _infinicore + +import infinicore + + +@pytest.mark.skipif( + _infinicore.get_device_count(_infinicore.Device.Type.METAX) == 0, + reason="A MetaX device is required.", +) +@pytest.mark.parametrize("allow_tf32", [None, "0"]) +def test_metax_gemm_precision(allow_tf32): + env = os.environ.copy() + if allow_tf32 is None: + env.pop("INFINIOP_METAX_ALLOW_TF32", None) + else: + env["INFINIOP_METAX_ALLOW_TF32"] = allow_tf32 + result = subprocess.run( + [sys.executable, __file__], env=env, capture_output=True, text=True + ) + assert result.returncode == 0, result.stdout + result.stderr + + +if __name__ == "__main__": + generator = torch.Generator().manual_seed(20260917) + for dtype in (torch.float32, torch.float16, torch.bfloat16): + for rows in (1, 3, 257): + x = torch.randn(rows, 768, generator=generator).to(dtype) + weight = torch.randn(96, 768, generator=generator).to(dtype) + expected = (x.double() @ weight.double().T).to(dtype) + x_device, weight_device = x.cuda(), weight.cuda().T + output = torch.empty((rows, 96), device="cuda", dtype=dtype) + torch.cuda.synchronize() + infinicore.matmul( + infinicore.from_torch(x_device), + infinicore.strided_from_blob( + weight_device.data_ptr(), + list(weight_device.shape), + list(weight_device.stride()), + dtype=infinicore.utils.to_infinicore_dtype(dtype), + device=infinicore.device("cuda", 0), + ), + out=infinicore.from_torch(output), + ) + infinicore.sync_device() + strict = ( + dtype == torch.float32 and os.getenv("INFINIOP_METAX_ALLOW_TF32") == "0" + ) + tolerance = (2e-4, 1e-5) if strict else (0.1, 1e-2) + torch.testing.assert_close( + output.cpu(), expected, atol=tolerance[0], rtol=tolerance[1] + )