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):