#!/usr/bin/env python3 """ Generate reference inputs and outputs for Python vs Rust comparison testing. This script creates: 1. A deterministic input tensor 2. FNO model weights in SafeTensors format 3. Reference output from neuraloperator The Rust test can then load these files and verify numerical equivalence. Usage: python generate_reference_outputs.py --output weights/fno/fno2d_darcy/reference """ import argparse import json import os from pathlib import Path import numpy as np import torch from safetensors.torch import save_file, load_file # Check if neuraloperator is available NEURALOP_AVAILABLE = False try: from neuralop.models import FNO NEURALOP_AVAILABLE = True except ImportError: pass def create_grid_positional_encoding(height: int, width: int) -> torch.Tensor: """Create 2D grid positional encoding matching neuraloperator.""" # Create coordinate grids in [-1, 1] y = torch.linspace(-1, 1, height) x = torch.linspace(-1, 1, width) # Create meshgrid (neuraloperator uses 'ij' indexing) grid_y, grid_x = torch.meshgrid(y, x, indexing='ij') # Stack to [2, H, W] grid = torch.stack([grid_x, grid_y], dim=0) return grid def create_deterministic_input(batch_size: int, in_channels: int, height: int, width: int, seed: int = 42) -> torch.Tensor: """Create a deterministic input tensor for reproducible testing.""" torch.manual_seed(seed) np.random.seed(seed) # Create input with known pattern for easier debugging # Use a combination of smooth functions y = torch.linspace(0, 2 * np.pi, height) x = torch.linspace(0, 2 * np.pi, width) yy, xx = torch.meshgrid(y, x, indexing='ij') # Create smooth input pattern input_tensor = torch.zeros(batch_size, in_channels, height, width) for b in range(batch_size): for c in range(in_channels): # Different frequency for each channel freq = (c + 1) * (b + 1) input_tensor[b, c] = torch.sin(freq * xx) * torch.cos(freq * yy) return input_tensor def create_fno_weights_manual(config: dict, seed: int = 42) -> dict: """ Create FNO weights manually that match both neuraloperator and rtx-neural-operator. Returns weights in rtx-neural-operator naming convention. """ torch.manual_seed(seed) data_channels = config["data_channels"] out_channels = config["out_channels"] width = config["width"] n_modes = config["n_modes"] n_layers = config["n_layers"] weights = {} # Lifting MLP: (data_channels + 2) -> 2*width -> width lifting_in = data_channels + 2 # +2 for positional encoding hidden_dim = 2 * width # Use Xavier/Glorot initialization for better numerical stability scale_fc1 = np.sqrt(2.0 / (lifting_in + hidden_dim)) scale_fc2 = np.sqrt(2.0 / (hidden_dim + width)) weights["lifting.fcs.0.weight"] = torch.randn(hidden_dim, lifting_in) * scale_fc1 weights["lifting.fcs.0.bias"] = torch.zeros(hidden_dim) weights["lifting.fcs.1.weight"] = torch.randn(width, hidden_dim) * scale_fc2 weights["lifting.fcs.1.bias"] = torch.zeros(width) # Spectral convolution layers for i in range(n_layers): # SpectralConv weights: [in_ch, out_ch, modes_h, modes_w] # Note: neuraloperator uses [in_ch, out_ch, ...] convention, not PyTorch's [out_ch, in_ch, ...] # Use small initialization for spectral weights scale_spectral = 1.0 / (width * n_modes[0] * n_modes[1]) # Shape is [in_channels, out_channels, modes_h, modes_w] to match Rust implementation weights[f"spectral_conv.{i}.weights1_real"] = torch.randn(width, width, n_modes[0], n_modes[1]) * scale_spectral weights[f"spectral_conv.{i}.weights1_imag"] = torch.randn(width, width, n_modes[0], n_modes[1]) * scale_spectral weights[f"spectral_conv.{i}.weights2_real"] = torch.randn(width, width, n_modes[0], n_modes[1]) * scale_spectral weights[f"spectral_conv.{i}.weights2_imag"] = torch.randn(width, width, n_modes[0], n_modes[1]) * scale_spectral # Skip connection (1x1 conv as linear) scale_conv = np.sqrt(2.0 / (width + width)) weights[f"conv.{i}.weight"] = torch.randn(width, width) * scale_conv weights[f"conv.{i}.bias"] = torch.zeros(width) # Projection layers: width -> 128 -> out_channels proj_hidden = 128 scale_proj1 = np.sqrt(2.0 / (width + proj_hidden)) scale_proj2 = np.sqrt(2.0 / (proj_hidden + out_channels)) weights["projection.0.weight"] = torch.randn(proj_hidden, width) * scale_proj1 weights["projection.0.bias"] = torch.zeros(proj_hidden) weights["projection.1.weight"] = torch.randn(out_channels, proj_hidden) * scale_proj2 weights["projection.1.bias"] = torch.zeros(out_channels) return weights def run_fno_forward_manual(input_tensor: torch.Tensor, weights: dict, config: dict) -> torch.Tensor: """ Run FNO forward pass manually to match rtx-neural-operator implementation. This implements the same algorithm as the Rust code for direct comparison. """ batch_size, in_channels, height, width_dim = input_tensor.shape model_width = config["width"] n_modes = config["n_modes"] n_layers = config["n_layers"] # Step 1: Add positional encoding grid = create_grid_positional_encoding(height, width_dim) # [2, H, W] grid = grid.unsqueeze(0).expand(batch_size, -1, -1, -1) # [B, 2, H, W] x = torch.cat([input_tensor, grid], dim=1) # [B, in_ch+2, H, W] # Step 2: Lifting MLP # Reshape for linear: [B, C, H, W] -> [B, H, W, C] x = x.permute(0, 2, 3, 1) # FC1 with GELU fc1_w = weights["lifting.fcs.0.weight"] # [hidden, in] fc1_b = weights["lifting.fcs.0.bias"] # [hidden] x = torch.nn.functional.linear(x, fc1_w, fc1_b) x = torch.nn.functional.gelu(x) # FC2 fc2_w = weights["lifting.fcs.1.weight"] # [width, hidden] fc2_b = weights["lifting.fcs.1.bias"] # [width] x = torch.nn.functional.linear(x, fc2_w, fc2_b) # Reshape back: [B, H, W, C] -> [B, C, H, W] x = x.permute(0, 3, 1, 2) # Step 3: Fourier layers for i in range(n_layers): # Spectral convolution path x_ft = torch.fft.fft2(x) # [B, C, H, W] complex # Get spectral weights w1_real = weights[f"spectral_conv.{i}.weights1_real"] w1_imag = weights[f"spectral_conv.{i}.weights1_imag"] w2_real = weights[f"spectral_conv.{i}.weights2_real"] w2_imag = weights[f"spectral_conv.{i}.weights2_imag"] # Create complex weight tensors # Weight shape: [in_ch, out_ch, modes_h, modes_w] (neuraloperator convention) w1 = torch.complex(w1_real, w1_imag) w2 = torch.complex(w2_real, w2_imag) modes_h, modes_w = n_modes # Initialize output in frequency domain out_ft = torch.zeros_like(x_ft) # Apply weights to upper left corner (positive frequencies) # x_ft[:, :, :modes_h, :modes_w] with w1 x_upper = x_ft[:, :, :modes_h, :modes_w] # [B, in_ch, modes_h, modes_w] # Einstein sum: batch, in_ch, h, w with in_ch, out_ch, h, w -> batch, out_ch, h, w # Contract over in_ch, element-wise multiply h,w out_upper = torch.einsum('bihw,iohw->bohw', x_upper, w1) out_ft[:, :, :modes_h, :modes_w] = out_upper # Apply weights to lower left corner (negative frequencies in height) # x_ft[:, :, -modes_h:, :modes_w] with w2 x_lower = x_ft[:, :, -modes_h:, :modes_w] out_lower = torch.einsum('bihw,iohw->bohw', x_lower, w2) out_ft[:, :, -modes_h:, :modes_w] = out_lower # Inverse FFT x_spectral = torch.fft.ifft2(out_ft).real # Skip connection path (1x1 conv as pointwise linear) conv_w = weights[f"conv.{i}.weight"] # [out, in] conv_b = weights[f"conv.{i}.bias"] # [out] # Reshape for linear x_skip = x.permute(0, 2, 3, 1) # [B, H, W, C] x_skip = torch.nn.functional.linear(x_skip, conv_w, conv_b) x_skip = x_skip.permute(0, 3, 1, 2) # [B, C, H, W] # Combine and activate x = x_spectral + x_skip if i < n_layers - 1: # No activation on last layer x = torch.nn.functional.gelu(x) # Step 4: Projection x = x.permute(0, 2, 3, 1) # [B, H, W, C] # Projection layer 1 with GELU proj1_w = weights["projection.0.weight"] proj1_b = weights["projection.0.bias"] x = torch.nn.functional.linear(x, proj1_w, proj1_b) x = torch.nn.functional.gelu(x) # Projection layer 2 proj2_w = weights["projection.1.weight"] proj2_b = weights["projection.1.bias"] x = torch.nn.functional.linear(x, proj2_w, proj2_b) x = x.permute(0, 3, 1, 2) # [B, C, H, W] return x def main(): parser = argparse.ArgumentParser( description="Generate reference inputs/outputs for Python vs Rust comparison" ) parser.add_argument( "--output", type=str, default="./weights/fno/fno2d_darcy", help="Output directory for reference files" ) parser.add_argument( "--height", type=int, default=64, help="Input height (default: 64)" ) parser.add_argument( "--width", type=int, default=64, help="Input width (default: 64)" ) parser.add_argument( "--batch-size", type=int, default=1, help="Batch size (default: 1)" ) parser.add_argument( "--seed", type=int, default=42, help="Random seed (default: 42)" ) args = parser.parse_args() # Model configuration (matching fno2d_darcy) config = { "description": "FNO2d trained on Darcy Flow equation", "url": "https://github.com/neuraloperator/neuraloperator", "in_channels": 3, # data_channels + 2 (positional encoding) "data_channels": 1, # Raw data channels "out_channels": 1, "width": 32, "n_modes": [12, 12], "n_layers": 4, "pde_type": "darcy", } output_dir = Path(args.output) output_dir.mkdir(parents=True, exist_ok=True) print(f"Generating reference data with seed={args.seed}") print(f" Input shape: [{args.batch_size}, {config['data_channels']}, {args.height}, {args.width}]") print(f" Model: width={config['width']}, modes={config['n_modes']}, layers={config['n_layers']}") # Create deterministic input input_tensor = create_deterministic_input( args.batch_size, config["data_channels"], args.height, args.width, seed=args.seed ) print(f"\nInput tensor shape: {input_tensor.shape}") print(f"Input range: [{input_tensor.min():.4f}, {input_tensor.max():.4f}]") # Create weights weights = create_fno_weights_manual(config, seed=args.seed) print(f"\nCreated {len(weights)} weight tensors") # Run forward pass print("\nRunning forward pass...") with torch.no_grad(): output_tensor = run_fno_forward_manual(input_tensor, weights, config) print(f"Output tensor shape: {output_tensor.shape}") print(f"Output range: [{output_tensor.min():.4f}, {output_tensor.max():.4f}]") print(f"Output mean: {output_tensor.mean():.6f}") print(f"Output std: {output_tensor.std():.6f}") # Save weights weights_path = output_dir / "model.safetensors" save_file(weights, weights_path) print(f"\nSaved weights to {weights_path}") # Save config config_path = output_dir / "config.json" with open(config_path, "w") as f: json.dump(config, f, indent=2) print(f"Saved config to {config_path}") # Save input tensor input_path = output_dir / "reference_input.safetensors" save_file({"input": input_tensor}, input_path) print(f"Saved input to {input_path}") # Save output tensor output_path = output_dir / "reference_output.safetensors" save_file({"output": output_tensor}, output_path) print(f"Saved output to {output_path}") # Save metadata metadata = { "seed": args.seed, "batch_size": args.batch_size, "height": args.height, "width": args.width, "input_shape": list(input_tensor.shape), "output_shape": list(output_tensor.shape), "output_mean": float(output_tensor.mean()), "output_std": float(output_tensor.std()), "output_min": float(output_tensor.min()), "output_max": float(output_tensor.max()), } metadata_path = output_dir / "reference_metadata.json" with open(metadata_path, "w") as f: json.dump(metadata, f, indent=2) print(f"Saved metadata to {metadata_path}") print("\n" + "=" * 60) print("Reference data generation complete!") print("=" * 60) print(f"\nFiles created in {output_dir}:") print(" - model.safetensors (weights)") print(" - config.json (model config)") print(" - reference_input.safetensors (input tensor)") print(" - reference_output.safetensors (expected output)") print(" - reference_metadata.json (test metadata)") if __name__ == "__main__": main()