import Parameters; #ifndef WORKGROUP_SIZE #define WORKGROUP_SIZE 64 #endif // Num elements #define OCBT_NUM_ELEMENTS 1048576 // Tree sizes #define OCBT_TREE_SIZE_BITS (32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128 + 16 * 256 + 16 * 512 + 16 * 1024 + 16* 2048 + 16 * 4096 + 8 * 8192) #define OCBT_TREE_NUM_SLOTS (OCBT_TREE_SIZE_BITS / 32) #define OCBT_BITFIELD_NUM_SLOTS (OCBT_NUM_ELEMENTS / 64) #define OCBT_LAST_LEVEL_SIZE 8192 // Tree last level #define TREE_LAST_LEVEL 13 // First virtual level #define FIRST_VIRTUAL_LEVEL 14 // Leaf level #define LEAF_LEVEL 20 // per level offset static const uint32_t OCBT_depth_offset[21] = { 0, // Level 0 32 * 1, // level 1 32 * 1 + 32 * 2, // level 2 32 * 1 + 32 * 2 + 32 * 4, // level 3 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8, // Level 4 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16, // Level 5 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32, // Level 6 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64, // Level 7 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128, // Level 8 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128 + 16 * 256, // Level 9 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128 + 16 * 256 + 16 * 512, // Level 10 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128 + 16 * 256 + 16 * 512 + 16 * 1024, // Level 11 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128 + 16 * 256 + 16 * 512 + 16 * 1024 + 16 * 2048, // Level 12 32 * 1 + 32 * 2 + 32 * 4 + 32 * 8 + 32 * 16 + 32 * 32 + 32 * 64 + 16 * 128 + 16 * 256 + 16 * 512 + 16 * 1024 + 16 * 2048 + 16 * 4096, // Level 13 0, // Level 14 0, // Level 15 0, // Level 16 0, // Level 17 0, // Level 18 0, // Level 19 0, // Level 20 }; static const uint64_t OCBT_bit_mask[21] = { 0xffffffff, // Root 17 0xffffffff, // Level 16 0xffffffff, // level 15 0xffffffff, // level 14 0xffffffff, // level 13 0xffffffff, // level 12 0xffffffff, // level 11 0xffff, // level 10 0xffff, // level 9 0xffff, // level 8 0xffff, // level 8 0xffff, // level 8 0xffff, // level 8 0xff, // level 8 0xffffffffffffffff, // level 7 0xffffffff, // Level 6 0xffff, // level 5 0xff, // level 4 0xf, // level 3 0x3, // level 2 0x1, // level 1 }; static const uint32_t OCBT_bit_count[21] = { 32, // Root 17 32, // Level 16 32, // level 15 32, // level 14 32, // level 13 32, // level 12 32, // level 11 16, // level 10 16, // level 9 16, // level 8 16, // level 8 16, // level 8 16, // level 8 8, // level 8 64, // Level 5 32, // Level 5 16, // Level 4 8, // level 3 4, // level 2 2, // level 1 1, // level 0 }; // Define the remaining values #define BUFFER_ELEMENT_PER_LANE ((OCBT_TREE_NUM_SLOTS + WORKGROUP_SIZE - 1) / WORKGROUP_SIZE) #define BUFFER_ELEMENT_PER_LANE_NO_BITFIELD ((OCBT_TREE_NUM_SLOTS + WORKGROUP_SIZE - 1) / WORKGROUP_SIZE) #define BITFIELD_ELEMENT_PER_LANE ((OCBT_BITFIELD_NUM_SLOTS + WORKGROUP_SIZE - 1) / WORKGROUP_SIZE) #define WAVE_TREE_DEPTH uint(log2(OCBT_NUM_ELEMENTS)) uint32_t cbt_size() { return OCBT_NUM_ELEMENTS; } uint32_t last_level_offset() { return OCBT_depth_offset[TREE_LAST_LEVEL] / 32; } groupshared uint gs_cbtTree[OCBT_TREE_NUM_SLOTS]; // Function that sets a given bit void set_bit(uint bitID, bool state) { // Coordinates of the bit uint32_t slot = bitID / 64; uint32_t local_id = bitID % 64; if (state) pParams.bitFieldBuffer[slot] |= 1uLL << local_id; else pParams.bitFieldBuffer[slot] &= ~(1uLL << local_id); } void set_bit_atomic(uint bitID, bool state) { // Coordinates of the bit uint32_t slot = bitID / 64; uint32_t local_id = bitID % 64; if (state) InterlockedOr(pParams.bitFieldBuffer[slot], 1uLL << local_id); else InterlockedAnd(pParams.bitFieldBuffer[slot], ~(1uLL << local_id)); } uint get_bit(uint bitID) { uint32_t slot = bitID / 64; uint32_t local_id = bitID % 64; return uint((pParams.bitFieldBuffer[slot] & (1uLL << local_id)) >> local_id); } uint get_heap_element(uint id) { // Figure out the location of the first bit of this element uint32_t real_heap_id = id - 1; uint32_t depth = uint32_t(log2(real_heap_id + 1)); uint32_t level_first_element = (1u << depth) - 1; uint32_t id_in_level = real_heap_id - level_first_element; uint32_t first_bit = OCBT_depth_offset[depth] + OCBT_bit_count[depth] * id_in_level; if (depth < FIRST_VIRTUAL_LEVEL) { uint32_t slot = first_bit / 32; uint32_t local_id = first_bit % 32; uint32_t target_bits = (gs_cbtTree[slot] >> local_id) & uint32_t(OCBT_bit_mask[depth]); return (gs_cbtTree[slot] >> local_id) & uint32_t(OCBT_bit_mask[depth]); } else { uint32_t slot = first_bit / 64; uint32_t local_id = first_bit % 64; uint64_t target_bits = (pParams.bitFieldBuffer[slot] >> local_id) & OCBT_bit_mask[depth]; uint32_t high = uint(target_bits >> 32); uint32_t low = uint(target_bits); return countbits(high) + countbits(low); } } // Should not be called if depth > TREE_LAST_LEVEL void set_heap_element(uint id, uint value) { // Figure out the location of the first bit of this element uint real_heap_id = id - 1; uint depth = uint(log2(real_heap_id + 1)); uint level_first_element = (1u << depth) - 1; uint first_bit = OCBT_depth_offset[depth] + OCBT_bit_count[depth] * (real_heap_id - level_first_element); // Find the slot and the local first bit uint slot = first_bit / 32; uint local_id = first_bit % 32; // Extract the relevant bits gs_cbtTree[slot] &= ~(uint32_t(OCBT_bit_mask[depth]) << local_id); gs_cbtTree[slot] |= ((uint32_t(OCBT_bit_mask[depth]) & value) << local_id); } // Should not be called if depth > TREE_LAST_LEVEL void set_heap_element_atomic(uint id, uint value) { // Figure out the location of the first bit of this element uint real_heap_id = id - 1; uint depth = uint(log2(real_heap_id + 1)); uint level_first_element = (1u << depth) - 1; uint first_bit = OCBT_depth_offset[depth] + OCBT_bit_count[depth] * (real_heap_id - level_first_element); // Find the slot and the local first bit uint slot = first_bit / 32; uint local_id = first_bit % 32; // Extract the relevant bits InterlockedAnd(gs_cbtTree[slot], ~(uint32_t(OCBT_bit_mask[depth]) << local_id)); InterlockedOr(gs_cbtTree[slot], ((uint32_t(OCBT_bit_mask[depth]) & value) << local_id)); } // Function that returns the number of active bits uint bit_count() { return gs_cbtTree[0]; } uint bit_count(uint depth, uint element) { return get_heap_element((1u << depth) + element); } // decodes the position of the i-th one in the bitfield uint decode_bit(uint handle) { #if defined(NAIVE_DECODE) uint bitID = 1; for (uint currentDepth = 0; currentDepth < WAVE_TREE_DEPTH; ++currentDepth) { uint heapValue = get_heap_element(2 * bitID); uint b = handle < heapValue ? 0 : 1; bitID = 2 * bitID + b; handle -= heapValue * b; } return (bitID ^ OCBT_NUM_ELEMENTS); #else uint currentDepth = 0; uint heapElementID = 1u; for (currentDepth = 0; currentDepth < FIRST_VIRTUAL_LEVEL; ++currentDepth) { // Read the left element uint heapValue = get_heap_element(2u * heapElementID); // Does it fall in the right or left subtree? uint b = handle < heapValue ? 0u : 1u; // Pick a subtree heapElementID = 2u * heapElementID + b; // Move the iterator to exclude the right subtree if required handle -= heapValue * b; } // Align with the internal depth currentDepth++; // Ok we have our subtree, now we need to pick the right bit uint64_t heapValue = pParams.bitFieldBuffer[heapElementID - OCBT_LAST_LEVEL_SIZE * 2]; uint64_t mask = 0xffffffff; uint32_t bitCount = 32; for (; currentDepth < (WAVE_TREE_DEPTH + 1); ++currentDepth) { // Figure out the location of the first bit of this element uint real_heap_id = 2 * heapElementID - 1; uint level_first_element = (1u << currentDepth) - 1; uint id_in_level = real_heap_id - level_first_element; uint first_bit = bitCount * id_in_level; uint local_id = first_bit % 64; uint64_t target_bits = (heapValue >> local_id) & mask; uint32_t high = uint(target_bits >> 32); uint32_t low = uint(target_bits); uint heapValue = countbits(high) + countbits(low); // Does it fall in the right or left subtree? uint b = handle < heapValue ? 0u : 1u; // Pick a subtree heapElementID = 2u * heapElementID + b; // Move the iterator to exclude the right subtree if required handle -= heapValue * b; // Adjust the mask and bitcount bitCount /= 2; mask = mask >> bitCount; } return (heapElementID ^ OCBT_NUM_ELEMENTS); #endif } // decodes the position of the i-th zero in the bitfield uint decode_bit_complement(uint handle) { #if defined(NAIVE_DECODE) uint bitID = 1u; uint c = OCBT_NUM_ELEMENTS / 2u; while (bitID < OCBT_NUM_ELEMENTS) { uint heapValue = c - get_heap_element(2u * bitID); uint b = handle < heapValue ? 0u : 1u; bitID = 2u * bitID + b; handle -= heapValue * b; c /= 2u; } return (bitID ^ OCBT_NUM_ELEMENTS); #else uint heapElementID = 1u; uint c = OCBT_NUM_ELEMENTS / 2u; uint currentDepth = 0; for (currentDepth = 0; currentDepth < FIRST_VIRTUAL_LEVEL; ++currentDepth) { uint heapValue = c - get_heap_element(2u * heapElementID); uint b = handle < heapValue ? 0u : 1u; heapElementID = 2u * heapElementID + b; handle -= heapValue * b; c /= 2u; } // Align with the internal depth currentDepth++; // Ok we have our subtree, now we need to pick the right bit uint64_t heapValue = pParams.bitFieldBuffer[heapElementID - OCBT_LAST_LEVEL_SIZE * 2]; uint64_t mask = 0xffffffff; uint32_t bitCount = 32; for (; currentDepth < (WAVE_TREE_DEPTH + 1); ++currentDepth) { // Figure out the location of the first bit of this element uint real_heap_id = 2 * heapElementID - 1; uint level_first_element = (1u << currentDepth) - 1; uint id_in_level = real_heap_id - level_first_element; uint first_bit = bitCount * id_in_level; uint local_id = first_bit % 64; uint64_t target_bits = (heapValue >> local_id) & mask; uint32_t high = uint(target_bits >> 32); uint32_t low = uint(target_bits); uint heapValue = c - (countbits(high) + countbits(low)); uint b = handle < heapValue ? 0u : 1u; heapElementID = 2u * heapElementID + b; handle -= heapValue * b; c /= 2u; // Adjust the mask and bitcount bitCount /= 2; mask = mask >> bitCount; } return (heapElementID ^ OCBT_NUM_ELEMENTS); #endif } void reduce(uint groupIndex) { // First we do a reduction until each lane has exactly one element to process uint initial_pass_size = OCBT_NUM_ELEMENTS / WORKGROUP_SIZE; for (uint it = initial_pass_size / 64, offset = OCBT_NUM_ELEMENTS / 64; it > 0 ; it >>=1, offset >>=1) { uint minHeapID = offset + (groupIndex * it); uint maxHeapID = offset + ((groupIndex + 1) * it); for (uint heapID = minHeapID; heapID < maxHeapID; ++heapID) { set_heap_element(heapID, get_heap_element(heapID * 2) + get_heap_element(heapID * 2 + 1)); } } GroupMemoryBarrierWithGroupSync(); for(uint s = WORKGROUP_SIZE / 2; s > 0u; s >>= 1) { if (groupIndex < s) { uint v = s + groupIndex; set_heap_element(v, get_heap_element(v * 2) + get_heap_element(v * 2 + 1)); } GroupMemoryBarrierWithGroupSync(); } } void reduce_prepass(uint dispatchThreadID) { // Initialize the packed sum uint packedSum = 0; // Loop through the 4 pairs to process for (uint pairIdx = 0; pairIdx < 4; ++pairIdx) { // First element of the pair uint64_t target_bits = pParams.bitFieldBuffer[dispatchThreadID * 8 + 2 * pairIdx]; uint32_t high = uint(target_bits >> 32); uint32_t low = uint(target_bits); uint elementC = countbits(high) + countbits(low); // Second element of the pair target_bits = pParams.bitFieldBuffer[dispatchThreadID * 8 + 2 * pairIdx + 1]; high = uint(target_bits >> 32); low = uint(target_bits); elementC += countbits(high) + countbits(low); // Store in the right bits packedSum |= (elementC << pairIdx * 8); } // Offset of the last level of the tree const uint bufferOffset = last_level_offset(); // Store the result into the bitfield pParams.cbtBuffer[bufferOffset + dispatchThreadID] = packedSum; } void reduce_first_pass(uint dispatchThreadID, uint groupIndex) { // Load the lowest level (and only the last level) const uint level0Offset = OCBT_depth_offset[TREE_LAST_LEVEL] / 32; for (uint e = 0; e < 4; ++e) { uint target_element = 4 * dispatchThreadID + e; gs_cbtTree[level0Offset + target_element] = pParams.cbtBuffer[level0Offset + target_element]; } GroupMemoryBarrierWithGroupSync(); // First we do a reduction until each lane has exactly one element to process uint initial_pass_size = OCBT_LAST_LEVEL_SIZE / 2; uint it, offset; for (it = initial_pass_size / 512, offset = initial_pass_size; it > 1; it >>=1, offset >>=1) { uint minHeapID = offset + (dispatchThreadID * it); uint maxHeapID = offset + ((dispatchThreadID + 1) * it); for (uint heapID = minHeapID; heapID < maxHeapID; ++heapID) { set_heap_element(heapID, get_heap_element(heapID * 2) + get_heap_element(heapID * 2 + 1)); } } // Last pass needs to be atomic uint heapID = offset + (dispatchThreadID * it); set_heap_element_atomic(heapID, get_heap_element(heapID * 2) + get_heap_element(heapID * 2 + 1)); GroupMemoryBarrierWithGroupSync(); // Load the first reduced level const uint level2Offset = OCBT_depth_offset[TREE_LAST_LEVEL - 1] / 32; for (uint e = 0; e < 4; ++e) { uint target_element = 4 * dispatchThreadID + e; pParams.cbtBuffer[level2Offset + target_element] = gs_cbtTree[level2Offset + target_element]; } // Load the first reduced level const uint level3Offset = OCBT_depth_offset[TREE_LAST_LEVEL - 2] / 32; for (uint e = 0; e < 2; ++e) { uint target_element = 2 * dispatchThreadID + e; pParams.cbtBuffer[level3Offset + target_element] = gs_cbtTree[level3Offset + target_element]; } const uint level4Offset = OCBT_depth_offset[TREE_LAST_LEVEL - 3] / 32; pParams.cbtBuffer[level4Offset + dispatchThreadID] = gs_cbtTree[level4Offset + dispatchThreadID]; const uint level5Offset = OCBT_depth_offset[TREE_LAST_LEVEL - 4] / 32; if (groupIndex % 2 == 0) pParams.cbtBuffer[level5Offset + dispatchThreadID / 2] = gs_cbtTree[level5Offset + dispatchThreadID / 2]; } void reduce_second_pass(uint groupIndex) { // Load the lowest level (and only the last level) const uint level0Offset = OCBT_depth_offset[9] / 32; for (uint e = 0; e < 4; ++e) { uint target_element = 4 * groupIndex + e; gs_cbtTree[level0Offset + target_element] = pParams.cbtBuffer[level0Offset + target_element]; } GroupMemoryBarrierWithGroupSync(); // First we do a reduction until each lane has exactly one element to process uint initial_pass_size = 256; for (uint it = initial_pass_size / 64, offset = initial_pass_size; it > 0 ; it >>=1, offset >>=1) { uint minHeapID = offset + (groupIndex * it); uint maxHeapID = offset + ((groupIndex + 1) * it); for (uint heapID = minHeapID; heapID < maxHeapID; ++heapID) { set_heap_element(heapID, get_heap_element(heapID * 2) + get_heap_element(heapID * 2 + 1)); } } GroupMemoryBarrierWithGroupSync(); for(uint s = WORKGROUP_SIZE / 2; s > 0u; s >>= 1) { if (groupIndex < s) { uint v = s + groupIndex; set_heap_element(v, get_heap_element(v * 2) + get_heap_element(v * 2 + 1)); } GroupMemoryBarrierWithGroupSync(); } // Make sure all the previous operations are done GroupMemoryBarrierWithGroupSync(); // Load the bitfield to the LDS for (uint e = 0; e < 5; ++e) { uint target_element = 5 * groupIndex + e; if (target_element < 319) pParams.cbtBuffer[target_element] = gs_cbtTree[target_element]; } } void reduce_no_bitfield(uint groupIndex) { // First we do a reduction until each lane has exactly one element to process uint initial_pass_size = OCBT_NUM_ELEMENTS / WORKGROUP_SIZE; for (uint it = initial_pass_size / 128, offset = OCBT_NUM_ELEMENTS / 128; it > 0 ; it >>=1, offset >>=1) { uint minHeapID = offset + (groupIndex * it); uint maxHeapID = offset + ((groupIndex + 1) * it); for (uint heapID = minHeapID; heapID < maxHeapID; ++heapID) { set_heap_element(heapID, get_heap_element(heapID * 2) + get_heap_element(heapID * 2 + 1)); } } GroupMemoryBarrierWithGroupSync(); for(uint s = WORKGROUP_SIZE / 2; s > 0u; s >>= 1) { if (groupIndex < s) { uint v = s + groupIndex; set_heap_element(v, get_heap_element(v * 2) + get_heap_element(v * 2 + 1)); } GroupMemoryBarrierWithGroupSync(); } } void clear_cbt(uint groupIndex) { for (uint v = 0; v < BUFFER_ELEMENT_PER_LANE; ++v) { uint target_element = BUFFER_ELEMENT_PER_LANE * groupIndex + v; if (target_element < OCBT_TREE_NUM_SLOTS) gs_cbtTree[target_element] = 0; } for (uint b = 0; b < BITFIELD_ELEMENT_PER_LANE; ++b) { uint target_element = BITFIELD_ELEMENT_PER_LANE * groupIndex + b; if (target_element < OCBT_BITFIELD_NUM_SLOTS) pParams.bitFieldBuffer[target_element] = 0; } GroupMemoryBarrierWithGroupSync(); } // Importante note // Depending on your target GPU architecture, the pattern used to load has a different performance behavior // here is the best performant based on our tests: // NVIDIA uint target_element = groupIndex + WORKGROUP_SIZE * e; // AMD uint target_element = BUFFER_ELEMENT_PER_LANE * groupIndex + e; void load_buffer_to_shared_memory(uint groupIndex) { // Load the bitfield to the LDS for (uint e = 0; e < BUFFER_ELEMENT_PER_LANE; ++e) { #ifdef AMD uint target_element = BUFFER_ELEMENT_PER_LANE * groupIndex + e; #else uint target_element = groupIndex + WORKGROUP_SIZE * e; #endif if (target_element < OCBT_TREE_NUM_SLOTS) gs_cbtTree[target_element] = pParams.cbtBuffer[target_element]; } GroupMemoryBarrierWithGroupSync(); } void load_shared_memory_to_buffer(uint groupIndex) { // Make sure all the previous operations are done GroupMemoryBarrierWithGroupSync(); // Load the bitfield to the LDS for (uint e = 0; e < BUFFER_ELEMENT_PER_LANE; ++e) { #ifdef AMD uint target_element = BUFFER_ELEMENT_PER_LANE * groupIndex + e; #else uint target_element = groupIndex + WORKGROUP_SIZE * e; #endif if (target_element < OCBT_TREE_NUM_SLOTS) pParams.cbtBuffer[target_element] = gs_cbtTree[target_element]; } } void set_bit_atomic_buffer(uint bitID, bool state) { // Coordinates of the bit uint32_t slot = bitID / 64; uint32_t local_id = bitID % 64; if (state) InterlockedOr(pParams.bitFieldBuffer[slot], 1uLL << local_id); else InterlockedAnd(pParams.bitFieldBuffer[slot], ~(1uLL << local_id)); } uint32_t bit_count_buffer() { return pParams.cbtBuffer[0]; }