diff --git a/packages/map2loop/src/map2loop/project.py b/packages/map2loop/src/map2loop/project.py index 3cf355fdc..6093cd242 100644 --- a/packages/map2loop/src/map2loop/project.py +++ b/packages/map2loop/src/map2loop/project.py @@ -8,7 +8,14 @@ from .thickness_calculator import InterpolatedStructure, ThicknessCalculator from .throw_calculator import ThrowCalculator, ThrowCalculatorAlpha from .fault_orientation import FaultOrientation -from .sorter import Sorter, SorterAgeBased, SorterAlpha, SorterUseNetworkX, SorterUseHint +from .sorter import ( + Sorter, + SorterAgeBased, + SorterAlpha, + SorterHierarchical, + SorterUseNetworkX, + SorterUseHint, +) from .stratigraphic_column import StratigraphicColumn from .deformation_history import DeformationHistory from .topology import Topology @@ -603,21 +610,25 @@ def calculate_stratigraphic_order(self, take_best=False): self.stratigraphic_column.column = column else: logger.info(f'Calculating stratigraphic column using sorter {self.sorter.sorter_label}') - # Update sorter with current data based on what it needs - if hasattr(self.sorter, 'unit_relationships') and self.sorter.unit_relationships is None: - self.sorter.unit_relationships = self.topology.get_unit_unit_relationships() - if hasattr(self.sorter, 'contacts') and self.sorter.contacts is None: - self.sorter.contacts = self.contact_extractor.contacts - if hasattr(self.sorter, 'geology_data') and self.sorter.geology_data is None: - self.sorter.geology_data = self.map_data.get_map_data(Datatype.GEOLOGY) - if hasattr(self.sorter, 'structure_data') and self.sorter.structure_data is None: - self.sorter.structure_data = self.map_data.get_map_data(Datatype.STRUCTURE) - if hasattr(self.sorter, 'dtm_data') and self.sorter.dtm_data is None: - self.sorter.dtm_data = self.map_data.get_map_data(Datatype.DTM) - if hasattr(self.sorter, 'min_age_column') and self.sorter.min_age_column is None: - self.sorter.min_age_column = self.stratigraphic_column.get_min_age_column() - if hasattr(self.sorter, 'max_age_column') and self.sorter.max_age_column is None: - self.sorter.max_age_column = self.stratigraphic_column.get_max_age_column() + # Update sorter with current data based on what it needs. A hierarchical + # sorter gives the data to the sorter that it uses at each level. + sorter = self.sorter + if isinstance(sorter, SorterHierarchical): + sorter = sorter.sorter + if hasattr(sorter, 'unit_relationships') and sorter.unit_relationships is None: + sorter.unit_relationships = self.topology.get_unit_unit_relationships() + if hasattr(sorter, 'contacts') and sorter.contacts is None: + sorter.contacts = self.contact_extractor.contacts + if hasattr(sorter, 'geology_data') and sorter.geology_data is None: + sorter.geology_data = self.map_data.get_map_data(Datatype.GEOLOGY) + if hasattr(sorter, 'structure_data') and sorter.structure_data is None: + sorter.structure_data = self.map_data.get_map_data(Datatype.STRUCTURE) + if hasattr(sorter, 'dtm_data') and sorter.dtm_data is None: + sorter.dtm_data = self.map_data.get_map_data(Datatype.DTM) + if hasattr(sorter, 'min_age_column') and sorter.min_age_column is None: + sorter.min_age_column = self.stratigraphic_column.get_min_age_column() + if hasattr(sorter, 'max_age_column') and sorter.max_age_column is None: + sorter.max_age_column = self.stratigraphic_column.get_max_age_column() self.stratigraphic_column.column = self.sorter.sort( self.stratigraphic_column.stratigraphicUnits, diff --git a/packages/map2loop/src/map2loop/sorter.py b/packages/map2loop/src/map2loop/sorter.py index 123d357ce..c378e87ea 100644 --- a/packages/map2loop/src/map2loop/sorter.py +++ b/packages/map2loop/src/map2loop/sorter.py @@ -1,9 +1,10 @@ from abc import ABC, abstractmethod +import copy import beartype import pandas import numpy as np import math -from typing import Union, Optional, List +from typing import Callable, Dict, Union, Optional, List from map2loop.topology import Topology import geopandas from osgeo import gdal @@ -647,3 +648,349 @@ def sort(self, units: pandas.DataFrame) -> list: order = list(nx.dfs_preorder_nodes(dd, source=list(dd.nodes())[0])) logger.info(','.join(order)) return order + + +def _is_missing(value) -> bool: + """ + Check if a group or supergroup value is empty + + Args: + value: the value from the group or supergroup column + + Returns: + bool: True if the value is None, NaN, an empty string or "None"/"nan" + """ + if value is None: + return True + if isinstance(value, float) and math.isnan(value): + return True + return str(value).strip() in ("", "None", "nan") + + +def relabel_contacts( + contacts: pandas.DataFrame, + labels: Dict[str, str], + unitname1_column: str = 'UNITNAME_1', + unitname2_column: str = 'UNITNAME_2', +) -> pandas.DataFrame: + """ + Change the unit names of the contacts to labels (for example the group of each unit) + + A contact with a unit that is not in labels is removed. A contact between two units + with the same label is removed. The lengths of the contacts between the same two + labels are added together. + + Args: + contacts (pandas.DataFrame): the contacts, with a 'length' column or a geometry + labels (dict): the label of each unit name + unitname1_column (str): the name of the column with the first unit name + unitname2_column (str): the name of the column with the second unit name + + Returns: + pandas.DataFrame: the contacts between the labels, with the columns + [unitname1_column, unitname2_column, 'length'] + """ + columns = [unitname1_column, unitname2_column, 'length'] + if contacts is None or len(contacts) == 0: + return pandas.DataFrame(columns=columns) + if 'length' in contacts.columns: + lengths = contacts['length'].astype(float) + elif isinstance(contacts, geopandas.GeoDataFrame): + lengths = contacts.geometry.length + else: + lengths = pandas.Series(1.0, index=contacts.index) + label1 = contacts[unitname1_column].map(labels) + label2 = contacts[unitname2_column].map(labels) + keep = label1.notna() & label2.notna() & (label1 != label2) + pairs = pandas.DataFrame( + {unitname1_column: label1[keep], unitname2_column: label2[keep], 'length': lengths[keep]} + ) + if len(pairs) == 0: + return pandas.DataFrame(columns=columns) + # (a, b) and (b, a) are the same contact + swap = pairs[unitname1_column] > pairs[unitname2_column] + pairs.loc[swap, [unitname1_column, unitname2_column]] = pairs.loc[ + swap, [unitname2_column, unitname1_column] + ].values + return pairs.groupby([unitname1_column, unitname2_column], as_index=False)['length'].sum() + + +def relabel_unit_relationships( + unit_relationships: pandas.DataFrame, labels: Dict[str, str] +) -> pandas.DataFrame: + """ + Change the unit names of the unit relationships to labels (for example the group of each unit) + + A relationship with a unit that is not in labels is removed. A relationship between + two units with the same label is removed. The direction of each relationship is kept. + + Args: + unit_relationships (pandas.DataFrame): the relationships, with the columns + 'UNITNAME_1' and 'UNITNAME_2' + labels (dict): the label of each unit name + + Returns: + pandas.DataFrame: the relationships between the labels + """ + columns = ['UNITNAME_1', 'UNITNAME_2'] + if unit_relationships is None or len(unit_relationships) == 0: + return pandas.DataFrame(columns=columns) + label1 = unit_relationships['UNITNAME_1'].map(labels) + label2 = unit_relationships['UNITNAME_2'].map(labels) + keep = label1.notna() & label2.notna() & (label1 != label2) + relationships = pandas.DataFrame({'UNITNAME_1': label1[keep], 'UNITNAME_2': label2[keep]}) + return relationships.drop_duplicates().reset_index(drop=True) + + +class SorterHierarchical(Sorter): + """ + Sorter class which keeps the units of each supergroup and of each group together + + This sorter uses a different sorter (for example SorterAlpha) at each level: + 1. It sorts the supergroups. + 2. It sorts the groups in each supergroup. + 3. It sorts the units in each group. + + To sort the supergroups (or the groups), the data of the sorter (contacts, unit + relationships and geology) is changed so that each supergroup (or group) is one + unit. At each step, the sorter uses only the data of the units in that step. + Thus each supergroup has its own stratigraphic order, and the contacts between + two supergroups only set the order of the two supergroups. + + A group with no supergroup is one item at the supergroup level, the same as a + supergroup. A unit with no group (and no supergroup) is one item at that level. + """ + + required_arguments: List[str] = ['sorter'] + + def __init__( + self, + *, + sorter: Sorter, + group_column: Optional[str] = 'group', + supergroup_column: Optional[str] = 'supergroup', + postprocess: Optional[Callable[[list, Dict[str, str]], list]] = None, + ): + """ + Initialiser for hierarchical sorter + + Args: + sorter (Sorter): the sorter to use at each level + group_column (str, optional): the column of the units with the group of each + unit. Set to None to not use groups. Defaults to 'group'. + supergroup_column (str, optional): the column of the units with the + supergroup of each unit. Set to None to not use supergroups. + Defaults to 'supergroup'. + postprocess (callable, optional): a function that is applied to the result + of each step, postprocess(order, labels) -> order. labels is the label of + each unit name in that step (the group, the supergroup or the unit name). + For example, use it to repair the closed route of a travelling salesman + sorter. Defaults to None. + """ + super().__init__() + if isinstance(sorter, SorterHierarchical): + raise TypeError("sorter must not be a SorterHierarchical") + self.sorter = sorter + self.group_column = group_column + self.supergroup_column = supergroup_column + self.postprocess = postprocess + self.unit_name_column = getattr(sorter, 'unit_name_column', None) or 'name' + self.sorter_label = f"SorterHierarchical({sorter.sorter_label})" + + def sort(self, units: pandas.DataFrame) -> list: + """ + Execute sorter method takes unit data and returns the sorted unit names based on this algorithm. + + Args: + units (pandas.DataFrame): the data frame to sort + + Returns: + list: the sorted unit names + """ + if self.unit_name_column not in units.columns: + raise ValueError(f"Column {self.unit_name_column} must be present in units DataFrame") + units = units.drop_duplicates(subset=[self.unit_name_column]).reset_index(drop=True) + levels = [ + column + for column in (self.supergroup_column, self.group_column) + if column and column in units.columns + ] + if not levels: + logger.warning( + f"{self.sorter_label}: no group or supergroup column in the units, " + "so the units are sorted with no hierarchy" + ) + order = self._sort_level(units, levels) + logger.info("Stratigraphic order calculated using hierarchical sorting") + logger.info(','.join(order)) + return order + + def _sort_level(self, units: pandas.DataFrame, levels: List[str]) -> list: + """ + Sort the units with the first level in levels, then each part with the next levels + + Args: + units (pandas.DataFrame): the units to sort + levels (list): the group columns to use, from the highest level + + Returns: + list: the sorted unit names + """ + names = list(units[self.unit_name_column]) + if not levels: + return self._sort_step(self._units_table(units), {name: name for name in names}) + column, next_levels = levels[0], levels[1:] + + # Find the label of each unit at this level. A unit with no value in this + # column uses the value of the next level (or its name), so that a group + # with no supergroup is one item at the supergroup level. + level_values = { + str(value).strip() for value in units[column] if not _is_missing(value) + } + labels = {} + parts = {} + for _, row in units.iterrows(): + name = row[self.unit_name_column] + label_column, label = None, name + for level in levels: + if not _is_missing(row[level]): + label_column, label = level, str(row[level]).strip() + break + if label_column != column and label in level_values: + label = f"{label} ({label_column or 'unit'})" + labels[name] = label + parts.setdefault(label, []).append(name) + if len(parts) == 1: + return self._sort_level(units, next_levels) + + label_order = self._sort_step(self._label_table(units, labels), labels) + order = [] + for label in label_order: + part = units[units[self.unit_name_column].isin(parts[label])] + order += self._sort_level(part, next_levels) + return order + + def _units_table(self, units: pandas.DataFrame) -> pandas.DataFrame: + """ + Make the units table for the sorter, for the units in one step + + Args: + units (pandas.DataFrame): the units + + Returns: + pandas.DataFrame: the units, with 'layerId' equal to the index + """ + hierarchy_columns = [ + column for column in (self.group_column, self.supergroup_column) if column + ] + table = units.drop(columns=[c for c in hierarchy_columns if c in units.columns]) + table = table.reset_index(drop=True) + # SorterUseNetworkX reads units["name"][layerId] + table['layerId'] = table.index + if 'name' not in table.columns: + table['name'] = table[self.unit_name_column] + return table + + def _label_table(self, units: pandas.DataFrame, labels: Dict[str, str]) -> pandas.DataFrame: + """ + Make a units table for the sorter where each label is one unit + + The minimum age of a label is the minimum of the minimum ages of its units and + the maximum age is the maximum of the maximum ages. + + Args: + units (pandas.DataFrame): the units + labels (dict): the label of each unit name + + Returns: + pandas.DataFrame: one row for each label + """ + unit_labels = units[self.unit_name_column].map(labels) + table = pandas.DataFrame({self.unit_name_column: list(dict.fromkeys(unit_labels))}) + for attribute, default, aggregate in ( + ('min_age_column', 'minAge', 'min'), + ('max_age_column', 'maxAge', 'max'), + ): + column = getattr(self.sorter, attribute, None) or default + if column in units.columns: + ages = pandas.to_numeric(units[column], errors='coerce') + ages = ages.groupby(unit_labels).agg(aggregate) + table[column] = table[self.unit_name_column].map(ages) + return self._units_table(table) + + def _relabelled_sorter(self, labels: Dict[str, str]) -> Sorter: + """ + Make a copy of the sorter that uses only the units in labels, with their labels as the unit names + + Args: + labels (dict): the label of each unit name + + Returns: + Sorter: the copy of the sorter + """ + sorter = copy.copy(self.sorter) + contacts = getattr(sorter, 'contacts', None) + if contacts is not None: + unitname1_column = ( + getattr(sorter, 'unitname1_column', None) + or getattr(sorter, 'unit1name_column', None) + or 'UNITNAME_1' + ) + unitname2_column = ( + getattr(sorter, 'unitname2_column', None) + or getattr(sorter, 'unit2name_column', None) + or 'UNITNAME_2' + ) + sorter.contacts = relabel_contacts( + contacts, labels, unitname1_column, unitname2_column + ) + unit_relationships = getattr(sorter, 'unit_relationships', None) + if unit_relationships is not None: + sorter.unit_relationships = relabel_unit_relationships(unit_relationships, labels) + geology_data = getattr(sorter, 'geology_data', None) + if geology_data is not None and 'UNITNAME' in geology_data.columns: + geology_data = geology_data[geology_data['UNITNAME'].isin(list(labels))].copy() + geology_data['UNITNAME'] = geology_data['UNITNAME'].map(labels) + sorter.geology_data = geology_data.reset_index(drop=True) + # copy.copy shares lists, so do not add to the list of the original sorter + if isinstance(getattr(sorter, 'lines', None), list): + sorter.lines = [] + return sorter + + def _sort_step(self, table: pandas.DataFrame, labels: Dict[str, str]) -> list: + """ + Sort the labels in the table with a relabelled copy of the sorter + + If the sorter fails, the labels are kept in the order of the table. If the sorter + does not return some labels, they are added at the end. + + Args: + table (pandas.DataFrame): one row for each label + labels (dict): the label of each unit name + + Returns: + list: the sorted labels + """ + expected = list(table[self.unit_name_column]) + if len(expected) < 2: + return expected + try: + order = self._relabelled_sorter(labels).sort(table) + except Exception as e: + logger.warning( + f"{self.sorter_label}: could not sort {expected} ({e}). " + "The order of these is not changed." + ) + return expected + expected_set = set(expected) + order = [label for label in dict.fromkeys(order) if label in expected_set] + missing = [label for label in expected if label not in order] + if missing: + logger.warning( + f"{self.sorter_label}: the sorter did not give a position for {missing}. " + "They are added at the end." + ) + order += missing + if self.postprocess is not None: + order = list(self.postprocess(order, labels)) + return order diff --git a/packages/map2loop/tests/sorter/test_sorter_hierarchical.py b/packages/map2loop/tests/sorter/test_sorter_hierarchical.py new file mode 100644 index 000000000..c56f6ea4b --- /dev/null +++ b/packages/map2loop/tests/sorter/test_sorter_hierarchical.py @@ -0,0 +1,181 @@ +import pandas +import pytest + +from map2loop.sorter import ( + SorterAgeBased, + SorterAlpha, + SorterHierarchical, + SorterUseNetworkX, + relabel_contacts, + relabel_unit_relationships, +) + + +def make_units(rows): + units = pandas.DataFrame(rows, columns=["name", "minAge", "maxAge", "group", "supergroup"]) + units["layerId"] = units.index + return units + + +def assert_contiguous(order, units, column): + """Check that the units of each value of column are next to each other in order""" + values = units.set_index("name")[column] + sequence = [values[name] for name in order] + seen = [] + for value in sequence: + if seen and seen[-1] == value: + continue + assert value not in seen, f"{column} {value} is split in {order}" + seen.append(value) + + +def test_relabel_contacts_adds_lengths_and_removes_internal_contacts(): + contacts = pandas.DataFrame( + { + "UNITNAME_1": ["A", "B", "C", "A", "X"], + "UNITNAME_2": ["B", "C", "A", "C", "A"], + "length": [10.0, 5.0, 2.0, 3.0, 100.0], + } + ) + labels = {"A": "G1", "B": "G1", "C": "G2"} + result = relabel_contacts(contacts, labels) + assert len(result) == 1 + row = result.iloc[0] + assert {row["UNITNAME_1"], row["UNITNAME_2"]} == {"G1", "G2"} + # B-C, C-A and A-C; A-B is in G1 and X is not in labels + assert row["length"] == pytest.approx(10.0) + + +def test_relabel_unit_relationships_keeps_direction(): + relationships = pandas.DataFrame( + {"UNITNAME_1": ["A", "B", "A"], "UNITNAME_2": ["B", "C", "C"]} + ) + labels = {"A": "G1", "B": "G1", "C": "G2"} + result = relabel_unit_relationships(relationships, labels) + assert result.values.tolist() == [["G1", "G2"]] + + +def test_age_based_keeps_groups_together(): + units = make_units( + [ + ["A", 1.0, 2.0, "G1", "S1"], + ["B", 3.0, 4.0, "G1", "S1"], + ["C", 2.5, 2.6, "G2", "S1"], + ] + ) + # With no hierarchy, C is between A and B + assert SorterAgeBased().sort(units.drop(columns=["group"])) == ["A", "C", "B"] + order = SorterHierarchical(sorter=SorterAgeBased()).sort(units) + assert order == ["A", "B", "C"] + + +def test_two_supergroups_have_separate_orders(): + units = make_units( + [ + ["A", 0, 0, "", "S1"], + ["B", 0, 0, "", "S1"], + ["C", 0, 0, "", "S1"], + ["D", 0, 0, "", "S2"], + ["E", 0, 0, "", "S2"], + ] + ) + relationships = pandas.DataFrame( + { + "UNITNAME_1": ["A", "B", "D", "B"], + "UNITNAME_2": ["B", "C", "E", "D"], + } + ) + sorter = SorterHierarchical(sorter=SorterUseNetworkX(unit_relationships=relationships)) + order = sorter.sort(units) + assert order == ["A", "B", "C", "D", "E"] + # the original sorter data is not changed + assert len(sorter.sorter.unit_relationships) == 4 + + +def test_groups_in_supergroups_with_contacts(): + units = make_units( + [ + ["A", 0, 0, "G1", "S1"], + ["B", 0, 0, "G2", "S1"], + ["C", 0, 0, "G1", "S1"], + ["D", 0, 0, "G2", "S1"], + ["E", 0, 0, "G3", "S2"], + ["F", 0, 0, "G3", "S2"], + ] + ) + contacts = pandas.DataFrame( + { + "UNITNAME_1": ["A", "C", "B", "B", "D", "E", "A"], + "UNITNAME_2": ["C", "B", "D", "E", "F", "F", "B"], + "length": [100.0, 50.0, 80.0, 500.0, 30.0, 90.0, 400.0], + } + ) + order = SorterHierarchical(sorter=SorterAlpha(contacts=contacts)).sort(units) + assert sorted(order) == sorted(units["name"]) + assert_contiguous(order, units, "supergroup") + assert_contiguous(order, units, "group") + + +def test_unit_with_no_group_is_one_item(): + units = make_units( + [ + ["A", 1.0, 2.0, "G1", ""], + ["B", 3.0, 4.0, "G1", ""], + ["C", 2.5, 2.6, None, ""], + ["D", 5.0, 6.0, "G2", ""], + ] + ) + order = SorterHierarchical(sorter=SorterAgeBased()).sort(units) + assert order == ["A", "B", "C", "D"] + + +def test_group_with_no_supergroup_is_one_item(): + units = make_units( + [ + ["A", 1.0, 2.0, "G1", "S1"], + ["B", 9.0, 10.0, "G2", "S1"], + ["C", 4.0, 5.0, "G3", ""], + ["D", 5.0, 6.0, "G3", ""], + ["E", 3.0, 4.0, "G4", "S2"], + ] + ) + order = SorterHierarchical(sorter=SorterAgeBased()).sort(units) + # S1 spans 1 to 10 (mean 5.5), G3 spans 4 to 6 (mean 5), S2 spans 3 to 4 + assert order == ["E", "C", "D", "A", "B"] + + +def test_sorter_failure_keeps_all_units(): + units = make_units( + [ + ["A", 0, 0, "G1", ""], + ["B", 0, 0, "G1", ""], + ["C", 0, 0, "G2", ""], + ] + ) + # no contacts in G1, so SorterAlpha can not sort A and B + contacts = pandas.DataFrame( + {"UNITNAME_1": ["B"], "UNITNAME_2": ["C"], "length": [10.0]} + ) + order = SorterHierarchical(sorter=SorterAlpha(contacts=contacts)).sort(units) + assert sorted(order) == ["A", "B", "C"] + assert_contiguous(order, units, "group") + + +def test_postprocess_is_applied_to_each_step(): + units = make_units( + [ + ["A", 1.0, 2.0, "G1", ""], + ["B", 3.0, 4.0, "G1", ""], + ["C", 5.0, 6.0, "G2", ""], + ["D", 7.0, 8.0, "G2", ""], + ] + ) + calls = [] + + def reverse(order, labels): + calls.append(sorted(set(labels.values()))) + return list(reversed(order)) + + order = SorterHierarchical(sorter=SorterAgeBased(), postprocess=reverse).sort(units) + assert order == ["D", "C", "B", "A"] + assert ["G1", "G2"] in calls