Skip to content

Commit e63c437

Browse files
Update implementation of dpnp.putmask (#3014)
This PR proposes a new implementation of `dpnp.putmask` replacing the legacy `dpnp_putmask` implementation with dedicated SYCL kernels : a vectorized contiguous kernel and a strided kernel for F-contiguous/transposed arrays. It also fully reworks the `putmask` tests by adding `TestPutMask`
1 parent 147a64a commit e63c437

12 files changed

Lines changed: 951 additions & 196 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,7 @@ This release is compatible with NumPy 2.5.
4949
* Reworked the ASV benchmarks and added end-to-end workload benchmarks derived from dpBench [#2996](https://github.com/IntelPython/dpnp/pull/2996)
5050
* Reduced allocations in `dpnp.linalg.norm` by reusing the reduction result as the `sqrt` output buffer in the 2-norm and Frobenius-norm branches [#3062](https://github.com/IntelPython/dpnp/pull/3062)
5151
* Avoided a copy of `dpnp.einsum` result into C-order by building the product in the requested layout directly [#3069](https://github.com/IntelPython/dpnp/pull/3069)
52+
* Updated the implementation of `dpnp.putmask` by adding dedicated contiguous and strided SYCL kernels [#3014](https://github.com/IntelPython/dpnp/pull/3014)
5253

5354
### Deprecated
5455

‎dpnp/backend/extensions/indexing/CMakeLists.txt‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
set(python_module_name _indexing_impl)
3131
set(_module_src
3232
${CMAKE_CURRENT_SOURCE_DIR}/choose.cpp
33+
${CMAKE_CURRENT_SOURCE_DIR}/putmask.cpp
3334
${CMAKE_CURRENT_SOURCE_DIR}/indexing_py.cpp
3435
)
3536

‎dpnp/backend/extensions/indexing/indexing_py.cpp‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,8 +33,10 @@
3333
#include <pybind11/pybind11.h>
3434

3535
#include "choose.hpp"
36+
#include "putmask.hpp"
3637

3738
PYBIND11_MODULE(_indexing_impl, m, py::mod_gil_not_used())
3839
{
3940
dpnp::extensions::indexing::init_choose(m);
41+
dpnp::extensions::indexing::init_putmask(m);
4042
}
Lines changed: 284 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,284 @@
1+
//*****************************************************************************
2+
// Copyright (c) 2026, Intel Corporation
3+
// All rights reserved.
4+
//
5+
// Redistribution and use in source and binary forms, with or without
6+
// modification, are permitted provided that the following conditions are met:
7+
// - Redistributions of source code must retain the above copyright notice,
8+
// this list of conditions and the following disclaimer.
9+
// - Redistributions in binary form must reproduce the above copyright notice,
10+
// this list of conditions and the following disclaimer in the documentation
11+
// and/or other materials provided with the distribution.
12+
// - Neither the name of the copyright holder nor the names of its contributors
13+
// may be used to endorse or promote products derived from this software
14+
// without specific prior written permission.
15+
//
16+
// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
17+
// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
18+
// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
19+
// ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE
20+
// LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
21+
// CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
22+
// SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
23+
// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN
24+
// CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
25+
// ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF
26+
// THE POSSIBILITY OF SUCH DAMAGE.
27+
//*****************************************************************************
28+
29+
#include <algorithm>
30+
#include <cstddef>
31+
#include <tuple>
32+
#include <utility>
33+
#include <vector>
34+
35+
#include <sycl/sycl.hpp>
36+
37+
#include <pybind11/pybind11.h>
38+
#include <pybind11/stl.h>
39+
40+
#include "dpnp4pybind11.hpp"
41+
42+
#include "kernels/indexing/putmask.hpp"
43+
44+
// dpnp tensor headers
45+
#include "utils/memory_overlap.hpp"
46+
#include "utils/offset_utils.hpp"
47+
#include "utils/output_validation.hpp"
48+
#include "utils/sycl_alloc_utils.hpp"
49+
#include "utils/type_dispatch.hpp"
50+
51+
// utils extension headers
52+
#include "ext/common.hpp"
53+
#include "ext/validation_utils.hpp"
54+
55+
namespace py = pybind11;
56+
namespace td_ns = dpnp::tensor::type_dispatch;
57+
58+
using dpnp::tensor::usm_ndarray;
59+
60+
using ext::validation::array_names;
61+
using ext::validation::check_c_contig;
62+
using ext::validation::check_has_dtype;
63+
using ext::validation::check_num_dims;
64+
using ext::validation::check_queue;
65+
using ext::validation::check_same_dtype;
66+
using ext::validation::check_same_size;
67+
using ext::validation::check_writable;
68+
69+
namespace dpnp::extensions::indexing
70+
{
71+
using ext::common::init_dispatch_vector;
72+
73+
typedef sycl::event (*putmask_strided_fn_ptr_t)(
74+
sycl::queue &,
75+
const int, // nd
76+
const std::size_t, // nelems
77+
const py::ssize_t *, // shape_strides
78+
char *, // dst
79+
py::ssize_t, // dst_offset
80+
const char *, // mask
81+
py::ssize_t, // mask_offset
82+
const char *, // values
83+
const std::size_t, // values_size
84+
const std::vector<sycl::event> &);
85+
86+
template <typename T>
87+
sycl::event putmask_strided_call(sycl::queue &q,
88+
const int nd,
89+
const std::size_t nelems,
90+
const py::ssize_t *shape_strides,
91+
char *dst_p,
92+
py::ssize_t dst_offset,
93+
const char *mask_p,
94+
py::ssize_t mask_offset,
95+
const char *values_p,
96+
const std::size_t values_size,
97+
const std::vector<sycl::event> &depends)
98+
{
99+
return dpnp::kernels::putmask::putmask_strided_impl<T>(
100+
q, nd, nelems, shape_strides, dst_p, dst_offset, mask_p, mask_offset,
101+
values_p, values_size, depends);
102+
}
103+
104+
typedef sycl::event (*putmask_contig_fn_ptr_t)(
105+
sycl::queue &,
106+
const std::size_t, // nelems
107+
char *, // dst
108+
const char *, // mask
109+
const char *, // values
110+
const std::size_t, // values_size
111+
const std::vector<sycl::event> &);
112+
113+
template <typename T>
114+
sycl::event putmask_contig_call(sycl::queue &q,
115+
const std::size_t nelems,
116+
char *dst_p,
117+
const char *mask_p,
118+
const char *values_p,
119+
const std::size_t values_size,
120+
const std::vector<sycl::event> &depends)
121+
{
122+
return dpnp::kernels::putmask::putmask_contig_impl<T>(
123+
q, nelems, dst_p, mask_p, values_p, values_size, depends);
124+
}
125+
126+
putmask_strided_fn_ptr_t putmask_strided_dispatch_vector[td_ns::num_types];
127+
putmask_contig_fn_ptr_t putmask_contig_dispatch_vector[td_ns::num_types];
128+
129+
std::pair<sycl::event, sycl::event>
130+
py_putmask(const usm_ndarray &dst,
131+
const usm_ndarray &mask,
132+
const usm_ndarray &values,
133+
sycl::queue &exec_q,
134+
const std::vector<sycl::event> &depends = {})
135+
{
136+
array_names names = {{&dst, "dst"}, {&mask, "mask"}, {&values, "values"}};
137+
138+
check_same_dtype(&dst, &values, names);
139+
check_has_dtype(&mask, td_ns::typenum_t::BOOL, names);
140+
141+
// TODO: redundant with the shape check below;
142+
// use `check_same_shape` later
143+
check_same_size({&dst, &mask}, names);
144+
const int nd = dst.get_ndim();
145+
check_num_dims({&mask}, nd, names);
146+
147+
check_queue({&dst, &mask, &values}, names, exec_q);
148+
check_writable({&dst}, names);
149+
150+
// values must be C-contiguous
151+
check_c_contig({&values}, names);
152+
153+
const auto &overlap = dpnp::tensor::overlap::MemoryOverlap();
154+
if (overlap(dst, mask) || overlap(dst, values)) {
155+
throw py::value_error("Arrays have overlapping segments of memory");
156+
}
157+
158+
auto types = td_ns::usm_ndarray_types();
159+
// dst_typeid == values_typeid (check_same_dtype(&dst, &values, names))
160+
int dst_values_typeid = types.typenum_to_lookup_id(dst.get_typenum());
161+
162+
const py::ssize_t *dst_shape = dst.get_shape_raw();
163+
const py::ssize_t *mask_shape = mask.get_shape_raw();
164+
bool shapes_equal(true);
165+
std::size_t nelems(1);
166+
167+
for (int i = 0; i < std::max(nd, 1); ++i) {
168+
const py::ssize_t d = (nd == 0 ? 1 : dst_shape[i]);
169+
const py::ssize_t m = (nd == 0 ? 1 : mask_shape[i]);
170+
nelems *= static_cast<std::size_t>(d);
171+
shapes_equal = shapes_equal && (d == m);
172+
}
173+
if (!shapes_equal) {
174+
throw py::value_error("`mask` and `dst` shapes must match");
175+
}
176+
177+
const std::size_t values_size = values.get_size();
178+
179+
// empty output or empty `values` is a no-op
180+
if (nelems == 0 || values_size == 0) {
181+
return {sycl::event(), sycl::event()};
182+
}
183+
184+
dpnp::tensor::validation::AmpleMemory::throw_if_not_ample(dst, nelems);
185+
186+
char *dst_p = dst.get_data();
187+
const char *mask_p = mask.get_data();
188+
const char *values_p = values.get_data();
189+
190+
// the contig kernel cycles `values` by the memory-linear index, which
191+
// matches numpy's C-order `values.flat` only for C-contiguous data
192+
// (`values` is already checked to be C-contiguous above)
193+
const bool all_c_contig = dst.is_c_contiguous() && mask.is_c_contiguous();
194+
195+
if (all_c_contig) {
196+
auto contig_fn = putmask_contig_dispatch_vector[dst_values_typeid];
197+
198+
auto comp_ev = contig_fn(exec_q, nelems, dst_p, mask_p, values_p,
199+
values_size, depends);
200+
sycl::event ht_ev = dpnp::utils::keep_args_alive(
201+
exec_q, {dst, mask, values}, {comp_ev});
202+
203+
return std::make_pair(ht_ev, comp_ev);
204+
}
205+
206+
// strided path: the iteration space is intentionally not simplified, so
207+
// the kernel's linear index stays equal to the C-order flat index used to
208+
// cycle `values` (simplify_iteration_space may reorder axes and break it)
209+
const auto &dst_strides = dst.get_strides_vector();
210+
const auto &mask_strides = mask.get_strides_vector();
211+
212+
// 0-d arrays go through the contig path, so here nd >= 1
213+
using shT = std::vector<py::ssize_t>;
214+
shT common_shape(dst_shape, dst_shape + nd);
215+
shT s_dst_strides = dst_strides;
216+
shT s_mask_strides = mask_strides;
217+
218+
// trivial offsets: shape and strides are passed without simplification
219+
constexpr py::ssize_t dst_off = 0;
220+
constexpr py::ssize_t mask_off = 0;
221+
222+
auto strided_fn = putmask_strided_dispatch_vector[dst_values_typeid];
223+
224+
using dpnp::tensor::offset_utils::device_allocate_and_pack;
225+
226+
std::vector<sycl::event> host_tasks;
227+
host_tasks.reserve(2);
228+
229+
auto pack = device_allocate_and_pack<py::ssize_t>(
230+
exec_q, host_tasks, common_shape, s_dst_strides, s_mask_strides);
231+
232+
auto shape_strides_owner = std::move(std::get<0>(pack));
233+
const py::ssize_t *shape_strides_dev = shape_strides_owner.get();
234+
const sycl::event &cpy_ev = std::get<2>(pack);
235+
236+
std::vector<sycl::event> all_deps = depends;
237+
all_deps.push_back(cpy_ev);
238+
239+
sycl::event comp_ev =
240+
strided_fn(exec_q, nd, nelems, shape_strides_dev, dst_p, dst_off,
241+
mask_p, mask_off, values_p, values_size, all_deps);
242+
243+
sycl::event cleanup_ev = dpnp::tensor::alloc_utils::async_smart_free(
244+
exec_q, {comp_ev}, shape_strides_owner);
245+
host_tasks.push_back(cleanup_ev);
246+
247+
sycl::event ht_ev =
248+
dpnp::utils::keep_args_alive(exec_q, {dst, mask, values}, host_tasks);
249+
250+
return std::make_pair(ht_ev, comp_ev);
251+
}
252+
253+
template <typename fnT, typename T>
254+
struct PutMaskStridedFactory
255+
{
256+
fnT get() { return putmask_strided_call<T>; }
257+
};
258+
259+
template <typename fnT, typename T>
260+
struct PutMaskContigFactory
261+
{
262+
fnT get() { return putmask_contig_call<T>; }
263+
};
264+
265+
static void populate_putmask_dispatch_vectors()
266+
{
267+
init_dispatch_vector<putmask_strided_fn_ptr_t, PutMaskStridedFactory>(
268+
putmask_strided_dispatch_vector);
269+
init_dispatch_vector<putmask_contig_fn_ptr_t, PutMaskContigFactory>(
270+
putmask_contig_dispatch_vector);
271+
}
272+
273+
void init_putmask(py::module_ &m)
274+
{
275+
populate_putmask_dispatch_vectors();
276+
277+
m.def("_putmask", &py_putmask, "", py::arg("dst"), py::arg("mask"),
278+
py::arg("values"), py::arg("sycl_queue"),
279+
py::arg("depends") = py::list());
280+
281+
return;
282+
}
283+
284+
} // namespace dpnp::extensions::indexing
Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
//*****************************************************************************
2+
// Copyright (c) 2026, Intel Corporation
3+
// All rights reserved.
4+
//
5+
// Redistribution and use in source and binary forms, with or without
6+
// modification, are permitted provided that the following conditions are met:
7+
// - Redistributions of source code must retain the above copyright notice,
8+
// this list of conditions and the following disclaimer.
9+
// - Redistributions in binary form must reproduce the above copyright notice,
10+
// this list of conditions and the following disclaimer in the documentation
11+
// and/or other materials provided with the distribution.
12+
// - Neither the name of the copyright holder nor the names of its contributors
13+
// may be used to endorse or promote products derived from this software
14+
// without specific prior written permission.
15+
//
16+
// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
17+
// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
18+
// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
19+
// ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE
20+
// LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
21+
// CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
22+
// SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
23+
// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN
24+
// CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
25+
// ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF
26+
// THE POSSIBILITY OF SUCH DAMAGE.
27+
//*****************************************************************************
28+
29+
#pragma once
30+
31+
#include <pybind11/pybind11.h>
32+
33+
namespace py = pybind11;
34+
35+
namespace dpnp::extensions::indexing
36+
{
37+
void init_putmask(py::module_ &m);
38+
} // namespace dpnp::extensions::indexing

0 commit comments

Comments
 (0)