Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions gempy_engine/API/interp_single/_aux_faults_ops.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
6 changes: 6 additions & 0 deletions gempy_engine/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'
2 changes: 2 additions & 0 deletions gempy_engine/core/data/interp_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
1 change: 1 addition & 0 deletions gempy_engine/core/data/output/blocks_value_type.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
65 changes: 60 additions & 5 deletions gempy_engine/core/data/raw_arrays_solution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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)

Expand All @@ -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

1 change: 1 addition & 0 deletions gempy_engine/core/data/scalar_field_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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([
Expand All @@ -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():
Expand Down Expand Up @@ -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
Expand All @@ -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):
Expand Down