words, this is a SIMD analog of `out[i] = table[index[i]]`.
We will use the `vqtbl1q_u8` instruction:
| part | meaning |
|:----:|------------------------------------------|
| v | vector intrinsic |
| q | table consists of 128-bit registers |
| tbl | table lookup |
| 1 | number of registers in table |
| q | result and indices are 128-bit registers |
| u8 | elements of table are `uint8_t` |
`tbl` permutes bytes, but we need to select floats (4 bytes). So, we will create `index` in blocks of 4 bytes:
to select the second (0-based) float of the register, `index` will contain its bytes [8, 9, 10, 11] (the second element
starts at an offset of `2 * sizeof(float) = 8`).
Computing `index` every time is slow. There are 16 variants in total (4 elements to take/drop), so we will precompute all the `index` variants.
But to select the `index` using the mask, we need to convert the mask to a number (call it `idx`):
### mask → idx
The mask consists of 4 elements, each either 0x00000000 (false) or 0xFFFFFFFF (true).
If the i-th element is true, we want to set the i-th bit in idx.
Trick: `mask & [1, 2, 4, 8]`. Because `0xFFFFFFFF & x = x`, the true elements keep their weight (1/2/4/8), while the false ones become 0.
We add all elements together and get a number between 0 and 15.
```cpp
std::array<uint32_t, 4> weights{1, 2, 4, 8};
size_t idx = vaddvq_u32(vandq_u32(mask, vld1q_u32(
weights.data())));
```
- `vld1q_u32(
weights.data())` - load 4 values from memory at address `
weights.data()` into a register (ld - load)
- `vandq_u32` - elementwise & (and)
- `vaddvq_u32` - sum of all the elements in the register (addv - add across vector)
### Precompute the `index` table
There is no way to compute registers at compile time, so instead of `uint8x16_t` (register of 16 `uint8_t`) we will store `std::array<uint8_t, 16>`.
For each `idx` we will go through the 4 elements of `mask`. If the element is selected,
we append the indices of its 4 bytes into `index` at the cursor position and advance the cursor by 4.
```cpp
consteval auto make_index_table() {
std::array<std::array<uint8_t, 16>, 16> index{};
for (size_t idx = 0; idx < 16; ++idx) { // iterate over all masks
size_t j = 0; // j is the cursor
for (size_t i = 0; i < 4; ++i) // iterate over the mask's elements
if (idx & (1 << i)) // if the i-th element is selected
for (size_t k = 0; k < 4; ++k) // iterate over its bytes
index[idx][j++] = i * 4 + k; // store the indices of its bytes
}
return index;
}
```
The `j` cursor advances only on selected elements, so their bytes are placed in `index` consecutively. `tbl` with that `index`
collects floats into a register. Unused positions in `index` are zeros, so in the tail, after `count` elements, there will be garbage.
### The `count` table
Next we need to compute the number of elements we select. Similarly we can precompute a table for this:
```cpp
consteval auto make_count_table() {
std::array<uint8_t, 16> count{};
for (size_t idx = 0; idx < 16; ++idx)
for (size_t i = 0; i < 4; ++i)
if (idx & 1 << i)
++count[idx];
return count;
}
```
### The full `compress`
```cpp
auto compress(uint32x4_t mask, float32x4_t a) {
static constexpr std::array<uint32_t, 4> weights{1, 2, 4, 8};
const size_t idx = vaddvq_u32(vandq_u32(mask, vld1q_u32(
weights.data())));
static constexpr auto count = make_count_table();
static constexpr auto index_table = make_index_table();
const auto index = vld1q_u8(index_table[idx].data()); // at runtime, loads only one row of the table into a register
return std::pair{vreinterpretq_f32_u8(vqtbl1q_u8(vreinterpretq_u8_f32(a), index)), count[idx]};
}
```
Because `tbl` works only with u8, we need to cast `a` to u8 and then cast the result back to f32.
We write full