From e8f580ce6ba3953da2e568bf808ccc1ecdaadc4b Mon Sep 17 00:00:00 2001 From: Miguel de la Varga Date: Tue, 25 Aug 2026 14:12:55 +0200 Subject: [PATCH] [ENH] Add support for finite-fault scalar field handling Introduces finite-fault scalar fields in interpolation workflows with conditional inclusion based on environment variables. Adds new data structures, methods for scalar field computation, and integration with stack outputs and raw arrays. Updates fault stack wiring, test cases for scalar validation, and introduces `finite_fault_scalar` value type. Refactors relevant methods to handle scalar field serialization and dense grid operations. # Conflicts: # tests/test_common/test_api/test_faults/test_finite_fault_stack_wiring.py --- .../API/interp_single/_aux_faults_ops.py | 3 + gempy_engine/config.py | 6 ++ gempy_engine/core/data/interp_output.py | 2 + .../core/data/output/blocks_value_type.py | 1 + gempy_engine/core/data/raw_arrays_solution.py | 65 +++++++++++++++++-- gempy_engine/core/data/scalar_field_output.py | 1 + .../test_finite_fault_stack_wiring.py | 62 +++++++++++++++++- 7 files changed, 133 insertions(+), 7 deletions(-) diff --git a/gempy_engine/API/interp_single/_aux_faults_ops.py b/gempy_engine/API/interp_single/_aux_faults_ops.py index 93907fb4..87fb9278 100644 --- a/gempy_engine/API/interp_single/_aux_faults_ops.py +++ b/gempy_engine/API/interp_single/_aux_faults_ops.py @@ -1,6 +1,7 @@ import numpy as np from gempy_engine.core.backend_tensor import BackendTensor +from gempy_engine.config import include_raw_scalar_fields from gempy_engine.core.data.kernel_classes.faults import FaultsData from gempy_engine.core.data.options import InterpolationOptions from gempy_engine.core.data.scalar_field_output import ScalarFieldOutput @@ -79,4 +80,6 @@ def _modify_faults_values_output( finite_fault_scalar_np, dtype=shifted_vals.dtype, ) + if include_raw_scalar_fields(): + output.finite_fault_scalar = finite_fault_scalar return shifted_vals * finite_fault_scalar diff --git a/gempy_engine/config.py b/gempy_engine/config.py index 223679e0..cd3f3070 100644 --- a/gempy_engine/config.py +++ b/gempy_engine/config.py @@ -43,3 +43,9 @@ class DualContouringOverlap(Flag): is_tensorflow_installed = find_spec("tensorflow") is not None is_pytorch_installed = find_spec("torch") is_pykeops_installed = find_spec("pykeops") is not None + + +def include_raw_scalar_fields() -> bool: + return os.getenv('ONLY_LITH_SOLUTION', 'False') != 'True' or os.getenv( + 'SET_RAW_SCALAR_FIELDS_IN_SOLUTION', 'False' + ) == 'True' diff --git a/gempy_engine/core/data/interp_output.py b/gempy_engine/core/data/interp_output.py index 746f886d..c52ddad8 100644 --- a/gempy_engine/core/data/interp_output.py +++ b/gempy_engine/core/data/interp_output.py @@ -168,6 +168,8 @@ def get_block_from_value_type(self, value_type: ValueType, slice_: slice): block = self.values_block[0] case ValueType.scalar: block = self.exported_fields.scalar_field + case ValueType.finite_fault_scalar: + block = self.scalar_fields.finite_fault_scalar case ValueType.squeeze_mask: block = self.squeezed_mask_array case ValueType.mask_component: diff --git a/gempy_engine/core/data/output/blocks_value_type.py b/gempy_engine/core/data/output/blocks_value_type.py index e208b985..ac427a81 100644 --- a/gempy_engine/core/data/output/blocks_value_type.py +++ b/gempy_engine/core/data/output/blocks_value_type.py @@ -7,6 +7,7 @@ class ValueType(enum.Enum): litho_faults_block = enum.auto() scalar = enum.auto() + finite_fault_scalar = enum.auto() squeeze_mask = enum.auto() mask_component = enum.auto() diff --git a/gempy_engine/core/data/raw_arrays_solution.py b/gempy_engine/core/data/raw_arrays_solution.py index 3ad75951..b7f7706c 100644 --- a/gempy_engine/core/data/raw_arrays_solution.py +++ b/gempy_engine/core/data/raw_arrays_solution.py @@ -5,6 +5,8 @@ import numpy as np +from gempy_engine.config import include_raw_scalar_fields + from .transforms import Transform from ..backend_tensor import BackendTensor from .interp_output import InterpOutput @@ -32,6 +34,8 @@ class BlockSolutionType(enum.Enum): litho_faults_block: np.ndarray = field(default_factory=lambda: np.empty(0)) scalar_field_matrix: np.ndarray = field(default_factory=lambda: np.empty(0)) + finite_fault_scalar_field_matrix: np.ndarray = field(default_factory=lambda: np.empty(0)) + finite_fault_stack_indices: np.ndarray = field(default_factory=lambda: np.empty(0, dtype=np.int64)) block_matrix: np.ndarray = field(default_factory=lambda: np.empty(0)) mask_matrix: np.ndarray = field(default_factory=lambda: np.empty(0)) mask_matrix_squeezed: np.ndarray = field(default_factory=lambda: np.empty(0)) @@ -182,10 +186,6 @@ def ___get_regular_grid_values_for_all_structural_groups(octree_output: list[Oct level=None, value_type=ValueType.litho_faults_block )).astype("int32").ravel() - raw_arrays_solution.scalar_field_matrix = ___get_regular_grid_values_for_all_structural_groups( - octree_output=octrees_output, - scalar_type=ValueType.scalar - ) raw_arrays_solution.mask_matrix = ___get_regular_grid_values_for_all_structural_groups( octree_output=octrees_output, scalar_type=ValueType.mask_component @@ -195,6 +195,13 @@ def ___get_regular_grid_values_for_all_structural_groups(octree_output: list[Oct scalar_type=ValueType.squeeze_mask ) + if include_raw_scalar_fields(): + raw_arrays_solution.scalar_field_matrix = ___get_regular_grid_values_for_all_structural_groups( + octree_output=octrees_output, + scalar_type=ValueType.scalar + ) + _fill_finite_fault_scalar_fields_with_octree_output(octrees_output, raw_arrays_solution) + lith_block_temp = BackendTensor.t.to_numpy(get_regular_grid_value_for_level( octree_list=octrees_output, level=None, @@ -224,7 +231,9 @@ def ___get_regular_grid_values_for_all_structural_groups(stacks_output: list[Int raw_arrays_solution.block_matrix = ___get_regular_grid_values_for_all_structural_groups(stacks_output, ValueType.values_block) raw_arrays_solution.fault_block = collapsed_output.get_block_from_value_type(ValueType.faults_block, slice_=collapsed_output.grid.dense_grid_slice).astype("int8").ravel() raw_arrays_solution.litho_faults_block = collapsed_output.get_block_from_value_type(ValueType.litho_faults_block, slice_=collapsed_output.grid.dense_grid_slice).astype("int32").ravel() - raw_arrays_solution.scalar_field_matrix = ___get_regular_grid_values_for_all_structural_groups(stacks_output, ValueType.scalar) + if include_raw_scalar_fields(): + raw_arrays_solution.scalar_field_matrix = ___get_regular_grid_values_for_all_structural_groups(stacks_output, ValueType.scalar) + _fill_finite_fault_scalar_fields_with_dense_grid(stacks_output, raw_arrays_solution) raw_arrays_solution.mask_matrix = ___get_regular_grid_values_for_all_structural_groups(stacks_output, ValueType.mask_component) raw_arrays_solution.mask_matrix_squeezed = ___get_regular_grid_values_for_all_structural_groups(stacks_output, ValueType.squeeze_mask) @@ -233,3 +242,49 @@ def ___get_regular_grid_values_for_all_structural_groups(stacks_output: list[Int raw_arrays_solution.lith_block = lith_block_temp +def _finite_fault_stack_indices(stacks_output: list[InterpOutput]) -> np.ndarray: + return np.asarray([ + i for i, output in enumerate(stacks_output) + if output.scalar_fields.finite_fault_scalar is not None + ], dtype=np.int64) + + +def _fill_finite_fault_scalar_fields_with_octree_output( + octrees_output: list[OctreeLevel], + raw_arrays_solution: RawArraysSolution, +) -> None: + stack_indices = _finite_fault_stack_indices(octrees_output[0].outputs) + if stack_indices.size == 0: + return + + fields = [ + get_regular_grid_value_for_level( + octree_list=octrees_output, + level=None, + value_type=ValueType.finite_fault_scalar, + scalar_n=int(stack_index), + ).ravel() + for stack_index in stack_indices + ] + raw_arrays_solution.finite_fault_scalar_field_matrix = BackendTensor.t.to_numpy(BackendTensor.t.stack(fields)) + raw_arrays_solution.finite_fault_stack_indices = stack_indices + + +def _fill_finite_fault_scalar_fields_with_dense_grid( + stacks_output: list[InterpOutput], + raw_arrays_solution: RawArraysSolution, +) -> None: + stack_indices = _finite_fault_stack_indices(stacks_output) + if stack_indices.size == 0: + return + + fields = [ + stacks_output[int(stack_index)].get_block_from_value_type( + ValueType.finite_fault_scalar, + slice_=stacks_output[int(stack_index)].grid.dense_grid_slice, + ) + for stack_index in stack_indices + ] + raw_arrays_solution.finite_fault_scalar_field_matrix = np.vstack(fields) + raw_arrays_solution.finite_fault_stack_indices = stack_indices + diff --git a/gempy_engine/core/data/scalar_field_output.py b/gempy_engine/core/data/scalar_field_output.py index 92642f83..f50f77e5 100644 --- a/gempy_engine/core/data/scalar_field_output.py +++ b/gempy_engine/core/data/scalar_field_output.py @@ -20,6 +20,7 @@ class ScalarFieldOutput: values_block: Optional[np.ndarray] #: Final values ignoring unconformities _values_block: Optional[np.ndarray] = dataclasses.field(init=False, repr=False) + finite_fault_scalar: Optional[np.ndarray] = None @property def values_block(self) -> Optional[np.ndarray]: diff --git a/tests/test_common/test_api/test_faults/test_finite_fault_stack_wiring.py b/tests/test_common/test_api/test_faults/test_finite_fault_stack_wiring.py index 785e1cb1..9d5c7967 100644 --- a/tests/test_common/test_api/test_faults/test_finite_fault_stack_wiring.py +++ b/tests/test_common/test_api/test_faults/test_finite_fault_stack_wiring.py @@ -18,7 +18,8 @@ from gempy_engine.core.data.stacks_structure import StacksStructure -def test_modify_fault_values_applies_projected_taper(): +def test_modify_fault_values_applies_projected_taper(monkeypatch): + monkeypatch.setenv("SET_RAW_SCALAR_FIELDS_IN_SOLUTION", "True") finite_fault = FiniteFault(center=(0.0, 0.0, 0.0), strike_radius=1.0, dip_radius=1.0) fault_input = FaultsData.from_user_input(thickness=None, finite_fault=finite_fault) points = np.array([ @@ -45,6 +46,52 @@ def test_modify_fault_values_applies_projected_taper(): tapered_values = _modify_faults_values_output(fault_input, output, points) assert np.allclose(tapered_values, [[0.0, 2.0, 0.0]]) + assert np.allclose(output.finite_fault_scalar, [0.0, 1.0, 0.0]) + + +@pytest.mark.parametrize("use_gpu", [False, True]) +def test_modify_fault_values_preserves_pytorch_backend(use_gpu): + torch = pytest.importorskip("torch") + if use_gpu and not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + BackendTensor._change_backend(AvailableBackends.PYTORCH, use_gpu=use_gpu) + dtype = BackendTensor.dtype_obj + device = BackendTensor.device + + finite_fault = FiniteFault( + center=(0.0, 0.0, 0.0), + strike_radius=1.0, + dip_radius=1.0, + taper=TaperType.SPLINE, + ) + fault_input = FaultsData.from_user_input(thickness=None, finite_fault=finite_fault) + points = torch.tensor([ + [2.0, 0.0, 2.0], + [0.0, 0.0, 2.0], + [1.0, 0.0, 2.0], + ], dtype=dtype, device=device) + exported_fields = ExportedFields( + _scalar_field=torch.full((3,), 2.0, dtype=dtype, device=device), + _gx_field=torch.zeros(3, dtype=dtype, device=device), + _gy_field=torch.zeros(3, dtype=dtype, device=device), + _gz_field=torch.ones(3, dtype=dtype, device=device), + _grid_size=3, + _scalar_field_at_surface_points=torch.tensor([0.0], dtype=dtype, device=device), + ) + output = ScalarFieldOutput( + weights=None, + grid=None, + exported_fields=exported_fields, + stack_relation=StackRelationType.FAULT, + values_block=torch.tensor([[1.0, 3.0, 3.0]], dtype=dtype, device=device), + ) + + tapered_values = _modify_faults_values_output(fault_input, output, points) + + assert isinstance(tapered_values, torch.Tensor) + assert tapered_values.dtype == dtype + assert tapered_values.device == points.device + assert torch.allclose(tapered_values, torch.tensor([[0.0, 2.0, 0.0]], dtype=dtype, device=device)) def test_finite_fault_gradient_options_do_not_mutate_input(): @@ -87,7 +134,9 @@ def test_finite_fault_stack_is_isolated_before_dependent_stack(): assert all(0 not in chunk or chunk == [0] for chunk in chunks) -def test_finite_fault_is_wired_into_dependent_stack(one_fault_model): +def test_finite_fault_is_wired_into_dependent_stack(one_fault_model, monkeypatch): + monkeypatch.setenv("ONLY_LITH_SOLUTION", "True") + monkeypatch.setenv("SET_RAW_SCALAR_FIELDS_IN_SOLUTION", "True") interpolation_input, data_descriptor, options = copy.deepcopy(one_fault_model) options.evaluation_options.number_octree_levels = 1 options.evaluation_options.mesh_extraction = False @@ -114,13 +163,22 @@ def test_finite_fault_is_wired_into_dependent_stack(one_fault_model): finite = compute_model(interpolation_input, options, data_descriptor) finite_fault_output = finite.octrees_output[0].outputs[0].exported_fields + finite_fault_scalar = finite.octrees_output[0].outputs[0].scalar_fields.finite_fault_scalar baseline_dependent = baseline.octrees_output[0].outputs[2].exported_fields.scalar_field finite_dependent = finite.octrees_output[0].outputs[2].exported_fields.scalar_field assert finite_fault_output.gx_field is not None assert finite_fault_output.gy_field is not None assert finite_fault_output.gz_field is not None + assert finite_fault_scalar is not None assert options.evaluation_options.compute_scalar_gradient is False assert not np.allclose(finite_dependent, baseline_dependent) + assert finite.raw_arrays.block_matrix.size == 0 + assert finite.raw_arrays.mask_matrix.size == 0 + assert finite.raw_arrays.scalar_field_matrix.shape[0] == 3 + assert np.array_equal(finite.raw_arrays.finite_fault_stack_indices, [0]) + assert finite.raw_arrays.finite_fault_scalar_field_matrix.shape == (1, finite.raw_arrays.lith_block.size) + assert np.all((finite.raw_arrays.finite_fault_scalar_field_matrix >= 0) & + (finite.raw_arrays.finite_fault_scalar_field_matrix <= 1)) def test_finite_fault_flat_stack_matches_serial(one_fault_model, monkeypatch):