trueno_graph/gpu/
buffer.rs1use super::GpuDevice;
7use crate::storage::CsrGraph;
8use anyhow::Result;
9
10#[derive(Debug)]
18pub struct GpuCsrBuffers {
19 pub num_nodes: usize,
21
22 pub num_edges: usize,
24
25 pub row_offsets: wgpu::Buffer,
27
28 pub col_indices: wgpu::Buffer,
30
31 pub edge_weights: Option<wgpu::Buffer>,
33}
34
35impl GpuCsrBuffers {
36 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 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 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 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 #[must_use]
78 pub const fn num_nodes(&self) -> usize {
79 self.num_nodes
80 }
81
82 #[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 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 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 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()); }
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 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 let mut graph = CsrGraph::new();
184 for i in 1..10 {
186 graph.add_edge(NodeId(0), NodeId(i), i as f32).unwrap();
187 }
188 for i in 10..15 {
190 graph.add_edge(NodeId(1), NodeId(i), i as f32).unwrap();
191 }
192 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); assert!(buffers.edge_weights.is_some());
200 }
201}