use std::path::PathBuf;
use gdal::raster::Buffer;
use gdal::Dataset;
use gdal::DriverManager;
use gdal::raster::RasterBand;
use gdal::GeoTransform;
use gdal::GeoTransformEx;
use nalgebra::DVector;
use std::rc::Rc;
pub fn read_raster_coords(path: &PathBuf) -> gdal::errors::Result<Vec<(f64, f64)>> {
let dataset = Dataset::open(path)?;
let geo_transform = dataset.geo_transform()?;
let (width, height) = (dataset.raster_size().0, dataset.raster_size().1);
let mut coords = Vec::with_capacity(width * height);
for y in 0..height {
for x in 0..width {
let px = geo_transform[0] + (x as f64 + 0.5) * geo_transform[1];
let py = geo_transform[3] + (y as f64 + 0.5) * geo_transform[5];
coords.push((px, py));
}
}
Ok(coords)
}
pub fn export_predictions_to_raster(
template_tif_path: &PathBuf,
output_tif_path: &PathBuf,
predictions: &[f64],
) -> gdal::errors::Result<()> {
let template = Dataset::open(template_tif_path)?;
let (size_x, size_y) = (template.raster_size().0, template.raster_size().1);
assert_eq!(
predictions.len(),
(size_x * size_y) as usize,
"Prediction vector length doesn't match raster size"
);
let driver = DriverManager::get_driver_by_name("GTiff")?;
let mut output = driver.create_with_band_type::<f64, &PathBuf>(
output_tif_path,
size_x,
size_y,
1
)?;
output.set_projection(&template.projection())?;
output.set_geo_transform(&template.geo_transform()?)?;
let mut buffer = Buffer::<f64>::new((size_x, size_y), predictions.to_vec());
let mut band: RasterBand = output.rasterband(1)?; band.write((0, 0), (size_x, size_y), &mut buffer)?;
band.set_no_data_value(Some(f64::NAN))?;
Ok(())
}
pub struct RasterDrift {
_datasets: Vec<Rc<Dataset>>,
bands: Vec<RasterBand<'static>>,
geo_transform: GeoTransform,
inv_transform: GeoTransform,
coefficients: Option<DVector<f64>>
}
impl RasterDrift {
pub fn new(raster_paths: &[&str]) -> Self {
assert!(!raster_paths.is_empty(), "At least one raster is required");
let mut datasets = Vec::new();
let mut bands = Vec::new();
for path in raster_paths {
let ds = Rc::new(Dataset::open(path).expect("Failed to open raster"));
let leaked_ds: &'static Dataset = Box::leak(Box::new(Rc::clone(&ds)));
let static_band = leaked_ds.rasterband(1).unwrap();
datasets.push(ds);
bands.push(static_band);
}
let geo_transform = datasets[0].geo_transform().expect("Missing geo transform");
let inv_transform = geo_transform.invert().expect("Failed to invert transform");
Self {
_datasets: datasets, bands,
geo_transform,
inv_transform,
coefficients: None
}
}
pub fn get_coefficients(&self) -> Option<DVector<f64>> {
self.coefficients.clone()
}
pub fn set_coefficients(&mut self, coeffs: DVector<f64>) {
self.coefficients = Some(coeffs);
}
pub fn predict(&self, point: (f64, f64)) -> f64 {
let drift_vec = self.drift_functions(point);
let coeffs = self.coefficients.as_ref()
.expect("Drift coefficients not set");
drift_vec.iter().zip(coeffs.iter()).map(|(d, c)| d * c).sum()
}
pub fn drift_functions(&self, point: (f64, f64)) -> Vec<f64> {
let (x , y) = &self.inv_transform.apply(point.0, point.1);
let px = *x as isize;
let py = *y as isize;
if px < 0 || py < 0 || px >= self.bands[0].x_size() as isize || py >= self.bands[0].y_size() as isize {
return vec![0.0, 0.0, 0.0];
}
let mut drift = vec![1.0];
for band in &self.bands {
let value = band
.read_as::<f64>((px, py), (1, 1), (1, 1), None)
.expect("Failed to read raster value")
.into_iter()
.next()
.expect("No value found for the pixel");
drift.push(value);
}
drift
}
pub fn drift_dimension(&self) -> usize {
self.bands.len() + 1
}
}