Skip to main content

trueno_graph/gpu/
buffer.rs

1//! GPU buffer management for CSR graph data
2//!
3//! Handles uploading CSR (`row_offsets`, `col_indices`, `edge_weights`) to GPU
4//! and downloading results (distances, visited arrays) from GPU.
5
6use super::GpuDevice;
7use crate::storage::CsrGraph;
8use anyhow::Result;
9
10/// GPU buffers for CSR graph representation
11///
12/// Manages GPU-side storage of:
13/// - Row offsets (CSR row pointers)
14/// - Column indices (CSR column indices)
15/// - Edge weights (optional)
16/// - Auxiliary arrays (visited, distances, etc.)
17#[derive(Debug)]
18pub struct GpuCsrBuffers {
19    /// Number of nodes in the graph
20    pub num_nodes: usize,
21
22    /// Number of edges in the graph
23    pub num_edges: usize,
24
25    /// GPU buffer for `row_offsets` (size: `num_nodes` + 1)
26    pub row_offsets: wgpu::Buffer,
27
28    /// GPU buffer for `col_indices` (size: `num_edges`)
29    pub col_indices: wgpu::Buffer,
30
31    /// GPU buffer for `edge_weights` (size: `num_edges`, optional)
32    pub edge_weights: Option<wgpu::Buffer>,
33}
34
35impl GpuCsrBuffers {
36    /// Upload CSR graph to GPU
37    ///
38    /// # Errors
39    ///
40    /// Returns error if buffer creation fails
41    pub fn from_csr_graph(device: &GpuDevice, graph: &CsrGraph) -> Result<Self> {
42        let num_nodes = graph.num_nodes();
43        let num_edges = graph.num_edges();
44
45        // Create row_offsets buffer
46        let row_offsets_data = graph.row_offsets_slice();
47        let row_offsets = device.create_buffer_init(
48            "CSR row_offsets",
49            bytemuck::cast_slice(row_offsets_data),
50            wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
51        )?;
52
53        // Create col_indices buffer
54        let col_indices_data = graph.col_indices_slice();
55        let col_indices = device.create_buffer_init(
56            "CSR col_indices",
57            bytemuck::cast_slice(col_indices_data),
58            wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
59        )?;
60
61        // Create edge_weights buffer (optional for unweighted graphs)
62        let edge_weights_data = graph.edge_weights_slice();
63        let edge_weights = if edge_weights_data.is_empty() {
64            None
65        } else {
66            Some(device.create_buffer_init(
67                "CSR edge_weights",
68                bytemuck::cast_slice(edge_weights_data),
69                wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
70            )?)
71        };
72
73        Ok(Self { num_nodes, num_edges, row_offsets, col_indices, edge_weights })
74    }
75
76    /// Get number of nodes
77    #[must_use]
78    pub const fn num_nodes(&self) -> usize {
79        self.num_nodes
80    }
81
82    /// Get number of edges
83    #[must_use]
84    pub const fn num_edges(&self) -> usize {
85        self.num_edges
86    }
87}
88
89#[cfg(test)]
90mod tests {
91    use super::*;
92    use crate::NodeId;
93
94    #[tokio::test]
95    async fn test_upload_csr_to_gpu() {
96        if !GpuDevice::is_gpu_available().await {
97            eprintln!("⚠️  Skipping test_upload_csr_to_gpu: GPU not available");
98            return;
99        }
100
101        let device = GpuDevice::new().await.unwrap();
102
103        // Create simple graph: 0 -> 1 -> 2
104        let mut graph = CsrGraph::new();
105        graph.add_edge(NodeId(0), NodeId(1), 1.0).unwrap();
106        graph.add_edge(NodeId(1), NodeId(2), 1.0).unwrap();
107
108        // Upload to GPU
109        let buffers = GpuCsrBuffers::from_csr_graph(&device, &graph).unwrap();
110
111        assert_eq!(buffers.num_nodes(), 3);
112        assert_eq!(buffers.num_edges(), 2);
113    }
114
115    #[tokio::test]
116    async fn test_upload_empty_graph() {
117        if !GpuDevice::is_gpu_available().await {
118            eprintln!("⚠️  Skipping test_upload_empty_graph: GPU not available");
119            return;
120        }
121
122        let device = GpuDevice::new().await.unwrap();
123        let graph = CsrGraph::new();
124
125        let buffers = GpuCsrBuffers::from_csr_graph(&device, &graph).unwrap();
126
127        assert_eq!(buffers.num_nodes(), 0);
128        assert_eq!(buffers.num_edges(), 0);
129    }
130
131    #[tokio::test]
132    async fn test_upload_weighted_graph() {
133        if !GpuDevice::is_gpu_available().await {
134            eprintln!("⚠️  Skipping test_upload_weighted_graph: GPU not available");
135            return;
136        }
137
138        let device = GpuDevice::new().await.unwrap();
139
140        // Create weighted graph
141        let mut graph = CsrGraph::new();
142        graph.add_edge(NodeId(0), NodeId(1), 2.5).unwrap();
143        graph.add_edge(NodeId(1), NodeId(2), 3.7).unwrap();
144
145        let buffers = GpuCsrBuffers::from_csr_graph(&device, &graph).unwrap();
146
147        assert_eq!(buffers.num_nodes(), 3);
148        assert_eq!(buffers.num_edges(), 2);
149        assert!(buffers.edge_weights.is_some()); // Weighted graph
150    }
151
152    #[tokio::test]
153    async fn test_upload_large_graph() {
154        if !GpuDevice::is_gpu_available().await {
155            eprintln!("⚠️  Skipping test_upload_large_graph: GPU not available");
156            return;
157        }
158
159        let device = GpuDevice::new().await.unwrap();
160
161        // Create larger graph
162        let mut graph = CsrGraph::new();
163        for i in 0..100 {
164            graph.add_edge(NodeId(i), NodeId((i + 1) % 100), 1.0).unwrap();
165        }
166
167        let buffers = GpuCsrBuffers::from_csr_graph(&device, &graph).unwrap();
168
169        assert_eq!(buffers.num_nodes(), 100);
170        assert_eq!(buffers.num_edges(), 100);
171    }
172
173    #[tokio::test]
174    async fn test_buffer_with_complex_graph() {
175        if !GpuDevice::is_gpu_available().await {
176            eprintln!("⚠️  Skipping test_buffer_with_complex_graph: GPU not available");
177            return;
178        }
179
180        let device = GpuDevice::new().await.unwrap();
181
182        // Create graph with varying degrees
183        let mut graph = CsrGraph::new();
184        // Node 0: high degree
185        for i in 1..10 {
186            graph.add_edge(NodeId(0), NodeId(i), i as f32).unwrap();
187        }
188        // Node 1: medium degree
189        for i in 10..15 {
190            graph.add_edge(NodeId(1), NodeId(i), i as f32).unwrap();
191        }
192        // Node 2: low degree
193        graph.add_edge(NodeId(2), NodeId(15), 15.0).unwrap();
194
195        let buffers = GpuCsrBuffers::from_csr_graph(&device, &graph).unwrap();
196
197        assert!(buffers.num_nodes() >= 16);
198        assert_eq!(buffers.num_edges(), 15); // 9 + 5 + 1
199        assert!(buffers.edge_weights.is_some());
200    }
201}