MCPcopy Create free account
hub / github.com/ml-explore/mlx-data / uniq_t

Function uniq_t

mlx/data/core/Utils.cpp:10–58  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9template <class T>
10void uniq_t(
11 std::shared_ptr<Array> dst,
12 std::shared_ptr<Array> dst_length,
13 const std::shared_ptr<Array> src,
14 const std::shared_ptr<Array> src_length,
15 int dim,
16 double pad) {
17 int64_t stride = 1;
18 for (int d = src->ndim() - 1; d > dim; d--) {
19 stride *= src->shape(d);
20 }
21
22 int64_t niter = 1;
23 for (int d = 0; d < src->ndim(); d++) {
24 if (d != dim) {
25 niter *= src->shape(d);
26 }
27 }
28 auto dst_t = dst->data<T>();
29 auto src_t = src->data<T>();
30 auto src_length_t = src_length->data<int64_t>();
31 auto dst_length_t = dst_length->data<int64_t>();
32 auto max_sz = src->shape(dim);
33 int64_t iterstride = (dim == src->ndim() - 1 ? src->shape(-1) : 1);
34 for (int64_t iter = 0; iter < niter; iter++) {
35 int64_t offset = iter * iterstride;
36 int64_t sz = src_length_t[iter];
37 int64_t idx = 0;
38 if (sz > max_sz) {
39 throw std::runtime_error("uniq: provided length exceeds input shape");
40 }
41 if (sz > 0) {
42 T last = src_t[offset];
43 dst_t[offset] = last;
44 idx = 1;
45 for (int64_t i = 1; i < sz; i++) {
46 if (src_t[offset + i * stride] != last) {
47 last = src_t[offset + i * stride];
48 dst_t[offset + idx * stride] = last;
49 idx++;
50 }
51 }
52 }
53 dst_length_t[iter] = idx;
54 for (; idx < max_sz; idx++) {
55 dst_t[offset + idx * stride] = static_cast<T>(pad);
56 }
57 }
58}
59
60template <class T>
61void remove_t(

Callers

nothing calls this directly

Calls 2

ndimMethod · 0.80
shapeMethod · 0.45

Tested by

no test coverage detected