diff --git a/crates/ml-dqn/src/seg_tree_kernel.cu b/crates/ml-dqn/src/seg_tree_kernel.cu index 7e6274c59..d49a3789d 100644 --- a/crates/ml-dqn/src/seg_tree_kernel.cu +++ b/crates/ml-dqn/src/seg_tree_kernel.cu @@ -7,10 +7,11 @@ // Three kernels: // 1. seg_tree_update — priority update: compute (|td|^alpha + eps), // write to priorities[], compute priority^alpha, -// write to tree leaf, propagate sums to root. +// write to tree leaf, propagate DELTA to root +// via atomicAdd (race-free, O(log n) per thread). // 2. seg_tree_insert — insert: take raw priorities (max_priority fill), // compute priority^alpha, write to tree leaf, -// propagate sums to root. +// propagate DELTA to root via atomicAdd. // 3. seg_tree_sample — proportional sampling: parallel root-to-leaf // traversal with Philox RNG. Output i64 indices. @@ -50,13 +51,18 @@ extern "C" __global__ void seg_tree_update( // Compute priority^alpha for tree leaf (same as pow_alpha_f32) float pa = powf(new_prio, alpha); - // Write leaf and propagate sums to root + // Write leaf and propagate delta to root via atomicAdd. + // Each thread computes delta = new_leaf - old_leaf and adds it to every + // ancestor. atomicAdd is commutative+associative, so concurrent threads + // produce correct sums without synchronization barriers. int leaf = capacity + (int)idx; + float old_pa = tree[leaf]; tree[leaf] = pa; + float delta = pa - old_pa; int node = leaf >> 1; while (node >= 1) { - tree[node] = tree[2 * node] + tree[2 * node + 1]; + atomicAdd(&tree[node], delta); node >>= 1; } } @@ -82,13 +88,15 @@ extern "C" __global__ void seg_tree_insert( // Compute priority^alpha for tree leaf float pa = powf(priorities[i], alpha); - // Write leaf and propagate sums to root + // Write leaf and propagate delta to root via atomicAdd (race-free). int leaf = capacity + (int)idx; + float old_pa = tree[leaf]; tree[leaf] = pa; + float delta = pa - old_pa; int node = leaf >> 1; while (node >= 1) { - tree[node] = tree[2 * node] + tree[2 * node + 1]; + atomicAdd(&tree[node], delta); node >>= 1; } }