Files
rustytorch/demos/neural-operator-shared/src/config.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
Consistent formatting pass: line wrapping, import sorting, trailing
whitespace removal, let-chain indentation, merged derive attributes,
and unsafe block reformatting.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-04-12 07:01:58 -07:00

219 lines
5.9 KiB
Rust

//! PDE configuration types for neural operator demo
//!
//! Defines the supported PDE types and their configurations.
use serde::{Deserialize, Serialize};
/// Supported PDE types for the neural operator
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[derive(Default)]
pub enum PDEType {
/// Darcy flow in porous media: -∇·(a(x)∇u) = f
/// Input: permeability field a(x)
/// Output: pressure field u(x)
#[default]
DarcyFlow,
/// Steady-state heat equation: -∇²u = f
/// Input: source term f(x) or boundary conditions
/// Output: temperature field u(x)
HeatEquation,
/// Poisson equation: -∇²u = f
/// Input: charge density f(x)
/// Output: potential field u(x)
Poisson,
/// 2D Navier-Stokes (simplified): steady-state velocity field
/// Input: initial vorticity
/// Output: velocity magnitude
NavierStokes,
}
impl PDEType {
/// Returns a human-readable name for the PDE type
#[must_use]
pub const fn name(&self) -> &'static str {
match self {
Self::DarcyFlow => "Darcy Flow",
Self::HeatEquation => "Heat Equation",
Self::Poisson => "Poisson Equation",
Self::NavierStokes => "Navier-Stokes",
}
}
/// Returns a description of the PDE
#[must_use]
pub const fn description(&self) -> &'static str {
match self {
Self::DarcyFlow => "Flow through porous media. Input: permeability field.",
Self::HeatEquation => "Steady-state heat conduction. Input: heat sources.",
Self::Poisson => "Electrostatics potential. Input: charge density.",
Self::NavierStokes => "Fluid flow velocity. Input: initial vorticity.",
}
}
/// Returns the default number of modes for this PDE type
#[must_use]
pub const fn default_modes(&self) -> (u32, u32) {
match self {
Self::DarcyFlow | Self::HeatEquation | Self::Poisson => (12, 12),
Self::NavierStokes => (16, 16),
}
}
/// Returns the model weight filename for this PDE type
#[must_use]
pub const fn weight_filename(&self) -> &'static str {
match self {
Self::DarcyFlow => "fno_darcy.safetensors",
Self::HeatEquation => "fno_heat.safetensors",
Self::Poisson => "fno_poisson.safetensors",
Self::NavierStokes => "fno_navier_stokes.safetensors",
}
}
}
/// Configuration for PDE solving
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PDEConfig {
/// Type of PDE to solve
pub pde_type: PDEType,
/// Grid resolution (assumes square grid)
pub resolution: u32,
/// Number of Fourier modes (height, width)
pub n_modes: (u32, u32),
/// Model width (hidden dimension)
pub model_width: u32,
/// Number of Fourier layers
pub n_layers: u32,
/// Whether to run FEM baseline for comparison
pub run_fem_baseline: bool,
}
impl PDEConfig {
/// Creates a new PDE configuration
#[must_use]
pub fn new(pde_type: PDEType, resolution: u32) -> Self {
let n_modes = pde_type.default_modes();
Self {
pde_type,
resolution,
n_modes,
model_width: 32,
n_layers: 4,
run_fem_baseline: false,
}
}
/// Creates a Darcy flow configuration
#[must_use]
pub fn darcy(resolution: u32) -> Self {
Self::new(PDEType::DarcyFlow, resolution)
}
/// Creates a heat equation configuration
#[must_use]
pub fn heat(resolution: u32) -> Self {
Self::new(PDEType::HeatEquation, resolution)
}
/// Creates a Poisson equation configuration
#[must_use]
pub fn poisson(resolution: u32) -> Self {
Self::new(PDEType::Poisson, resolution)
}
/// Creates a Navier-Stokes configuration
#[must_use]
pub fn navier_stokes(resolution: u32) -> Self {
Self::new(PDEType::NavierStokes, resolution)
}
/// Sets the number of Fourier modes
#[must_use]
pub const fn with_modes(mut self, modes_h: u32, modes_w: u32) -> Self {
self.n_modes = (modes_h, modes_w);
self
}
/// Sets the model width
#[must_use]
pub const fn with_width(mut self, width: u32) -> Self {
self.model_width = width;
self
}
/// Sets the number of layers
#[must_use]
pub const fn with_layers(mut self, layers: u32) -> Self {
self.n_layers = layers;
self
}
/// Enables FEM baseline comparison
#[must_use]
pub const fn with_fem_baseline(mut self) -> Self {
self.run_fem_baseline = true;
self
}
/// Returns total grid points
#[must_use]
pub const fn total_points(&self) -> u32 {
self.resolution * self.resolution
}
}
impl Default for PDEConfig {
fn default() -> Self {
Self::darcy(64)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_pde_type_defaults() {
assert_eq!(PDEType::default(), PDEType::DarcyFlow);
}
#[test]
fn test_pde_config_builder() {
let config = PDEConfig::darcy(128)
.with_modes(16, 16)
.with_width(64)
.with_layers(6)
.with_fem_baseline();
assert_eq!(config.pde_type, PDEType::DarcyFlow);
assert_eq!(config.resolution, 128);
assert_eq!(config.n_modes, (16, 16));
assert_eq!(config.model_width, 64);
assert_eq!(config.n_layers, 6);
assert!(config.run_fem_baseline);
}
#[test]
fn test_total_points() {
let config = PDEConfig::darcy(64);
assert_eq!(config.total_points(), 4096);
}
#[test]
fn test_serialization() {
let config = PDEConfig::default();
let json = serde_json::to_string(&config).unwrap();
let deserialized: PDEConfig = serde_json::from_str(&json).unwrap();
assert_eq!(config, deserialized);
}
}