tvm
Loading...
Searching...
No Matches
random_engine.h
Go to the documentation of this file.
1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements. See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership. The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License. You may obtain a copy of the License at
9 *
10 * http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied. See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
24#ifndef TVM_S_TIR_RANDOM_ENGINE_H_
25#define TVM_S_TIR_RANDOM_ENGINE_H_
26
27#include <tvm/ffi/error.h>
28#include <tvm/ir/prim/expr.h>
29
30#include <cstdint>
31#include <random>
32
33namespace tvm {
34namespace s_tir {
35using namespace tvm::prim;
36
48 public:
53 static constexpr TRandState multiplier = 48271;
55 static constexpr TRandState increment = 0;
57 static constexpr TRandState modulus = 2147483647;
59 static constexpr result_type min() { return 0; }
61 static constexpr result_type max() { return modulus - 1; }
62
68 std::random_device rd;
69 return rd() % modulus;
70 }
71
83 (*rand_state_ptr_) = ((*rand_state_ptr_) * multiplier + increment) % modulus;
84 return *rand_state_ptr_;
85 }
91 static TRandState NormalizeSeed(TRandState rand_state) {
92 if (rand_state == -1) {
93 rand_state = DeviceRandom();
94 } else {
95 rand_state %= modulus;
96 }
97 if (rand_state == 0) {
98 rand_state = 1;
99 }
100 if (rand_state < 0) {
101 TVM_FFI_THROW(ValueError) << "Random seed must be non-negative";
102 }
103 return rand_state;
104 }
109 void Seed(TRandState rand_state) {
110 TVM_FFI_ICHECK(rand_state_ptr_ != nullptr);
111 *rand_state_ptr_ = NormalizeSeed(rand_state);
112 }
113
119 // In order for reproducibility, we compute the new seed using RNG's random state and a
120 // different set of parameters. Note that both 32767 and 1999999973 are prime numbers.
121 return ((*this)() * 32767) % 1999999973;
122 }
123
132 rand_state_ptr_ = rand_state_ptr;
133 }
134
135 private:
136 TRandState* rand_state_ptr_;
137};
138
139} // namespace s_tir
140} // namespace tvm
141
142#endif // TVM_S_TIR_RANDOM_ENGINE_H_
RAII wrapper function to enter and exit a context object similar to python's with syntax.
Definition with_context.h:59
This linear congruential engine is a drop-in replacement for std::minstd_rand. It strictly correspond...
Definition random_engine.h:47
static constexpr TRandState multiplier
The multiplier.
Definition random_engine.h:53
void Seed(TRandState rand_state)
Change the start random state of RNG with the seed of a new random state value.
Definition random_engine.h:109
LinearCongruentialEngine(TRandState *rand_state_ptr)
Construct a random number generator with a random state pointer.
Definition random_engine.h:131
TRandState ForkSeed()
Fork a new seed for another RNG from current random state.
Definition random_engine.h:118
static constexpr TRandState modulus
The modulus.
Definition random_engine.h:57
static TRandState DeviceRandom()
Get a device random state.
Definition random_engine.h:67
result_type operator()()
Operator to move the random state to the next and return the new random state. According to definitio...
Definition random_engine.h:82
static TRandState NormalizeSeed(TRandState rand_state)
Normalize the random seed to the range of [1, modulus - 1].
Definition random_engine.h:91
static constexpr result_type max()
The maximum possible value of random state here.
Definition random_engine.h:61
int64_t TRandState
Definition random_engine.h:49
static constexpr result_type min()
The minimum possible value of random state here.
Definition random_engine.h:59
static constexpr TRandState increment
The increment.
Definition random_engine.h:55
TIR expressions.
Definition builtin.h:25
An object that builds and maintains block scope and StmtSref mapping for Dependence analysis.
Definition analyzer.h:40