| 20 | |
| 21 | template <typename T, typename Context> |
| 22 | void ShuffleBatchKernel(const Context& dev_ctx, |
| 23 | const DenseTensor& x, |
| 24 | const DenseTensor& seed, |
| 25 | int startup_seed, |
| 26 | DenseTensor* out, |
| 27 | DenseTensor* shuffleidx, |
| 28 | DenseTensor* seed_out) { |
| 29 | auto x_embed_size = x.dims()[x.dims().size() - 1]; |
| 30 | int elem_size = 1; |
| 31 | for (auto i = 0; i < x.dims().size() - 1; i++) |
| 32 | elem_size *= static_cast<int>(x.dims()[i]); |
| 33 | |
| 34 | std::vector<int64_t> idx_vec; // record shuffled order |
| 35 | idx_vec.reserve(elem_size); |
| 36 | for (int i = 0; i < elem_size; i++) { |
| 37 | idx_vec.push_back(i); |
| 38 | } |
| 39 | int64_t seed_int = 0; |
| 40 | if (seed.initialized()) { |
| 41 | seed_int = *seed.data<int64_t>(); |
| 42 | } else { |
| 43 | seed_int = startup_seed; |
| 44 | } |
| 45 | std::default_random_engine engine; |
| 46 | engine.seed(seed_int); |
| 47 | |
| 48 | auto custom_random_shuffle = [&idx_vec]() { |
| 49 | std::random_device rnd; |
| 50 | int64_t seed_tmp = rnd(); |
| 51 | std::default_random_engine rng(seed_tmp); |
| 52 | const int n = static_cast<int>(idx_vec.size()); |
| 53 | std::vector<int> v(n); |
| 54 | std::iota(v.begin(), v.end(), 0); |
| 55 | std::vector<bool> visit(n, false); |
| 56 | while (!v.empty()) { |
| 57 | std::shuffle(v.begin(), v.end(), rng); |
| 58 | int tmp = v.back(); |
| 59 | v.pop_back(); |
| 60 | if (v.empty()) { |
| 61 | std::uniform_int_distribution<int> distr(0, n - 2); |
| 62 | idx_vec[tmp] = tmp; |
| 63 | std::swap(idx_vec[tmp], idx_vec[(distr(rng) + tmp + 1) % n]); |
| 64 | return; |
| 65 | } |
| 66 | visit[tmp] = true; |
| 67 | std::shuffle(v.begin(), v.end(), rng); |
| 68 | int curr = v.back(); |
| 69 | v.pop_back(); |
| 70 | v.push_back(tmp); |
| 71 | idx_vec[tmp] = curr; |
| 72 | while (!visit[curr]) { |
| 73 | visit[curr] = true; |
| 74 | std::shuffle(v.begin(), v.end(), rng); |
| 75 | idx_vec[curr] = v.back(); |
| 76 | v.pop_back(); |
| 77 | curr = static_cast<int>(idx_vec[curr]); |
| 78 | } |
| 79 | } |
nothing calls this directly
no test coverage detected