5 ms·
For 32 bit floats, you can skip the math and just test all of them. LLVM will vectorise and unroll this nicely. fn main() { let start = std::time::
by cormacrelf 3y ago
For 32 bit floats, you can skip the math and just test all of them. LLVM will vectorise and unroll this nicely.
fn main() {
let start = std::time::Instant::now();
let total = (0..=u32::MAX)
.filter(|&x| {
let f = f32::from_bits(x);
0. <= f && f <= 1.
})
.count();
println!("total {total} in {:?}", start.elapsed());
}
total 1065353218 in 1.364751583s
Edit: Apparently if you move the sum to its own function it runs in 500ms. A bit temperamental.
Edit 2: it's the size of the sum accumulator that makes it slow. The version above is like `.fold(0usize, |a, _| a + 1)`. When I moved it to another function, I cast the return value to u32, so LLVM saw basically `.fold(0u32, |a, _| a + 1)` and could use u32 throughout. Godbolt says the usize version ends up with floats in xmm* registers on x86, which fit 4 32-bit floats, but the u32 version ends up with floats in ymm* registers (8 32-bit floats) and similar half-as-wide behaviour on ARM.
- raphlinus 3y agoThis is 0x3f800002, which should look pretty familiar to people who work with floating point in hex; it's 2 + the representation of 1.0. To understand the lowest order bits, you've got -0.0 (0x80000000) in there, as well as everything from 0.0 (0x00000000) to 1.0 (0x3f800000) inclusive. This is a larger number than the one cited in the post because the latter only includes "normal" numbers. This includes the "denormals" as well. It's very easy to generalize this to 64 bits, there are 4,607,182,418,800,017,410 of them.
- pansa2 3y agoI did the same thing, recalling "There are Only Four Billion Floats–So Test Them All!" [0]. Except I used Python, so it took over 10 minutes: >>> import struct >>> count = 0 >>> for i in range(0x1_0000_0000): ... f = struct.unpack('<f', struct.pack('<I', i))[0] ... if 0 <= f <= 1: count += 1 ... >>> count 1065353218 [0] https://randomascii.wordpress.com/2014/01/27/theres-only-four-billion-floatsso-test-them-all/ https://randomascii.wordpress.com/2014/01/27/theres-only-fou...
- jherskovic 3y agoMy very-slightly-different Python implementation runs in 9m56s minutes on Python 3.11, but in 2m54s using pypy3.9 (all on an M1 Max) pypy never ceases to amaze me for low-effort performance gains on computational stuff. A naive implementation in C compiled with Apple's clang (-O2) takes 0.62 seconds, of course.
- eesmith 3y agoThis should be a bit faster, while still sticking to stock Python: import time, array n = 0 start_time = time.time() values = bytearray(b"".join(i.to_bytes(4, "big") for i in range(2**24))) for prefix in range(256): values[::4] = prefix.to_bytes(1, "big") * (2**24) arr = array.array("f", values) n += sum(1 for x in arr if 0.0 <= x <= 1.0) print(f"Found: {n} Time: {time.time() - start_time:.1f} seconds") It uses an array.array() to convert from a large number of concatenated 4 bytes to floats, rather than going through the struct module one-by-one. Rather than build all 0x1_0000_0000 values in memory, I process 0x100_0000 at a time, and accumulate the partial sums. On my laptop this reports "Found: 1065353218 Time: 228.5 seconds", so just under 4 minutes. A slightly improved algorithm in NumPy takes 5.4 seconds. import time, numpy as np start_time = time.time() arr = np.arange(2**24, dtype=np.uint32) incr = np.full(2**24, 2**24, dtype=arr.dtype) values = np.frombuffer(arr, np.float32) n = ((0.0 < values) & (values <= 1.0)).sum() for i in range(1, 256): arr += incr n += ((0.0 < values) & (values <= 1.0)).sum() print(f"Found: {n} Time: {time.time() - start_time:.1f} seconds")
- enriquto 3y agoIn standard C you have nextafterf(3) that gives the next float and allows to traverse them starting from 0: #include <math.h> // nextafterf #include <stdio.h> // printf int main() { long n = 1; float f = 0; while (f <= 1) { f = nextafterf(f, 2); n = n + 1; } printf("%ld\n", n); return 0; }
- flerchin 3y agoIt runs in ~300ms in java long start = System.currentTimeMillis(); int count = 1; for (int intBits = Float.floatToIntBits(0.0f); Float.intBitsToFloat(intBits) <= 1.0f; intBits ++) { count++; } System.out.println(count); System.out.println((System.currentTimeMillis() - start) + "ms");
- flerchin 3y agoAlternatively, it can be solved at compile time with the streams api long start = System.currentTimeMillis(); long count = IntStream.rangeClosed(Float.floatToIntBits(0.0f), Float.floatToIntBits(1.0f)).count(); System.out.println(count); System.out.println((System.currentTimeMillis() - start) + "ms"); 1065353217 0ms
- refulgentis 3y agoAt compile time!? Really cool trick! I can't tell that'd be possible - I see the arguments could be evaluated at compile time, but knowing IntStream.rangeClosed will be evaluated at compile time is a leap, to me. Did you know before you wrote the code that'd happen? Is it like C? You cross your fingers and hope the compiler unrolls?
- messe 3y agoThe key is that the length of IntStream.rangeClosed(a, b) is just going to be equal to (b - a) + 1. So there's actually no reason to even invoke the streams API there.
- SpaghettiCthulu 3y ago> So there's actually no reason to even invoke the streams API there. While technically correct, in practice it is not the case. First of all (and this relates to flerchin's comment), it's not actually done at "compile time" in the traditional sense (i.e. when you invoke `javac`). Rather, it is done at runtime when the method is compiled by the JIT. As for this comment, the call to `IntStream.rangeClosed` will not simply be reduced to `(b - a) + 1` as you suggest (and one reason for this is that streams are far more complex internally than regular iterators). In reality, it will just be (potentially) inlined and then further optimized alongside the rest of the method's code. Edit: the mentioned previous comment was from flerchin, not meese Edit 2: I might have misunderstood what you were getting at there. You are sort of correct in that the `count` operation on this stream is optimized, but it is still technically going through the streams API.
- version_five 3y ago> LLVM will vectorise and unroll this nicely Curious to know the total compile + run time under different compiler optimizations. For code you only need to run once, I don't see how having the compiler unroll the loops actually saves you any time.
- cormacrelf 3y agoUnrolling loops reduces the cost of loop termination checking. A tight loop like `for (i=0; i < n; i++) { acc += 1; }` is mostly overhead, you are constantly checking to see if you’re done or not. If you were given a loop like this in real life, like “come back to me and check if you’ve made enough doughnuts after every single individual doughnut you make, I’ll tell you when you’re done”, you would just quit your job in frustration, because it is a huge waste of your time. The compiler can unroll loops like that to `for (i=0; i < n; i+=8) { acc += 8; }` (ignoring the remainder of n/8), which does 8x fewer compare+conditional jump instructions. It has nothing to do with how many times you’ll be running the code; this kind of optimisation reduces the runtime of any single execution.
- version_five 3y agoThanks for the explanation! So how would that play out in optimizing the code you posted? It will run all 2^32 iterations without checking for completion? How does it know when it's finished, and is that the only "unrolling" optimization? I had thought that in cases like this (where the output can already be determined at compile time) the compiler could just run the loop to get the answer, therefore avoiding running it every time the code runs, which was what I was picturing when I wrote my earlier comment.
- cormacrelf 3y agoHere is an annotated version of rustc's `--release --target aarch64-apple-darwin` assembly output, for the fast version of my code that runs in 500ms. Annotated https://gist.github.com/cormacrelf/4b2b37d1d870377d548b3a32a3c7b26a https://gist.github.com/cormacrelf/4b2b37d1d870377d548b3a32a... and full Godbolt https://rust.godbolt.org/z/aKK4TEs3x https://rust.godbolt.org/z/aKK4TEs3x. You may be interested to see the LLVM Opt Pipeline, any of the green passes (non-zero diffs) that are called LoopUnrollPass or LoopVectorizePass. You can see that it does a couple of LoopUnrollPasses and LoopVectorizePasses. In the end it results in a transformation from something like this: for (uint32_t i = 0; i <= 0xffff; i++) { if (in range) acc += 1; } to u32x4 v1 = [1,2,3,4]; u32x4 v0 = [0,0,0,0], v5 = [0,0,0,0]; u32x4 v6, v7; uint32_t w8 = 0xfff8; do { v6 = v1 + [4,4,4,4]; v7 = v1 >= 0 && v1 <= 1; v6 = v6 >= 0 && v6 <= 1; v0 -= v7; // add 1 to each lane if lane was true in v7 v5 -= v6; // add 1 to each lane if lane was true in v6 v1 += [8,8,8,8]; w8 -= 8; } while (w8); return element_sum(v0 + v5); That is basically unrolled x8, and then the loop body is grouped into two groups of four vector ops. You also save a lot of time by vectorising the sum operation over 8 lanes. The only non-vector operations in the loop body are on the w8 counter. The ratio of real work to "increment checking" is 8x better, and because we unrolled multiple ops into the loop body, we were able to vectorise it and run it in parallel. Re "how does it know when it's finished?": it counts like normal, as you can see there, it uses a register to count iteration. Unrolling does not mean "completely eliminate the termination check" unless your heuristic/supplied unroll threshold is larger than the (known) number of iterations. It generally just does the termination check fewer times by taking bigger strides. It's like walking up stairs two at a time. The answer re compile-time-full-evaluation: compilers use heuristics and pattern matching, they don't automatically try compile time evaluation for things they don't recognise. Nobody would have bothered to write a heuristic to detect "count floating point values in range" and eval it. There is no compiler pass to arbitrarily try to execute eveything in your program at compile time that's remotely possible.
- fm77 3y agoIn Turbo Pascal :-) you have to filter for NaN or else you will end up with a Runtime Error. const NaN = $FF shl 23; var x, t: longint; f: single absolute x; begin t := 0; x := 0; repeat inc(t, ord((x and NaN <> NaN) and (0<=f) and (f<=1))); inc(x); until x = 0; WriteLn('total: ', t); end. total: 1065353218 (in ca. 40 seconds)
- zX41ZdbW 3y agoIf you write it as an SQL query in ClickHouse, it will take only 0.170 seconds: SELECT sum(reinterpretAsFloat32(number) BETWEEN 0.0 AND 1.0) FROM numbers_mt(0x100000000) 1 row in set. Elapsed: 0.170 sec. Processed 2.39 billion rows, 19.08 GB (14.02 billion rows/s., 112.13 GB/s.)