MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / ShuffleBatchKernel

Function ShuffleBatchKernel

paddle/phi/kernels/cpu/shuffle_batch_kernel.cc:22–105  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

20
21template <typename T, typename Context>
22void 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 }

Callers

nothing calls this directly

Calls 14

shuffleFunction · 0.85
swapFunction · 0.50
dimsMethod · 0.45
sizeMethod · 0.45
reserveMethod · 0.45
push_backMethod · 0.45
initializedMethod · 0.45
seedMethod · 0.45
beginMethod · 0.45
endMethod · 0.45
emptyMethod · 0.45
backMethod · 0.45

Tested by

no test coverage detected