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

Function RandpermKernel

paddle/phi/kernels/cpu/randperm_kernel.cc:101–149  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

99
100template <typename T, typename Context>
101void RandpermKernel(const Context& dev_ctx,
102 int n,
103 DataType dtype UNUSED,
104 DenseTensor* out) {
105 T* out_data = dev_ctx.template Alloc<T>(out);
106
107 if (FLAGS_use_accuracy_compatible_kernel) {
108 // MT19937 engine with that seed so the random sequence is identical.
109 uint64_t seed = dev_ctx.GetGenerator()->GetCurrentSeed();
110 TorchMT19937Engine engine(seed);
111
112 if (n < static_cast<int>(std::numeric_limits<uint32_t>::max() / 20)) {
113 // For small n: classic Fisher-Yates shuffle using 32-bit random values
114 for (int i = 0; i < n; ++i) {
115 out_data[i] = static_cast<T>(i);
116 }
117 for (int i = 0; i < n - 1; i++) {
118 int64_t z = engine() % (n - i);
119 T save = out_data[i];
120 out_data[i] = out_data[z + i];
121 out_data[z + i] = save;
122 }
123 } else {
124 // For large n: inside-out Fisher-Yates using 64-bit random values
125 for (int i = 0; i < n; i++) {
126 int64_t z = static_cast<int64_t>(engine.random64() % (i + 1));
127 out_data[i] = out_data[z];
128 out_data[z] = static_cast<T>(i);
129 }
130 }
131
132 // Advance the generator state so that successive randperm calls within the
133 // same run produce different results
134 dev_ctx.GetGenerator()->SetCurrentSeed(engine());
135 } else {
136 int seed = 0;
137 std::shared_ptr<std::mt19937_64> engine;
138 if (seed) {
139 engine = std::make_shared<std::mt19937_64>();
140 engine->seed(seed);
141 } else {
142 engine = dev_ctx.GetGenerator()->GetCPUEngine();
143 }
144 for (int i = 0; i < n; ++i) {
145 out_data[i] = static_cast<T>(i);
146 }
147 std::shuffle(out_data, out_data + n, *engine);
148 }
149}
150
151} // namespace phi
152

Callers

nothing calls this directly

Calls 8

shuffleFunction · 0.85
GetCurrentSeedMethod · 0.80
random64Method · 0.80
SetCurrentSeedMethod · 0.80
GetCPUEngineMethod · 0.80
maxFunction · 0.50
GetGeneratorMethod · 0.45
seedMethod · 0.45

Tested by

no test coverage detected