53 const std::uint8_t *mask_u8_ =
nullptr;
54 const T *vals_ =
nullptr;
55 const TwoOffsets_IndexerT indexer_;
56 const std::size_t values_size_ = 0;
62 const TwoOffsets_IndexerT &indexer,
63 const std::size_t values_size)
64 : dst_(dst), mask_u8_(
reinterpret_cast<const std::uint8_t *
>(mask)),
65 vals_(vals), indexer_(indexer), values_size_(values_size)
69 void operator()(sycl::id<1> wid)
const
71 const std::size_t lin = wid[0];
72 auto offset = indexer_(
static_cast<ssize_t
>(lin));
74 const dpnp::tensor::ssize_t dst_off = offset.get_first_offset();
75 const dpnp::tensor::ssize_t mask_off = offset.get_second_offset();
77 if (mask_u8_[mask_off]) {
78 const std::size_t vlin = lin % values_size_;
79 dst_[dst_off] = vals_[vlin];
92 const std::uint8_t *mask_u8_ =
nullptr;
93 const T *values_ =
nullptr;
94 std::size_t nelems_ = 0;
95 std::size_t val_size_ = 0;
101 const std::size_t nelems,
102 const std::size_t val_size)
103 : dst_(dst), mask_u8_(
reinterpret_cast<const std::uint8_t *
>(mask)),
104 values_(values), nelems_(nelems), val_size_(val_size)
108 void operator()(sycl::nd_item<1> ndit)
const
110 const bool values_no_repeat = (val_size_ >= nelems_);
112 constexpr std::uint8_t elems_per_wi = n_vecs * vec_sz;
116 using dpnp::tensor::type_utils::is_complex_v;
117 if constexpr (enable_sg_loadstore && !is_complex_v<T>) {
118 auto sg = ndit.get_sub_group();
119 const std::uint32_t sgSize = sg.get_max_local_range()[0];
120 const std::size_t lane_id = sg.get_local_id()[0];
122 const std::size_t base =
123 elems_per_wi * (ndit.get_group(0) * ndit.get_local_range(0) +
124 sg.get_group_id()[0] * sgSize);
126 if (base + elems_per_wi * sgSize <= nelems_) {
127 using dpnp::tensor::sycl_utils::sub_group_load;
128 using dpnp::tensor::sycl_utils::sub_group_store;
131 for (std::uint8_t it = 0; it < elems_per_wi; it += vec_sz) {
132 const std::size_t offset = base + it * sgSize;
134 auto dst_multi_ptr = sycl::address_space_cast<
135 sycl::access::address_space::global_space,
136 sycl::access::decorated::yes>(&dst_[offset]);
137 auto mask_multi_ptr = sycl::address_space_cast<
138 sycl::access::address_space::global_space,
139 sycl::access::decorated::yes>(&mask_u8_[offset]);
141 const sycl::vec<T, vec_sz> dst_vec =
142 sub_group_load<vec_sz>(sg, dst_multi_ptr);
143 const sycl::vec<std::uint8_t, vec_sz> mask_vec =
144 sub_group_load<vec_sz>(sg, mask_multi_ptr);
146 sycl::vec<T, vec_sz> val_vec;
148 if (values_no_repeat) {
149 auto values_multi_ptr = sycl::address_space_cast<
150 sycl::access::address_space::global_space,
151 sycl::access::decorated::yes>(&values_[offset]);
153 val_vec = sub_group_load<vec_sz>(sg, values_multi_ptr);
156 const std::size_t idx = offset + lane_id;
158 for (std::uint8_t k = 0; k < vec_sz; ++k) {
159 const std::size_t g =
160 idx +
static_cast<std::size_t
>(k) * sgSize;
161 val_vec[k] = values_[g % val_size_];
165 sycl::vec<T, vec_sz> out_vec;
167 for (std::uint8_t vec_id = 0; vec_id < vec_sz; ++vec_id) {
168 out_vec[vec_id] = (mask_vec[vec_id]) ? val_vec[vec_id]
172 sub_group_store<vec_sz>(sg, out_vec, dst_multi_ptr);
176 const std::size_t lane_id = sg.get_local_id()[0];
177 for (std::size_t k = base + lane_id; k < nelems_; k += sgSize) {
179 const std::size_t v =
180 values_no_repeat ? k : (k % val_size_);
181 dst_[k] = values_[v];
187 const std::size_t gid = ndit.get_global_linear_id();
188 const std::size_t gws = ndit.get_global_range(0);
190 for (std::size_t offset = gid; offset < nelems_; offset += gws) {
191 if (mask_u8_[offset]) {
192 const std::size_t v =
193 values_no_repeat ? offset : (offset % val_size_);
194 dst_[offset] = values_[v];