Skip to main content

ursus_core/render/
graph.rs

1use crate::assets::storage::GpuAssetServer;
2use crate::render::gfx::descriptor::ImageUsage;
3use crate::render::gfx::types::format::ImageLayout;
4use crate::render::gfx::types::{DescriptorSetId, SamplerId};
5use crate::render::gfx::CommandEncoder;
6use crate::render::resource::{
7    DescriptorBinding, DescriptorBindingRegistry, DescriptorImageType, FlushReason, LayoutTracker, ResourceHandle,
8    ResourcePool,
9};
10use crate::render::world::RWorld;
11use crate::vulkan::core::debug::{cmd_begin_label, cmd_end_label};
12use crate::vulkan::core::DeviceContext;
13use crate::vulkan::timestamps::{GpuFrameTimes, GpuTimestampPool};
14use ash::ext::debug_utils;
15use ash::vk;
16use std::collections::{HashMap, HashSet, VecDeque};
17use std::sync::Arc;
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum AccessType {
21    Read,
22    Write,
23    ReadWrite,
24}
25
26#[derive(Debug, Clone)]
27pub struct PassAccess {
28    pub handle: ResourceHandle,
29    pub access: AccessType,
30    pub layout: vk::ImageLayout,
31}
32
33impl PassAccess {
34    pub fn read(handle: ResourceHandle, layout: vk::ImageLayout) -> Self {
35        Self { handle, access: AccessType::Read, layout }
36    }
37
38    pub fn write(handle: ResourceHandle, layout: vk::ImageLayout) -> Self {
39        Self { handle, access: AccessType::Write, layout }
40    }
41
42    pub fn read_write(handle: ResourceHandle, layout: vk::ImageLayout) -> Self {
43        Self { handle, access: AccessType::ReadWrite, layout }
44    }
45}
46pub type RecordFn = Box<dyn FnMut(&mut CommandEncoder<'_>, &RWorld, &GpuAssetServer) -> anyhow::Result<()> + Send>;
47
48pub struct PassNode {
49    pub name: String,
50    pub accesses: Vec<PassAccess>,
51    pub record: RecordFn,
52    pub enabled: bool,
53    pub depends_on: Vec<PassHandle>,
54}
55
56impl PassNode {
57    pub fn new(name: impl Into<String>, accesses: Vec<PassAccess>, record: RecordFn) -> Self {
58        Self { name: name.into(), accesses, record, enabled: true, depends_on: Vec::new() }
59    }
60}
61
62impl std::fmt::Debug for PassNode {
63    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
64        f.debug_struct("PassNode")
65            .field("name", &self.name)
66            .field("enabled", &self.enabled)
67            .field("accesses", &self.accesses.len())
68            .finish()
69    }
70}
71
72#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
73pub struct PassHandle(pub(crate) u32);
74
75#[derive(Clone)]
76struct CompiledBarrier {
77    handle: ResourceHandle,
78    new_layout: vk::ImageLayout,
79}
80
81struct CompiledPass {
82    node_index: usize,
83    barriers: Vec<CompiledBarrier>,
84}
85
86pub struct RenderGraph {
87    pub pool: ResourcePool,
88    nodes: Vec<PassNode>,
89    sorted_order: Vec<usize>,
90    tracker: LayoutTracker,
91    bindings: DescriptorBindingRegistry,
92    debug_utils: Option<Arc<debug_utils::Device>>,
93
94    internal_resolution: (u32, u32),
95    output_resolution: (u32, u32),
96
97    compiled: bool,
98    allocated: bool,
99
100    compiled_passes: Vec<CompiledPass>,
101    compiled_finals: Vec<CompiledBarrier>,
102
103    timestamps: Option<GpuTimestampPool>,
104    pub last_frame_times: Option<GpuFrameTimes>,
105    current_frame: usize,
106    frames_in_flight: usize,
107}
108
109impl RenderGraph {
110    pub fn new(
111        pool: ResourcePool,
112        device: ash::Device,
113        internal_resolution: (u32, u32),
114        output_resolution: (u32, u32),
115        debug_utils: Option<Arc<debug_utils::Device>>,
116    ) -> Self {
117        Self {
118            bindings: DescriptorBindingRegistry::new(device),
119            pool,
120            nodes: Vec::new(),
121            sorted_order: Vec::new(),
122            tracker: LayoutTracker::new(),
123            internal_resolution,
124            output_resolution,
125            compiled: false,
126            allocated: false,
127            compiled_passes: Vec::new(),
128            compiled_finals: Vec::new(),
129            debug_utils,
130            timestamps: None,
131            last_frame_times: None,
132            current_frame: 0,
133            frames_in_flight: 0,
134        }
135    }
136
137    pub fn enable_timestamps(
138        &mut self,
139        ctx: DeviceContext,
140        frames_in_flight: u32,
141        command_pool: vk::CommandPool,
142        queue: vk::Queue,
143    ) -> anyhow::Result<()> {
144        assert!(self.compiled, "enable_timestamps вызван до compile()");
145
146        let pass_names = self.nodes.iter().map(|n| n.name.clone()).collect();
147
148        self.timestamps = Some(GpuTimestampPool::new(
149            ctx.device,
150            ctx.physical_device,
151            ctx.instance,
152            frames_in_flight,
153            pass_names,
154            command_pool,
155            queue,
156        )?);
157        self.frames_in_flight = frames_in_flight as usize;
158        Ok(())
159    }
160
161    pub fn disable_timestamps(&mut self) {
162        self.timestamps = None;
163        self.last_frame_times = None;
164    }
165
166    pub fn add_pass(&mut self, node: PassNode) -> PassHandle {
167        for access in &node.accesses {
168            let extra = match access.layout {
169                vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL => ImageUsage::SAMPLED,
170                vk::ImageLayout::COLOR_ATTACHMENT_OPTIMAL => ImageUsage::COLOR_ATTACHMENT,
171                vk::ImageLayout::DEPTH_ATTACHMENT_OPTIMAL => ImageUsage::DEPTH_ATTACHMENT,
172                _ => ImageUsage::empty(),
173            };
174            self.pool.add_usage(access.handle, extra);
175        }
176
177        let handle = PassHandle(self.nodes.len() as u32);
178        self.nodes.push(node);
179        self.compiled = false;
180        handle
181    }
182
183    pub fn compile(&mut self) -> anyhow::Result<()> {
184        let n = self.nodes.len();
185        let mut adj: Vec<HashSet<usize>> = vec![HashSet::new(); n];
186        let mut in_degree = vec![0usize; n];
187
188        build_resource_edges(&self.nodes, &mut adj, &mut in_degree);
189        build_explicit_edges(&self.nodes, &mut adj, &mut in_degree);
190
191        self.sorted_order = topological_sort(n, adj, in_degree)?;
192        self.compiled = true;
193
194        self.build_compiled_passes();
195
196        log::info!(
197            "RenderGraph скомпилирован: {} пассов -> {:?}",
198            n,
199            self.nodes.iter().enumerate().map(|(i, p)| format!("[{}]{}", i, p.name)).collect::<Vec<_>>()
200        );
201        Ok(())
202    }
203
204    fn build_compiled_passes(&mut self) {
205        let mut sim_layouts: HashMap<ResourceHandle, vk::ImageLayout> = HashMap::new();
206
207        self.compiled_passes.clear();
208        self.compiled_finals.clear();
209
210        for &idx in &self.sorted_order {
211            let node = &self.nodes[idx];
212            if !node.enabled {
213                self.compiled_passes.push(CompiledPass { node_index: idx, barriers: Vec::new() });
214                continue;
215            }
216
217            let mut barriers = Vec::new();
218            for access in &node.accesses {
219                let old = sim_layouts.get(&access.handle).copied().unwrap_or(vk::ImageLayout::UNDEFINED);
220                if old != access.layout {
221                    barriers.push(CompiledBarrier { handle: access.handle, new_layout: access.layout });
222                    sim_layouts.insert(access.handle, access.layout);
223                }
224            }
225
226            self.compiled_passes.push(CompiledPass { node_index: idx, barriers });
227        }
228
229        for handle in self.pool.external_handles().collect::<Vec<_>>() {
230            if let Some(final_layout) = self.pool.external_final_layout(handle) {
231                self.compiled_finals.push(CompiledBarrier { handle, new_layout: final_layout });
232            }
233        }
234    }
235
236    pub fn bind_resource(&mut self, binding: DescriptorBinding) {
237        self.bindings.register(binding);
238    }
239
240    pub fn bind_sampled(&mut self, resource: ResourceHandle, set: DescriptorSetId, binding: u32, sampler: SamplerId) {
241        self.bindings.register(DescriptorBinding {
242            resource,
243            set,
244            binding,
245            array_element: 0,
246            image_type: DescriptorImageType::CombinedImageSampler(sampler),
247            image_layout: vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL,
248        });
249    }
250
251    pub fn allocate(&mut self, gpu: &GpuAssetServer) -> anyhow::Result<()> {
252        self.pool.allocate(self.internal_resolution, self.output_resolution)?;
253        self.bindings.flush_all(&self.pool, gpu, FlushReason::InitialAllocation);
254        self.allocated = true;
255        Ok(())
256    }
257
258    pub fn reset_external_layouts(&mut self) {
259        for handle in self.pool.external_handles().collect::<Vec<_>>() {
260            if let Some(initial) = self.pool.external_initial_layout(handle) {
261                self.tracker.set(handle, initial);
262            }
263        }
264    }
265
266    pub fn execute(
267        &mut self,
268        device: &ash::Device,
269        cmd: vk::CommandBuffer,
270        rw: &RWorld,
271        gpu_assets: &GpuAssetServer,
272    ) -> anyhow::Result<()> {
273        assert!(self.compiled, "RenderGraph::compile() не был вызван");
274
275        if let Some(ts) = &mut self.timestamps {
276            ts.read_and_reset(self.current_frame, cmd);
277        }
278
279        for (order_idx, cp) in self.compiled_passes.iter().enumerate() {
280            let node = &mut self.nodes[cp.node_index];
281            if !node.enabled {
282                continue;
283            }
284
285            if let Some(du) = &self.debug_utils {
286                cmd_begin_label(du, cmd, &node.name);
287            }
288
289            let barriers =
290                self.tracker.plan_transition(&self.pool, cp.barriers.iter().map(|cb| (cb.handle, cb.new_layout)));
291            if !barriers.is_empty() {
292                unsafe {
293                    device.cmd_pipeline_barrier2(cmd, &vk::DependencyInfo::default().image_memory_barriers(barriers));
294                }
295            }
296
297            if let Some(ts) = &self.timestamps {
298                ts.begin_pass(cmd, self.current_frame, order_idx);
299            }
300
301            {
302                let mut encoder = CommandEncoder::new(device, cmd, &self.pool, gpu_assets.pipeline_cache(), gpu_assets);
303                (node.record)(&mut encoder, rw, gpu_assets)?;
304            }
305
306            if let Some(ts) = &self.timestamps {
307                ts.end_pass(cmd, self.current_frame, order_idx);
308            }
309
310            if let Some(du) = &self.debug_utils {
311                cmd_end_label(du, cmd);
312            }
313        }
314
315        let barriers =
316            self.tracker.plan_transition(&self.pool, self.compiled_finals.iter().map(|cb| (cb.handle, cb.new_layout)));
317        if !barriers.is_empty() {
318            unsafe {
319                device.cmd_pipeline_barrier2(cmd, &vk::DependencyInfo::default().image_memory_barriers(barriers));
320            }
321        }
322
323        if let Some(ts) = &self.timestamps {
324            self.last_frame_times = Some(ts.last_frame.clone());
325        }
326
327        self.current_frame = (self.current_frame + 1) % self.frames_in_flight.max(1);
328        Ok(())
329    }
330
331    pub fn mark_submitted(&mut self) {
332        if let Some(ts) = &mut self.timestamps {
333            ts.mark_submitted(self.current_frame);
334        }
335    }
336
337    pub fn resize_output(&mut self, new_output: (u32, u32), gpu: &GpuAssetServer) -> anyhow::Result<()> {
338        self.output_resolution = new_output;
339        self.pool.resize_output(self.internal_resolution, new_output)?;
340        let affected: Vec<ResourceHandle> = self.pool.output_handles().collect();
341        self.bindings.flush(&self.pool, &affected, gpu, FlushReason::Resize);
342        self.tracker.invalidate(&affected);
343        self.build_compiled_passes();
344        Ok(())
345    }
346
347    pub fn resize_internal(&mut self, new_internal: (u32, u32), gpu: &GpuAssetServer) -> anyhow::Result<()> {
348        self.internal_resolution = new_internal;
349        self.pool.resize_internal(new_internal, self.output_resolution)?;
350        let affected: Vec<ResourceHandle> = self.pool.internal_handles().collect();
351        self.bindings.flush(&self.pool, &affected, gpu, FlushReason::Resize);
352        self.tracker.invalidate(&affected);
353        Ok(())
354    }
355
356    pub fn internal_resolution(&self) -> (u32, u32) {
357        self.internal_resolution
358    }
359    pub fn output_resolution(&self) -> (u32, u32) {
360        self.output_resolution
361    }
362
363    pub fn pass_mut(&mut self, handle: PassHandle) -> &mut PassNode {
364        &mut self.nodes[handle.0 as usize]
365    }
366}
367
368fn build_resource_edges(nodes: &[PassNode], adj: &mut [HashSet<usize>], in_degree: &mut [usize]) {
369    let mut last_writer: HashMap<ResourceHandle, usize> = HashMap::new();
370    let mut last_readers: HashMap<ResourceHandle, Vec<usize>> = HashMap::new();
371
372    for (i, node) in nodes.iter().enumerate() {
373        for access in &node.accesses {
374            match access.access {
375                AccessType::Read => {
376                    add_edge(adj, in_degree, last_writer.get(&access.handle).copied(), i);
377                    last_readers.entry(access.handle).or_default().push(i);
378                }
379                AccessType::Write | AccessType::ReadWrite => {
380                    add_edge(adj, in_degree, last_writer.get(&access.handle).copied(), i);
381                    add_edges_from_readers(adj, in_degree, &last_readers, access.handle, i);
382                    last_writer.insert(access.handle, i);
383                    last_readers.remove(&access.handle);
384                }
385            }
386        }
387    }
388}
389
390fn build_explicit_edges(nodes: &[PassNode], adj: &mut [HashSet<usize>], in_degree: &mut [usize]) {
391    for (i, node) in nodes.iter().enumerate() {
392        for &dep_handle in &node.depends_on {
393            add_edge(adj, in_degree, Some(dep_handle.0 as usize), i);
394        }
395    }
396}
397
398fn add_edge(adj: &mut [HashSet<usize>], in_degree: &mut [usize], from: Option<usize>, to: usize) {
399    let Some(from) = from else { return };
400    if from != to && !adj[from].contains(&to) {
401        adj[from].insert(to);
402        in_degree[to] += 1;
403    }
404}
405
406fn add_edges_from_readers(
407    adj: &mut [HashSet<usize>],
408    in_degree: &mut [usize],
409    last_readers: &HashMap<ResourceHandle, Vec<usize>>,
410    handle: ResourceHandle,
411    to: usize,
412) {
413    let Some(readers) = last_readers.get(&handle) else {
414        return;
415    };
416    for &reader in readers {
417        add_edge(adj, in_degree, Some(reader), to);
418    }
419}
420
421fn topological_sort(n: usize, adj: Vec<HashSet<usize>>, mut in_degree: Vec<usize>) -> anyhow::Result<Vec<usize>> {
422    let mut queue: VecDeque<usize> = (0..n).filter(|&i| in_degree[i] == 0).collect();
423    let mut order = Vec::with_capacity(n);
424
425    while let Some(idx) = queue.pop_front() {
426        order.push(idx);
427        for &dep in &adj[idx] {
428            in_degree[dep] -= 1;
429            if in_degree[dep] == 0 {
430                queue.push_back(dep);
431            }
432        }
433    }
434
435    if order.len() != n {
436        anyhow::bail!("RenderGraph: обнаружен цикл в графе пассов");
437    }
438    Ok(order)
439}
440
441pub struct PassBuilder {
442    name: String,
443    accesses: Vec<PassAccess>,
444    deferred_bindings: Vec<DescriptorBinding>,
445    explicit_deps: Vec<PassHandle>,
446}
447
448impl PassBuilder {
449    pub fn new(name: impl Into<String>) -> Self {
450        Self { name: name.into(), accesses: Vec::new(), deferred_bindings: Vec::new(), explicit_deps: Vec::new() }
451    }
452
453    pub fn read(mut self, handle: ResourceHandle, layout: ImageLayout) -> Self {
454        self.accesses.push(PassAccess::read(handle, layout.to_vk()));
455        self
456    }
457
458    pub fn write(mut self, handle: ResourceHandle, layout: ImageLayout) -> Self {
459        self.accesses.push(PassAccess::write(handle, layout.to_vk()));
460        self
461    }
462
463    pub fn read_write(mut self, handle: ResourceHandle, layout: ImageLayout) -> Self {
464        self.accesses.push(PassAccess::read_write(handle, layout.to_vk()));
465        self
466    }
467
468    pub fn bind_resource(mut self, binding: DescriptorBinding) -> Self {
469        self.deferred_bindings.push(binding);
470        self
471    }
472
473    pub fn bind_sampled(
474        mut self,
475        resource: ResourceHandle,
476        set: DescriptorSetId,
477        binding: u32,
478        sampler: SamplerId,
479    ) -> Self {
480        self.deferred_bindings.push(DescriptorBinding {
481            resource,
482            set,
483            binding,
484            array_element: 0,
485            image_type: DescriptorImageType::CombinedImageSampler(sampler),
486            image_layout: vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL,
487        });
488        self
489    }
490
491    pub fn bind_sampled_at(
492        mut self,
493        resource: ResourceHandle,
494        set: DescriptorSetId,
495        binding: u32,
496        array_element: u32,
497        sampler: SamplerId,
498        image_layout: vk::ImageLayout,
499    ) -> Self {
500        self.deferred_bindings.push(DescriptorBinding {
501            resource,
502            set,
503            binding,
504            array_element,
505            image_type: DescriptorImageType::CombinedImageSampler(sampler),
506            image_layout,
507        });
508        self
509    }
510
511    pub fn after(mut self, handle: PassHandle) -> Self {
512        self.explicit_deps.push(handle);
513        self
514    }
515
516    pub fn record<F>(self, f: F) -> PassNodeReady
517    where
518        F: FnMut(&mut CommandEncoder<'_>, &RWorld, &GpuAssetServer) -> anyhow::Result<()> + Send + 'static,
519    {
520        PassNodeReady {
521            node: PassNode { depends_on: self.explicit_deps, ..PassNode::new(self.name, self.accesses, Box::new(f)) },
522            deferred_bindings: self.deferred_bindings,
523        }
524    }
525}
526
527pub struct PassNodeReady {
528    node: PassNode,
529    deferred_bindings: Vec<DescriptorBinding>,
530}
531
532impl PassNodeReady {
533    pub fn build(self, graph: &mut RenderGraph, gpu: &GpuAssetServer) -> PassHandle {
534        for b in self.deferred_bindings {
535            let resource = b.resource;
536            graph.bindings.register(b);
537            if graph.allocated {
538                graph.bindings.flush(&graph.pool, &[resource], gpu, FlushReason::LateBinding);
539            }
540        }
541        graph.add_pass(self.node)
542    }
543}
544
545pub fn pass(name: impl Into<String>) -> PassBuilder {
546    PassBuilder::new(name)
547}