正如聊天中所宣布的,我添加了一个精致的答案。它包含三个部分,每个部分后面都有对该部分的描述。
第 1 部分 get.h 是我的解决方案,但很笼统,有一处更正。
第二部分 got.h 是问题中发布的原始算法,也适用于任何无符号类型的任何 STL 容器。
第 3 部分 main.cpp 包含用于验证正确性和衡量性能的单元测试。
#include <cstddef>
using std::size_t;
template < typename C >
typename C::value_type get ( C const &container, size_t lo, size_t hi )
{
typedef typename C::value_type item; // a container entry
static unsigned const bits = (unsigned)sizeof(item)*8u; // bits in an item
static size_t const mask = ~(size_t)0u/bits*bits; // huge multiple of bits
// everthing above has been computed at compile time. Now do some work:
size_t lo_adr = (lo ) / bits; // the index in the container of ...
size_t hi_adr = (hi-(hi>0)) / bits; // ... the lower or higher item needed
// we read container[hi_adr] first and possibly delete the highest bits:
unsigned hi_shift = (unsigned)(mask-hi)%bits;
item hi_val = container[hi_adr] << hi_shift >> hi_shift;
// if all bits are in the same item, we delete the lower bits and are done:
unsigned lo_shift = (unsigned)lo%bits;
if ( hi_adr <= lo_adr ) return (hi_val>>lo_shift) * (lo<hi);
// else we have to read the lower item as well, and combine both
return ( hi_val<<bits-lo_shift | container[lo_adr]>>lo_shift );
}
第一部分,上面的 get.h,是我最初的解决方案,但可以泛化为与任何无符号整数类型的 STL 容器一起使用。因此,您也可以使用和测试 32 位整数或 128 位整数。对于非常小的数字,我仍然使用 unsigned,但您也可以将它们替换为 size_t。该算法几乎没有变化,只是稍作修正——如果 lo 是容器中的总位数,我之前的 get() 将访问容器大小上方的项目。现在已修复。
#include <cstddef>
using std::size_t;
// Shift right, but without being undefined behaviour if n >= 64
template < typename val_type >
val_type shiftr(val_type val, size_t n)
{
val_type good = n < sizeof(val_type)*8;
return good * (val >> (n * good));
}
// Shift left, but without being undefined behaviour if n >= 64
template < typename val_type >
val_type shiftl(val_type val, size_t n)
{
val_type good = n < sizeof(val_type)*8;
return good * (val << (n * good));
}
// Mask the word preserving only the lower n bits.
template < typename val_type >
val_type lowbits(val_type val, size_t n)
{
val_type mask = shiftr<val_type>((val_type)(-1), sizeof(val_type)*8 - n);
return val & mask;
}
// Struct for return values of locate()
struct range_location_t {
size_t lindex; // The word where is located the 'begin' position
size_t hindex; // The word where is located the 'end' position
size_t lbegin; // The position of 'begin' into its word
size_t llen; // The length of the lower part of the word
size_t hlen; // The length of the higher part of the word
};
// Locate the one or two words that will make up the result
template < typename val_type >
range_location_t locate(size_t begin, size_t end)
{
size_t lindex = begin / (sizeof(val_type)*8);
size_t hindex = end / (sizeof(val_type)*8);
size_t lbegin = begin % (sizeof(val_type)*8);
size_t hend = end % (sizeof(val_type)*8);
size_t len = (end - begin) * size_t(begin <= end);
size_t hlen = hend * (hindex > lindex);
size_t llen = len - hlen;
range_location_t l = { lindex, hindex, lbegin, llen, hlen };
return l;
}
// Main function.
template < typename C >
typename C::value_type got ( C const&container, size_t begin, size_t end )
{
typedef typename C::value_type val_type;
range_location_t loc = locate<val_type>(begin, end);
val_type low = lowbits<val_type>(container[loc.lindex] >> loc.lbegin, loc.llen);
val_type high = shiftl<val_type>(lowbits<val_type>(container[loc.hindex], loc.hlen), loc.llen);
return high | low;
}
上面的第二部分 got.h 是问题中的原始算法,也被概括为接受任何无符号整数类型的任何 STL 容器。与 get.h 一样,该版本除了定义容器类型的单个模板参数外不使用任何定义,因此可以轻松测试其他项目大小或容器类型。
#include <vector>
#include <cstddef>
#include <stdint.h>
#include <stdio.h>
#include <sys/time.h>
#include <sys/resource.h>
#include "get.h"
#include "got.h"
template < typename Container > class Test {
typedef typename Container::value_type val_type;
typedef val_type (*fun_type) ( Container const &, size_t, size_t );
typedef void (Test::*fun_test) ( unsigned, unsigned );
static unsigned const total_bits = 256; // number of bits in the container
static unsigned const entry_bits = (unsigned)sizeof(val_type)*8u;
Container _container;
fun_type _function;
bool _failed;
void get_value ( unsigned lo, unsigned hi ) {
_function(_container,lo,hi); // we call this several times ...
_function(_container,lo,hi); // ... because we measure ...
_function(_container,lo,hi); // ... the performance ...
_function(_container,lo,hi); // ... of _function, ....
_function(_container,lo,hi); // ... not the performance ...
_function(_container,lo,hi); // ... of get_value and ...
_function(_container,lo,hi); // ... of the loop that ...
_function(_container,lo,hi); // ... calls get_value.
}
void verify ( unsigned lo, unsigned hi ) {
val_type value = _function(_container,lo,hi);
if ( lo < hi ) {
for ( unsigned i=lo; i<hi; i++ ) {
val_type val = _container[i/entry_bits] >> i%entry_bits & 1u;
if ( val != (value&1u) ) {
printf("lo=%d hi=%d [%d] is'nt %d\n",lo,hi,i,(unsigned)val);
_failed = true;
}
value >>= 1u;
}
}
if ( value ) {
printf("lo=%d hi=%d value contains high bits set to 1\n",lo,hi);
_failed = true;
}
}
void run ( fun_test fun ) {
for ( unsigned lo=0; lo<total_bits; lo++ ) {
unsigned h0 = 0;
if ( lo > entry_bits ) h0 = lo - (entry_bits+1);
unsigned h1 = lo+64;
if ( h1 > total_bits ) h1 = total_bits;
for ( unsigned hi=h0; hi<=h1; hi++ ) {
(this->*fun)(lo,hi);
}
}
}
static uint64_t time_used ( ) {
struct rusage ru;
getrusage(RUSAGE_THREAD,&ru);
struct timeval t = ru.ru_utime;
return (uint64_t) t.tv_sec*1000 + t.tv_usec/1000;
}
public:
Test ( fun_type function ): _function(function), _failed() {
val_type entry;
unsigned index = 0; // position in the whole bit array
unsigned value = 0; // last value assigned to a bit
static char const entropy[] = "The quick brown Fox jumps over the lazy Dog";
do {
if ( ! (index%entry_bits) ) entry = 0;
entry <<= 1;
entry |= value ^= 1u & entropy[index/7%sizeof(entropy)] >> index%7;
++index;
if ( ! (index%entry_bits) ) _container.push_back(entry);
} while ( index < total_bits );
}
bool correctness() {
_failed = false;
run(&Test::verify);
return !_failed;
}
void performance() {
uint64_t t1 = time_used();
for ( unsigned i=0; i<999; i++ ) run(&Test::get_value);
uint64_t t2 = time_used();
printf("used %d ms\n",(unsigned)(t2-t1));
}
void operator() ( char const * name ) {
printf("testing %s\n",name);
correctness();
performance();
}
};
int main()
{
typedef typename std::vector<uint64_t> Container;
Test<Container> test(get<Container>); test("get");
Test<Container> tost(got<Container>); tost("got");
}
上面的第 3 部分 main.cpp 包含一类单元测试并将它们应用于 get.h 和 got.h,即应用于我的解决方案和问题的原始代码,稍作修改。单元测试验证正确性并测量速度。他们通过创建一个 256 位的容器来验证正确性,用一些数据填充它,读取所有可能的位部分(最多可容纳到容器条目中的位数以及许多病理情况),并验证每个结果的正确性。他们通过再次经常阅读相同的部分并报告线程在用户空间中使用的时间来衡量速度。