#ifndef BOOST_COMPUTE_RANDOM_THREEFRY_HPP
#define BOOST_COMPUTE_RANDOM_THREEFRY_HPP
#include <algorithm>
#include <boost/compute/types.hpp>
#include <boost/compute/buffer.hpp>
#include <boost/compute/kernel.hpp>
#include <boost/compute/context.hpp>
#include <boost/compute/program.hpp>
#include <boost/compute/command_queue.hpp>
#include <boost/compute/algorithm/transform.hpp>
#include <boost/compute/detail/iterator_range_size.hpp>
#include <boost/compute/utility/program_cache.hpp>
#include <boost/compute/container/vector.hpp>
#include <boost/compute/iterator/discard_iterator.hpp>
namespace boost {
namespace compute {
template<class T = uint_>
class threefry_engine
{
public:
typedef T result_type;
static const ulong_ default_seed = 0UL;
explicit threefry_engine(command_queue &queue,
ulong_ value = default_seed)
: m_key(value),
m_counter(0),
m_context(queue.get_context())
{
load_program();
}
threefry_engine(const threefry_engine<T> &other)
: m_key(other.m_key),
m_counter(other.m_counter),
m_context(other.m_context),
m_program(other.m_program)
{
}
threefry_engine<T>& operator=(const threefry_engine<T> &other)
{
if(this != &other){
m_key = other.m_key;
m_counter = other.m_counter;
m_context = other.m_context;
m_program = other.m_program;
}
return *this;
}
~threefry_engine()
{
}
void seed(ulong_ value, command_queue &queue)
{
(void) queue;
m_key = value;
m_counter = 0;
}
void seed(command_queue &queue)
{
seed(default_seed, queue);
}
template<class OutputIterator>
void generate(OutputIterator first, OutputIterator last, command_queue &queue)
{
const size_t size = detail::iterator_range_size(first, last);
kernel fill_kernel(m_program, "fill");
fill_kernel.set_arg(0, first.get_buffer());
fill_kernel.set_arg(1, static_cast<const uint_>(size));
fill_kernel.set_arg(2, m_key);
fill_kernel.set_arg(3, m_counter);
queue.enqueue_1d_range_kernel(fill_kernel, 0, (size + 1)/2, 0);
discard(size, queue);
}
void generate(discard_iterator first, discard_iterator last, command_queue &queue)
{
(void) queue;
ulong_ offset = std::distance(first, last);
m_counter += offset;
}
template<class OutputIterator, class Function>
void generate(OutputIterator first, OutputIterator last, Function op, command_queue &queue)
{
vector<T> tmp(std::distance(first, last), queue.get_context());
generate(tmp.begin(), tmp.end(), queue);
::boost::compute::transform(tmp.begin(), tmp.end(), first, op, queue);
}
void discard(size_t z, command_queue &queue)
{
generate(discard_iterator(0), discard_iterator(z), queue);
}
private:
void load_program()
{
boost::shared_ptr<program_cache> cache =
program_cache::get_global_cache(m_context);
std::string cache_key =
std::string("__boost_threefry_engine_32x2");
const char source[] =
"#define THREEFRY2x32_DEFAULT_ROUNDS 20\n"
"#define SKEIN_KS_PARITY_32 0x1BD11BDA\n"
"enum r123_enum_threefry32x2 {\n"
" R_32x2_0_0=13,\n"
" R_32x2_1_0=15,\n"
" R_32x2_2_0=26,\n"
" R_32x2_3_0= 6,\n"
" R_32x2_4_0=17,\n"
" R_32x2_5_0=29,\n"
" R_32x2_6_0=16,\n"
" R_32x2_7_0=24\n"
"};\n"
"static uint RotL_32(uint x, uint N)\n"
"{\n"
" return (x << (N & 31)) | (x >> ((32-N) & 31));\n"
"}\n"
"struct r123array2x32 {\n"
" uint v[2];\n"
"};\n"
"typedef struct r123array2x32 threefry2x32_ctr_t;\n"
"typedef struct r123array2x32 threefry2x32_key_t;\n"
"threefry2x32_ctr_t threefry2x32_R(unsigned int Nrounds, threefry2x32_ctr_t in, threefry2x32_key_t k)\n"
"{\n"
" threefry2x32_ctr_t X;\n"
" uint ks[3];\n"
" uint i; \n"
" ks[2] = SKEIN_KS_PARITY_32;\n"
" for (i=0;i < 2; i++) {\n"
" ks[i] = k.v[i];\n"
" X.v[i] = in.v[i];\n"
" ks[2] ^= k.v[i];\n"
" }\n"
" X.v[0] += ks[0]; X.v[1] += ks[1];\n"
" if(Nrounds>0){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_0_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>1){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_1_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>2){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_2_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>3){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_3_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>3){\n"
" X.v[0] += ks[1]; X.v[1] += ks[2];\n"
" X.v[1] += 1;\n"
" }\n"
" if(Nrounds>4){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_4_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>5){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_5_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>6){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_6_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>7){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_7_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>7){\n"
" X.v[0] += ks[2]; X.v[1] += ks[0];\n"
" X.v[1] += 2;\n"
" }\n"
" if(Nrounds>8){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_0_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>9){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_1_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>10){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_2_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>11){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_3_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>11){\n"
" X.v[0] += ks[0]; X.v[1] += ks[1];\n"
" X.v[1] += 3;\n"
" }\n"
" if(Nrounds>12){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_4_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>13){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_5_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>14){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_6_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>15){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_7_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>15){\n"
" X.v[0] += ks[1]; X.v[1] += ks[2];\n"
" X.v[1] += 4;\n"
" }\n"
" if(Nrounds>16){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_0_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>17){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_1_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>18){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_2_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>19){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_3_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>19){\n"
" X.v[0] += ks[2]; X.v[1] += ks[0];\n"
" X.v[1] += 5;\n"
" }\n"
" if(Nrounds>20){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_4_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>21){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_5_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>22){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_6_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>23){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_7_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>23){\n"
" X.v[0] += ks[0]; X.v[1] += ks[1];\n"
" X.v[1] += 6;\n"
" }\n"
" if(Nrounds>24){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_0_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>25){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_1_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>26){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_2_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>27){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_3_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>27){\n"
" X.v[0] += ks[1]; X.v[1] += ks[2];\n"
" X.v[1] += 7;\n"
" }\n"
" if(Nrounds>28){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_4_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>29){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_5_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>30){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_6_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>31){ X.v[0] += X.v[1]; X.v[1] = RotL_32(X.v[1],R_32x2_7_0); X.v[1] ^= X.v[0]; }\n"
" if(Nrounds>31){\n"
" X.v[0] += ks[2]; X.v[1] += ks[0];\n"
" X.v[1] += 8;\n"
" }\n"
" return X;\n"
"}\n"
"__kernel void fill(__global uint * output,\n"
" const uint output_size,\n"
" const uint2 key,\n"
" const uint2 counter)\n"
"{\n"
" uint gid = get_global_id(0);\n"
" threefry2x32_ctr_t c;\n"
" c.v[0] = counter.x + gid;\n"
" c.v[1] = counter.y + (c.v[0] < counter.x ? 1 : 0);\n"
"\n"
" threefry2x32_key_t k = { {key.x, key.y} };\n"
"\n"
" threefry2x32_ctr_t result;\n"
" result = threefry2x32_R(THREEFRY2x32_DEFAULT_ROUNDS, c, k);\n"
"\n"
" if(gid < output_size/2)\n"
" {\n"
" output[2 * gid] = result.v[0];\n"
" output[2 * gid + 1] = result.v[1];\n"
" }\n"
" else if(gid < (output_size+1)/2)\n"
" output[2 * gid] = result.v[0];\n"
"}\n";
m_program = cache->get_or_build(cache_key, std::string(), source, m_context);
}
ulong_ m_key; ulong_ m_counter;
context m_context;
program m_program;
};
} }
#endif