OSCR

Three-dimensional inversion of gravity data using implicit neural representations and scientific machine learning.

Code ↔ Paper

15 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 15 matches
  1. [1] § Methods › Implicit neural representation and spatial encodings ↔ 001-EncodingComparisons_blocky.py, lines 120–245 · score 0.76 · learnable feature vectors, trilinear interpolation, hash encoding, concatenated, corner, embeddings
  2. [2] § Methods › Implicit neural representation and spatial encodings ↔ 001-EncodingComparisons_blocky.py, lines 118–243 · score 0.76 · learnable feature vectors, trilinear interpolation, hash encoding, concatenated, corner, embeddings
  3. [3] § Results › Spectral bias and spatial encoding comparison ↔ v01/05-TestingNoiseSensitivity.py, lines 1–54 · score 0.74 · rectangular prism kernel, stationary Gaussian, Gaussian noise, gravity inversion, whitened, surface
  4. [4] § Methods › Forward modelling of gravity data ↔ v01/04-BlockModel.py, lines 109–130 · score 0.69 · vertical component, rectangular prism, sensitivity matrix, cells, gravity, model
  5. [5] § Results › Network capacity and grid refinement ↔ 002a-BlockNetworkSizeComparison.py, lines 268–399 · score 0.66 · model error, training loss, voxel model, ratio, INR parameter, largest
  6. [6] § Methods › Forward modelling of gravity data ↔ v01/05-TestingNoiseSensitivity.py, lines 76–90 · score 0.64 · vertical component, rectangular prism, arctan, sensitivity, gravity
  7. [7] § Results › Implicit versus explicit regularisation ↔ v01/04-BlockModel.py, lines 232–281 · score 0.64 · smallness term, linear system, depth weighted, gradient, smoothness, noise
  8. [8] § Results › Spectral bias and spatial encoding comparison ↔ 001-EncodingComparisons_blocky.py, lines 120–245 · score 0.63 · finest resolution, noise floor, hash encoding, encoding comparison, geometric, blocky
  9. [9] § Results › Spectral bias and spatial encoding comparison ↔ 001-EncodingComparisons_blocky.py, lines 118–243 · score 0.63 · finest resolution, noise floor, hash encoding, encoding comparison, geometric, blocky
  10. [10] § Results › Implicit versus explicit regularisation ↔ 003-ImplictvsExplicitRegularisation.py, lines 1–21 · score 0.60 · explicit smoothness, hidden layers, frequency bands, positional encoding, implicit
  11. [11] § Results › Spectral bias and spatial encoding comparison ↔ 001-EncodingComparisons_blocky.py, lines 73–117 · score 0.60 · overly smooth, sinusoidal positional encoding, lowest, blocky, MLP, bias
  12. [12] § Results › Spectral bias and spatial encoding comparison ↔ 001-EncodingComparisons_blocky.py, lines 71–115 · score 0.60 · overly smooth, sinusoidal positional encoding, lowest, blocky, MLP, bias
  13. [13] § Results › Network capacity and grid refinement ↔ 002b-BlockGridSizeComparison.py, lines 1–20 · score 0.53 · voxel model parameter, INR architectures, INR parameter, gravity inversion, block model, discretisation
  14. [14] § Results › Network capacity and grid refinement ↔ 004-ModelEnsambles.py, lines 1–23 · score 0.51 · kept fixed, INR architecture, block model, medium, positional encoding, Network
  15. [15] § Results › Network capacity and grid refinement ↔ 002b-BlockGridSizeComparison.py, lines 1–20 · score 0.50 · voxel model parameter, INR parameter, refined, architecture, discretisation, inversion

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 759 lines · 30 KB · MIT · 3 matches

  1. import os
  2. import sys
  3. import time
  4. import random
  5. import numpy as np
  6. import torch
  7. import torch.nn as nn
  8. import matplotlib
  9. import matplotlib.pyplot as plt
  10. # --- Random seed -------------------------------------------------------
  11. SEED = 42 #10,42, 40, 41
  12. DATA_SEED = SEED
  13. # --- Grid / domain -----------------------------------------------------
  14. DX = 50.0 # cell size in x (m)
  15. DY = 50.0 # cell size in y (m)
  16. DZ = 50.0 # cell size in z (m)
  17. X_MAX = 1000.0 # domain extent in x (m)
  18. Y_MAX = 1000.0 # domain extent in y (m)
  19. Z_MAX = 500.0 # domain extent in z (m)
  20. # --- Block model -------------------------------------------------------
  21. RHO_BG = 0.0 # background density contrast (kg/m³)
  22. RHO_BLK = 400.0 # block density contrast (kg/m³)
  23. # --- Noise level -------------------------------------------------------
  24. NOISE_LEVEL = 0.01 # fraction of gz_true std
  25. # --- Training / optimisation -------------------------------------------
  26. GAMMA = 1.0 # data-term weight
  27. EPOCHS = 500 # number of training epochs
  28. LR = 1e-2 # Adam learning rate
  29. # --- Early stopping ----------------------------------------------------
  30. # Stop near the expected noise floor to reduce overfitting to 1% noise.
  31. USE_EARLY_STOPPING = True
  32. EARLY_STOP_MIN_EPOCHS = 100
  33. EARLY_STOP_PATIENCE = 25
  34. EARLY_STOP_TARGET = 1.0
  35. EARLY_STOP_TOL = 0.05
  36. EARLY_STOP_OVERFIT_PATIENCE = 5
  37. # --- INR network -------------------------------------------------------
  38. HIDDEN = 256 # hidden-layer width
  39. DEPTH = 4 # number of hidden layers
  40. RHO_ABS_MAX = 600.0 # tanh output scaling (kg/m³)
  41. # --- Encoding strategy -------------------------------------------------
  42. # Compare only the two encoded models in this script.
  43. ENCODING_ORDER = ('positional', 'hash')
  44. ENCODING_CONFIGS = {
  45. 'positional': dict(num_freqs=2),
  46. 'hash': dict(n_levels=2, n_features_per_level=2,
  47. log2_hashmap_size=17, base_resolution=4,
  48. finest_resolution=128),
  49. }
  50. # --- Plotting -----------------------------------------------------------
  51. CMAP = 'Spectral_r' # colormap for all plots
  52. INV_VMAX = 250 # fixed colorbar max for inverted model
  53. FIG_DPI = 300
  54. TITLE_FONTSIZE = 14
  55. LABEL_FONTSIZE = 12
  56. TICK_FONTSIZE = 11
  57. LEGEND_FONTSIZE = 10
  58. class PositionalEncoding(nn.Module):
  59. """Sinusoidal positional encoding (NeRF-style Fourier features).
  60. Intuition
  61. ---------
  62. Plain MLPs are biased toward learning low-frequency functions — they
  63. tend to produce overly smooth outputs and struggle to represent sharp
  64. edges or fine detail ("spectral bias"). Positional encoding fights
  65. this by lifting each input coordinate into a set of sine and cosine
  66. waves at *exponentially increasing* frequencies:
  67. γ(p) = [sin(2⁰πp), cos(2⁰πp), sin(2¹πp), cos(2¹πp), …]
  68. Think of it as giving the network a set of "rulers" at different
  69. scales: the lowest frequency captures the overall trend, while
  70. higher frequencies let the network resolve increasingly fine spatial
  71. variation. The resulting feature vector is fixed (no learnable
  72. parameters in the encoding itself).
  73. Trade-offs
  74. ----------
  75. • Simple and deterministic — no extra learnable parameters.
  76. • The frequency ladder is fixed (powers of 2), so the spectrum can
  77. have gaps and may not cover all scales equally well.
  78. • Too few frequencies → overly smooth model; too many → potential
  79. noise fitting and slower convergence.
  80. • Works best with a sufficiently deep/wide MLP.
  81. Args:
  82. num_freqs: Number of frequency octaves (L). Output grows
  83. as input_dim × (1 + 2L) with include_input.
  84. include_input: Whether to prepend the raw (x, y, z) values.
  85. input_dim: Spatial dimensionality (default 3).
  86. """
  87. def __init__(self, num_freqs=8, include_input=True, input_dim=3):
  88. super().__init__()
  89. self.include_input = include_input
  90. self.register_buffer('freqs', 2.0 ** torch.arange(0, num_freqs))
  91. self.out_dim = input_dim * (1 + 2 * num_freqs) if include_input else input_dim * 2 * num_freqs
  92. def forward(self, x):
  93. parts = [x] if self.include_input else []
  94. for f in self.freqs:
  95. parts += [torch.sin(f * x), torch.cos(f * x)]
  96. return torch.cat(parts, dim=-1)
  97. class HashEncoding(nn.Module):
  98. """Multi-resolution hash encoding (Müller et al., 2022 / Instant-NGP).
  99. Intuition
  100. ---------
  101. Imagine overlaying your 3-D domain with a stack of voxel grids, from
  102. very coarse (e.g. 4³) to very fine (e.g. 128³). At each resolution
  103. level, every voxel stores a small learnable feature vector. To
  104. encode a point, you look up its 8 surrounding voxel corners at every
  105. level, trilinearly interpolate, and concatenate across levels.
  106. The "hash" trick makes this memory-efficient: instead of allocating a
  107. full 3-D grid (which grows as O(R³)), each level maps voxel corners
  108. to a fixed-size hash table. Collisions (two corners sharing one
  109. slot) are resolved implicitly by the gradient-based optimisation —
  110. the network learns to disentangle them.
  111. The result is an encoding that is:
  112. • **Adaptive**: features are *learned*, so detail concentrates where
  113. the data demands it (unlike fixed Fourier features).
  114. • **Multi-scale**: coarse levels capture large-scale trends; fine
  115. levels capture sharp boundaries and small anomalies.
  116. • **Memory-bounded**: hash table size is constant regardless of
  117. how fine the resolution is — crucial for large 3-D domains.
  118. Trade-offs
  119. ----------
  120. • Most flexible encoding — can represent both smooth and sharp models.
  121. • Hash collisions at fine levels can introduce small artifacts or a
  122. noise floor if the hash table is too small (increase
  123. log2_hashmap_size to mitigate).
  124. • Many hyperparameters to tune (n_levels, base/finest resolution,
  125. hash table size).
  126. • Learnable parameters mean more total parameters to optimise, and
  127. they may overfit noisy data without appropriate regularisation.
  128. Args:
  129. n_levels: Number of resolution levels. More levels
  130. give smoother multi-scale interpolation.
  131. n_features_per_level: Feature dimensions stored per hash entry.
  132. Typically 2; larger values add capacity.
  133. log2_hashmap_size: Log₂ of hash-table size per level.
  134. 2^19 ≈ 500 k entries is a common default.
  135. base_resolution: Coarsest voxel-grid resolution (e.g. 4).
  136. finest_resolution: Finest voxel-grid resolution (e.g. 128).
  137. Intermediate levels are spaced
  138. geometrically between base and finest.
  139. input_dim: Spatial dimensionality (default 3).
  140. """
  141. def __init__(self, n_levels=16, n_features_per_level=2,
  142. log2_hashmap_size=19, base_resolution=16,
  143. finest_resolution=512, input_dim=3):
  144. super().__init__()
  145. self.n_levels = n_levels
  146. self.n_features_per_level = n_features_per_level
  147. self.input_dim = input_dim
  148. self.out_dim = n_levels * n_features_per_level
  149. self.hashmap_size = 2 ** log2_hashmap_size
  150. if n_levels > 1:
  151. self.growth_factor = np.exp(
  152. (np.log(finest_resolution) - np.log(base_resolution))
  153. / (n_levels - 1))
  154. else:
  155. self.growth_factor = 1.0
  156. self.base_resolution = base_resolution
  157. # Learnable hash tables (one per level)
  158. self.hash_tables = nn.ModuleList([
  159. nn.Embedding(self.hashmap_size, n_features_per_level)
  160. for _ in range(n_levels)
  161. ])
  162. for tbl in self.hash_tables:
  163. nn.init.uniform_(tbl.weight, -1e-4, 1e-4)
  164. # Large primes for the spatial hash
  165. self.register_buffer(
  166. 'primes', torch.tensor([1, 2654435761, 805459861], dtype=torch.long))
  167. def _hash(self, coords_int):
  168. """Spatial hash: integer grid coords -> hash-table index."""
  169. result = torch.zeros(coords_int.shape[:-1],
  170. dtype=torch.long, device=coords_int.device)
  171. for d in range(self.input_dim):
  172. result ^= coords_int[..., d] * self.primes[d]
  173. return result % self.hashmap_size
  174. def forward(self, x):
  175. # Normalise to [0, 1] using per-batch bounds
  176. x_min = x.min(dim=0, keepdim=True).values
  177. x_max = x.max(dim=0, keepdim=True).values
  178. x_scaled = (x - x_min) / (x_max - x_min + 1e-8)
  179. outputs = []
  180. for level in range(self.n_levels):
  181. resolution = int(self.base_resolution * (self.growth_factor ** level))
  182. x_grid = x_scaled * resolution # (N, 3)
  183. x_floor = torch.floor(x_grid).long() # voxel origin
  184. x_frac = x_grid - x_floor.float() # interpolation weight
  185. # Eight voxel corners
  186. corners = []
  187. for dz in (0, 1):
  188. for dy in (0, 1):
  189. for dx in (0, 1):
  190. corners.append(
  191. x_floor + torch.tensor([dx, dy, dz], device=x.device))
  192. corners = torch.stack(corners, dim=1) # (N, 8, 3)
  193. indices = self._hash(corners) # (N, 8)
  194. features = self.hash_tables[level](indices) # (N, 8, F)
  195. # Trilinear interpolation weights
  196. wx, wy, wz = (x_frac[:, 0:1],
  197. x_frac[:, 1:2],
  198. x_frac[:, 2:3])
  199. weights = torch.stack([
  200. (1-wx)*(1-wy)*(1-wz), wx*(1-wy)*(1-wz),
  201. (1-wx)* wy *(1-wz), wx* wy *(1-wz),
  202. (1-wx)*(1-wy)* wz , wx*(1-wy)* wz ,
  203. (1-wx)* wy * wz , wx* wy * wz ,
  204. ], dim=1) # (N, 8, 1)
  205. outputs.append((weights * features).sum(dim=1)) # (N, F)
  206. return torch.cat(outputs, dim=-1) # (N, n_levels*F)
  207. def create_encoding(encoding_type, **kwargs):
  208. """Factory: build an encoding module by name.
  209. Supported types
  210. ---------------
  211. positional – Sinusoidal Fourier features (Mildenhall et al., 2020
  212. / NeRF). Fixed frequencies on a power-of-2 ladder.
  213. Simple, no learnable params; good general-purpose
  214. baseline. Tune `num_freqs` for resolution.
  215. hash – Multi-resolution hash tables (Müller et al., 2022 /
  216. Instant-NGP). Learnable feature grids at multiple
  217. resolutions compressed via spatial hashing. Most
  218. flexible; best for sharp boundaries. More params &
  219. hyperparameters.
  220. """
  221. if encoding_type == 'positional':
  222. enc = PositionalEncoding(
  223. num_freqs=kwargs.get('num_freqs', 8),
  224. include_input=kwargs.get('include_input', True))
  225. return enc
  226. if encoding_type == 'hash':
  227. return HashEncoding(
  228. n_levels=kwargs.get('n_levels', 16),
  229. n_features_per_level=kwargs.get('n_features_per_level', 2),
  230. log2_hashmap_size=kwargs.get('log2_hashmap_size', 19),
  231. base_resolution=kwargs.get('base_resolution', 16),
  232. finest_resolution=kwargs.get('finest_resolution', 512))
  233. raise ValueError(f"Unknown encoding type: {encoding_type}. Expected 'positional' or 'hash'.")
  234. class DensityContrastINR(nn.Module):
  235. """INR density-contrast model with pluggable spatial encoding."""
  236. def __init__(self, encoding_type='positional', hidden=256, depth=5,
  237. rho_abs_max=600.0, **encoding_kwargs):
  238. super().__init__()
  239. self.pe = create_encoding(encoding_type, **encoding_kwargs)
  240. in_dim = self.pe.out_dim
  241. layers = []
  242. h = hidden
  243. layers += [nn.Linear(in_dim, h), nn.LeakyReLU(0.01)]
  244. for _ in range(depth - 1):
  245. layers += [nn.Linear(h, h), nn.LeakyReLU(0.01)]
  246. layers += [nn.Linear(h, 1)]
  247. self.net = nn.Sequential(*layers)
  248. self.rho_abs_max = float(rho_abs_max)
  249. def forward(self, x):
  250. z = self.pe(x)
  251. out = self.net(z)
  252. return self.rho_abs_max * torch.tanh(out)
  253. # ──────────────────────────────────────────────────────────────────────
  254. # UTILITY FUNCTIONS
  255. # ──────────────────────────────────────────────────────────────────────
  256. def set_seed(seed: int = 42):
  257. random.seed(seed)
  258. np.random.seed(seed)
  259. torch.manual_seed(seed)
  260. if torch.cuda.is_available():
  261. torch.cuda.manual_seed(seed)
  262. torch.cuda.manual_seed_all(seed)
  263. torch.backends.cudnn.deterministic = True
  264. torch.backends.cudnn.benchmark = False
  265. os.environ['PYTHONHASHSEED'] = str(seed)
  266. print(f"Seed = {seed}")
  267. def capture_rng_state():
  268. state = {
  269. 'python': random.getstate(),
  270. 'numpy': np.random.get_state(),
  271. 'torch': torch.get_rng_state(),
  272. }
  273. if torch.cuda.is_available():
  274. state['cuda'] = torch.cuda.get_rng_state_all()
  275. return state
  276. def restore_rng_state(state):
  277. random.setstate(state['python'])
  278. np.random.set_state(state['numpy'])
  279. torch.set_rng_state(state['torch'])
  280. if torch.cuda.is_available() and 'cuda' in state:
  281. torch.cuda.set_rng_state_all(state['cuda'])
  282. def A_integral_torch(x, y, z):
  283. eps = 1e-20
  284. r = torch.sqrt(x**2 + y**2 + z**2).clamp_min(eps)
  285. return -(x * torch.log(torch.abs(y + r) + eps) +
  286. y * torch.log(torch.abs(x + r) + eps) -
  287. z * torch.atan2(x * y, z * r + eps))
  288. @torch.inference_mode()
  289. def construct_sensitivity_matrix_G_torch(cell_grid, data_points, d1, d2, device):
  290. Gamma = 6.67430e-11
  291. cx = cell_grid[:, 0].unsqueeze(0)
  292. cy = cell_grid[:, 1].unsqueeze(0)
  293. cz = cell_grid[:, 2].unsqueeze(0)
  294. czh = cell_grid[:, 3].unsqueeze(0)
  295. ox = data_points[:, 0].unsqueeze(1)
  296. oy = data_points[:, 1].unsqueeze(1)
  297. oz = data_points[:, 2].unsqueeze(1)
  298. x2, x1 = (cx + d1/2) - ox, (cx - d1/2) - ox
  299. y2, y1 = (cy + d2/2) - oy, (cy - d2/2) - oy
  300. z2, z1 = (cz + czh) - oz, (cz - czh) - oz
  301. A = (A_integral_torch(x2, y2, z2) - A_integral_torch(x2, y2, z1) -
  302. A_integral_torch(x2, y1, z2) + A_integral_torch(x2, y1, z1) -
  303. A_integral_torch(x1, y2, z2) + A_integral_torch(x1, y2, z1) +
  304. A_integral_torch(x1, y1, z2) - A_integral_torch(x1, y1, z1))
  305. return (Gamma * A).to(device)
  306. def generate_grf_torch(nx, ny, nz, dx, dy, dz, lam, nu, sigma, device):
  307. kx = torch.fft.fftfreq(nx, d=dx, device=device) * 2 * torch.pi
  308. ky = torch.fft.fftfreq(ny, d=dy, device=device) * 2 * torch.pi
  309. kz = torch.fft.fftfreq(nz, d=dz, device=device) * 2 * torch.pi
  310. Kx, Ky, Kz = torch.meshgrid(kx, ky, kz, indexing='ij')
  311. k2 = Kx**2 + Ky**2 + Kz**2
  312. P = (k2 + (1/lam**2))**(-nu - 1.5)
  313. P[0, 0, 0] = 0
  314. noise = torch.randn(nx, ny, nz, dtype=torch.complex64, device=device)
  315. f = noise * torch.sqrt(P)
  316. m = torch.real(torch.fft.ifftn(f))
  317. m = (m - m.mean()) / (m.std() + 1e-9)
  318. return sigma * m
  319. def train_inr(model, opt, coords_norm, G, gz_obs, Wd, Nx, Ny, Nz, dx, dy, dz, cfg):
  320. history = {"total": [], "gravity": []}
  321. use_early_stopping = cfg.get('use_early_stopping', False)
  322. min_epochs = cfg.get('early_stop_min_epochs', 0)
  323. patience = cfg.get('early_stop_patience', 0)
  324. target = cfg.get('early_stop_target', 1.0)
  325. tol = cfg.get('early_stop_tol', 0.0)
  326. overfit_patience = cfg.get('early_stop_overfit_patience', 0)
  327. best_gap = float('inf')
  328. best_epoch = -1
  329. best_weighted_mse = None
  330. best_state = None
  331. in_band_count = 0
  332. overfit_count = 0
  333. for ep in range(cfg['epochs']):
  334. opt.zero_grad()
  335. m_pred = model(coords_norm).view(-1)
  336. gz_pred = torch.matmul(G, m_pred.unsqueeze(1)).squeeze(1)
  337. residual = gz_pred - gz_obs
  338. data_term = cfg['gamma'] * torch.mean((Wd * residual) ** 2)
  339. loss = data_term
  340. loss.backward()
  341. opt.step()
  342. history['gravity'].append(float(data_term.item()))
  343. history['total'].append(float(loss.item()))
  344. if use_early_stopping:
  345. weighted_mse = float(data_term.item() / cfg['gamma'])
  346. gap = abs(weighted_mse - target)
  347. if gap < best_gap:
  348. best_gap = gap
  349. best_epoch = ep
  350. best_weighted_mse = weighted_mse
  351. best_state = {
  352. name: value.detach().cpu().clone()
  353. for name, value in model.state_dict().items()
  354. }
  355. if ep + 1 >= min_epochs:
  356. if target - tol <= weighted_mse <= target + tol:
  357. in_band_count += 1
  358. else:
  359. in_band_count = 0
  360. if weighted_mse < target - tol:
  361. overfit_count += 1
  362. else:
  363. overfit_count = 0
  364. if in_band_count >= patience:
  365. print(
  366. f"Early stopping at epoch {ep:4d} | "
  367. f"reason = near noise floor | "
  368. f"loss = {weighted_mse:.3f} | "
  369. f"best epoch = {best_epoch}"
  370. )
  371. break
  372. if overfit_count >= overfit_patience:
  373. print(
  374. f"Early stopping at epoch {ep:4d} | "
  375. f"reason = below noise floor | "
  376. f"loss = {weighted_mse:.3f} | "
  377. f"best epoch = {best_epoch}"
  378. )
  379. break
  380. if ep % 50 == 0 or ep == cfg['epochs'] - 1:
  381. print(f"Epoch {ep:4d} | loss {history['gravity'][-1]:.3e}")
  382. if use_early_stopping and best_state is not None:
  383. model.load_state_dict(best_state)
  384. history['best_epoch'] = best_epoch
  385. history['best_weighted_mse'] = best_weighted_mse
  386. else:
  387. history['best_epoch'] = len(history['gravity']) - 1
  388. history['best_weighted_mse'] = history['gravity'][-1] / cfg['gamma']
  389. return history
  390. def make_block_model(Nx, Ny, Nz, dx, dy, dz, rho_bg=0.0, rho_blk=400.0):
  391. m = torch.full((Nx, Ny, Nz), rho_bg)
  392. for i in range(7):
  393. z_idx = 1 + i
  394. y_start, y_end = 11 - i, 16 - i
  395. x_start, x_end = 7, 13
  396. if 0 <= z_idx < Nz:
  397. ys, ye = max(0, y_start), min(Ny, y_end)
  398. xs, xe = max(0, x_start), min(Nx, x_end)
  399. m[xs:xe, ys:ye, z_idx] = rho_blk
  400. return m.view(-1), m
  401. def get_block_boundaries(Nx, Ny, Nz):
  402. boundaries = []
  403. for i in range(7):
  404. z_idx = 1 + i
  405. y_start, y_end = 11 - i, 16 - i
  406. x_start, x_end = 7, 13
  407. if 0 <= z_idx < Nz:
  408. ys, ye = max(0, y_start), min(Ny, y_end)
  409. xs, xe = max(0, x_start), min(Nx, x_end)
  410. boundaries.append((xs, xe, ys, ye, z_idx))
  411. return boundaries
  412. def add_xy_outline(ax, boundaries, x, y, dx, dy, iz):
  413. boundary_for_z = next((b for b in boundaries if b[4] == iz), None)
  414. if boundary_for_z:
  415. xs, xe, ys, ye, _ = boundary_for_z
  416. rect = plt.Rectangle((x[xs] - dx / 2, y[ys] - dy / 2),
  417. (xe - xs) * dx,
  418. (ye - ys) * dy,
  419. edgecolor='white', facecolor='none', linewidth=2)
  420. ax.add_patch(rect)
  421. def add_xz_outline(ax, boundaries, x, z, dx, dz, iy):
  422. z_indices_in_slice = []
  423. x_range = None
  424. for xs, xe, ys, ye, z_idx in boundaries:
  425. if ys <= iy < ye:
  426. z_indices_in_slice.append(z_idx)
  427. if x_range is None:
  428. x_range = (xs, xe)
  429. if z_indices_in_slice and x_range:
  430. min_z_idx, max_z_idx = min(z_indices_in_slice), max(z_indices_in_slice)
  431. xs, xe = x_range
  432. rect = plt.Rectangle((x[xs] - dx / 2, z[min_z_idx] - dz / 2),
  433. (xe - xs) * dx,
  434. (max_z_idx - min_z_idx + 1) * dz,
  435. edgecolor='white', facecolor='none', linewidth=2)
  436. ax.add_patch(rect)
  437. def add_yz_outline(ax, boundaries, y, z, dy, dz, ix):
  438. for xs, xe, ys, ye, z_idx in boundaries:
  439. if xs <= ix < xe:
  440. rect = plt.Rectangle((y[ys] - dy / 2, z[z_idx] - dz / 2),
  441. (ye - ys) * dy,
  442. dz,
  443. edgecolor='white', facecolor='none', linewidth=2)
  444. ax.add_patch(rect)
  445. def style_axes(ax, xlabel, ylabel):
  446. ax.set_xlabel(xlabel, fontsize=LABEL_FONTSIZE)
  447. ax.set_ylabel(ylabel, fontsize=LABEL_FONTSIZE)
  448. ax.tick_params(labelsize=TICK_FONTSIZE)
  449. def encoding_label(name):
  450. return 'Positional Encoding' if name == 'positional' else 'Hash Encoding'
  451. def print_runtime_info(device):
  452. print('Runtime info:')
  453. print(f" Python = {sys.version.split()[0]}")
  454. print(f" NumPy = {np.__version__}")
  455. print(f" Matplotlib = {matplotlib.__version__}")
  456. print(f" PyTorch = {torch.__version__}")
  457. print(f" Device = {device}")
  458. print(f" CUDA = {torch.cuda.is_available()}")
  459. def run():
  460. set_seed(DATA_SEED)
  461. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  462. print_runtime_info(device)
  463. os.makedirs('plots', exist_ok=True)
  464. dx, dy, dz = DX, DY, DZ
  465. x = np.arange(0.0, X_MAX + dx, dx)
  466. y = np.arange(0.0, Y_MAX + dy, dy)
  467. z = np.arange(0.0, Z_MAX + dz, dz)
  468. Nx, Ny, Nz = len(x), len(y), len(z)
  469. Xc = x.astype(float)
  470. Yc = y.astype(float)
  471. Zc = z.astype(float)
  472. X3, Y3, Z3 = np.meshgrid(Xc, Yc, Zc, indexing='ij')
  473. grid_coords = np.stack([X3.ravel(), Y3.ravel(), Z3.ravel()], axis=1)
  474. c_mean = grid_coords.mean(axis=0, keepdims=True)
  475. c_std = grid_coords.std(axis=0, keepdims=True)
  476. coords_norm = (grid_coords - c_mean) / (c_std + 1e-12)
  477. coords_norm = torch.tensor(coords_norm, dtype=torch.float32, device=device, requires_grad=True)
  478. dz_half = dz / 2.0
  479. cell_grid = np.hstack([grid_coords, np.full((grid_coords.shape[0], 1), dz_half)])
  480. cell_grid = torch.tensor(cell_grid, dtype=torch.float32, device=device)
  481. XX, YY = np.meshgrid(x, y, indexing='ij')
  482. obs = np.column_stack([XX.ravel(), YY.ravel(), -np.ones(XX.size)])
  483. obs = torch.tensor(obs, dtype=torch.float32, device=device)
  484. print("Assembling sensitivity G ...")
  485. t0 = time.time()
  486. G = construct_sensitivity_matrix_G_torch(cell_grid, obs, dx, dy, device)
  487. G = G.clone().detach().requires_grad_(False)
  488. print(f"G shape = {tuple(G.shape)}, time = {time.time() - t0:.2f}s")
  489. rho_true_vec, rho_true_3d = make_block_model(Nx, Ny, Nz, dx, dy, dz, rho_bg=RHO_BG, rho_blk=RHO_BLK)
  490. rho_true_vec = rho_true_vec.to(device)
  491. with torch.no_grad():
  492. gz_true = (G @ rho_true_vec.unsqueeze(1)).squeeze(1)
  493. sigma = NOISE_LEVEL * gz_true.std()
  494. noise = sigma * torch.randn_like(gz_true)
  495. gz_obs = gz_true + noise
  496. Wd = 1.0 / sigma
  497. model_rng_state = capture_rng_state()
  498. cfg = dict(gamma=GAMMA, epochs=EPOCHS, lr=LR)
  499. cfg.update(
  500. use_early_stopping=USE_EARLY_STOPPING,
  501. early_stop_min_epochs=EARLY_STOP_MIN_EPOCHS,
  502. early_stop_patience=EARLY_STOP_PATIENCE,
  503. early_stop_target=EARLY_STOP_TARGET,
  504. early_stop_tol=EARLY_STOP_TOL,
  505. early_stop_overfit_patience=EARLY_STOP_OVERFIT_PATIENCE,
  506. )
  507. results = {}
  508. for encoding_type in ENCODING_ORDER:
  509. restore_rng_state(model_rng_state)
  510. print(f"\n▶ Encoding strategy: {encoding_type}")
  511. model = DensityContrastINR(
  512. encoding_type=encoding_type, hidden=HIDDEN, depth=DEPTH,
  513. rho_abs_max=RHO_ABS_MAX,
  514. **ENCODING_CONFIGS[encoding_type]
  515. ).to(device)
  516. opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'])
  517. hist = train_inr(model, opt, coords_norm, G, gz_obs, Wd, Nx, Ny, Nz, dx, dy, dz, cfg)
  518. with torch.no_grad():
  519. m_inv = model(coords_norm).view(-1)
  520. gz_pred = (G @ m_inv.unsqueeze(1)).squeeze(1)
  521. results[encoding_type] = {
  522. 'm_inv': m_inv.detach().cpu().numpy().reshape(Nx, Ny, Nz),
  523. 'gz_pred': gz_pred.detach().cpu().numpy(),
  524. 'hist': hist,
  525. 'rms_rho': torch.sqrt(torch.mean((m_inv - rho_true_vec.to(device)) ** 2)).item(),
  526. 'rms_gz': torch.sqrt(torch.mean((gz_pred - gz_obs) ** 2)).item() * 1e5,
  527. }
  528. def get_axes_coords():
  529. x1d = grid_coords[:, 0].reshape(Nx, Ny, Nz)[:, 0, 0]
  530. y1d = grid_coords[:, 1].reshape(Nx, Ny, Nz)[0, :, 0]
  531. z1d = grid_coords[:, 2].reshape(Nx, Ny, Nz)[0, 0, :]
  532. return x1d, y1d, z1d
  533. x1d, y1d, z1d = get_axes_coords()
  534. block_boundaries = get_block_boundaries(Nx, Ny, Nz)
  535. ix, iy, iz = Nx // 2, Ny // 2, min(Nz - 1, 5)
  536. tru = rho_true_3d.cpu().numpy()
  537. inv_pos = results['positional']['m_inv']
  538. inv_hash = results['hash']['m_inv']
  539. tru_max = tru.max()
  540. inv_max = INV_VMAX
  541. fig1, axes = plt.subplots(4, 3, figsize=(16, 20))
  542. x1d, y1d, z1d = get_axes_coords()
  543. # cell-edge limits
  544. x_edge_min, x_edge_max = x1d[0] - dx/2, x1d[-1] + dx/2
  545. y_edge_min, y_edge_max = y1d[0] - dy/2, y1d[-1] + dy/2
  546. z_edge_min, z_edge_max = z1d[0] - dz/2, z1d[-1] + dz/2
  547. # use edges for all extents
  548. extent_xy = [x_edge_min, x_edge_max, y_edge_min, y_edge_max]
  549. # for depth plots keep depth increasing downward by reversing z limits
  550. extent_xz = [x_edge_min, x_edge_max, z_edge_max, z_edge_min]
  551. extent_yz = [y_edge_min, y_edge_max, z_edge_max, z_edge_min]
  552. im = axes[0, 0].imshow(tru[:, :, iz].T, origin='lower', extent=extent_xy,
  553. aspect='auto', vmin=0, vmax=tru_max, cmap=CMAP)
  554. axes[0, 0].set_title(f"True Model XY @ z≈{z1d[iz]:.0f} m", fontsize=TITLE_FONTSIZE)
  555. cbar = fig1.colorbar(im, ax=axes[0, 0], label='kg/m³', fraction=0.046, pad=0.04)
  556. cbar.ax.tick_params(labelsize=TICK_FONTSIZE)
  557. im = axes[0, 1].imshow(tru[:, iy, :].T, origin='upper', extent=extent_xz,
  558. aspect='auto', vmin=0, vmax=tru_max, cmap=CMAP)
  559. axes[0, 1].set_title(f"True Model XZ @ y≈{y1d[iy]:.0f} m", fontsize=TITLE_FONTSIZE)
  560. im = axes[0, 2].imshow(tru[ix, :, :].T, origin='upper', extent=extent_yz,
  561. aspect='auto', vmin=0, vmax=tru_max, cmap=CMAP)
  562. axes[0, 2].set_title(f"True Model YZ @ x≈{x1d[ix]:.0f} m", fontsize=TITLE_FONTSIZE)
  563. model_rows = [
  564. ('positional', inv_pos, 1),
  565. ('hash', inv_hash, 2),
  566. ]
  567. for encoding_type, inv, row in model_rows:
  568. title = encoding_label(encoding_type)
  569. im = axes[row, 0].imshow(inv[:, :, iz].T, origin='lower', extent=extent_xy,
  570. aspect='auto', vmin=0, vmax=inv_max, cmap=CMAP)
  571. axes[row, 0].set_title(f"Recovered Model XY, {title}", fontsize=TITLE_FONTSIZE)
  572. add_xy_outline(axes[row, 0], block_boundaries, x, y, dx, dy, iz)
  573. cbar = fig1.colorbar(im, ax=axes[row, 0], label='kg/m³', fraction=0.046, pad=0.04)
  574. cbar.ax.tick_params(labelsize=TICK_FONTSIZE)
  575. im = axes[row, 1].imshow(inv[:, iy, :].T, origin='upper', extent=extent_xz,
  576. aspect='auto', vmin=0, vmax=inv_max, cmap=CMAP)
  577. axes[row, 1].set_title(f"Recovered Model XZ, {title}", fontsize=TITLE_FONTSIZE)
  578. add_xz_outline(axes[row, 1], block_boundaries, x, z, dx, dz, iy)
  579. im = axes[row, 2].imshow(inv[ix, :, :].T, origin='upper', extent=extent_yz,
  580. aspect='auto', vmin=0, vmax=inv_max, cmap=CMAP)
  581. axes[row, 2].set_title(f"Recovered Model YZ, {title}", fontsize=TITLE_FONTSIZE)
  582. add_yz_outline(axes[row, 2], block_boundaries, y, z, dy, dz, ix)
  583. for ax in [axes[0, 1], axes[0, 2], axes[1, 1], axes[1, 2], axes[2, 1], axes[2, 2]]:
  584. ax.set_aspect(1.0)
  585. for ax in [axes[0, 0], axes[1, 0], axes[2, 0]]:
  586. ax.set_aspect(1.0)
  587. for row in range(3):
  588. style_axes(axes[row, 0], 'x (m)', 'y (m)')
  589. style_axes(axes[row, 1], 'x (m)', 'Depth (m)')
  590. style_axes(axes[row, 2], 'y (m)', 'Depth (m)')
  591. def to_mgal(g):
  592. return 1e5 * g.detach().cpu().numpy()
  593. obs_mgal = to_mgal(gz_obs)
  594. res_pos_mgal = obs_mgal - 1e5 * results['positional']['gz_pred']
  595. res_hash_mgal = obs_mgal - 1e5 * results['hash']['gz_pred']
  596. obs_x = obs[:, 0].cpu().numpy()
  597. obs_y = obs[:, 1].cpu().numpy()
  598. axes[3, 0].plot(results['positional']['hist']['gravity'], label='Positional', color='tab:red')
  599. axes[3, 0].plot(results['hash']['hist']['gravity'], label='Hash', color='black')
  600. axes[3, 0].set_title('Training Convergence by Encoding', fontsize=TITLE_FONTSIZE)
  601. axes[3, 0].set_yscale('log')
  602. axes[3, 0].grid(True, which='both', ls='--', alpha=0.3)
  603. axes[3, 0].legend(fontsize=LEGEND_FONTSIZE)
  604. style_axes(axes[3, 0], 'Epoch', 'Loss = mean((residual / sigma)^2)')
  605. vmax_res = max(np.abs(res_pos_mgal).max(), np.abs(res_hash_mgal).max())
  606. sc = axes[3, 1].scatter(obs_x, obs_y, c=res_pos_mgal, s=80, cmap=CMAP,
  607. vmin=-vmax_res, vmax=vmax_res, marker='o', edgecolors='none')
  608. axes[3, 1].set_title(
  609. f"Data Misfit, Positional Encoding (RMS = {np.sqrt(np.mean(res_pos_mgal**2)):.3f} mGal)",
  610. fontsize=TITLE_FONTSIZE,
  611. )
  612. cbar = fig1.colorbar(sc, ax=axes[3, 1], fraction=0.046, pad=0.04)
  613. cbar.ax.tick_params(labelsize=TICK_FONTSIZE)
  614. sc = axes[3, 2].scatter(obs_x, obs_y, c=res_hash_mgal, s=80, cmap=CMAP,
  615. vmin=-vmax_res, vmax=vmax_res, marker='o', edgecolors='none')
  616. axes[3, 2].set_title(
  617. f"Data Misfit, Hash Encoding (RMS = {np.sqrt(np.mean(res_hash_mgal**2)):.3f} mGal)",
  618. fontsize=TITLE_FONTSIZE,
  619. )
  620. cbar = fig1.colorbar(sc, ax=axes[3, 2], fraction=0.046, pad=0.04)
  621. cbar.ax.tick_params(labelsize=TICK_FONTSIZE)
  622. for ax in [axes[3, 1], axes[3, 2]]:
  623. style_axes(ax, 'x (m)', 'y (m)')
  624. ax.set_aspect('equal')
  625. fig1.tight_layout()
  626. fig1.savefig('plots/BlockModel_EncodingComparison.png', dpi=FIG_DPI)
  627. plt.close(fig1)
  628. for encoding_type in ENCODING_ORDER:
  629. label = encoding_label(encoding_type)
  630. hist = results[encoding_type]['hist']
  631. print(
  632. f"{label} RMS density-contrast error ≈ {results[encoding_type]['rms_rho']:.2f} kg/m^3"
  633. )
  634. print(
  635. f"{label} RMS data misfit ≈ {results[encoding_type]['rms_gz']:.3f} mGal | "
  636. f"best epoch = {hist['best_epoch']} | "
  637. f"best loss = {hist['best_weighted_mse']:.3f}"
  638. )
  639. if __name__ == '__main__':
  640. run()

001-EncodingComparisons_blocky.py at commit 3d5a05d, under MIT · at the source

Overview

Authors: Pankaj K Mishra1, Sanni Laaksonen1, Jochen Kamm1, Anand Singh2
ORCID iDs: Pankaj K Mishra
  1. Geological Survey of Finland,Vuorimiehentie 5, Espoo, Finland
  2. Indian Institute of Technology Bombay,Mumbai, India
Journal: Scientific reports, volume 16, issue 1, article 25630
Dates: received 14 October 2025; accepted 28 May 2026; published online 4 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41598-026-55960-5 · PMID 42243533 · PMCID PMC13478334 · OpenAlex W4416054661
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: systems (subfield)
Methods: Machine learning
Keywords: Neural fields, Gravity, Physics-based deep learning, Scientific machine learning, Inversion, Mathematics and computing, Solid Earth sciences
Topic: Geophysical and Geoelectrical Methods (Geophysics, Earth and Planetary Sciences), according to OpenAlex
Citations: not cited yet (Europe PMC); 39 references in the paper

Abstract

Inversion of gravity data is an important method for investigating subsurface density variations relevant to mineral exploration, geothermal assessment, carbon storage, natural hydrogen, groundwater resources, and tectonic evolution. Here we present a scientific machine-learning approach for three-dimensional gravity inversion that represents subsurface density as a continuous field using an implicit neural representation (INR). The method trains a deep neural network directly through a physics-based forward-model loss, mapping spatial coordinates to a continuous density field without predefined meshes or discretisation. Spatial encoding enhances the network’s capacity to capture sharp contrasts and short-wavelength features that conventional coordinate-based networks tend to oversmooth due to spectral bias. We demonstrate the approach on synthetic examples including smooth models, representing realistic geological complexity, and a dipping block model to assess recovery of structures at different depths. The INR framework reconstructs detailed structure and geologically plausible boundaries without explicit regularisation or depth weighting, while reducing the number of inversion parameters as the problem size grows bigger. These results highlight the potential of implicit representations to enable scalable, flexible, and interpretable large-scale geophysical inversion. This framework could generalise to other geophysical methods and for joint/multiphysics inversion.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repositories

Its files are read in the Code ↔ Paper reader above, with 15 matches between paragraphs and lines of code.

Zenodo 19440024

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: “Data availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (12 files), NumPy (12 files), PyTorch (12 files)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
14 files

pankajkmishra/inrgravity3dinv

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 3d5a05d1b1c59d909be7dc765f888ea5f2247849, 13 April 2026
Languages: Python (12)
Size: 16 files, 12 scripts
Software Heritage: not archived
Found in: the Zenodo archive record
Holds: README, license file, environment (requirements.txt)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: Matplotlib (12 files), NumPy (12 files), PyTorch (12 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
14 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 24 scripts, each with its path and the digest of its content;
  • 15 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data availability

The codes for reproducing the results can be found at https://zenodo.org/records/19440024.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 7 keywords, 1 funder, 23 references.

Cite

This paper

Mishra, P. K., Laaksonen, S., Kamm, J., & Singh, A. (2026). Three-dimensional inversion of gravity data using implicit neural representations and scientific machine learning. Scientific reports, 16(1), 25630. https://doi.org/10.1038/s41598-026-55960-5

BibTeX

@article{mishra2026three,
author = {Mishra, Pankaj K and Laaksonen, Sanni and Kamm, Jochen and Singh, Anand},
title = {{Three-dimensional inversion of gravity data using implicit neural representations and scientific machine learning}},
journal = {Scientific reports},
year = {2026},
month = jun,
volume = {16},
number = {1},
pages = {25630},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/s41598-026-55960-5},
url = {https://doi.org/10.1038/s41598-026-55960-5},
pmid = {42243533},
pmcid = {PMC13478334}
}

RIS

TY - JOUR
AU - Mishra, Pankaj K
AU - Laaksonen, Sanni
AU - Kamm, Jochen
AU - Singh, Anand
TI - Three-dimensional inversion of gravity data using implicit neural representations and scientific machine learning
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/06/04
VL - 16
IS - 1
SP - 25630
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-55960-5
UR - https://doi.org/10.1038/s41598-026-55960-5
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-55960-5",
"type": "article-journal",
"title": "Three-dimensional inversion of gravity data using implicit neural representations and scientific machine learning",
"container-title": "Scientific reports",
"author": [
{
"family": "Mishra",
"given": "Pankaj K"
},
{
"family": "Laaksonen",
"given": "Sanni"
},
{
"family": "Kamm",
"given": "Jochen"
},
{
"family": "Singh",
"given": "Anand"
}
],
"container-title-short": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "25630",
"DOI": "10.1038/s41598-026-55960-5",
"PMID": "42243533",
"PMCID": "PMC13478334",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-55960-5",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
4
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s41467-026-74092-y [code]
High-rate phase association with travel time neural fields.
Journal: Nature communications
In common: PyTorch, Matplotlib, NumPy, 1 reference
[2] doi:10.3390/jimaging12040170
ARS-GS: Anisotropic Reflective Spherical 3D Gaussian Splatting.
Journal: Journal of imaging
In common: 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.