In the last post we got to a basic implementation of the radix 2 decimation-in-time Cooley-Tukey algorithm. But there are still other additions and optimizations we can make before
getting into the way FFT-20XX works. We’ll start by diving into the “radix 2” part of the algorithm: we can make
improvements there.
Butterflies
Before we go into the specifics, I’m finally going to bring in a classic FFT diagram: the butterfly, an FFT
data flow diagram. The most basic one is for a FFT of length 2, which has only a single step over
two inputs. Here’s a diagram for the radix-2 decimation-in-time 2-length FFT:
In this diagram, the inputs are on the left, each computation (just one, in this example) is represented as a small circle,
and the outputs are on the right. The computation (for a decimation-in-time FFT) rotates the bottom input by the given
rotation value; in this diagram it’s just a , but in general the value there will be some where the actual
rotation angle is clockwise turns ( radians).
Note
The formatting of these butterfly diagrams is a little different than a standard butterfly diagram, which tend to
look more like the following:
These diagrams are more comprehensive (the value makes the whole turn fraction of clockwise turns clear,
where in my diagrams the “” is implied), it has a multiplier on the lower input going into the sum circle
(to signify that the bottom “sum” is a difference), and the sums are each on their own line as opposed to a central
circle that represents the add/subtract pair, but in practice – especially on lower-resolution displays – my variant
is, imo, less visually noisy.
Plus, once you get into the higher radixes, even the standard diagrams start to be drawn more simply.
Now here’s one for the decimation-in-time 4-length FFT (where the bit reverse is performed on the input):
There are two “phases” in this radix-4 FFT (two runs of the outer loop of the code), which are separated via the
vertical dotted lines. Really, though, this is four of the previous (length 2 FFT) butterflies bolted together as two
separate 2-element computations (the first pass) followed by a single 4-element computation loop (the second pass,
itself just two 2-element computations stacked on top of each other wearing a trenchcoat).
Finally, here’s a butterfly diagram for a full 16-element FFT:
There are four phases in this one, and note that you can still see the structure discussed in the last post of 8
2-element computations in the first phase, then four 4-element computations in the next, then two 8-element
computations, then a single 16-element computation. Also note the rotation angles in each phase: the last phase has
rotations 0-7 (corresponding to to turns), the previous one has just the even numbers, the one before
that just 0 and 4 (every other even value), down to no rotations in the first phase.
Note
The first radix-2 phase of any decimation-in-time FFT will always have no rotations, and the second phase will only have
0- or 90-degree rotations (a in the above diagram, as it’s turns, or ). When it comes time to start
optimizing, these first two phases can be done with no multiplications at all (obvious when there’s no rotation, but
a 90 degree (clockwise) rotation is just taking and changing it to ), and can typically be special-
cased for a significant speed bost.
Doubling the Radix
We went with a radix 2 breakdown in the last post because it’s the most straightforward, but there are others. The one
we’re going to focus on here is radix 4, which, in contrast to the above 16-element butterfly diagram, would instead
look like the following:
Don’t worry about the specifics of the rotation values or what this four-in/four-out computation node
is doing yet – we’ll get to that shortly. For now, the important thing is that there are now half as many passes over
the data, and each computation handles four values at a time instead of two.
This will end up being more efficient in a couple different ways. The main efficiency gain is that having half as many
passes means it halves the number of reads from and writes to memory, which as the FFT lengths get larger can make
a big difference in how often there are cache misses. Second, it is possible to reduce the number of rotations
needed by , which can also make a difference on some CPUs (but not all, which we’ll get to).
To derive a radix 4 solution, you could go back to the original equation, split the inputs into four interleaved
sections (where is , , , and respectively), then split the outputs into quarters, and redo all
of the rotation shenanigans that we did when deriving radix-2.
Nah. I’m not doing all that.
Instead, we’re going to start with the radix-2 diagram and build up to what that four-in/four-out computation looks
like. Every pair of radix-2 passes can be turned into a single radix-4 pass, so let’s start by highlighting a couple
such groupings in our radix-2 diagram:
Each of the groupings of “two pairs radix-2 passes that all use the same 4 values” can be represented generically as:
The input/output indices are relative to the position in the FFT, but note that the middle two indices are in the
opposite order as the output indices: this will always be true when we’ve done the bit reverse at the outset, and
remains true no matter which intermediate step we’re on (i.e. it just automatically happens that way).
Also note that the angles themselves are relative, too: there will be some base angle (in the first highlighed group
at input indices 0, 8, 4, and 12, is . In the second at output indices “2, 6, 10, and 14”, is ) that the
other angles are relative to, and we’ll set it as the first angle in the second radix-2 pass. Both angles in the
first radix-2 pass are double this value (both labeled ), and the second angle in the last pass is always a
quarter turn more than (labeled ).
To get from here to a full radix-4, we can rearrange the angles a bit. If we separate the from the in
the last angle (taking advantage of the fact that, since we’re using complex numbers, Rotate(a + b) is the same as
Rotate(a) * Rotate(b)), there is a common factor to both of the second radix-2 angles, and we can distribute that
back to the latter two inputs to the first radix-2 step (the first of which has no rotation initially so it just becomes
a rotation of , and the latter of which becomes effectively Rotate(2t) * Rotate(t) which, as just explained,
becomes Rotate(2t + t) or Rotate(3t):
As mentioned in an above note, a quarter turn (90-degree clockwise rotation) can be done without any sine/cosine
multiplications at all (it transforms into ), so this is mathematically more efficient: it does
3 rotations on 3 inputs at the start of the process instead of 4 rotations (split between the first and second
inner radix-2 phases). The number of adds/subtractions remains the same, but that’s okay: there’s no way to reduce it
any further.
This, then, is what becomes the computation node of the radix-4 butterfly diagram we made above! This, then, is
the equivalent radix 4 diagram to the previous diagram (which hides the internal complexities, including the
effectively-free quarter turn):
Note
An implementation detail worth mentioning here: on some systems using the first 4-element breakdown from above (the
one with the term) can be more efficient! There are systems that have FMA instructions (the
fused multiply-add)
which have a cost close to doing either a multiply or an add by itself (x64 CPUs with the FMA extension – like
most with AVX2 instructions – have this property). Doing it that way (which is sometimes referred to as radix-),
even with what seems like an extra complex multiply, can end up being cheaper due to less overall operations and
one less required angle (the angle isn’t required, using and is sufficient).
Here’s the 16-element radix-4 diagram we used earlier, with our same two radix-4 sections highlighted:
Note
This formulation of radix-4 is slightly different than the canonical one, which has a different “scrambled” input
or output order from the bit reverse that we’re using. In that case, the reordered indices are reversed pairs of
bits: for instance, an 8-bit binary value with bits 01234567 would becomes 67452301 instead of 76543210.
Instead, we’ve kept the radix-2 style bit reverse, the consequence of which is the reversal of the middle 2 inputs
in the above diagrams (with reversed bit pairs, those would be in proper order, and the angles would instead be ,
, and , respectively).
Now, hopefully, it’s clear why the angle values are the way they are in this diagram. The first phase has no rotations
(except for the free 90-degree rotation), and the second phase has four computations with increasing rotation values of
the form , , .
The only issue with using the radix-4 algorithm is that it only fully works with FFTs with power-of-four lengths
(because it splits in fours each time, so it works for length 4 but not 8, 16 but not 32, etc). The good news is that
the way we’ve formulated it (by effectively smooshing together two radix-2 passes into a single radix-4 pass), you can
just put a single radix-2 pass in there to make up the difference (and it doesn’t even matter where!)
Here is an 8-length FFT where the first pass is still radix 2 but the second is radix 4:
And here’s the same 8-length FFT the other way around: radix-4 then radix-2:
For longer FFTs, the radix-2 pass can be anywhere in the middle, as well - whatever is most efficient for your
implementation! I found it easiest to place it right after the very first radix-4 pass (i.e. like that last diagram),
mostly for simplicity.
Other Radix Flavors
So if radix 4 is more efficient than radix 2, that raises the question: are there other, even more efficient radixes?
radix-8? radix-16? How far is too far?
Radix-4 is the largest power-of-two radix that has “free” rotations in the middle (that quarter turn). Radix-8 has some
45-degree-multiple rotations in the middle, which still require some multiplication. However, because both cosine and
sine of 45 degrees (and related friends) are both , the multiply can be done as a common factor so
it still is somewhat more efficient.
The problems with radix-8 tend to have to do with CPU register limits: you need to load 8 complex values (so 16 scalars),
then (similar to how radix-2 needs 1 and radix-4 needs 3) it tends to need 7 complex twiddles plus the
internal value (so another 15 scalars) and you tend to also need some scratch registers as well to
do computations. AVX and AVX2 only have 16 SIMD registers so even before you get to the angles you’ve already run out of
register space. AARCH64 has 32, so you could maybe squeeze in there, but it would be tight. Plus you’re going to be
up against CPU cache associativity for systems with eight-way
associative caches, because you’re loading from 8 locations but also need to be loading the angles from somewhere.
Past radix-8, you’re out of register space, hitting CPU cache associativity limits, and also have more internal
rotations, so it’s going to perform way worse under most circumstances.
Note
One place where it’s really tempting to use radix-8, however, is on the very first decimation-in-time step, which
has no twiddles (all seven of the twiddle angles are 0), and thus only has those nice-to-compute internal 45-degree
rotations. However, the register limits are still an issue and, depending on how you have to load your values (i.e.
for SIMD purposes you may want to do a deinterleave to separate the real and imaginary components), you may still not
have enough space. But it’s worth considering, if you can.
There is another radix option worth mentioning: the split-radix FFT algorithm, which I’m not going to detail here. In terms of required math operations, the split-radix
algorithm gets you to some of the lowest-known possible operation counts. Unfortunately, in practice, I found that it
breaks up the nice regularity of using just a standard radix-4 algorithm and ultimately – at least in my attempts at
implementing it – had worse performance overall.
Okay, we’re through all of the radix shenanigans!
Lookup Tables
The last thing to touch on in this post is doing lookup tables for the angles. These tend to be exactly what they
sound like: you precompute a table for your required FFT length that contains all of the sin/cos values needed for that
length.
The simplest lookup table is one where you have cos/sin values ( turns through
turns), in order. However, there are some potential ways to make the table a little more efficient (with size,
cache-friendliness, or both):
You don’t need both and values (unless you’re doing SIMD with interleaved
real and imaginary components): because , you can store scalar
values and start the cosine lookup partway into the table, and use less memory.
if you’re doing radix-4, you can avoid storing the angles and instead compute them using the and rotations,
which saves a bunch of table space. This may sound like you’re adding a complex multiply back in after eliminating
it, but you can move that computation into an outer loop so it isn’t an issue.
It can also be convenient – for SIMD loads and CPU cache-friendliness – to store not only the angles in order for
size , but to store all of the table sizes from some minimum length (like 4 or 8) up to your max length.
Doing this is nice because it means you can build a single table for your maximum required FFT length and it will
perform just as well for any smaller FFT lengths as well).
Also, the earlier passes of the algorithm can take advantage of these smaller tables (the second radix-2 pass only
needs the angles for a length-4 FFT, and the next only needs the angles for a length-8 FFT, etc). Loading from them
can help with CPU cache coherence (as all of the required angles at that level are contiguous), plus SIMD loading
if you need multiple angles within a single pass.
Note
It’s also worth noting, for some algorithmic variants (i.e. if you do the bit reverse at the end, or if you are doing a decimation-in-frequency implementation), it can be better to store the table in bit-reversed order (for example, for a
16-length FFT you’d store the rotations in order 0 4 2 6 1 5 3 7)! This has a neat property: if you make a table for a
length- FFT, the first entries are also exactly what would be in a table for a length- FFT! This means
the table can be built for the maximum required FFT length, but unlike the last bullet point above, with no extra
memory footprint.
For FFT-20XX I found all of the above bullet points to be useful.
What Now?
No code this time, but using the above tricks (namely: using radix-4 instead of radix-2 where possible and switching
to lookup tables for the rotations) gets you to a somewhat decent baseline FFT implementation! There are many more
places to optimize: figuring out SIMD support and optimizing the bit reversal pass (or – spoiler alert – eliminating
it altogether). We’ll get to those once we finally get into FFT-20XX’s algorithms.
Before that, next time, we’ll touch on the decimation-in-frequency version of the algorithm on our way to finally
talking about the inverse FFT and how to compute it!
Part 1 went through the core of how the
discrete Fourier transform (DFT) works to turn a series of
samples in time into a set of samples in frequency, and ended up with the following algoritm:
voidCalculateDFT(complex[] inputs,complex[] outputs,int length){for(outIdx =0; outIdx < length; outIdx++){// Start the sum at 0.float sum =0// This is the angular frequency (in radians/sample) for// the given output sample.float angularFreq =-2* pi * outIdx / length
for(inIdx =0; inIdx < length; inIdx++){// Get the rotation value for the current input index// then use that to calculate the resulting rotation.float radians = inIdx * angularFreq
complex rotation =cos(radians)+ i*sin(radians)// A complex multiply with a unit complex vector is// a 2D rotation.
sum += inputs[inIdx]* rotation
}
outputs[outIdx]= sum
}}
Note
A reminder: we’re going to be talking about these values as complex numbers (of the form ), but if you prefer,
just think of them as a 2D vector instead ( where is the x coordinate and is the y coordinate).
We’ll be doing complex multiplication of these vectors with others, but in all cases a multiply like will
be a rotation of point by the angle that represents (where ).
This algorithm is , but there’s a much more efficient, way: the
fast Fourier transform (FFT) – and we’re going to derive it.
Note
The algorithm we end up deriving here is going to work specifically for FFT lengths that are powers of two,
because it’s the easiest to implement (plus, to be honest, power-of-two lengths were sufficient for FFT-20XX so that’s all I
implemented). The following techniques can be adapted for other lengths (i.e. if you need a length that is a mutliple of 3
or 5), plus there are ways to do this for arbitrary lengths (for instance:
Bluestein’s algorithm). This series isn’t going to touch on any of those.
Halving the Work
The first step of deriving the FFT algoritm is figuring out where we can halve the amount of work. Since the DFT boils down
to “every output takes every input multiplied by a rotation”, we can visualize all of the rotations as a 2D grid, where
each column is an input and each row is an output.
Here’s a visualization of the rotations for an
FFT of length 8 (adding a dividing line between the top and bottom halves, which should make sense in a moment):
in[0]
[0]
in[1]
[1]
in[2]
[2]
in[3]
[3]
in[4]
[4]
in[5]
[5]
in[6]
[6]
in[7]
[7]
outO[0] =
+
+
+
+
+
+
+
outO[1] =
+
+
+
+
+
+
+
outO[2] =
+
+
+
+
+
+
+
outO[3] =
+
+
+
+
+
+
+
outO[4] =
+
+
+
+
+
+
+
outO[5] =
+
+
+
+
+
+
+
outO[6] =
+
+
+
+
+
+
+
outO[7] =
+
+
+
+
+
+
+
Each rotation (as per the above pseudocode) for a given input index and output index is turns (where
is 8, the FFT length). But rotations are cyclical, and, for instance, turning a full turn (360 degrees) is
equivalent to not turning at all.
We can use this cyclicity to our advantage! If you look at the even inputs (columns indexed 0, 2, 4, and 6), you’ll note that the
rotations in the bottom half are identical to the top half. For instance, with input 2, outputs 0 and 4 are both
unrotated, outputs 1 and 5 both rotate a quarter turn, and so on.
For the odd inputs (columns 1, 3, 5, 7), the angles aren’t the same. However, they’re still related: the
bottom half rotations are always 180 degrees from the corresponding top half rotations.
As an example, looking at the column for input 1, we see that output 4’s rotation is a half turn, where output 0 is
unrotated – they’re facing away from each other. The other odd outputs are similar, with the bottom half rotations
always pointing the opposite direction as their top-half counterparts.
In other words, the bottom half rotations are negated versions of the top half (because
with 2D vector A, Rotate(A, 180 degrees) is the same as the 2D vector -A).
Let’s rearrange the columns of the diagram to group the even and odd elements together (evens to the left, odds to the right):
in[0]
[0]
in[2]
[2]
in[4]
[4]
in[6]
[6]
in[1]
[1]
in[3]
[3]
in[5]
[5]
in[7]
[7]
outO[0] =
+
+
+
+
+
+
+
outO[1] =
+
+
+
+
+
+
+
outO[2] =
+
+
+
+
+
+
+
outO[3] =
+
+
+
+
+
+
+
outO[4] =
+
+
+
+
+
+
+
outO[5] =
+
+
+
+
+
+
+
outO[6] =
+
+
+
+
+
+
+
outO[7] =
+
+
+
+
+
+
+
We can then rewrite the bottom-half odd inputs as negations of the top-half ones (note the minus signs in place of the
usual plus signs):
in[0]
[0]
in[2]
[2]
in[4]
[4]
in[6]
[6]
in[1]
[1]
in[3]
[3]
in[5]
[5]
in[7]
[7]
outO[0] =
+
+
+
+
+
+
+
outO[1] =
+
+
+
+
+
+
+
outO[2] =
+
+
+
+
+
+
+
outO[3] =
+
+
+
+
+
+
+
outO[4] =
+
+
+
-
-
-
-
outO[5] =
+
+
+
-
-
-
-
outO[6] =
+
+
+
-
-
-
-
outO[7] =
+
+
+
-
-
-
-
Now all of the corresponding top and bottom half outputs use the same corresponding rotations per input; the
only difference is that in the bottom half we subtract the rotated odd inputs instead of adding them.
Our code, then, can switch to summing up all of the even and odd elements separately, then computing two outputs at once
by doing evenSum + oddSum for the first (top) half of the outputs and evenSum - oddSum for the second (bottom) half,
which means we’re now doing half the work as before!
Note
If you want this (and the next section) done algebraically instead of visually, the wiki page for the
Cooley-Tukey algorithm has a fairly clear breakdown
of the derivation of the whole thing.
Here’s what the code looks like if we do that:
voidCalculateDFT(complex[] inputs,complex[] outputs,int length){// NOTE: Now looping over half the lengthfor(outIdx =0; outIdx < length /2; outIdx++){// Start the sums at 0.float evenSum =0float oddSum =0// angular frequency for the given output// (in radians/sample)float angularFreq =-2* pi * outIdx / length
// NOTE: now incrementing by 2 to do odds and evens// separately.for(inIdx =0; inIdx < length; inIdx +=2){// Get the rotation value for the current input index// then use that to calculate the resulting rotation.float evenRads = inIdx * angularFreq
complex evenRot =cos(radians)+ i*sin(radians)
evenSum += inputs[inIdx +0]* evenRot
// Do the same with the next index (inIdx + 1), for// the odd sum.float oddRads =(inIdx +1)* angularFreq
complex oddRot =cos(radians)+ i*sin(radians)
oddSum += inputs[inIdx +1]* oddRot
}// Computing 2 outputs in the same loop!
outputs[outIdx]= evenSum + oddSum
outputs[outIdx + length/2]= evenSum - oddSum
}}
If you compare this code to what we started with, it does half the rotations and half the adds (counting
subtractions) as the original: the inner loop does effectively the same amount of work in both, but the outer loop runs
half as much. Unfortunately, it’s not quite to yet, so we need a way to effectively do this halving of
work recursively.
Halving Work Again (and Again and…)
Let’s look again at just the first half of the outputs from our 8-length example:
in[0]
[0]
in[2]
[2]
in[4]
[4]
in[6]
[6]
in[1]
[1]
in[3]
[3]
in[5]
[5]
in[7]
[7]
outO[0] =
+
+
+
+
+
+
+
outO[1] =
+
+
+
+
+
+
+
outO[2] =
+
+
+
+
+
+
+
outO[3] =
+
+
+
+
+
+
+
If we look at just the even inputs for a moment, the pattern of rotations in that 4×4 block of the diagram ends up
being the exact same pattern of rotations as for a standard 4-element FFT (half the size of our 8-element example):
in[0]
in[1]
in[2]
in[3]
out[0] =
+
+
+
out[1] =
+
+
+
out[2] =
+
+
+
out[3] =
+
+
+
Note
If you care about the algebra: this is because for each even rotation (), is a multiple of 2. So if
we say , then the rotation becomes which is equivalent to , which is exactly how
the angles for an FFT of length would be specified.
This means that we can calculate the even sums recursively, computing the even sums as if they were an
-length DFT, which means we can do the same even/odd dance to halve its work as well.
That’s great for the even inputs, but what about the odd inputs? They have a similar pattern, but it’s not quite the
same:
in[1]
in[3]
in[5]
in[7]
odd[0] =
+
+
+
odd[1] =
+
+
+
odd[2] =
+
+
+
odd[3] =
+
+
+
The first row is the same as in the 4-element DFT: no rotation. The second row, however, is different: each angle in the
row is 1/8th of a turn rotated relative to the corresponding row of the 4-element DFT. The third and fourth rows are
also different, with an extra 2/8th and 3/8th turn added to the angles in each row, respectively.
In other words, any given output ’s odd inputs will have their rotations be turns more than doing a
standard -length FFT with those inputs.
Since “rotating every vector by t then summing” is the same as “summing then rotating by t” (that is, Rotate(A, t) + Rotate(B, t) equals
Rotate(A + B, t)), we can calculate each odd sum by doing the work as a half-length DFT (like we can with
the evens), then adjust each sum by rotating by its corresponding -turn angle (called a twiddle factor, which is a fun
name to say), like so:
in[1]
in[3]
in[5]
in[7]
odd[0] =
×
+
+
+
odd[1] =
×
+
+
+
odd[2] =
×
+
+
+
odd[3] =
×
+
+
+
Using that, we can now recursively calculate the even and odd sums using half-length FFTs, apply the twiddles, then do
the add and subtract to get the final outputs.
Let’s write this as code again, this time recursively! To do this, for now we’ll use some temporary storage for simplicity:
// Helper function to build a complex// rotation valuecomplexTwiddle(int k,int len){float angularFreq =-2* pi * k / len
returncos(angularFreq)+ i*sin(angularFreq)}voidCalculateDFT(complex[] inputs,complex[] outputs,int len){if(len ==1){// To end the recursion, special-case// the 1-length FFT, which sums one input,// unrotated (i.e. it does nothing).
outputs[0]= inputs[0]return}// Temp storage for the recursion
complex[length /2] evens
complex[length /2] odds
// Separate the even and odd inputsfor(k =0; k < len /2; k++){
evens[k]= inputs[2* k]
odds[k]= inputs[2* k +1]}// Recurse into the even and odd FFTs// (each using the same array as input and// output, which is fine here)CalculateDFT(evens, evens, len/2)CalculateDFT(odds, odds, len/2)for(k =0; k < len /2; k++){// Get the corresponding recursively-calcualted// sum from the relevant arrays.complex evenSum = evens[k]complex oddSum = odds[k]// The odd sum needs to be rotated by this// output's twiddle factor.
oddSum *=Twiddle(k, len)// Now, as before, calculate two outputs from// these sums.
outputs[k]= evenSum + oddSum
outputs[k + len/2]= evenSum - oddSum
}}
This is now truly, algorithmically ! We halve the work at each level of the FFT, which does way fewer
rotations and adds than the initial DFT algorithm (for instance, the original algorithm for a 1024-length DFT would do
around a million rotations, but this one does around five thousand).
But it’s not ideal: it’s recursive, and uses temporary storage per recursion. We can do better! But now we get
into the fun that is…
Deinterleaving and Bit Reversal
To do all of this without recursion or scratch space, we need a way to do each level of FFT in-place. So let’s look at
how data is flowing through the system.
Here’s a diagram of the routine (at a single recursion level) where, from left to right:
The inputs are split into evens and odds.
Each of those two groups has a half-length sub-DFT performed on it.
Those results are then used in pairs by each step of the computation loop to write corresponding pairs of outputs.
If you look at the flow through this diagram, what’s interesting is that in the box on the right (representing the loop)
each computation pulls in from the same row that it writes out to (for instance, DFT result indices ‘0’ and ‘1’ are lined
up perfectly with final outputs ‘0’ and ‘8’).
This means that, theoretically, we could compute the entire FFT in-place (where the inputs and outputs are the same
array), by doing the following:
Deinterleave the data in-place (separating the evens and odds).
Calculate the sub-DFTs recursively (again, in-place), using the first and last halves of the data as the even and odd
sub-DFTs, respectively.
Then, run the computation loop over the results, each iteration of which reads from and writes to the same elements,
so no rearranging needs to take place.
There’s one problem with this: it turns out, efficiently deinterleaving an array in-place is one of those problems
that sounds like it would be super easy, but is instead kind of a nightmare!
To try to avoid that, let’s keep investigating the pattern of data moves. We know that, in terms of data flow, the data into the last
step comes in from two half-length sub-FFTs:
What if we track this pattern of deinterleaves all the way back to the lowest level (the two-element FFTs), where the
indices from right to left (from output to input) get deinterleaved at each step? Here’s the ordering we end up with:
At a glance, the pattern of inputs on the left isn’t necessarily obvious, but it turns out it’s straightforward: the
values are the original indices with their bits reversed!
For a 16-element FFT (like the diagrams above), we have 4 bits worth of index (since ), so here’s
the table of reversals:
0
0000
↔
0000
0
1
0001
↔
1000
8
2
0010
↔
0100
4
3
0011
↔
1100
12
4
0100
↔
0010
2
5
0101
↔
1010
10
6
0110
↔
0110
6
7
0111
↔
1110
14
8
1000
↔
0001
1
9
1001
↔
1001
9
10
1010
↔
0101
5
11
1011
↔
1101
13
12
1100
↔
0011
3
13
1101
↔
1011
11
14
1110
↔
0111
7
15
1111
↔
1111
15
As you can see, the ordering on the right of this table (0, 8, 4, 12, etc) is the same as the input ordering from the
previous diagram and, in their bit representations, you can see that each index is just the 4 bits of the other reversed.
Here is a simple routine that does a bit-reversed copy from one array to another:
voidBitReverseCopy(complex[] inputs,complex[] outputs,int len){for(int k =0; k < len; k++)
outputs[k]= inputs[BitRev(k, len)]}
Note
The above code doesn’t have an implementation of BitRev, because an efficient implementation on many systems is sorta
unreadable; it turns into a sequence of swapping increasingly-large groups of bits until you reach the size of the int.
For a C++20 implementaton, here’s a compiler explorer link.
Additionally, the above routine doesn’t work in-place. It is totally possible to do this bit reversal in place: it
can be written efficiently as a square matrix transpose (or two) if you get creative with the indexing of
the “matrix” rows. However, the final FFT-20XX algorithms do not have an explicit bit reversal step, so I’m not going to
go into those details. Sorry!
Going Iterative
We now have all of the pieces we need to assemble the final algorithm. Let’s go back to this earlier diagram of the data
flow for a 16-element FFT:
This diagram takes the inputs in what we now know to be bit-reversed order, and does four () passes on the inputs:
The first pass does eight 2-element computation loops
The next pass does four 4-element loops
Next is two 8-element loops
Finally, a single 16-element loop to get the final outputs
Note
In each of the above passes, the number of computation loops multiplied by the number of elements in the loop
is 16 (, , , and ), which means each pass processes elements. Since there
are passes of elements each, this algorithm – like the recursive one – is our target .
This process can (finally!) be done iteratively instead of recursively, by getting the elements into the expected
bit-reversed order, then modifying each element in place, in passes:
voidCalculateDFT(complex[] inputs,complex[] outputs,int len){if(len ==1){// Special case length 1.
outputs[0]= inputs[0]return}// Do the bit reversal so we can do the rest// iteratively in the outputs, in-place.BitReverseCopy(inputs, outputs, len)// Outer loop iterates all of the subFFT lengths:// starting at 2, then doubling until we reach the// full length ... this is log2(len) iterations.for(int subLen =2; subLen <= len; subLen *=2){// This loop iterates over each subsection, as the starting// index into the data.for(int subStart =0; subStart < len; subStart += subLen){// The data for this subsection starts at subStart.var*data =&outputs[subStart]// Finally, do the computation loop for this subsection.// The logic here is unchanged from the previous code// examples, but the reads and writes are to the same// pairs of elements.for(int k =0; k < subLen /2; k++){int kOdd = k + subLen /2complex evenSum = data[k]complex oddSum = data[kOdd]
oddSum *=Twiddle(k, subLen)
data[k]= evenSum + oddSum
data[kOdd]= evenSum - oddSum
}}}}
Note
The bit reversal does not necessarily have to come at the beginning: the algorithm could instead be written to take the
inputs in natural order and end up with the outputs in bit-reversed order, and do the bit reversal in-place as the
last step. The indexing within the passes would be different, but the algorithm would otherwise work the same.
In fact, some applications can just deal with the frequency-space outputs being out of order, and so can skip the
bit-reversal step entirely! This is really nice when it works out.
What we’ve ended up with is the radix-2 decimation in time
Cooley-Tukey algorithm. It’s radix-2 because
each “recursion” splits into two sub-FFTs, and it’s decimation in time because each pass deinterleaves the input,
time domain samples.
Note
If you go back up to the rotation diagrams again, you might note that the rotations are diagonally symmetrical, which
means you could instead split on even/odd output for the first and last half of the inputs (instead of the other
way around, like we did). This will end up as a similar but different algorithm, and is the decimation in frequency
version of an FFT.
Room for Improvement
While we’ve now gotten to the core FFT algorithm, this still isn’t anywhere near an optimal implementation, and
there are a number of things we’ll want to do to get this even faster. For instance:
As mentioned above, this is a radix 2 algorithm, but there are other radixes
(such as radix-4, where each FFT splits into 4 sub-FFTs) that are more efficient and worth investigating.
The sines and cosines could be done via a lookup table instead of calculating them at runtime, since for a given FFT
length there are only so many unique angles that are required.
Doing the bit reverse pass to get things into the correct order (either at the start or end) can be a sizeable chunk
of time, and depending on implementation can wreak havoc on your CPU’s caches.
And then there’s doing SIMD optimizations and otherwise trying to make the memory accesses CPU-cache-friendly.
We’ll touch on the first couple of those next time!
Because the world clearly doesn’t already have enough of them, I recently released a new FFT library:
FFT-20XX. To the best of my ability to tell, it uses an algorithm that is
at least slightly novel, and I thought it might be interesting to write about it.
This is the first of a series of posts about FFTs and the specific algorithms that FFT-20XX uses:
Any explanation of a Fourier transform would be remiss to not include a link to
this fabulous 3Blue1Brown video that nicely illustrates how the Fourier
transform works. I highly recommend watching it!
This post, however, is specifically about the discrete Fourier transform (referred to as the DFT), which is
a Fourier transform that operates on sampled data (like, say, digital audio or images) instead of continuous data.
The FFT (fast Fourier transform) is the standard algorithm used
to compute the result of a DFT – we’ll get to it in the next post.
What Does the DFT Even Do?
Any periodic signal (i.e. a signal that repeats, such as looped audio) can be represented as the sum of a set of
sine waves at different frequencies. Here is an example of a simple signal that can be broken down into three waves:
=
+
+
This is true of any periodic signal whether it is analog or sampled – no matter how complex it may seem.
Note
The “periodic signal” bit is important, but there
are many applications that use the DFT on non-periodic signals (like small blocks of a longer piece
of audio) and there are a few tricks that can be done (like windowing
the signal) to effectively “pretend” that it is periodic when it’s not.
The DFT answers a simple question: given a set of samples that represents a periodic signal, which frequencies
are represented in that signal?
To put it another way: the DFT converts a signal that is made up of samples in time (where each sample
represents a point on the signal at a given time) into the same signal represented in frequency (where each sample
represents the amplitude (how tall the wave is) and phase (the side-to-side position of the wave) of its
corresponding frequency). These two representations are frequently referred to as time domain and frequency
domain, and both encode the same signal in different ways.
Note
There is also an inverse DFT (the IDFT) which converts the other way – from the frequency domain back to the
time domain. Many signal processing applications convert from time to frequency space to do some operation, then use
the inverse DFT to convert back.
For a given repeating signal that is N samples long, the DFT returns contribution values for N frequencies, each
with one more cycle than the previous. To put it a different way: for each frequency with index i
(starting at 0 and ending at N - 1) there are i cycles of a wave (so index 0 has 0 cycles, index 1 has 1 cycle, etc).
Here are the 8 waves for a DFT with 8 samples: (the red wave with circles is the cosine and the green wave with squares
is the sine)
0
1
2
3
4
5
6
7
Ultimately, a DFT will give you a set of values that tell you how much each of these frequencies contributes to the signal.
Specifically, each frequency will have a value that represents both its amplitude and its phase.
Note
The frequency at index 0 always has cosine values of 1 and sine values of 0, and this ends up measuring how off-center
the signal is. This is known as the DC bias or DC offset.
Let’s figure out how to calculate the DFT for a signal!
Frequency Testing
The first step of computing the DFT is to figure out how to get a single frequency out of a signal. Ignoring phase for
a moment, this can be done in a surprisingly straightforward (to me, at least) manner: if we multiply each sample in the
signal by the cosine value of the test frequency and sum the results, you end up with a value that represents the
contribution of that frequency to the signal (which is 0 if the frequency is not present in the signal).
Here is an interactive example where you can set the signal and test frequencies and see the result of this process.
For purposes of illustration, the final sum in the diagram is instead an average (otherwise it would be too large to fit
in the diagram):
Signal Amplitude:
1.0
Signal Cycles:
2
Test Cycles:
2
Note that when the frequencies match, the resulting average scales with the amplitude of the signal, otherwise
(with one exception we’ll touch on in a moment) it is 0. This makes the average effectively a frequency detector
(again, ignoring phase for now)! This is still true even when there are multiple frequencies represented within the
signal: this will result in a zero value for frequenices that are not in the signal and non-zero values for those that
are.
Note
The mechanism by which this works is that the product of two waves with mismatched frequencies (over a span where they
have no fractional cycles) will have the same amount of area above and beneath the curve, and thus the average area is
zero.
This is hard to illustrate now, but it will be clearer in a later diagram.
As noted, there’s an exception here, where the multiply-and-sum seems to register a match when it shouldn’t.
For example: if you set the Signal Cycles to 1 and the Test Cycles to 15 (or vice versa), it will have a non-zero
average even though the frequencies don’t match. In fact, with the exception of frequencies 0 and 8, there are
two matches for every other frequency. Why is that?
Well, now it’s time to talk about…
Aliasing and Negative Frequencies
Let’s take a closer look at those two frequencies’ cosine values:
1
15
Note that while the waves are different frequencies, the cosine samples are the same! This is because of
aliasing, which is where a frequency above the maximum representable frequency
ends up sampling as if it were a different, lower frequency.
The frequency above which aliasing occurs is called the Nyquist frequency,
which is the highest frequency where there can be both a high and low sample for every cycle.
In a signal with N samples, that frequency has N / 2 cycles. For example, with 16 samples
the Nyquist frequency has 8 cycles (which you can see have a clear up then down pattern):
Above this frequency, the samples end up aliasing to a lower frequency and cannot be differentiated anymore.
Here are some more example pairs of aliased cosines from our 16-sample example:
2
14
3
13
7
9
Note that each pair has matching cosine samples, despite the different wave. Specifically: with a repeating sample count
of N, any frequency with cycles has the same cosine values as the frequency with cycles (that is 1 matches 15, 2 matches 14, etc).
That’s neat! But it’s not the whole picture. We can’t just look at the cosine values, we need to also look at what
happens to the sine values as well:
1
15
2
14
3
13
7
9
While the cosines are the same in these pairs of frequencies, the sines are negated and appear vertically
flipped! This means that a given frequency with cycles has negated sine values from the frequency with cycles.
The relationship here is that the aliased frequencies in each pair are negated versions of the other, since
and . Thus, frequency has the same angle as the hypothetical frequency at :
Cosines:
-1
15
Sines:
-1
15
In practice, when computing the DFT, these frequencies are referred to via their negative indices instead of as the positive
aliased ones.
This explains why all of the frequencies (except DC and Nyquist) have two matches: the cosines match at both the
positive and negative versions of the angle. This may seem redundant, but it isn’t always – we’ll get into why in a
bit.
Note
It’s also worth noting that the multiply-then-average in the example diagram results in half the signal amplitude for
each match, and when both are summed, the result represents the full amplitude!
However, DC and Nyquist – because they each only have a single match – have an average that matches the amplitude exactly.
Measuring Phase
So far we have been looking at a signal that contains a cosine, comparing it with a test wave that is also a cosine.
But the original description of the DFT mentions both amplitude and phase as being represented, so what happens
if we let the phase of the signal shift around relative to the test?
Signal Phase:
0.0
Oh no! Sliding the phase around causes resulting average to oscillate! While we get the expected 0.5 average when
the phase is 0, it is -1 at a phase of 0.5 (a half-wavelength offset) and 0 at phases of
0.25 and 0.75 (quarter- and three-quarter-wavelength offsets)! Clearly there is something else we need to consider
here…and it turns out it’s the sine of the test frequency.
Instead of just multiplying by the cosine of the test value, we can instead multiply by both cosine and sine to get two values per sample.
If we then treat those two resulting values as a 2D coordinate (where cosine corresponds to x and sine to y), we can average
them all to get our final result, which will be at the origin for non-matching frequencies and at some non-origin position
for matches.
Note
“Multiply by both sine and cosine” is a simplification for purposes of illustration. It works for now but is
not the general form, which we’ll get to in a bit.
Signal Phase:
Signal Cycles:
Test Cycles:
Some observations from this diagram:
If the cycle value magnitudes don’t match, the points are all evenly distributed around the center
(and thus average to the origin). This makes it much easier to visualize why we end up with no measured amplitude for
mismatched waves.
If the cycle value magnitudes do match (i.e. a signal of 2
cycles with either 2 or -2 test cycles), the points are off-center, and thus their average is also off-center.
Increasing the phase causes the resulting average to spin clockwise when the test cycle count is positive, but
counter-clockwise when the test cycle count is negative. The magnitude of the average remains the same no matter
the phase.
At Nyquist, the resulting points always lie on the x axis, as the test wave’s sine values are all zeros. It turns out:
a phase-shifted wave at Nyquist is indistinguishable from a non-shifted wave with scaled amplitude.
This gives us a new way to think about this computation: rather than thinking of the operation as a multiplication,
instead we can think of each sample of the signal being rotated by the angle in the test wave at that sample
position to get a new 2D position. Those 2D positions are then averaged to get a value that corresponds to the
frequency’s contribution.
The computed amplitude (the strength of this frequency’s contribution) is the magnitude of the
resulting average, and the phase is the angle of the average (computed, in this case, using atan2).
2D Coordinates and Complex Numbers
Each result of this multiplication with cosine and sine ends up as a 2D coordinate. Canonically, the Fourier transform
represents these as a complex number (of the form x + i*y). Why complex numbers? Honestly I think it’s a hack to get
a 2D coordinate that is simple to represent with standard math notation (mathemeticians, don’t @ me).
But, to be fair, complex numbers do have a nice property that standard 2D vectors do not: multiplying a complex number by another is equivalent to scaling the first value
by the magnitude of the second then rotating it by the angle of the second.
Thus, multiplying a sample from the signal (whether it is complex or real-valued) by cos(t) + i*sin(t) is equivalent to
rotating that sample by the angle t
(the rotation value has a magnitude of 1 so no scaling occurs in this case).
Note
This also finally gives us an answer to the earlier question of why we need both the positive and negative versions of a given
frequency: if the input samples are complex instead of purely real-valued, the contributions to these paired frequencies
can be different, and both are needed to fully reconstruct the signal.
For real-valued inputs, however, the paired results will be complex conjugates,
where the result (a + i*b) for frequency k is the conjugate of the result for frequency -k: (a - i*b). These
results are, in fact, redundant – and we will absolutely take that redundancy into account when implementing the FFT
for real-valued inputs.
If you’re more familiar with matrices and linear algebra than complex numbers, the reason this works is that this complex multiply:
gives exactly the same result as multiplying a 2D vector (a, b) with a 2x2 rotation matrix representing the same angle:
Note
It’s worth pointing out something that shows up in the standard formula for the DFT: mathematical notation has an
absolutely wild shorthand for :
That’s right, taking the mathematical constant
to an imaginary power is the same as a rotation around the origin. Why? Well, 3blue1brown has
an explanation if you’re curious.
Personally I don’t like this shorthand because it obscures the meaning of the equation, especially for
non-mathemeticians (which is, if I’m doing my calculations right, most people).
Calculating the DFT
Okay, we have all of the pieces of how to calculate a DFT, it’s time to put them together.
A DFT takes a set of complex input samples in the time domain and produces the same number of complex output
samples in the frequency domain, where the first output sample (at index 0) represents the DC offset, the middle one (at index N/2)
represents the Nyquist frequency, and the rest are either the positive or negative angles in between.
There are a couple minor (but important) differences in the value calculation done in the above interactive diagrams vs
how the DFT does them:
Typically the DFT rotates the opposite direction as the above diagrams (i.e. using instead of ), so
the rotations actually spin the opposite way (clockwise becomes counter-clockwise and vice versa).
Rather than averaging the results of these rotations, the DFT simply sums them – no division by N is performed.
Otherwise, the process described above can be used to get the contribution of any single frequency within the signal,
so then we just need to do that for every possible frequency in the signal.
Thus, here we are, finally at some pseudocode to calculate the DFT:
voidCalculateDFT(complex[] inputs,complex[] outputs,int length){for(outIdx =0; outIdx < length; outIdx++){// Start the sum at 0.float sum =0// This is the angular frequency (in radians/sample) for// the given output sample (the "test frequency" from the// earlier examples).float angularFreq =-2* pi * outIdx / length
for(inIdx =0; inIdx < length; inIdx++){// Get the rotation value for the current input index// then use that to calculate the resulting rotation.float radians = inIdx * angularFreq
complex rotation =cos(radians)+ i*sin(radians)// A complex multiply with a unit complex vector is// a rotation.
sum += inputs[inIdx]* rotation
}
outputs[outIdx]= sum;}}
Note
As mentioned in a previous note, the standard formulation of the DFT uses an exponential shorthand (which I do not
like) to represent the rotation. Thus, the DFT is written in standard mathematical notation as:
…which is equivalent to the above pseudocode.
We’ve finally done it! We have some code that we can run to calculate the DFT for a given signal.
The above algorithm is , which may seem like the best that can be done. However, there’s a bit of
algorithmic trickery that can be done to get that down to , which is considerably more efficient. That
trickery is the fast Fourier transform (FFT), and the next post will go into the details of how that works!
At some point, if you’ve done multithreaded programming, you’ve probably used a mutex (or some other locking mechanism). Locks are relatively straightforward to understand and use (“lock this thing before accessing your data and unlock it when you’re done”), but they do have their issues. The most commonly-discussed issue with using locks is the dreaded deadlock, scourge of many a poor soul who needed to hold multiple locks at once for something.
But we’re not here to talk about deadlocks…instead, I’d like to focus on a problem that I have encountered much more frequently: accidentally reading or modifying protected data without having acquired the lock.
That’s right, we’re once again going to try to protect ourselves from our worst enemy: ourselves.
Note
As with the previous entry in this semi-series, the examples and implementation are in C++, but you can probably build something similar in your language of choice (unless it doesn’t need it).
An Easy Mistake
It turns out it’s surprisingly easy to not grab a lock before reading or writing things that need to be synchronized. Here’s a toy example:
In the real world, “oops I accessed a thing outside of the lock” bugs tend to be more subtle - it can be especially hard to notice when a thing that should be present isn’t. Since all of the accesses are to normal members of your class/struct/global scope/whatever, it’s easy to slip and put an access to a value where it wasn’t intended, and wasn’t propertly protected.
But what if instead you could do something more like the following:
// Note: this is pseudocode, not actual C++classFoo{// ... class stuff ...// Mysterious m_state object that protects everything inside of it with a mutex
MutexProtected m_state
{int m_thing =0;};};voidFoo::IncrementThing(int c){lock(m_state)// Lock here{
m_thing += c;// Only accessible in the lock}// m_thing += c; // Won't compile: Can't access outside of lock}
Basically, what if you could have some sort of protected wrapper around the state that walls it off and makes it inaccessible except when the lock is actually held? With something like that it would become much more difficult to get at values when it shouldn’t be allowed.
Additionally, it would help both with organization (forcing the mutex-protected values together in the code) as well as making the intent clear: values that are protected sit inside the protected block, clarifying which values are (and are not) intended to be accessed solely through the mutex, even to folks who were not the original author (the list of which stealthily includes the original author, one month in the future).
With C++ it’s possible to get quite close to the above:
classFoo{// ... class stuff ...// State structure for the protected member(s)structState{int m_thing =0;};// The mutex and the state it's protecting
MutexProtected<State> m_state;};voidFoo::IncrementThing(int c){// lock the mutex and get the inner stateauto state = m_state.Lock();// Use the lock to access the values within.
state->m_thing += c;}
The nice thing is, a basic implementation of this is fairly compact. There are just two classes involved: MutexProtected<T> and its lock object, MutexLocked<T>.
EDIT (2025-01-25): People have pointed out two existing versions of this: boost has boost::synchronized_value, and there’s cs_libguarded which looks like it has a few variants on this concept, too. So if you don’t feel like rolling your own you can always check one of those out instead!
MutexProtected<T>
The outer class, MutexProtected<T>, has very few moving parts:
template<typenameT>classMutexProtected{public:
MutexLocked<T>Lock(){return{&m_t, m_mutex };}private:
std::mutex m_mutex;
T m_t;};
There’s a lock function that returns a MutexLocked<T> (which we’ll get to in a moment), and it contains both a mutex as well as a T (the templated type), which contains the data to be protected by the mutex. In the earlier example, this was the State struct that contained the protected m_thing value, but it could contain any number of values/objects that should only be accessed from within the same lock.
Note
You may be wondering “there’s a Lock, so why is there no corresponding Unlock function?” - this is because we’re taking advantage of the classic C++ idiom of RAII, so the returned MutexLocked<T> object holds the lifetime of the lock, and unlocks the mutex when it goes out of scope. As such, there’s no need to manually unlock the mutex; it will happen automatically.
MutexLocked<T>
But what is the thing that the Lock function returns? Well, that’s the MutexLocked<T> class and it looks like this:
template<typenameT>classMutexLocked{public:// Add the standard pointer-like accessors.
T *operator->()const{return m_p;}
T &operator*()const{return*m_p;}private:// Construct this with a pointer to the state and the mutex to lock.MutexLocked(T *p, std::mutex &mutex): m_lock{std::lock_guard(mutex)}, m_p{p}{}
std::lock_guard<std::mutex> m_lock;
T *m_p =nullptr;template<typenameT>friendclassMutexProtected;};
As you can see, it also doesn’t have much to it! It constructs with a pointer to the state object (the m_t in the MutexProtected<T>) and the mutex (which it immediately locks in the constructor), and so all it does is hold onto the lock until destruction time, while giving the user a way to access the state object (via the two public operators).
So, using this setup in full:
Call Lock on the MutexProtected<T> to get a MutexLocked<T>
Access the protected state through that acquired object (using the -> operator) to do whatever needs to be done within the lock
Let the MutexLocked<T> leave scope, at which point the lock is released
A nice little feature of this is if you do have a single thing that you’re doing in the lock (Say, inserting a value into a protected queue), the above steps can even fit nicely on a single line (while still being perfectly readable), thanks to the rules of C++ temporary object lifetimes:
voidFoo::EnqueueItem(Item *i){
state.Lock()->queue.Enqueue(i);// Lock/Update/Unlock// Do more stuff here outside of the lock, it's guaranteed to be released}
Sub-Functions That Require A Lock
This type of object also helps with a secondary part of this problem, which is when your class has functions (that are likely private or protected) that require the lock to already be acquired before you call it:
// This function is intended to only be called while the lock is heldvoidFoo::AdjustStateWithLockHeld(int delta){ m_thing += delta;}
Void Foo::IncrementThing(int delta){
std::lock_guard lock { m_mutex };// Call this function while the lock is held onlyAdjustStateWithLockHeld(delta);}
Void Foo::DecrementThing(int delta){// Oh no I have once again forgotten to lock the mutex before doing the thingAdjustStateWithLockHeld(-delta);}
In the above example, AdjustStateWithLockHeld is assuming that the lock is being held by its caller and that it’s free to manipulate the state within it. It doesn’t make much sense in this example, but if you have mutex-protected state that has a complex update process (such as multiple things needing to be kept in sync), it can be nice to move such logic to a subroutine.
Thankfully, using the MutexProtected<T> object, the state is only accessible through a MutexLocked<T> object. In order for the sub-function to be able to do anything with the state, it would need to additionally take a reference to the MutexLocked<T> for the state in question, thus effectively ensuring that the mutex is locked by the caller:
// Now the function takes the locked State object as a parameter:voidFoo::AdjustStateWithLockHeld(
MutexLocked<State>&state,int delta){ state->m_thing += delta;}
Void Foo::DecrementThing(int delta){// Now state has to be locked to get the object to pass to the sub-function!auto state = m_state.Lock();AdjustStateWithLockHeld(state,-delta);}
A Variation
For completeness, I also wanted to mention an alternate form of this that I thought of while designing it: where, rather than Lock returning an object that represents the scoped lock, you instead pass a lambda (or function) that gets all of the state values in it:
structState{int a, b;};
MutexProtected<State> state;
state.Lock([&](int&a,int&b){// "a" and "b" correspond to the two values of the same name in the state structure.// All the things that need to modify state values must, then, happen in this lambda,// as there is no way to access the struct members directly.
a += someVariable;
b -=2;});
The main advantage this has over the other form is that it makes it considerably more difficult to accidentally (or intentionally?) grab a reference to the internal state structure that could then persist outside of the lock. However, in practice I felt like that kind of mistake is not going to be super common relative to the problem being solved, but also should be easy to catch during code review. Plus, there are additional downsides to this version that I didn’t like:
It’s easy to get the order or names of the lambda paramters wrong and end up doing the wrong things with the wrong values since they’d effectively have to match the order in the state.
It gets harder to call sub-functions that need the lock (you’d have to pass all of the state objects that need updating which can be a pain in practice).
It’s worse to debug, since you have to step into the Lock call instead of just over it, and then into the lambda from there.
This is also the reason I didn’t consider another variant where instead of a parameter per state object it’s a single reference to State and you access it that way.
All said, I felt like having a lock object that provides access to the inner state as-is (via the -> operator) was cleaner in practice, and easier to step through in the debugger.
Limitations
This, of course, isn’t a perfect solution:
There are absolutely cases where some things need to be accessible outside of the mutex lock (i.e. an atomic which can be safely read at any time but only gets updated from within the mutex due to sequencing issues), and as such those values couldn’t live inside of the inner state object.
Also, this is C++ so there’s nothing preventing someone from grabbing a reference or pointer to the state object from the lock and holding onto it until after the lock ends, then partying on the data. However, that kind of code is more likely to get caught at review time.
There are likely other cases (multiple locks, perhaps, which I haven’t fully thought through with this because I haven’t needed to) where this would present problems. This is definitely primarily intended for the “I have a set of data that should only ever be accessed during a lock” case.
Potential Improvements
Also, there are some improvements that could be made:
Perhaps instead of using a lock_guard you could use a unique_lock which would let you wait on a critical section using the lock (if you need to wait for the state to be in some specific configuration), which could even be built into the MutexProtected<T> class (or a similar one) if the critical section is a core part of the usage.
It’s a good idea to declare a custom move constructor and move assignment operator for MutexLocked<T> that nulls the m_t pointer of the object being moved from so that you can’t do a move of it to some other location (which subsequently releases the lock) and then still party on the internal pointer. (I also have asserts in my production version in the pointer-access operators that assert the pointer is non-null)
Similarly, it may be nice to have a constructor for MutexProtected<T> that takes constructor arguments for the contained state structure (similar to, say std::vector::emplace_back), especially if you are going to have state structures that cannot default construct.
There are also other flavors of locking that might occur: for instance, a TryLock function that returns a std::optional<MutexLocked<T>>, and only locks the lock (and returns the locked state) if there is no contention on the mutex.
Closing Time
All in all, having an abstraction like this makes it way more difficult to party all over internal state without properly locking the mutex first. Switching some old code to use this actually found a couple places where I’d done things incorrectly (reading values that should have been mutex protected on read, in those cases).
So, yeah, by protecting our data from being accessed when it shouldn’t be, we’re also protecting ourselves from ourselves.
Imagine this: you’ve got some value, x, that you want to ensure is at least 1. That is to say, you want to ensure its minimum value is 1. So, being the smart, experienced programmer that you are, you write the following:
x =Min(x,1);
You give yourself the small, satisfied nod of a job well done and run the program and then it all goes immediately sideways because that should have been Max and not Min.
If you’ve been writing code for basically any length of time, the above was probably less an “imagine this” and more a “remember this” because if you’re anything like me, you’ve done this over. and over. and over.
Inspired by once again mistakenly using Min instead of Max to limit the minimum allowed value of something, I’ve decided to start a little series (will it have more than one entry? who knows!) called Protecting Coders From Ourselves, in which we rework some bit of API surface to make it less error-prone. We’re going to deal with the “I chose the wrong Min/Max again” problem, but first I want to talk about Clamp and Lerp.
Note
The examples here are in C++, but the concepts should be relevant to basically any language.
Clamp and Lerp
Clamp
Clamp is a simple enough function: Take some value v and make sure it is no less than min and no greater than max. Almost every clamp function in every library I’ve seen has three parameters, one of which represents the value to be clamped, and two of which represent the range that it should be clamped within:
Clamp(a, b, c)
The question, of course, is “which parameter is which?” Many languages (ex: C++, C#, HLSL, and Rust) have the following arrangement in their standard libraries:
Clamp(valueToClamp, min, max)
where the first parameter is the value being clamped and the last two are the min and max ends of the range.
But I’ve also seen this one (looking at you, CSS):
Clamp(min, valueToClamp, max)
This one puts the value to clamp in the middle of the range (which, honestly, is conceptually where it belongs).
Lerp
Another common function that takes a value and a range is Lerp, which uses a value in the range [0, 1] to linearly interpolate between two endpoint values. Most lerp functions that I’ve seen (ex: C++, C#, and HLSL) have their parameters ordered as follows:
Lerp(a, b, t)
where a and b are the range endpoints and t is the interpolating value.
Depending on the projects you work on, you maybe don’t write code that uses Lerp very often (or ever), but I do, and for me the combination of Lerp and Clamp are a source of constant, mild confusion.
Parameter Confusion
Both of these functions take three parameters, two of which represent a range and one which is a value that is either limited by or used to interpolate within the range. In isolation it’s easy to rationalize the order of the parameters for each:
Clamp(v, min, max): Clamp v to be within the range [min, max].
Clamp(min, v, max): Clamp such that min <= v <= max.
Lerp(a, b, t): Get a linearly-interpolated value between a and b using t.
…but in combination I am constantly second-guessing which order the parameters need to go in. Sometimes, for instance, I’ll write the equivalent of Lerp(t, a, b) and wonder why nothing is working the way I expect.
That brings us (finally) to the question of the article: how can we make it clear which parameters are which in these functions?
An obvious way to do this, given language support, is to make use of named parameters when calling the function:
Clamp(v=value, min=0, max=5)
but not all languages (looking at you, C++) support named parameters, so what then?
Grouping the Range Values
What if instead of the above, calls looked more like this:
Clamp(a,{b, c});Lerp({a, b}, c);
With this added structure, it’s clear even without reasonable variable names which part is the range and which is the value, and it’s much more difficult to accidentally call them with the parameters in the wrong order, since, for instance, Lerp(t,{a, b}) wouldn’t even compile.
To do this, we need a simple range structure. in C++ it could look something like this:
template<typenameT>structValueRange{// Use a constructor to ensure both endpoints are required.ValueRange(T a_, T b_): a{a_}, b{b_}{}
T a;
T b;};
Using this range, then, you could define Clamp and Lerp as follows:
template<typenameT>
T Clamp(T v, ValueRange<T> range){returnMax(range.a,Min(range.b, v));}template<typenameT>
T Lerp(ValueRange<T> range, T t){return range.b * t + range.a *(1- t);}
Now instead of three parameters, they take two, which matches how they work conceptually: with a range and a value in some order.
Once you start grouping your input parameters , you may start seeing other places to do it, like IsInRange<Inclusive>(v,{0,20}).
“Almost every time I use [min or max], I think very carefully and then pick the wrong one.”
This tends to happen because you think “I need to make sure the max value of x is 10” and it just feels right to turn that into Max(x,10)…that’s how it gets you. Or, well, me. That’s how it gets me. Basically every time I have to do this, I will choose the wrong one on the first try, have my code explode on me, and then go back and headdesk at it until it turns into the correct one.
Basically. Every. Time.
Reframing The Problem
But what if you looked at it a different way? What if you instead thought of it as “I want to clamp x so that it’s no larger than 10”? You want something that just clamps one end of it, in a clear way. What if you could declare a one-sided clamp:
// Using "Open" to declare a side of the range is open. Same as:// x = Min(x, 10);
x =Clamp(x,{Open,10});// Another alternative, ensure that y is no smaller than 2, same as: // y = Max(y, 2);
y =Clamp(y,{2, Open});
To do this efficiently, we’ll have multiple overloads of Clamp, and define two additional “range” structures, each of which takes an OpenEnded_t:
// Declare this as a nice, type-safe enum classenumclassOpenEnded_t{
Open,};// But make "Open" easy to reach using C++20's "using enum" feature.usingenumOpenEnded_t;// This is a "range" where only the "a" value is specified, the "b" end is opentemplate<typenameT>structValueRangeOpenB{ValueRangeOpenB(T a_, OpenEnded_t):a(a_){}
T a;};// Like the above, but it's the "b" end that's specified while "a" is opentemplate<typenameT>structValueRangeOpenA{ValueRangeOpenA(OpenEnded_t, T b_):b(b_){}
T b;};
Note
You can, of course, call Open whatever you’d prefer: I considered many options (including Unbounded, Infinite, OpenEnded, and None), but Open was short and, to my mind, clear.
If you have a global Open function (or you’re in a class that has a function named Open), this likely won’t work. You could add a second enum value called OpenEnded that could be used interchangeably with Open for that case, or just specify it fully qualified, or just pick a less-inconvenient name.
Once you have these structures, you can define two additional overloads of Clamp:
template<typenameT>
T Clamp(T v, ValueRangeOpenB<T> range){returnMax(v, range.a);}template<typenameT>
T Clamp(T v, ValueRangeOpenA<T> range){returnMin(v, range.b);}
These just turn into the correct call to Min or Max, but now you can think of it in terms of limiting one side of its range or the other, rather than trying to A Beautiful Mind your way into picking the correct function right off the bat.
Now, finally, you can limit a value in multiple ways using the same concept, which can make it easier to reason about when you’re writing the code, and also easier to understand when you’re reading it a month later.
// Keep within a range:
x =Clamp(x,{1,5});// Limit the lower bound:
y =Clamp(y,{1, Open});// Limit the upper bound:
z =Clamp(z,{Open,5});
Final Thoughts
These are ideas I first proposed at my job, and got immediate buy-in from the dev team, because we all kept making the same kinds of mistakes with these functions. There are, of course, times when Min and Max are still the appropriate function to use (like when you’re thinking “I need the minimum of these values”) - but when you’re trying to limit the range of something, Clamp is a clearer declaration of intent.
There are ways to improve these functions:
In C++ I highly recommend making all of this constexpr (including the constructors) so that you can use these functions at compile time as well.
Depending on your codebase it may be desirable to additionally mark them [[nodiscard]] and noexcept.
Also, restricting the template types using C++20 concepts can help give you better error messages if you try to compile it with something that it can’t work with.
The full-range Clamp function may want to have some validation that b >= a (perhaps an assert).
A while back I was working on a programming language idea and while I haven’t made any progress on it in ages, I really liked the string design that I came up with. I don’t know that any of the ideas are original, but I haven’t seen anything exactly like it so I figure I’d throw the idea out into the ether in case anyone else happens to do something similar in their own personal language that they definitely shouldn’t be making 😆
(I’ll note upfront: this design uses dollar signs ($) and backticks(`). There are many languages that do so (like Javascript!) but these keys are not universally on all keyboards internationally so for those locales it may be more difficult to type these out…my language design was pretty much just for me with my standard US keyboard so I didn’t take this into account)
Basic Strings
There’s nothing fancy about these, they’re just like most other languages’ basic strings:
"This is a string""This is a string with a newline at the end: \n""Quotes? \"Escape them\""
"Oops forgot to end this one, it's a compiler error
Just a pair of double quotes with everything between being non-quotes (or escaped quotes), contained within a single line.
Raw Strings
Raw strings in my language design are kind of a blend between C++11-style raw strings and C#11 raw strings (the elevens!), using a different delimiting character variation than I’ve seen elsewhere. In this case, the simplest one starts and ends with a pair of backticks:
``This is a single raw string
it is multiple lines long``
It can contain any character sequence except the delimiting sequence (again, at simplest a pair of backticks: ``)
But what if your string needs to contain a consecutive pair of backticks? This is where the C++11 raw string inspiration comes in: you can put any string of characters between the ticks (excluding ticks, obviously, or newlines), and then the start and end have to match.
`uniqueString`This is a single string
it contains backticks without terminating: `` ... see?
This is the last line and ends here:`uniqueString`
This one starts with `uniqueString`, and so the only thing that will terminate it is that same sequence: `uniqueString` (with the tick marks around it).
To add to this, cribbing from C#11 it will:
Trim the very first newline if there is one
Also trim the last newline if there is one
Unindent every line of it based on the indentation of the final quote sequence:
myString =``
This is actually the first line of the string, the newline was ignored
{
indented further
}
``;// Note that this is indended 2 spaces
which turns into the string (note the lack of being completely indented:
This is actually the first line, the newline was ignored
{
indented further
}
This makes it easier to generate code (or text files or whatever) that are properly indented, without having to make the indenting of the string in your code all weird.
Interpolated Strings
I’m additionally adding string interpolated string support (which is a weird term), using a mix of C# and Javascript’s setup:
$"This string has a ${value} in it"
If a string starts with $, it’s treated as an interpolated string. A string-convertible expression can be inserted in-place in the string within ${}.
But what if you need to have the character sequence ${ in your string? Add more dollar signs to the start, and you need that many dollar signs before a { to enter the Interpolation Zone:
$$"This string has a $${value} in it, ${but this isn't one}"
Raw strings can also be used as interpolated strings (making for some nice codegen), same rules apply:
$$``
Interpolated string with
multiple lines and a $${value} in it.
${this is not a value because only one $}
``
If value == 5 this would turn into the following string (upon formatting):
Interpolated string with
multiple lines and a 5 in it.
${this is not a value because only one $}
I Just Think They’re Neat
Anyway, I think this is a really nice combination of properties that make it easy to format strings nicely without being overly-complicated to actually use (unlike C++'s raw strings, which I have to look up literally every time I need to use one). Need a hardcoded regex or path with backslashes? In most cases, just use a raw string with `` on either end:
``C:\Path\With\Single\Backslashes``
or
``^[\r\n \t]*Hi[\r\n \t]*$``
Hope this was at least mildly interesting to someone!
(This post follows Part 1: 32-bit floats and will make very little sense without having read that one first. Honestly, it might make little sense having read that one first, I dunno!)
Last time we went over how to calculate the results of the FMAdd instruction (a fused-multiply-add calculated as if it had infinite internal precision) for 32-bit single-precision float values:
Calculate the double-precision product of a and b
Add this product to c to get a double-precision sum
Calculate the error of the sum
Use the error to odd-round the sum
Round the double-precision sum back down to single precision
This requires casting up to a 64-bit double-precision float to get extra bits of precision. But what if you can’t do that? What if you’re using doubles? You can’t just (in most cases) cast up to a quad-precision float. So what do you do?
To do this natively as doubles, we need to invent a new operation: MulWithError. This is the multiplication equivalent of the AddWithError function from the 32-bit solution:
(double prod,double err)MulWithError(double x,double y){double prod = x * y;double err =// ??? how do we do thisreturn(prod, err);}
We’ll get to how to implement that in a moment, but first we’ll walk through how to use that function to calculate a proper FMAdd.
We need to do the following:
Calculate the product of a and b and the error of that product
Calculate the sum of that product and c (giving us a * b + c) and the error of this sum
We’re not using the error of the product … yet
Add the two error terms (product error and sum error) together, rounding the result to odd
Add this summed error term to our actual result, which will round normally.
In code, that looks like this:
// Start with an "OddRoundedToAdd" helper since we do// this operation frequentlydoubleOddRoundedAdd(double x,double y){(double sum,double err)=AddwithError(x, y);returnRoundToOdd(sum, err);}doubleFMAdd(double a,double b,double c){(double ab,double abErr)=MulWithError(a, b);(double abc,double abcErr)=AddWithError(ab, c);// Odd-round the sum of the two errors before // adding it in to the final result.double err =OddRoundedAdd(abErr, abcErr);return abc + err;}
By keeping the error terms from both the product and the sum, we have kept all of the exact result. That is, we can assemble the mathematically-exact result given enough precision by doing abc + abErr + abcErr.
But we can’t do infinite-precision addition of three values. However, we can odd-round an intermediate result, the same way we did with the single-precision case.
In this case, we know that abErr and abcErr both (necessarily) have much lower magnitudes than the final result, as each error value’s highest bit is lower than the lowest bit of the mantissa of their respective operations. So, if we odd-round the sum of these two values, it actually effectively fulfills the condition of having more bits of precision than the final result. Thus, if we add the error terms together with odd rounding, the odd-rounded fake-sticky final digit will be taken into account by the actual sticky bit used when doing the final sum of the result and error terms.
So how do we calculate the error term of a 64-bit multiply? We can’t use 128-bit values, but what we can do is break each 64-bit value up into two values, each with less bits of precision.
We’ll break x and y (our two multiplicands) up into high and low values, where:
x = xh + xl;
y = yh + yl;
We do this by breaking the double’s mantissa up:
xh contains the top 25 bits of x’s mantissa (plus its implied 1, giving it 26 bits of precision).
xl contains the bottom 27 bits of x’s mantissa. The highest-set ‘1’ in this mantissa will become the implied 1 bit that’s part of the standard floating-point format, so this value will have 27 bits of precision, max.
yh and yl are the same, but for y.
(double h,double l)Split(double v){// In C++ this Zero function can be implemented by masking off// the bottom 27 bits by casting to a 64-bit int:// constexpr uint64_t Mask = ~0x07ff'ffff;// double h = std::bit_cast<double>(// std::bit_cast<uint64_t>(v) & Mask);double h =ZeroBottom27BitsOfMantissa(v);// We can get the lower bits of the mantissa (correctly normalized and// with correct signs) by subtracting the extracted upper bits from // the original value.double l = v - h;return(h, l);}
What does this split give us? Well, we can now break the multiplication up into a sum of multiplies that now each have enough bits of precision to be exactly representable (27BitValueA * 27BitValueB == 54BitValue, which fits perfectly in a double (with the implied 1 bit), using our old friend from Algebra, FOIL:
x * y
=(xh + xl)*(yh + yl)= xh*yh + xh*yl + xl*yh + xl*yl;
We can’t actually do those adds directly, but what we can do is similar to how we did AddWithError: use a sequence of precision-preserving operations to calculate the difference between that idealized result and our rounded product:
(double prod,double err)MulWithError(double x,double y){double prod = x * y;(xh, xl)=Split(x);(yh, yl)=Split(y);// Parentheses to demonstrate the precise order these // operations must occur indouble err =(((xh*yh - prod)+ xh*yl)+ xl*yh)+ xl*yl;return(prod, err);}
It works like this:
Calculate the (rounded) product of x and y
Subtract that rounded product from the product of xh and yh
These should have roughly the same magnitude (and definitely the same sign) so this is a precision-preserving subtraction.
Since |xh * yh| <= rounded(|x * y|) (because xh and yh are truncated versions of x and y and thus have lower magnitudes) this is a smaller - larger operation and we’ll get a result with a sign opposite that of the final product.
Keep adding in next-lower-magnitudes of values, which will continue to preserve precision
(because we have a value that is opposite-sign these are effectively subtractions, in the same way that a + -b is)
It’s also worth noting here that xh*yl and xl*yh will have equivalent magnitudes so the order that you add them in doesn’t matter, as long as they’re both after xh*yh and before xl*yl
Once you’ve done that, you have the computed product as well as the error term, and we can then follow our FMAdd algorithm above to calculate the FMAdd.
So, that’s it, we’re done, right?
Edge Cases
Nope! Well, yes if you just wanted the gist, but now it’s time to get into all those annoying implementation details that the papers this is based on completely glossed over. Here’s where it gets ugly (unless you thought it was already ugly, in which case, sorry, it’s about to get worse somehow).
In our single-precision case, we didn’t have to worry about exponent overflow or underflow because we were using double-precision intermediates, which not only have additional mantissa range, but also additional exponent range.
It’s possible that the product of a * b (an intermediate value in our calculation) goes out of range of what a double can represent, but that the addition of c might bring the final result back into range (which can happen when the sign of c is opposite the sign of a * b). This causes a different set of errors on either end:
If a*b is too large to be represented, it turns into infinity which means that adding c in will just leave it as infinity even though the final result should have been a representable value (albeit one with a very large magnitude)
If a*b is too small to be represented, it will go subnormal which means bits of the intermediate result will slide off the bottom of the mantissa and we lose bits of information, which can causes us to round incorrectly at our final result.
To solve this, we’ll introduced a bias into the calculation, for when the value goes very small or very large:
doubleCalculateFMAddBias(double a,double b,double c){// Calculate what our final result would be if we just did it normallydouble testResult =Abs(a * b + c);if(testResult <Pow2(-500)&&Max(a, b, c)<Pow2(800)){// Our result is very small and our maximum value is not so large // that we'll blow up with a bias, so bias our values up to// ensure we don't go subnormal in our intermediate resultreturnPow2(110);}elseif(IsInfinite(testResult)){// We hit infinity, but that might be due to exponent overflow,// so bias everything down (this may cause c to go subnormal, // but if that's the case then a*b on its own is infinity and// so it won't affect the final result in any way)returnPow2(-55);}else{// No bias neededreturn1.0;}}
For any results that aren’t extreme, the bias will remain 1.0 but, for values at the extremes, we’ll scale our intermediates down (using powers of 2 which only affect the exponent and not the mantissa) into a range such that we can’t temporarily poke outside of range. Also note that my choices of powers of 2 are not perfectly chosen, I didn’t bother trying to figure out the exact right biases/thresholds so I just picked ones that I knew were good enough.
So then we do our FMAdd calculation as before, but with the bias introduced (and then backed out at the end):
// Do our multiplication with the bias applied to 'a'// (the choice of applying it to 'a' vs 'b' is completely// arbitrary)(double ab,double abErr)=MulWithError(a * bias, b);// Then the sum with the bias applied to 'c'(double abc,double abcErr)=AddWithError(ab, c * bias);
err =OddRoundedAdd(abErr, abcErr);// Calculate our final result then un-bias the result.return(abc + err)/ bias;
Alright, we’ve avoided both overflow and underflow and everything is great, right?
Two (Point Five) Last Annoying Implementation Details
Nope, sorry again! It turns out there are still two cases we need to deal with.
Case 1: Infinity or NaN even with the bias
If our result (without error applied) hits infinity even with the avoid-infinity bound, then we should just go ahead and return now to avoid Causing Problems Later (that is, turning what should be infinity into a NaN). And if it’s already NaN we can just return now because it’s going to be NaN forever.
Except, there’s one additional necessary check here, for a case caught by dzaima over on Bluesky): in the event a and b are finite numbers but a * b blows up to infinity, and then c is the opposite infinity, the correct return value is whichever infinity (positive or negative) c is, so in our early-out check has to catch that case as well:
if(IsInfiniteOrNaN(abc)){// If we got NaN (or Inf, which won't affect the output) and // a and b are both finite but c is infinite, return c (without// this check, we will incorrectly return NaN instead of -Inf// for FMAdd(1e200, 1e200, -Infinity))if(IsInfinite(c)&&!IsInfiniteOrNaN(a)&&!IsInfiniteOrNaN(b)){return c;}// Otherwise, return whichever Inf or NaN we got directly;return abc;}
Case 2: Subnormal Results
If our result is subnormal (after the bias is backed out), then it’s going to lose bits of precision as it shifts down (because the exponent can’t go any lower so instead the value itself shifts down the mantissa), which means whoops here’s another rounding step, and the dreaded double-rounding returns.
In this case we need to actually odd-round the addition of the error term as well, so that when the bias is backed out and it rounds, it does the correct thing:
// Multiply the smallest-representable normalized value by our avoid-// subnormal bias. Any (biased) value below this will go subnormal.// (In production code it'd be nicer to use something like// std::numeric_limits instead of hard-coding -1022)constdouble SubnormThreshold =Pow2(-1022)* AvoidDenormalBias;if(bias == AvoidSubnormalBias &&Abs(abc)< SubnormThreshold){// Odd-round the addition of the error so that the rounding that // happens on the divide by the bias is correct.(double finalSum, finalSumErr)=AddWithError(abc, err);
finalSum =RoundToOdd(finalSum, finalSumErr);return finalSum / bias;}
And this almost works, except there’s one more annoying case, and that’s where our result is going subnormal, but only by exactly one bit. Remember that the odd-rounding trick only works if we have two or more bits so that the final rounding works properly, but in this case we’re truncating the mantissa by exactly one bit, so we have to do even more work:
Split the value that will be shifting down into a high and low part (same as we did for the multiply)
Add our error term to the low part of it
This preserves additional bits of the error term since we gave ourselves more headroom by removing the upper half of its mantissa
Remove the bias from both the high and low parts separately
Removing the bias from the high part doesn’t round since we know the lowest bit is 0
Removing the bias from the low part applies the actual final rounding (correctly) since we gave ourselves more bits to work with
Sum the halves back together and return that as our final result
This sum is (thankfully) perfectly representable by the final precision and doesn’t introduce any additional error.
constdouble OneBitSubnormalThreshold =
OneBitSubnormalThreshold *0.5;if(Abs(finalResult.result)>= k_oneBitDenormThreshold){// Split into halves(rh, rl)=Split(finalSum);// Add the error term into the low part of the split
rl =OddRoundedAddition(rl, finalSumErr);// Scale them both down by the bias. Note that // the rh division cannot round since the lowest bit// is 0
rh /= bias;// This division is what actually introduces the final// rounding (correctly, since we gave ourselves more// bits to work with)
rl /= bias;// This sum is perfectly representable by the final// precision and will not introduce additional error.return rh + rl;}
OMG Are We Done Now?
As far as I’m aware, those are all the implementation details to doing a 64-bit double-precision FMAdd implementation. It’s conceptually not that much more complicated than the 32-bit one, but mechanically it’s worse, plus there are those fun extra edge cases to consider.
Here’s the final code:
(double h,double l)Split(double v){double h =ZeroBottom27BitsOfMantissa(v);double l = v - h;return(h, l);}(double prod,double err)MulWithError(double x,double y){double prod = x * y;(xh, xl)=Split(x);(yh, yl)=Split(y);double err =(((xh*yh - prod)+ xh*yl)+ xl*yh)+ xl*yl;return(prod, err);}doubleOddRoundedAdd(double x,double y){(double sum,double err)=AddwithError(x, y);returnRoundToOdd(sum, err);}doubleFMAdd(double a,double b,double c){constdouble AvoidSubnormalBias =Pow2(110);double bias =1.0;{// Calculate our final result as if done normallydouble testResult =Abs(a * b + c);// Bias if the result goes too low or too highif(testResult <Pow2(-500)&&Max(a, b, c)<Pow2(800)){ bias = AvoidSubnormalBias;}// too lowelseif(IsInfinite(testResult)){ bias =Pow2(-55);}// too high}// Calculate using our bias(double ab,double abErr)=MulWithError(a * bias, b);(double abc,double abcErr)=AddWithError(ab, c * bias);// Check for infinity or NaN and return earlyif(IsInfiniteOrNaN(abc)){// Handle the case of "a multiply of two finite values hit infinity// even *with* the bias, but c is the opposite infinity" case and// return the correct result of "c"if(IsInfinite(c)&&!IsInfiniteOrNaN(a)&&!IsInfiniteOrNaN(b)){return c;}// Otherwise just return the inf or nan directlyreturn abc;}// Odd-round the intermediate error resulttdouble err =OddRoundedAdd(abErr, abcErr);// Multiply the smallest-representable normalized value by our avoid-// subnormal bias. Any (biased) value below this will go subnormalconstdouble SubnormThreshold =Pow2(-1022)* AvoidSubnormalBias;if(bias == AvoidSubnormalBias &&Abs(abc)< SubnormThreshold){(double finalSum, finalSumErr)=AddWithError(abc, err);// This is half of SubnormThreshold. Any value between SubnormThresold// and this value will only lose a single bit of precision when// the bias is removed, which requires some extra careconstdouble OneBitSubnormalThreshold =
OneBitSubnormalThreshold *0.5;if(Abs(finalSum)>= OneBitSubnormalThreshold){// Split into halves(rh, rl)=Split(finalSum);// Add the error term into the LOW part of our split value
rl =OddRoundedAdd(rl, finalSumErr);// Divide out the bias from both halves (which will cause rl to// round to its final, correctly-rounded value) then sum them // together (which is perfectly representable).
rh /= bias;
rl /= bias;return rh + rl;}else{// For more-than-one-bit subnormals, we do an odd-rounded addition of// the error term and then divide out the bias, doing full rounding// just once.
finalSum =RoundToOdd(finalSum, finalSumErr);return finalSum / bias;}}else{// Not subnormal, so we can calculate our final result normally and un-// bias the result.return(abc + err)/ bias;}}
Compare that to the 32-bit version and you can see why this one got its own post:
A thing that I had to do at work is write an emulation of the FMAdd (fused multiply-add) instruction for hardware where it wasn’t natively supported (specifically I was writing a SIMD implementation, but the idea is the same), and so I thought I’d share a little bit about how FMAdd works, since I’ve already been posting about how float rounding works.
So, screw it, here we go with another unnecessarily technical, mathy post!
What is the FMAdd Instruction?
A fused multiply-add is basically doing a multiply and an add as a single operation, and it gives you the result as if it were computed with infinite precision and then rounded down at the final result. FMAdd computes (a * b) + c without intermediate floating-point error being introduced:
floatFMAdd(float a,float b,float c){// ??? Somehow do this with no intermediate roundingreturn(a * b)+ c;}
Computing it normally (using the code above) for some values will get you double rounding (explained in a moment) which means you might be an extra bit off (or, more formally, one ULP) from where your actual result should be. An extra bit doesn’t sound like a lot, but it can add up over many operations.
Fused multiply-add avoids this extra rounding, making it more accurate than a multiply followed by a separate add, which is great! (It can also be faster if it’s supported by hardware but, as you’ll see, computing it without a dedicated instruction on the CPU is actually surprisingly spendy, especially once you get into doing it for 64-bit floats, but sometimes you need precision instead of performance).
Double Rounding
Double rounding happens when the intermediate value rounds down (or up), then the final result also rounds in the same direction - but because of the first rounding, actually overshoots the correctly-rounded final value by a bit.
Here’s an example using two successive sums of some 4-bit float values. We’ll do the following sum (in top-down order):
1.000*2^4+1.001*2^0+1.100*2^0
The first sum, done with “infinite” internal precision, looks like this:
If we were to then use that result directly (with no intermediate rounding) and do the second sum, only rounding the final result:
1.0001001*2^4+0.0001100*2^4----------=1.0010101*2^4->1.001*2^4// Rounds down
The final result rounds (to nearest) to 1.001.
However, if we were to round that intermediate value to 4 bits first, we’d get this:
1.0001001*2^4->1.001*2^4// Rounded up+0.0001100*2^4----------1.0011100*2^4->1.010*2^4// Up again
In this one, we end up with 1.010 instead of 1.001 because of the intermediate rounding, which pushed us past the correctly-rounded final result.
How to Pretend That You Have Infinite Precision
Okay, for FMAdd we want to calculate a multiply, and then somehow throw an add in there and have it act as if we didn’t lose any precision on the multiply.
First we’re going to handle the case of 32-bit floats (singles) because it’s a wildly simpler case on CPUs that have 64-bit floats (doubles).
(also, sorry in advance, the term “double” for a “double-precision float” and the “double” in “double rounding” are two different instances of “double” but I’ve written so much of this post and like hell am I changing it now so hopefully it’s not too confusing)
The immediately obvious thing to try to get an accurate single-precision FMA is “hey, what if we do the multiply and add as doubles and then round the result back down to a single”:
floatFMAdd(float a,float b,float c){// Do the math as 64-bit floats and truncate at the end. // Surely that's good enough, right?returnfloat((double(a)*double(b))+double(c));}
While that gives a much better result than doing it as pure 32-bits, it actually can still have double rounding. But where does the extra rounding come from, in this case?
The multiply itself isn’t the source of the first rounding: Surprisingly (to me, at least): casting two singles to doubles and multiplying those together always results in an exact answer - this is because each of the single-precision values has 24 bits of precision, but a double can store 53 bits of precision, which is more than enough to store the result of multipling two singles (2 * 24 bits of precision max). Since floats are stored as:
sign *1.mantissa *2^(exponent)
…it means we’re multiplying two numbers of the form 1.xxxxxxxxxx and 1.yyyyyyyyy together then adding the exponents together to get the new number, so unlike addition and subtraction (where, say, 1 + 1*10^60 requires a ton of extra precision), if two float numbers have wildly different exponents it doesn’t actually matter because the exponents and significand values are handled separately.
To illustrate this, let’s pretend we have two 4-digit (base 10) numbers and we multiply them and store the result using 8 digits (double precision):
(1.234*2^1)*(1.457*2^100)->(1.2340000*1.4570000)*(2^1*2^100)=1.7979380*2^101// no rounding here!
Great, so the double-precision multiply is fine and introduces no rounding at all. So then how do we get double rounding?
As mentioned above, an add (or subtract) can introduce rounding:
(1.234*10^0)+(1.457*10^9)->(1.2340000*10^0)+(1.4570000*10^9)=(1.45700001234*10^9)// Too many digits!->1.45700000// Rounded to nearest here
This rounding happens at double precision (so well below the threshold of our target 32-bit result), but there’s still rounding, and then the value is rounded again when converted back down to single precision. That’s the double rounding and the source of a potential error.
Okay, so, double rounding is bad? Kinda! But it turns out there is a way to introduce a new rounding mode to use for the first rounding that, in the right situations, does not introduce any additional error and ensures that your final result is correct.
The key to eliminating the extra precision loss is by using a non-standard rounding mode: rounding to odd. Standard floating point rounding calculates results with some additional bits of precision (three bits, to be precise), and then rounds based on the result (usually using “round to nearest with round to even on ties”, although that detail doesn’t end up mattering here - this technique works with any standard rounding mode).
So, assume that we have some way of calculating a double precision addition and also having access to the error between the calculated result and the mathematically exact result. Given those two values we can perform a Round To Odd step:
doubleRoundToOdd(double value,double errorTerm){if(errorTerm !=0.0// if the result is not exact&&LowestBitOfMantissa(value)==0)//and mantissa is even{// We need to round, so round either up or down to oddif(errorTerm >0){// Round up to an odd value
value =AddOneBitToMantissa(value);}else// (errorTerm < 0){// Round down to an odd value
value =SubtractOneBitFromMantissa(value);}}}
Basically: if we have any error at all, and the mantissa is currently even, either add or subtract a single bit’s worth of mantissa, based on the sign of the error.
(In practice, I found that I also had to ensure the result was not Infinity before doing this operation, since I implemented this using some bitwise shenanigans that would end up “rounding” Infinity to NaN, so, you know, watch out for that).
Why Does Odd-Rounding the Intermediate Value Work?
Round to odd works as long as we have more bits of value than the final result - specifically we need at least two extra bits. Standard float rounding makes use of something called a “sticky” bit - basically the lowest bit of the extra precision is a 1 if any of the bits below it would have been 1.
And, hey, that is basically what “round to odd” does!
If the mantissa is odd, regardless of whether there’s error or not the lowest bit is already odd.
If the error was positive and the mantissa was even, we set the lower bit to 1 anyway, effectively stickying (yeah that’s a word now) all the error bits below it.
If the error was negative and the mantissa was even, we subtract 1 from the mantissa, making the lower bit odd, and effectively sticky since some of the digits below it are also 1s.
Effectively, round to odd is just “emulate having a sticky bit at the bottom of your intermediate result” - that way, you have a guaranteed tiebreaker for the final rounding step.
But note that I said it requires you to have at least two extra bits. In the case of our using-doubles-instead-of-singles intermediate addition, good news: we have way more than two extra bits - our intermediate value is a whole-ass double-precision float, so we have 29 extra bits vs. our single-precision final value and (mathematically speaking) 29 is greater than 2.
So, for the true single-precision FMAdd instruction we need to do the following:
floatFMAdd(float a,float b,float c){double product =double(a)*double(b);// No rounding here// Calculate our sum, but somehow get the error along with it(double sum,double err)=AddWithError(product, c);// Round our intermediate value to odd
sum =RoundToOdd(sum, err);// Final rounding here, which now does the correct thing and gives us// a properly-rounded final result (as if we'd used infinite bits)returnfloat(sum);}
That’s it! …wait, what’s that AddWithError function, we haven’t even–
Calculating An Exact Addition Result
Right, we need to calculate that intermediate addition along with some accurate error term. It turns out it’s possible to calculate a set of numbers, sum and error where mathematicallyExactSum = sum + error.
Calculating the error term of adding two numbers (I’ll use x and y) is relatively straightforward if |x| > |y|:
sum = x + y;
err = y -(sum - x);
(this is equation 4.14 in the linked paper)
This is just a different ordering of (x + y) - sum that preserves accuracy: due to the nature of the values involved in these subtractions (sum’s value is directly related to those of x and y, and y is smaller than x), it turns out that each of those subtractions is an exact result (the paper has a proof of this, and it’s a lot so I’m not going to expand on that here), so we get the precise difference between the calculated sum and the real sum.
But this only works if you know that x’s magnitude is larger than (or equal to) y’s. If you don’t know which of the two values has a larger magnitude, you can do a bit more work and end up with:
(double sum,double err)AddWithError(double x,double y){double sum = a + b;double intermediate = sum - x;double err1 = y - intermediate;double err2 = x -(sum - intermediate);return(sum, err1 + err2);}
(This is effectively the expanded version of listing 4.16 from the linked paper)
err1 here is the same as the value in the first version we calculated (a precision-preserving rewrite of (x + y) - sum)
err2 is, mathematically, x - (sum - (sum - x)) or 0; its goal is to calculate the error involved in calculating err1, since without the |x| > |y| guarantee those subtractions might NOT be exact … but these ones will be.
Thus, summing these two error terms together gives us a final, precise error term.
(More details in the paper, hopefully this isn’t too glossed over that it loses any meaning)
Finally, the End (For Single-Precision Floats)
So, yeah, that’s how you implement the FMAdd instruction for single-precision floats on a machine that has double-precision support:
Calculate the double-precision product of a and b
Add this product to c to get a double-precision sum
Calculate the error of the sum
Use the error to odd-round the sum
Round the double-precision sum back down to single precision
But what if you have to calculate FMAdd for double-precision floats? You can’t easily just cast up to, like, quad-precision floats and do the work there, so what now? Can you still do this?
The answer is yes, but it’s a lot more work, and that’s what the next post is about.
C++17 added support for hex float literals, so you can put more bit-accurate floating point values into your code. They’re handy to have, and I wanted to be able to parse them from a text file in a C# app I was writing.
I had a bit of a mental block on this number format for a bit - like, what does it even mean to have fractional hex digits? But it turns out it’s a concept that we already use all the time and my brain just needed some prodding to make the connection.
With our standard base 10 numbers, moving the decimal point left one digit means dividing the number by 10:
12.3==1.23*10^1==0.123*10^2
Hex floats? Same deal, just in 16s instead:
0x1B.C8==0x1.BC8*16^1==0x0.1BC8*16^2
Okay, so now what’s the “p” part in the number? Well, that’s the start of the exponent. A standard float has an exponent starting with ‘e’:
1.3e2==1.3*10^2
But ‘e’ is a hex digit, so you can’t use ‘e’ anymore as the exponent starter, so they chose ‘p’ instead (why not ‘x’, the second letter? Probably because a hex number starts with ‘0x’, so ‘x’ also already has a use - but ‘p’ is free so it wins)
The exponent for a hex float is in powers of 2 (so it corresponds perfectly to the exponent as it is stored in the value), so:
0x1.ABp3==0x1.AB*2^3
So that’s how a hex float literal works! Here’s a quick breakdown:
C++ conveniently has functions to parse these (std::strtod/strtof handle this nicely). However, if you’re (hypothetically) making a parser, and you happen to be writing it in C# which does not have an inbuilt way to parse these, then you’ll have to parse your own.
It ended up being a little more complicated than I thought for a couple reasons.
Arbitrarily long hex strings seem to be supported, which means you need to track both where the top-most non-zero hex value starts (i.e. skip any leading zeros) but also properly handle bits that are extra tiny which may affect rounding.
In order to properly handle said rounding, running through the hex digits in reverse and pushing them in from the top ends up being a nice strategy, because float rounding works via a sticky bit that stays set as things right-shift down through it.
Ultimately I ended up with the following algorithm (in C#) to parse a hex float, which I’m definitely sure is ~perfect~ and has absolutely no bugs whatsoever. It is also absolutely the most efficient version of this possible, with no room for improvement. Yep.
I’m throwing it in here in case anyone ever finds it useful.
staticboolIsDigit(char c)=>(c >='0'&& c <='9');staticboolIsHexDigit(char c)=>IsDigit(c)||(c >='A'&& c <='F')||(c >='a'&& c <='f');doubleParseHexFloat(string s){// This doesn't handle a negative sign, if only because the parser I have only// needed to support positive values, but it'd be easy to addif(s.Length <2|| s[0]!='0'||char.ToLowerInvariant(s[1])!='x'){thrownewFormatException("Missing 0x prefix");}if(s.Length <3||!IsHexDigit(s[2])){thrownewFormatException("Hex float literal must contain at least one whole part digit");}int i =2;int decimalPointIndex =-1;int firstNonZeroHexDigitIndex =-1;// Scan through our digits, looking for the index of the first set (non-zero)// hex value and the decimal point (if we have one).while(i < s.Length &&(IsHexDigit(s[i])|| s[i]=='.')){if(s[i]=='.'){// Found the decimal point! Hopefully there wasn't already one!if(decimalPointIndex >=0){thrownewFormatException("Too many decimal points");}
decimalPointIndex = i;}elseif(s[i]!='0'&& firstNonZeroHexDigitIndex <0){ firstNonZeroHexDigitIndex = i;}// Here's our top-most set hex value.
i++;}// Also make a note of where our last hex digit was (usually the digit before// the 'p' that we should be at right now)int lastHexDigitIndex = i -1;// ... but if the previous character was the decimal point, the last hex// digit is before that.if(lastHexDigitIndex == decimalPointIndex){ lastHexDigitIndex--;}// If we didn't find a decimal point, it's EFFECTIVELY here, at the 'p'if(decimalPointIndex <0){ decimalPointIndex = i;}// Validate and skip the 'p' characterif(i >= s.Length ||char.ToLowerInvariant(s[i])!='p'){thrownewFormatException("Missing exponent 'p'");}
i++;// Grab the sign if we have onebool negativeExponent =false;if(i < s.Length &&(s[i]=='+'|| s[i]=='-')){
negativeExponent =(s[i]=='-');
i++;}if(i >= s.Length){thrownewFormatException("Missing exponent digits");}// Parse the exponent!int exponent =0;while(i < s.Length &&IsDigit(s[i])){if(int.MaxValue /10< exponent){thrownewFormatException("Exponent overflow");}
exponent *=10;
exponent +=(int)(s[i]-'0');
i++;}if(negativeExponent){ exponent =-exponent;}// If we had no non-zero hex digits, there's no point in continuing, it's// zero. if(firstNonZeroHexDigitIndex <0){return0.0;}if(i != s.Length){thrownewFormatException("Unexpected characters at end of string");}// We have the supplied exponent, but we want to massage it a bit. In a// (IEEE) floating-point value, the mantissa is entirely fractional - that// is, the value is 1.mantissa * 2^(exponent) - there's an implied 1// (excluding subnormal floats, which we'll handle properly, but we can// ignore for the moment). Two things need to happen here:// 1. We need to adjust the exponent based on the position of the first// non-zero hex digit, to match the fact that we're parsing hex digits// such that the top hex digit is sitting in the top 4 bits of our 64-// bit int.// 2. But we EXPECT a single bit to be above the mantissa (the implied// 1) so subtract 1 from our adjustment to take into account that// there will be 4 bits in that hex, so if we had parsed a single "1"// (from, say, 0x1p0, which just equals 1.0) our effective exponent// should be 3 (which we will later shift back down to 0 to position// the 1s bit at the very top)if(decimalPointIndex >= firstNonZeroHexDigitIndex){ exponent +=((decimalPointIndex - firstNonZeroHexDigitIndex)*4)-1;}else{ exponent +=((decimalPointIndex - firstNonZeroHexDigitIndex +1)*4)-1;}// Now that we have the exponent and know the bounds of our hex digits,// we can parse backwards through the hex digits, shifting them in from// the top. We do this so that we can easily handle the rounding to the// final 53 bits of significand (by ensuring that we don't ever shift// any 1s off the bottom)ulong mantissa =0;for(i = lastHexDigitIndex; i >= firstNonZeroHexDigitIndex; i--){// Skip the '.' if there was one.if(s[i]=='.'){continue;}char c =char.ToLowerInvariant(s[i]);ulong v =(c >='a'&& c <='f')?(ulong)(c -'a'+10):(ulong)(c -'0');// Shift the mantissa down, but keep any 1s that happen to be in the bottom// 4 bits (this is a reasonably-efficient emulation of the "sticky bit"// that is used to round a floating point number properly.
mantissa =(mantissa >>4)|(mantissa &0xf);// Add in our parsed hex value, putting its 4 bits at the very top of the// mantissa ulong.
mantissa |=(v <<60);}// We know the mantissa is non-zero (checked earlier), and we want to position// the highest set bit at the top of our ulong so shift up until the top bit// is set (and adjust our exponent down 1 to compensate).while((mantissa &0x8000_0000_0000_0000ul)==0){
mantissa <<=1;
exponent--;}constint DoubleExponentBias =1023;constint MaxBiasedDoubleExponent =1023+ DoubleExponentBias;constulong MantissaMask =0x000f_ffff_ffff_fffful;constulong ImpliedOneBit =0x0010_0000_0000_0000ul;constint ExponentShift =52;constint MantissaShiftRight =sizeof(double)*8- ExponentShift -1;// Exponents are stored in a biased form (they can't go negative) so add our// bias now.
exponent += DoubleExponentBias;if(exponent <=0){// We have a subnormal value, which means there is no implied 1, so first// we need to shift our mantissa down by one to get rid of the implied 1.// (note that we're not letting any 1s shift off the bottom, keeping them// sticky)
mantissa =(mantissa >>1)|(mantissa &1);// Continue to denormalize the mantissa until our exponent reaches zerowhile(exponent <0){
mantissa =(mantissa >>1)|(mantissa &1);
exponent++;}}// Now to do the actual rounding of the mantissa. Sometimes floating point// rounding needs 3 bits (guard, round, sticky) to do rounding, but in our// case, two will suffice: one bit that represents the uppermost bit that// shifts right off of the edge of the mantissa (i.e. the "0.5" bit) and// then "literally any 1 bit underneath that" (which is why we've been// holding on to extra 1s when shifting right) that is the tiebreakerbool roundBit =(mantissa &0b10000000000)!=0;bool tiebreakerBit =(mantissa &0b01111111111)!=0;// Now that we have those bits, we can shift our mantissa down into its// proper place (as the lower 53 bits of our 64-bit ulong).
mantissa >>= MantissaShiftRight;if(roundBit){// If there's a tiebreaker, we'll increment the mantissa. Otherwise,// if there's a tie (could round either way), we round so that the// mantissa value is even (lowest bit in the double is 0)if(tiebreakerBit ||(mantissa &1)!=0){ mantissa++;}// If we have a subnormal float we may have overflowed into the implied 1// bit, otherwise we might have overflowed into the ... I guess the// implied *2* bit?ulong overflowedMask = ImpliedOneBit <<((exponent ==0)?0:1);if((mantissa & overflowedMask)!=0){// Shift back down one. This is not going to drop a 1 off the bottom// because if we overflowed it means we were odd, and added one to// become even.
exponent++;
mantissa >>=1;}}// It's possible that the truncation we ended up with a 0 mantissa after all,// so our final value has rounded allll the way down to 0.if(mantissa ==0){return0.0;}// If our exponent is too large to be represented, this value is infinity.if(exponent > MaxBiasedDoubleExponent){returndouble.PositiveInfinity;}// Mask off the implied one bit (if we have one)
mantissa &=~ImpliedOneBit;// Alright assemble the final double's bits, which means shifting and// adding the exponent into its proper place.// (if we had a sign to apply we'd apply it to the top bit). ulong assembled = mantissa |(((ulong)exponent)<< ExponentShift);double result = BitConverter.UInt64BitsToDouble(assembled);return result;}