//! Fast inference engine for temperature field visualization //! //! Provides optimized methods for querying the trained PINN //! at many points for real-time visualization. use crate::network::{Coordinate4D, ThermalPinn}; use bioheat_shared::{BoundingBox3D, Point3D, SliceAxis, SliceData, TemperatureField}; use serde::{Deserialize, Serialize}; /// Inference engine for fast temperature queries #[derive(Debug, Clone)] pub struct InferenceEngine { /// The trained network pinn: ThermalPinn, /// Domain bounds bounds: BoundingBox3D, } impl InferenceEngine { /// Create a new inference engine from a trained PINN #[must_use] pub fn new(pinn: ThermalPinn, bounds: BoundingBox3D) -> Self { Self { pinn, bounds } } /// Query temperature at a single point #[must_use] pub fn query_point(&self, point: &Point3D, t: f32) -> f32 { let coord = Coordinate4D::from_point_and_time(*point, t); self.pinn.forward(&coord) as f32 } /// Query temperature at multiple points (same time) #[must_use] pub fn query_points(&self, points: &[Point3D], t: f32) -> Vec { points.iter().map(|p| self.query_point(p, t)).collect() } /// Generate a 3D temperature field on a regular grid #[must_use] pub fn generate_field(&self, resolution: (usize, usize, usize), t: f32) -> TemperatureField { let (nx, ny, nz) = resolution; let size = self.bounds.size(); let dx = size.x / (nx - 1).max(1) as f32; let dy = size.y / (ny - 1).max(1) as f32; let dz = size.z / (nz - 1).max(1) as f32; let mut values = Vec::with_capacity(nx * ny * nz); for k in 0..nz { for j in 0..ny { for i in 0..nx { let point = Point3D::new( self.bounds.min.x + i as f32 * dx, self.bounds.min.y + j as f32 * dy, self.bounds.min.z + k as f32 * dz, ); values.push(self.query_point(&point, t)); } } } TemperatureField { resolution, values, bounds: self.bounds, } } /// Generate a 2D slice at the specified position #[must_use] pub fn generate_slice( &self, axis: SliceAxis, position: f32, resolution: (usize, usize), t: f32, ) -> SliceData { let (nu, nv) = resolution; let size = self.bounds.size(); // Determine the slice bounds based on axis let (u_range, v_range, pos_clamped) = match axis { SliceAxis::X => { let pos = position.clamp(self.bounds.min.x, self.bounds.max.x); ( (self.bounds.min.y, self.bounds.max.y), (self.bounds.min.z, self.bounds.max.z), pos, ) } SliceAxis::Y => { let pos = position.clamp(self.bounds.min.y, self.bounds.max.y); ( (self.bounds.min.x, self.bounds.max.x), (self.bounds.min.z, self.bounds.max.z), pos, ) } SliceAxis::Z => { let pos = position.clamp(self.bounds.min.z, self.bounds.max.z); ( (self.bounds.min.x, self.bounds.max.x), (self.bounds.min.y, self.bounds.max.y), pos, ) } }; let du = (u_range.1 - u_range.0) / (nu - 1).max(1) as f32; let dv = (v_range.1 - v_range.0) / (nv - 1).max(1) as f32; let mut values = Vec::with_capacity(nu * nv); for iv in 0..nv { for iu in 0..nu { let u = u_range.0 + iu as f32 * du; let v = v_range.0 + iv as f32 * dv; let point = match axis { SliceAxis::X => Point3D::new(pos_clamped, u, v), SliceAxis::Y => Point3D::new(u, pos_clamped, v), SliceAxis::Z => Point3D::new(u, v, pos_clamped), }; values.push(self.query_point(&point, t)); } } // Calculate index from position let index = match axis { SliceAxis::X => ((pos_clamped - self.bounds.min.x) / size.x * 100.0) as usize, SliceAxis::Y => ((pos_clamped - self.bounds.min.y) / size.y * 100.0) as usize, SliceAxis::Z => ((pos_clamped - self.bounds.min.z) / size.z * 100.0) as usize, }; SliceData { axis, position: pos_clamped, index, resolution, values, bounds_2d: (u_range.0, u_range.1, v_range.0, v_range.1), } } /// Generate all three orthogonal slices at a point #[must_use] pub fn generate_orthogonal_slices( &self, center: &Point3D, resolution: (usize, usize), t: f32, ) -> (SliceData, SliceData, SliceData) { let x_slice = self.generate_slice(SliceAxis::X, center.x, resolution, t); let y_slice = self.generate_slice(SliceAxis::Y, center.y, resolution, t); let z_slice = self.generate_slice(SliceAxis::Z, center.z, resolution, t); (x_slice, y_slice, z_slice) } /// Find the location of maximum temperature #[must_use] pub fn find_max_temperature( &self, resolution: (usize, usize, usize), t: f32, ) -> (Point3D, f32) { let field = self.generate_field(resolution, t); let (nx, ny, _nz) = resolution; let mut max_temp = f32::NEG_INFINITY; let mut max_idx = 0; for (idx, &temp) in field.values.iter().enumerate() { if temp > max_temp { max_temp = temp; max_idx = idx; } } // Convert index back to (i, j, k) let i = max_idx % nx; let j = (max_idx / nx) % ny; let k = max_idx / (nx * ny); let point = field.coords_at(i, j, k); (point, max_temp) } /// Query temperature along a line (for 1D profile plots) #[must_use] pub fn query_line( &self, start: &Point3D, end: &Point3D, num_points: usize, t: f32, ) -> Vec<(f32, f32)> { let mut results = Vec::with_capacity(num_points); for i in 0..num_points { let frac = i as f32 / (num_points - 1).max(1) as f32; let point = Point3D::new( start.x + frac * (end.x - start.x), start.y + frac * (end.y - start.y), start.z + frac * (end.z - start.z), ); let distance = start.distance(&point); let temp = self.query_point(&point, t); results.push((distance, temp)); } results } /// Get the bounds #[must_use] pub fn bounds(&self) -> &BoundingBox3D { &self.bounds } } /// Performance metrics for inference #[derive(Debug, Clone, Serialize, Deserialize)] pub struct InferenceMetrics { /// Number of points queried pub num_points: usize, /// Total time in milliseconds pub total_time_ms: f32, /// Time per point in microseconds pub time_per_point_us: f32, /// Throughput in points per second pub throughput: f32, } impl InferenceMetrics { /// Create from measurement #[must_use] pub fn from_measurement(num_points: usize, elapsed_ms: f32) -> Self { let time_per_point_us = if num_points > 0 { (elapsed_ms * 1000.0) / num_points as f32 } else { 0.0 }; let throughput = if elapsed_ms > 0.0 { num_points as f32 / (elapsed_ms / 1000.0) } else { 0.0 }; Self { num_points, total_time_ms: elapsed_ms, time_per_point_us, throughput, } } } #[cfg(test)] mod tests { use super::*; use crate::config::NetworkConfig; fn create_test_engine() -> InferenceEngine { let pinn = ThermalPinn::new(NetworkConfig { fourier_features: 16, hidden_layers: vec![32, 32], ..Default::default() }); let bounds = BoundingBox3D::from_dimensions(0.1, 0.1, 0.1); InferenceEngine::new(pinn, bounds) } #[test] fn test_query_point() { let engine = create_test_engine(); let temp = engine.query_point(&Point3D::origin(), 0.0); assert!(temp.is_finite()); } #[test] fn test_generate_field() { let engine = create_test_engine(); let field = engine.generate_field((8, 8, 8), 100.0); assert_eq!(field.len(), 512); } #[test] fn test_generate_slice() { let engine = create_test_engine(); let slice = engine.generate_slice(SliceAxis::Z, 0.0, (32, 32), 50.0); assert_eq!(slice.resolution, (32, 32)); assert_eq!(slice.values.len(), 1024); } #[test] fn test_orthogonal_slices() { let engine = create_test_engine(); let (x_slice, y_slice, z_slice) = engine.generate_orthogonal_slices(&Point3D::origin(), (16, 16), 0.0); assert_eq!(x_slice.axis, SliceAxis::X); assert_eq!(y_slice.axis, SliceAxis::Y); assert_eq!(z_slice.axis, SliceAxis::Z); } #[test] fn test_query_line() { let engine = create_test_engine(); let start = Point3D::new(-0.05, 0.0, 0.0); let end = Point3D::new(0.05, 0.0, 0.0); let profile = engine.query_line(&start, &end, 10, 0.0); assert_eq!(profile.len(), 10); // First point should be at distance 0 assert!(profile[0].0.abs() < 1e-6); } #[test] fn test_find_max() { let engine = create_test_engine(); let (point, temp) = engine.find_max_temperature((8, 8, 8), 0.0); assert!(engine.bounds().contains(&point)); assert!(temp.is_finite()); } #[test] fn test_metrics() { let metrics = InferenceMetrics::from_measurement(1000, 100.0); assert_eq!(metrics.num_points, 1000); assert!((metrics.throughput - 10000.0).abs() < 100.0); } }