Skip to main content

custom/
custom.rs

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    // println!("{model:?}");
20
21    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    // 处理输出张量
44    let out = out.squeeze(0).unwrap(); // 移除批次维度 [c, h, w]
45
46    let (_, height, width) = out.dims3().unwrap();
47
48
49    let image_data: Vec<u8> = out
50        .permute((1, 2, 0))
51        .unwrap() // [H, W, C] 符合图像布局
52        .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    // 保存图像
61    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}