Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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. 生成九齿算子(可选)
Expand Down
1 change: 1 addition & 0 deletions include/infinicore/ops.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
11 changes: 11 additions & 0 deletions include/infinicore/ops/mamba2_scan.hpp
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions include/infiniop.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
25 changes: 25 additions & 0 deletions include/infiniop/ops/mamba2_scan.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
#ifndef INFINIOP_MAMBA2_SCAN_API_H_
#define INFINIOP_MAMBA2_SCAN_API_H_
#include "../operator_descriptor.h"
#include <stddef.h>

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
2 changes: 2 additions & 0 deletions python/infinicore/nn/functional/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -88,6 +89,7 @@
"interpolate",
"log_softmax",
"mamba_selective_scan",
"mamba2_scan",
"moe_fused_dense",
"upsample_nearest",
"triplet_margin_with_distance_loss",
Expand Down
43 changes: 43 additions & 0 deletions python/infinicore/nn/functional/mamba2_scan.py
Original file line number Diff line number Diff line change
@@ -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,
)
)
21 changes: 21 additions & 0 deletions src/infinicore/ops/mamba2_scan/mamba2_scan.cc
Original file line number Diff line number Diff line change
@@ -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
25 changes: 25 additions & 0 deletions src/infinicore/ops/mamba2_scan/mamba2_scan_infiniop.cc
Original file line number Diff line number Diff line change
@@ -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> 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<PlannedMeta *>(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<PlannedMeta **>(planned_meta);
*planned_meta = nullptr;
}
INFINICORE_GRAPH_OP_REGISTER_ALLDEVICE(Mamba2Scan, &plan, &run, &cleanup);
} // namespace infinicore::op::mamba2_scan_impl::infiniop
2 changes: 2 additions & 0 deletions src/infinicore/pybind11/ops.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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);
Expand Down
8 changes: 8 additions & 0 deletions src/infinicore/pybind11/ops/mamba2_scan.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
#pragma once
#include "infinicore/ops/mamba2_scan.hpp"
#include <pybind11/pybind11.h>
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
10 changes: 8 additions & 2 deletions src/infiniop/ops/gemm/metax/gemm_metax.cc
Original file line number Diff line number Diff line change
@@ -1,11 +1,14 @@
#include "gemm_metax.h"
#include "../../../devices/metax/metax_common.h"
#include "../../../devices/metax/metax_handle.h"
#include <cstdlib>
#include <cstring>

namespace op::gemm::metax {

struct Descriptor::Opaque {
std::shared_ptr<device::metax::Handle::Internal> internal;
bool allow_tf32;
};

Descriptor::~Descriptor() {
Expand All @@ -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;
}
Expand Down Expand Up @@ -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:
Expand Down
141 changes: 141 additions & 0 deletions src/infiniop/ops/mamba2_scan/cuda/kernel.cuh
Original file line number Diff line number Diff line change
@@ -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 <typename T>
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<size_t>(blockIdx.x) * blockDim.x + threadIdx.x;
if (i < tokens * heads) {
const float step = softplus(static_cast<float>(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 <typename T, bool Summary, bool Chunked>
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<size_t>(end) > info.tokens
|| src < 0 || static_cast<size_t>(src) >= info.pool_size
|| dst <= 0 || static_cast<size_t>(dst) >= info.pool_size) {
return;
}
const size_t first = start + (Chunked ? blockIdx.z * info.chunk_size : 0);
if (first >= static_cast<size_t>(end)) {
return;
}
const size_t last = Chunked ? min(first + info.chunk_size, static_cast<size_t>(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<size_t>(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<float>(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<float>(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<float>(b[bc_base + n]));
if constexpr (!Summary) {
y = fmaf(values[j], static_cast<float>(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<T>(y + d[h] * input);
}
}
}
if constexpr (Summary || !Chunked) {
const size_t target_base = Summary ? chunk_base : static_cast<size_t>(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<size_t>(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<size_t>(end) > info.tokens
|| src < 0 || static_cast<size_t>(src) >= info.pool_size
|| dst <= 0 || static_cast<size_t>(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<size_t>(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<size_t>(dst) * (info.heads * info.head_dim * info.state_size) + element] = value;
}

} // namespace op::mamba2_scan::cuda
Loading