| 99 | |
| 100 | template <typename T, typename Context> |
| 101 | void 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 |
nothing calls this directly
no test coverage detected