Skip to content

Commit 0f1ae44

Browse files
committed
feat: support graph on moore
1 parent d3551f3 commit 0f1ae44

6 files changed

Lines changed: 141 additions & 91 deletions

File tree

‎python/infinicore/ops/moore_mate_flash_attn.py‎

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -141,7 +141,8 @@ def moore_mate_flash_attn_decode(
141141
scale: float,
142142
block_size: int,
143143
max_seq_len: int,
144-
) -> torch.Tensor:
144+
return_state: bool = False,
145+
) -> torch.Tensor | tuple[torch.Tensor, ...]:
145146
"""
146147
Decode entry point with native flash_attn KV cache layout (B, P, H, D).
147148
No layout conversion is performed.
@@ -174,7 +175,7 @@ def moore_mate_flash_attn_decode(
174175
pack_gqa=pack_gqa,
175176
)
176177

177-
out, *_ = flash_attn_with_kvcache(
178+
result = flash_attn_with_kvcache(
178179
q=q,
179180
k_cache=k_cache,
180181
v_cache=v_cache,
@@ -188,7 +189,18 @@ def moore_mate_flash_attn_decode(
188189
pack_gqa=pack_gqa,
189190
return_softmax_lse=True,
190191
)
191-
return out
192+
if not return_state:
193+
return result[0]
194+
195+
# Mate allocates output/workspace and scheduler tensors while the graph is
196+
# captured. Keep every referenced allocation alive for graph replay.
197+
return (
198+
*result,
199+
cache_seqlens,
200+
page_table,
201+
cu_seqlens_q,
202+
metadata,
203+
)
192204

193205

194206
# =============================================================================

‎src/infinicore/ops/mha_kvcache/moore/mha_kvcache_flashattn_moore.cc‎

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
#include <pybind11/pybind11.h>
99
#include <stdexcept>
1010
#include <torch/csrc/utils/pybind.h>
11+
#include <vector>
1112

1213
namespace infinicore::op::mha_kvcache_impl::flashattn_moore {
1314

@@ -37,6 +38,7 @@ struct PlannedMeta {
3738
graph::GraphTensor out, q, k_cache, v_cache, seqlens_k, block_table;
3839
std::optional<graph::GraphTensor> alibi_slopes;
3940
float scale;
41+
std::vector<at::Tensor> mate_state;
4042
};
4143

4244
void *plan(Tensor out,
@@ -89,20 +91,26 @@ void run(void *planned_meta) {
8991
py::object py_seqlens_k = py::cast(seqlens_k);
9092
py::object py_blk_tbl = py::cast(block_table);
9193

92-
py::object result = wrapper.attr("moore_mate_flash_attn_decode")(
94+
py::tuple result = wrapper.attr("moore_mate_flash_attn_decode")(
9395
py_q,
9496
py_k_cache,
9597
py_v_cache,
9698
py_blk_tbl,
9799
py_seqlens_k,
98100
p->scale,
99101
block_size,
100-
max_seq_len);
102+
max_seq_len,
103+
true);
101104

102-
at::Tensor result_t = result.cast<at::Tensor>();
105+
at::Tensor result_t = result[0].cast<at::Tensor>();
103106
out_tensor.copy_(result_t.unsqueeze(1));
104107

105-
result = py::none();
108+
p->mate_state.clear();
109+
p->mate_state.reserve(result.size());
110+
for (py::handle item : result) {
111+
p->mate_state.push_back(item.cast<at::Tensor>());
112+
}
113+
106114
py_q = py_k_cache = py_v_cache = py_seqlens_k = py_blk_tbl = py::none();
107115
} catch (const py::error_already_set &e) {
108116
throw std::runtime_error(

‎src/infiniop/elementwise/moore/elementwise_moore.h‎

Lines changed: 64 additions & 60 deletions
Original file line numberDiff line numberDiff line change
@@ -6,12 +6,19 @@
66
#include "../../devices/moore/moore_kernel_common.h"
77
#include "elementwise_moore_api.h"
88

9+
#include <algorithm>
10+
911
namespace op::elementwise::moore {
1012
template <typename T>
1113
__device__ __forceinline__ const T *typedInputPtr(const void *ptr) {
1214
return reinterpret_cast<const T *>(ptr);
1315
}
1416

17+
template <size_t N>
18+
struct InputPointerArray {
19+
const void *values[N];
20+
};
21+
1522
__device__ __forceinline__ size_t getOutputIndex(size_t idx, bool is_contiguous, size_t ndim,
1623
const size_t *shape, const ptrdiff_t *strides) {
1724
return is_contiguous ? idx : device::moore::indexToOffset(idx, ndim, shape, strides);
@@ -50,20 +57,20 @@ INFINIOP_MOORE_KERNEL elementwiseKernel(
5057
const ptrdiff_t *__restrict__ output_strides,
5158
const ptrdiff_t *__restrict__ input_strides,
5259
Tdata *output,
53-
const void *const *inputs,
60+
InputPointerArray<N> inputs,
5461
size_t offset,
5562
Args... args) {
5663

5764
size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset;
5865

5966
if (idx < output_size) {
60-
const Tdata *const *typed_inputs = reinterpret_cast<const Tdata *const *>(inputs);
6167
size_t out_idx = getOutputIndex(idx, output_contiguous, ndim, output_shape, output_strides);
6268
InputIndexer indexer{idx, ndim, input_contiguous, input_broadcasted, input_shapes, input_strides, output_strides};
6369

6470
unpackInputsAndApply(
6571
[&](auto... Is) {
66-
output[out_idx] = Op{}(typed_inputs[Is.value][indexer(Is.value)]..., std::forward<Args>(args)...);
72+
output[out_idx] = Op{}(
73+
typedInputPtr<Tdata>(inputs.values[Is.value])[indexer(Is.value)]..., std::forward<Args>(args)...);
6774
},
6875
std::make_index_sequence<N>{});
6976
}
@@ -81,7 +88,7 @@ INFINIOP_MOORE_KERNEL elementwiseKernel(
8188
const ptrdiff_t *__restrict__ output_strides,
8289
const ptrdiff_t *__restrict__ input_strides,
8390
Tout *output,
84-
const void *const *__restrict__ inputs,
91+
InputPointerArray<sizeof...(Tin)> inputs,
8592
size_t offset) {
8693

8794
size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset;
@@ -93,17 +100,56 @@ INFINIOP_MOORE_KERNEL elementwiseKernel(
93100
unpackInputsAndApply(
94101
[&](auto... Is) {
95102
output[out_idx] = Op{}.template operator()<Tout, Tin...>(
96-
(typedInputPtr<Tin>(inputs[Is.value])[indexer(Is.value)])...);
103+
(typedInputPtr<Tin>(inputs.values[Is.value])[indexer(Is.value)])...);
97104
},
98105
std::index_sequence_for<Tin...>{});
99106
}
100107
}
101108

102109
struct DeviceImpl::Opaque {
103110
std::shared_ptr<device::moore::Handle::Internal> internal;
111+
void *device_meta = nullptr;
112+
const bool *input_contiguous = nullptr;
113+
const bool *input_broadcasted = nullptr;
114+
const size_t *output_shape = nullptr;
115+
const ptrdiff_t *output_strides = nullptr;
116+
const size_t *input_shapes = nullptr;
117+
const ptrdiff_t *input_strides = nullptr;
118+
infiniStatus_t init_status = INFINI_STATUS_SUCCESS;
119+
120+
Opaque(const std::shared_ptr<device::moore::Handle::Internal> &internal_,
121+
const op::elementwise::ElementwiseInfo &info)
122+
: internal(internal_), init_status(initialize(info)) {}
123+
124+
~Opaque() {
125+
if (device_meta != nullptr) {
126+
musaFree(device_meta);
127+
}
128+
}
129+
130+
infiniStatus_t initialize(const op::elementwise::ElementwiseInfo &info) {
131+
const auto meta_size = info.getMetaMemSize();
132+
if (meta_size == 0) {
133+
return INFINI_STATUS_SUCCESS;
134+
}
135+
136+
CHECK_MOORE(musaMalloc(&device_meta, meta_size));
137+
CHECK_MOORE(musaMemcpy(device_meta,
138+
info.getMetaStart(),
139+
meta_size,
140+
musaMemcpyHostToDevice));
141+
142+
const auto ndim = info.getNdim();
143+
const auto input_size = info.getInputSize();
144+
output_shape = reinterpret_cast<const size_t *>(device_meta);
145+
output_strides = reinterpret_cast<const ptrdiff_t *>(output_shape + ndim);
146+
input_shapes = reinterpret_cast<const size_t *>(output_strides + ndim);
147+
input_strides = reinterpret_cast<const ptrdiff_t *>(input_shapes + input_size * ndim);
148+
input_contiguous = reinterpret_cast<const bool *>(input_strides + input_size * ndim);
149+
input_broadcasted = input_contiguous + input_size;
104150

105-
Opaque(const std::shared_ptr<device::moore::Handle::Internal> &internal)
106-
: internal(internal) {}
151+
return INFINI_STATUS_SUCCESS;
152+
}
107153

108154
template <uint32_t BLOCK_SIZE, size_t N, typename Op, typename Tdata, typename... Args>
109155
infiniStatus_t calculateImpl(const op::elementwise::ElementwiseInfo &info,
@@ -136,42 +182,6 @@ struct DeviceImpl::Opaque {
136182
}
137183

138184
private:
139-
template <size_t N>
140-
infiniStatus_t infoToDevice(
141-
const op::elementwise::ElementwiseInfo &info,
142-
void *workspace,
143-
const void *const *h_inputs_arr,
144-
const void **&d_inputs_arr,
145-
const bool *&d_input_contiguous,
146-
const bool *&d_input_broadcasted,
147-
const size_t *&d_output_shape,
148-
const ptrdiff_t *&d_output_strides,
149-
const size_t *&d_input_shapes,
150-
const ptrdiff_t *&d_input_strides,
151-
musaStream_t stream) const {
152-
153-
constexpr auto input_size = N;
154-
const auto ndim = info.getNdim();
155-
constexpr auto input_arr_size = N * sizeof(*h_inputs_arr);
156-
const int8_t *info_meta_start = info.getMetaStart();
157-
const int8_t *d_meta_start = reinterpret_cast<int8_t *>(workspace) + input_arr_size;
158-
159-
// copy the input pointer array and meta to device
160-
CHECK_MOORE(musaMemcpyAsync(workspace, h_inputs_arr, input_arr_size, musaMemcpyHostToDevice, stream));
161-
CHECK_MOORE(musaMemcpyAsync((void *)d_meta_start, info_meta_start, info.getMetaMemSize(), musaMemcpyHostToDevice, stream));
162-
163-
// offset/assign the pointers
164-
d_inputs_arr = reinterpret_cast<const void **>(workspace);
165-
d_output_shape = reinterpret_cast<const size_t *>(d_meta_start);
166-
d_output_strides = reinterpret_cast<const ptrdiff_t *>(d_output_shape + ndim);
167-
d_input_shapes = reinterpret_cast<const size_t *>(d_output_strides + ndim);
168-
d_input_strides = reinterpret_cast<const ptrdiff_t *>(d_input_shapes + input_size * ndim);
169-
d_input_contiguous = reinterpret_cast<const bool *>(d_input_strides + input_size * ndim);
170-
d_input_broadcasted = reinterpret_cast<const bool *>(d_input_contiguous + input_size);
171-
172-
return INFINI_STATUS_SUCCESS;
173-
}
174-
175185
template <uint32_t BLOCK_SIZE, size_t N, typename KernelFunc, typename Tout, typename... Args>
176186
infiniStatus_t launchElementwiseKernel(
177187
const op::elementwise::ElementwiseInfo &info,
@@ -187,19 +197,10 @@ struct DeviceImpl::Opaque {
187197
return INFINI_STATUS_SUCCESS;
188198
}
189199

190-
// Device pointers
191-
const void **d_inputs_arr = nullptr;
192-
const bool *d_input_contiguous = nullptr;
193-
const bool *d_input_broadcasted = nullptr;
194-
const size_t *d_output_shape = nullptr;
195-
const ptrdiff_t *d_output_strides = nullptr;
196-
const size_t *d_input_shapes = nullptr;
197-
const ptrdiff_t *d_input_strides = nullptr;
198-
199-
CHECK_STATUS(infoToDevice<N>(info, workspace, inputs.data(), d_inputs_arr,
200-
d_input_contiguous, d_input_broadcasted,
201-
d_output_shape, d_output_strides,
202-
d_input_shapes, d_input_strides, stream));
200+
(void)workspace;
201+
CHECK_OR_RETURN(inputs.size() == N, INFINI_STATUS_BAD_PARAM);
202+
InputPointerArray<N> input_ptrs{};
203+
std::copy_n(inputs.begin(), N, input_ptrs.values);
203204

204205
dim3 blockDims(std::min(BLOCK_SIZE, static_cast<uint32_t>(internal->maxThreadsPerBlock())));
205206
dim3 gridDims(std::min(uint32_t(CEIL_DIV(output_size, blockDims.x)), static_cast<uint32_t>(internal->gridSizeX())));
@@ -208,10 +209,10 @@ struct DeviceImpl::Opaque {
208209
for (size_t i = 0; i < output_size; i += step) {
209210
kernel_func<<<gridDims, blockDims, 0, stream>>>(
210211
output_size, info.getNdim(), info.isOutputContiguous(),
211-
d_input_contiguous, d_input_broadcasted,
212-
d_output_shape, d_input_shapes,
213-
d_output_strides, d_input_strides,
214-
output, reinterpret_cast<const void **>(d_inputs_arr),
212+
input_contiguous, input_broadcasted,
213+
output_shape, input_shapes,
214+
output_strides, input_strides,
215+
output, input_ptrs,
215216
i, std::forward<Args>(args)...);
216217
}
217218

@@ -222,6 +223,9 @@ struct DeviceImpl::Opaque {
222223
template <typename... Args>
223224
utils::Result<DeviceImpl *> DeviceImpl::create(Args &&...args) {
224225
auto opaque = std::make_shared<Opaque>(std::forward<Args>(args)...);
226+
if (opaque->init_status != INFINI_STATUS_SUCCESS) {
227+
return opaque->init_status;
228+
}
225229
return utils::Result<DeviceImpl *>(new DeviceImpl(opaque));
226230
}
227231

‎src/infiniop/elementwise/moore/elementwise_moore_api.h‎

Lines changed: 16 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -38,22 +38,22 @@ class DeviceImpl final {
3838
Args &&...args);
3939
};
4040
} // namespace op::elementwise::moore
41-
#define CREATE_ELEMENTWISE_MOORE_DESCRIPTOR(HANDLE, DTYPE, OUT_DESC, INPUT_DESC_VEC) \
42-
\
43-
auto info_result = op::elementwise::ElementwiseInfo::create(OUT_DESC, INPUT_DESC_VEC); \
44-
CHECK_RESULT(info_result); \
45-
auto info = info_result.take(); \
46-
auto workspace_size = info.getMetaMemSize() + info.getInputSize() * sizeof(void *); \
47-
\
48-
auto device_impl_result = op::elementwise::moore::DeviceImpl::create(HANDLE->internal()); \
49-
CHECK_RESULT(device_impl_result); \
50-
\
51-
*desc_ptr = new Descriptor( \
52-
DTYPE, \
53-
std::move(info), \
54-
std::move(device_impl_result.take()), \
55-
workspace_size, \
56-
HANDLE->device, \
41+
#define CREATE_ELEMENTWISE_MOORE_DESCRIPTOR(HANDLE, DTYPE, OUT_DESC, INPUT_DESC_VEC) \
42+
\
43+
auto info_result = op::elementwise::ElementwiseInfo::create(OUT_DESC, INPUT_DESC_VEC); \
44+
CHECK_RESULT(info_result); \
45+
auto info = info_result.take(); \
46+
auto workspace_size = 0; \
47+
\
48+
auto device_impl_result = op::elementwise::moore::DeviceImpl::create(HANDLE->internal(), info); \
49+
CHECK_RESULT(device_impl_result); \
50+
\
51+
*desc_ptr = new Descriptor( \
52+
DTYPE, \
53+
std::move(info), \
54+
std::move(device_impl_result.take()), \
55+
workspace_size, \
56+
HANDLE->device, \
5757
HANDLE->device_id);
5858

5959
#endif // __INFINIOP_ELEMENTWISE_MOORE_API_H__

‎src/infiniop/ops/hardtanh/moore/hardtanh_moore.mu‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -94,9 +94,9 @@ infiniStatus_t Descriptor::create(
9494
auto info_result = op::elementwise::ElementwiseInfo::create(out_desc, input_desc_vec);
9595
CHECK_RESULT(info_result);
9696
auto info = info_result.take();
97-
auto workspace_size = info.getMetaMemSize() + info.getInputSize() * sizeof(void *);
97+
auto workspace_size = 0;
9898

99-
auto device_impl_result = op::elementwise::moore::DeviceImpl::create(handle->internal());
99+
auto device_impl_result = op::elementwise::moore::DeviceImpl::create(handle->internal(), info);
100100
CHECK_RESULT(device_impl_result);
101101

102102
*desc_ptr = new Descriptor(

0 commit comments

Comments
 (0)