99 lines
4.1 KiB
C++
99 lines
4.1 KiB
C++
// mars::rng tests — reference vectors + state (de)serialization.
|
|
#include <cstdint>
|
|
#include <cstdio>
|
|
#include <cstring>
|
|
|
|
#include "mars/rng/mt19937.h"
|
|
|
|
using mars::rng::MT19937;
|
|
|
|
static int fails = 0;
|
|
#define CHECK(cond) \
|
|
do { \
|
|
if (!(cond)) { \
|
|
std::printf("FAIL %s:%d: %s\n", __FILE__, __LINE__, #cond); \
|
|
++fails; \
|
|
} \
|
|
} while (0)
|
|
|
|
int main() {
|
|
// --- standard MT19937 outputs for init_genrand(5489) ------------------------
|
|
{
|
|
MT19937 r(5489u);
|
|
static const uint32_t expect[10] = {3499211612u, 581869302u, 3890346734u, 3586334585u, 545404204u,
|
|
4161255391u, 3922919429u, 949333985u, 2715962298u, 1323567403u};
|
|
for (uint32_t e : expect) CHECK(r.next_u32() == e);
|
|
// the 10000th output of the 5489 stream (well-known check value)
|
|
MT19937 r2(5489u);
|
|
uint32_t v = 0;
|
|
for (int i = 0; i < 10000; ++i) v = r2.next_u32();
|
|
CHECK(v == 4123659995u);
|
|
}
|
|
// --- a fresh generator has twisted once: left == N ------------------------
|
|
{
|
|
MT19937 r(1u);
|
|
CHECK(r.left() == MT19937::N);
|
|
CHECK(r.index() == 0);
|
|
r.next_u32();
|
|
CHECK(r.left() == MT19937::N - 1);
|
|
for (int i = 1; i < MT19937::N; ++i) r.next_u32();
|
|
CHECK(r.left() == 0); // block exhausted; the next draw twists
|
|
r.next_u32();
|
|
CHECK(r.left() == MT19937::N - 1);
|
|
}
|
|
// --- float mapping: y * 2^-32 narrowed to float -----------------------------
|
|
{
|
|
MT19937 r(5489u);
|
|
float f = r.next_float();
|
|
float expect = static_cast<float>(3499211612.0 / 4294967296.0);
|
|
CHECK(f == expect);
|
|
CHECK(f >= 0.f && f <= 1.f);
|
|
}
|
|
// --- next_int: in range, and consumes exactly one word when mask == n-1 ----
|
|
{
|
|
MT19937 r(7u);
|
|
for (int i = 0; i < 1000; ++i) CHECK(r.next_int(10) < 10);
|
|
MT19937 a(9u), b(9u);
|
|
uint32_t x = a.next_int(256);
|
|
CHECK(x == (b.next_u32() & 255u));
|
|
}
|
|
// --- save_state / load_state round trip, blob layout mt[624] + left --------
|
|
{
|
|
MT19937 r(123456u);
|
|
for (int i = 0; i < 700; ++i) r.next_u32(); // past one twist
|
|
uint8_t blob[MT19937::kStateBytes];
|
|
r.save_state(blob);
|
|
CHECK(MT19937::kStateBytes == 0x9c4);
|
|
uint32_t left_in_blob = uint32_t(blob[2496]) | (uint32_t(blob[2497]) << 8) | (uint32_t(blob[2498]) << 16) |
|
|
(uint32_t(blob[2499]) << 24);
|
|
CHECK(int(left_in_blob) == r.left());
|
|
CHECK(std::memcmp(blob, r.state(), 4) == 0 || true); // first word is mt[0] little-endian
|
|
uint32_t w0 = uint32_t(blob[0]) | (uint32_t(blob[1]) << 8) | (uint32_t(blob[2]) << 16) | (uint32_t(blob[3]) << 24);
|
|
CHECK(w0 == r.state()[0]);
|
|
|
|
MT19937 s(1u);
|
|
CHECK(s.load_state(blob, sizeof blob));
|
|
CHECK(s.left() == r.left());
|
|
for (int i = 0; i < 2000; ++i) CHECK(s.next_u32() == r.next_u32());
|
|
|
|
// truncated / out-of-range blobs are rejected
|
|
CHECK(!s.load_state(blob, sizeof blob - 1));
|
|
uint8_t bad[MT19937::kStateBytes];
|
|
std::memcpy(bad, blob, sizeof bad);
|
|
bad[2496] = 0x71; // left = 625 > N
|
|
bad[2497] = 0x02;
|
|
CHECK(!s.load_state(bad, sizeof bad));
|
|
}
|
|
// --- load_state(mt, left) positions the next word at mt[N - left] -----------
|
|
{
|
|
MT19937 r(42u);
|
|
uint32_t st[MT19937::N];
|
|
std::memcpy(st, r.state(), sizeof st);
|
|
MT19937 s(0u);
|
|
s.load_state(st, 5);
|
|
for (int i = 0; i < MT19937::N - 5; ++i) r.next_u32();
|
|
for (int i = 0; i < 100; ++i) CHECK(s.next_u32() == r.next_u32());
|
|
}
|
|
std::printf("test_rng: %s\n", fails ? "FAILED" : "ok");
|
|
return fails ? 1 : 0;
|
|
}
|