1use baracuda_cusolver::{geqrf, DnHandle};
16use baracuda_driver::{Context, Device, DeviceBuffer};
17
18fn main() -> Result<(), Box<dyn std::error::Error>> {
19 baracuda_driver::init()?;
20 let device = Device::get(0)?;
21 println!("device: {}", device.name()?);
22 let ctx = Context::new(&device)?;
23 let handle = DnHandle::new()?;
24
25 let m: i32 = 3;
31 let n: i32 = 3;
32 let lda: i32 = m;
33 let a_host: [f32; 9] = [12.0, 6.0, -4.0, -51.0, 167.0, 24.0, 4.0, -68.0, -41.0];
34
35 let mut a = DeviceBuffer::from_slice(&ctx, &a_host)?;
36 let mut tau: DeviceBuffer<f32> = DeviceBuffer::new(&ctx, m.min(n) as usize)?;
37 let mut info: DeviceBuffer<i32> = DeviceBuffer::new(&ctx, 1)?;
38
39 geqrf::<f32>(&handle, m, n, &mut a, lda, &mut tau, &mut info)?;
40
41 let mut info_host = [42i32];
42 info.copy_to_host(&mut info_host)?;
43 assert_eq!(info_host[0], 0, "geqrf info = {}", info_host[0]);
44
45 let mut a_out = [0.0f32; 9];
46 a.copy_to_host(&mut a_out)?;
47 let mut tau_out = vec![0.0f32; tau.len()];
48 tau.copy_to_host(&mut tau_out)?;
49
50 println!("R (upper triangle, in column-major positions):");
52 for i in 0..m as usize {
53 for j in 0..n as usize {
54 if i <= j {
55 print!("{:10.4} ", a_out[j * lda as usize + i]);
56 } else {
57 print!("{:>10} ", "·");
58 }
59 }
60 println!();
61 }
62 println!("tau = {tau_out:?}");
63
64 let r00 = a_out[0].abs();
66 assert!((r00 - 14.0).abs() < 1e-3, "R[0,0]={r00}, expected 14");
67 println!("OK (|R[0,0]| = {r00})");
68 Ok(())
69}