| 8 | |
| 9 | template <class T> |
| 10 | void 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 | |
| 60 | template <class T> |
| 61 | void remove_t( |