Skip to content

Core Concepts

Overview

Opifex provides a unified framework for scientific machine learning, combining traditional numerical methods with modern deep learning approaches. Built on JAX and FLAX NNX, it offers high-performance, differentiable computing for scientific applications with full physics-informed capabilities.

Framework Architecture

JAX Ecosystem Foundation

Opifex is built entirely on the JAX ecosystem for maximum performance and scientific computing capabilities:

import jax
import jax.numpy as jnp
import flax.nnx as nnx

# Configure JAX for scientific computing
jax.config.update("jax_enable_x64", True)
print(f"Available devices: {[str(d) for d in jax.devices()]}")
print(f"Backend: {jax.default_backend()}")
print(f"64-bit precision: {jax.config.read('jax_enable_x64')}")

Key Benefits:

  • Automatic Differentiation: Forward and reverse mode AD for gradients
  • JIT Compilation: XLA optimization for high-performance execution
  • Multi-Device Support: Seamless CPU/GPU/TPU execution
  • Functional Programming: Pure functions for reproducible computations
  • 64-bit Precision: Scientific accuracy with configurable precision

FLAX NNX Integration

Modern neural network framework with stateful transforms:

from flax import nnx
from opifex.neural.base import StandardMLP

# Create RNG for reproducible initialization
rngs = nnx.Rngs(jax.random.PRNGKey(42))

# Build neural network with modern FLAX NNX
model = StandardMLP(
    layer_sizes=[2, 64, 64, 1],
    activation="swish",
    use_bias=True,
    rngs=rngs
)

# Forward pass
x = jax.random.normal(jax.random.PRNGKey(0), (32, 2))
output = model(x)
print(f"Input shape: {x.shape}, Output shape: {output.shape}")

Key Components

1. Problems (opifex.core.problems)

Define scientific problems with full specification capabilities:

from opifex.core.problems import create_pde_problem
from opifex.core.conditions import DirichletBC
from opifex.geometry import Rectangle

# Define heat equation residual
def heat_equation(x, u, u_derivatives):
    """Heat equation: du/dt - alpha * (d2u/dx2 + d2u/dy2) = 0"""
    u_t = u_derivatives['t']
    u_xx = u_derivatives['xx']
    u_yy = u_derivatives['yy']
    alpha = 0.01
    return u_t - alpha * (u_xx + u_yy)

# Define geometry and boundary conditions
geometry = Rectangle(center=jnp.array([0.5, 0.5]), width=1.0, height=1.0)

problem = create_pde_problem(
    geometry=geometry,
    equation=heat_equation,
    boundary_conditions=[
        DirichletBC(boundary="left", value=0.0),
        DirichletBC(boundary="right", value=1.0),
    ],
    parameters={"diffusivity": 0.01}
)

2. Neural Networks (opifex.neural)

Specialized architectures for scientific computing:

Available Architectures:

  • StandardMLP: Multi-layer perceptrons with scientific activations
  • AtomisticModel: Machine-learning interatomic potentials for molecular and materials systems (opifex.neural.atomistic)
  • FourierNeuralOperator: Learn mappings between function spaces
  • DeepONet: Deep operator networks for operator learning
  • PhysicsInformedOperator: Physics-aware neural operators
from opifex.neural.operators import FourierNeuralOperator

# Create Fourier Neural Operator
fno = FourierNeuralOperator(
    in_channels=2,
    out_channels=1,
    hidden_channels=64,
    modes=16,
    num_layers=4,
    rngs=rngs
)

# Process spatial data
spatial_data = jax.random.normal(jax.random.PRNGKey(1), (8, 64, 64, 2))
operator_output = fno(spatial_data)
print(f"Operator: {spatial_data.shape} -> {operator_output.shape}")

3. Training (opifex.training)

Physics-aware training procedures with advanced optimization:

from opifex.training.basic_trainer import ModularTrainer
from opifex.core.training.config import TrainingConfig

# Configure full training
config = TrainingConfig(
    num_epochs=5000,
    batch_size=128,
    learning_rate=1e-3,
    validation_frequency=100,
    checkpoint_frequency=500
)

# Create modular trainer with error recovery
trainer = ModularTrainer(
    model=model,
    config=config,
    rngs=rngs
)

# Train with automatic error recovery and optimization
trained_model, history = trainer.train(
    train_data=(x_train, y_train),
    val_data=(x_val, y_val)
)

4. Geometry (opifex.geometry)

Full geometric modeling with CSG operations:

from opifex.geometry import Rectangle, Circle, union, intersection
from opifex.geometry.manifolds import SphericalManifold

# Create 2D shapes
rect = Rectangle(center=jnp.array([0.0, 0.0]), width=2.0, height=1.5)
circle = Circle(center=jnp.array([1.0, 0.5]), radius=0.8)

# CSG operations
combined = union(rect, circle)
overlap = intersection(rect, circle)

# Sample boundary points
key = jax.random.PRNGKey(42)
boundary_points = rect.sample_boundary(n_points=100, key=key)

# Work with manifolds
sphere = SphericalManifold(dimension=2)
manifold_points = sphere.sample_points(n_points=50, key=key)

Scientific ML Paradigms

Physics-Informed Neural Networks (PINNs)

Neural networks that incorporate physical laws as soft constraints during training:

from opifex.neural.pinns.multi_scale import MultiScalePINN
from opifex.core.physics.losses import PhysicsInformedLoss, PhysicsLossConfig

# Create multi-scale PINN
pinn = MultiScalePINN(
    input_dim=2,
    output_dim=1,
    scales=[1, 2, 4],
    hidden_dims=[50, 50, 50],
    rngs=rngs
)

# Configure physics-informed loss
physics_config = PhysicsLossConfig(
    physics_loss_weight=1.0,
    boundary_loss_weight=10.0,
    data_loss_weight=1.0
)
physics_loss = PhysicsInformedLoss(
    config=physics_config,
    equation_type="heat",
    domain_type="rectangular"
)

Key Features:

  • Residual Computation: Automatic PDE residual calculation
  • Boundary Enforcement: Strong and weak boundary condition enforcement
  • Multi-Scale Training: Handle problems across different scales
  • Adaptive Weighting: Dynamic loss weight adjustment

Neural Operators

Learn mappings between function spaces, enabling generalization across different problem parameters:

from opifex.neural.operators import (
    FourierNeuralOperator,
    DeepONet,
    AdaptiveDeepONet,
    OperatorNetwork
)

# Fourier Neural Operator for PDEs
fno = FourierNeuralOperator(
    in_channels=2, out_channels=1,
    hidden_channels=64, modes=16,
    rngs=rngs
)

# Deep Operator Network
deeponet = DeepONet(
    branch_sizes=[100, 128, 128],
    trunk_sizes=[2, 128, 128],
    rngs=rngs
)

Operator Types Available:

  • FNO: Fourier Neural Operators with spectral convolutions
  • DeepONet: Deep operator networks with branch-trunk architecture
  • U-NO: U-Net style neural operators
  • GINO: Graph-informed neural operators
  • DISCO: Discrete-continuous convolutions

Atomistic Machine-Learning Potentials

Predict molecular and materials properties (energy, forces, stress) with a machine-learning interatomic potential assembled from a backbone and typed heads:

from flax import nnx

from opifex.core.quantum.molecular_system import create_water_molecule
from opifex.core.quantum.protocols import RadiusNeighborList
from opifex.core.quantum.registry import BackboneRegistry
from opifex.neural.atomistic import AtomisticModel
from opifex.neural.atomistic.heads import EnergyHead, ForcesHead

# Importing the backbones package registers "schnet" / "painn" / "nequip".
import opifex.neural.atomistic.backbones  # noqa: F401

rngs = nnx.Rngs(0)
backbone = BackboneRegistry().require("schnet")(rngs=rngs)
model = AtomisticModel(
    backbone=backbone,
    heads={"energy": EnergyHead(feature_dim=64, rngs=rngs), "forces": ForcesHead()},
    neighbor_list=RadiusNeighborList(cutoff=5.0),
    max_edges=64,
)

prediction = model(create_water_molecule())
print(f"Energy: {prediction['energy']}; forces shape: {prediction['forces'].shape}")

See the Atomistic Potentials guide for the full backbone/head design.

Probabilistic Numerics

Uncertainty quantification in scientific computations:

from opifex.uncertainty.aggregators import UncertaintyQuantifier

# Uncertainty quantification interface
uq = UncertaintyQuantifier(
    num_samples=100,
    confidence_level=0.95
)

# Decompose uncertainty from model predictions (samples x batch x output)
predictions = jax.random.normal(jax.random.PRNGKey(3), (100, 50, 1))
components = uq.decompose_uncertainty(predictions)
print(f"Epistemic uncertainty: {jnp.mean(components.epistemic):.3f}")
print(f"Aleatoric uncertainty: {jnp.mean(components.aleatoric):.3f}")

Advanced Features

Multi-Device Support

Seamless scaling across hardware:

# Check available devices
devices = jax.devices()
print(f"Available devices: {[str(d) for d in devices]}")

# Automatic device placement
if len(jax.devices('gpu')) > 0:
    print("🎮 GPU acceleration enabled")
else:
    print("💻 Running on CPU")

Checkpointing and Persistence

Robust model saving and loading:

from opifex.core.training.config import TrainingConfig

config = TrainingConfig(
    checkpoint_frequency=100,
    checkpoint_config={
        "save_directory": "./checkpoints",
        "max_to_keep": 5,
        "save_best_only": True
    }
)

Performance Optimization

Built-in performance monitoring and optimization:

# JIT compilation for performance
@jax.jit
def optimized_forward(model, x):
    return model(x)

# Vectorized operations
batch_output = jax.vmap(optimized_forward, in_axes=(None, 0))(model, batch_data)

Design Principles

1. Composability

Modular components that can be combined flexibly:

# Compose different components
from opifex.training.basic_trainer import ModularTrainer
from opifex.core.training.components.recovery import ErrorRecoveryManager

trainer = ModularTrainer(
    model=model,
    config=config,
    components={
        "error_recovery": ErrorRecoveryManager(config={}),
        "custom_component": CustomTrainingComponent()
    }
)

2. Performance

Optimized for scientific computing workloads:

  • JAX transformations (jit, vmap, pmap)
  • XLA compilation for optimal performance
  • Memory-efficient implementations
  • GPU/TPU acceleration

3. Extensibility

Easy to add new methods and approaches:

  • Protocol-based interfaces
  • Modular architecture
  • Plugin system for custom components
  • Clear extension points

4. Reproducibility

Deterministic computations with proper seeding:

# Reproducible random number generation
key = jax.random.PRNGKey(42)
rngs = nnx.Rngs(key)

# Deterministic model initialization
model = StandardMLP(layer_sizes=[2, 64, 1], rngs=rngs)

# Reproducible training
trainer = ModularTrainer(model=model, config=config, rngs=rngs)

Getting Started

Quick Example

import jax
import jax.numpy as jnp
from flax import nnx
from opifex.neural.base import StandardMLP
from opifex.training.basic_trainer import ModularTrainer
from opifex.core.training.config import TrainingConfig

# 1. Setup
key = jax.random.PRNGKey(42)
rngs = nnx.Rngs(key)

# 2. Create model
model = StandardMLP(layer_sizes=[2, 64, 64, 1], activation="swish", rngs=rngs)

# 3. Generate data
x = jax.random.uniform(key, (1000, 2), minval=-2, maxval=2)
y = jnp.sin(jnp.pi * x[:, 0]) * jnp.cos(jnp.pi * x[:, 1])

# 4. Configure training
config = TrainingConfig(num_epochs=1000, learning_rate=1e-3)

# 5. Train
trainer = ModularTrainer(model=model, config=config, rngs=rngs)
trained_model, history = trainer.train(train_data=(x, y))

print("✅ Opifex training complete!")

This complete framework enables researchers and practitioners to tackle complex scientific machine learning problems with advanced methods and high-performance computing capabilities.