1use dehazing::{model::DehazeNet, Module as _, Tensor, VarBuilder};
2
3fn main() {
4 let device = candle_core::Device::cuda_if_available(0).unwrap();
5
6 let base_dir = env!("CARGO_MANIFEST_DIR");
7 let weight_path = format!("{base_dir}/dehazer.safetensors");
8 let vb = unsafe {
9 VarBuilder::from_mmaped_safetensors(
10 &[&weight_path],
11 candle_core::DType::F32,
12 &device,
13 )
14 .unwrap()
15 };
16
17 let model = DehazeNet::new(vb).unwrap();
18
19 let img = image::open(format!("{base_dir}/testdata/test2.png")).unwrap();
22
23 let raw = img.to_rgb8().into_vec();
24 let data = Tensor::from_vec(
25 raw,
26 (img.height() as usize, img.width() as usize, 3),
27 &device,
28 )
29 .unwrap()
30 .to_dtype(candle_core::DType::F32)
31 .unwrap()
32 .broadcast_div(&Tensor::new(255f32, &device).unwrap())
33 .unwrap()
34 .permute((2, 0, 1))
35 .unwrap()
36 .unsqueeze(0)
37 .unwrap();
38
39 println!("{data:?}");
40
41 let out = model.forward(&data).unwrap();
42
43 let out = out.squeeze(0).unwrap(); let (_, height, width) = out.dims3().unwrap();
47
48
49 let image_data: Vec<u8> = out
50 .permute((1, 2, 0))
51 .unwrap() .flatten_all()
53 .unwrap()
54 .to_vec1::<f32>()
55 .unwrap()
56 .iter()
57 .map(|&v| (v.clamp(0.0, 1.0) * 255.0) as u8)
58 .collect();
59
60 let img_out =
62 image::RgbImage::from_raw(width as u32, height as u32, image_data).expect("创建图像失败");
63
64 img_out.save("dehazed_output.jpg").expect("保存图像失败");
65 println!("去雾结果已保存为 dehazed_output.jpg");
66}