git.lucas.co / cce-ui
GPU-accelerated UI toolkit (Vulkan)
git clone https://git.lucas.co/cce-ui.git

src/vk/rt_bvh.wgsl (3.5K)

  1 // rt_bvh.wgsl — tier-1 intersect_scene: a CPU-built binned-SAH BVH (rt.rs)
  2 // traversed with an explicit stack. Pure compute; runs on any device.
  3 // Concatenated after rt_common.wgsl at pipeline creation.
  4 
  5 // 32-byte BVH node (rt.rs `GpuBvhNode`): count > 0 marks a leaf over
  6 // tris[left_first .. left_first+count]; otherwise children are at
  7 // left_first and left_first + 1.
  8 struct Node {
  9     bmin: vec3<f32>,
 10     left_first: u32,
 11     bmax: vec3<f32>,
 12     count: u32,
 13 }
 14 @group(0) @binding(1) var<storage, read> nodes: array<Node>;
 15 
 16 // Slab test: entry distance, or 1e30 on miss / beyond the current hit.
 17 fn intersect_aabb(
 18     ro: vec3<f32>,
 19     inv_rd: vec3<f32>,
 20     bmin: vec3<f32>,
 21     bmax: vec3<f32>,
 22     t_limit: f32,
 23 ) -> f32 {
 24     let t1 = (bmin - ro) * inv_rd;
 25     let t2 = (bmax - ro) * inv_rd;
 26     let lo = min(t1, t2);
 27     let hi = max(t1, t2);
 28     let tn = max(max(lo.x, lo.y), lo.z);
 29     let tf = min(min(hi.x, hi.y), hi.z);
 30     if tf >= max(tn, 0.0) && tn < t_limit {
 31         return tn;
 32     }
 33     return 1e30;
 34 }
 35 
 36 // Möller–Trumbore, two-sided (scene triangles have no guaranteed winding).
 37 fn intersect_tri(ro: vec3<f32>, rd: vec3<f32>, i: u32, t_limit: f32) -> f32 {
 38     let tri = tris[i];
 39     let e1 = tri.p1.xyz - tri.p0.xyz;
 40     let e2 = tri.p2.xyz - tri.p0.xyz;
 41     let h = cross(rd, e2);
 42     let a = dot(e1, h);
 43     if abs(a) < 1e-8 {
 44         return 1e30;
 45     }
 46     let f = 1.0 / a;
 47     let s = ro - tri.p0.xyz;
 48     let u = f * dot(s, h);
 49     if u < 0.0 || u > 1.0 {
 50         return 1e30;
 51     }
 52     let q = cross(s, e1);
 53     let v = f * dot(rd, q);
 54     if v < 0.0 || u + v > 1.0 {
 55         return 1e30;
 56     }
 57     let t = f * dot(e2, q);
 58     if t > 1e-4 && t < t_limit {
 59         return t;
 60     }
 61     return 1e30;
 62 }
 63 
 64 // Ordered BVH traversal with an explicit stack.
 65 fn intersect_scene(ro: vec3<f32>, rd: vec3<f32>) -> HitInfo {
 66     var hit = HitInfo(1e30, 0u);
 67     if arrayLength(&nodes) == 0u {
 68         return hit;
 69     }
 70     let inv_rd = vec3<f32>(1.0, 1.0, 1.0) / rd;
 71     var stack: array<u32, 32>;
 72     var sp: u32 = 0u;
 73     var node_idx: u32 = 0u;
 74     if intersect_aabb(ro, inv_rd, nodes[0].bmin, nodes[0].bmax, hit.t) >= 1e30 {
 75         return hit;
 76     }
 77     loop {
 78         let node = nodes[node_idx];
 79         if node.count > 0u {
 80             for (var i: u32 = 0u; i < node.count; i = i + 1u) {
 81                 let tri_idx = node.left_first + i;
 82                 let t = intersect_tri(ro, rd, tri_idx, hit.t);
 83                 if t < hit.t {
 84                     hit.t = t;
 85                     hit.tri = tri_idx;
 86                 }
 87             }
 88             if sp == 0u {
 89                 break;
 90             }
 91             sp = sp - 1u;
 92             node_idx = stack[sp];
 93             continue;
 94         }
 95         // Internal: visit the nearer child first, defer the farther one.
 96         var near = node.left_first;
 97         var far = node.left_first + 1u;
 98         var t_near = intersect_aabb(ro, inv_rd, nodes[near].bmin, nodes[near].bmax, hit.t);
 99         var t_far = intersect_aabb(ro, inv_rd, nodes[far].bmin, nodes[far].bmax, hit.t);
100         if t_far < t_near {
101             let tmp_i = near;
102             near = far;
103             far = tmp_i;
104             let tmp_t = t_near;
105             t_near = t_far;
106             t_far = tmp_t;
107         }
108         if t_near >= 1e30 {
109             if sp == 0u {
110                 break;
111             }
112             sp = sp - 1u;
113             node_idx = stack[sp];
114             continue;
115         }
116         if t_far < 1e30 && sp < 32u {
117             stack[sp] = far;
118             sp = sp + 1u;
119         }
120         node_idx = near;
121     }
122     return hit;
123 }