From ff3403176d19c23603717e2af03d6c12478c6f46 Mon Sep 17 00:00:00 2001 From: lachlangrose Date: Mon, 5 Oct 2026 15:37:10 -0300 Subject: [PATCH] feat(map2loop): add SorterHierarchical to keep groups and supergroups together SorterHierarchical sorts the supergroups, then the groups in each supergroup, then the units in each group, with a different sorter at each level. The sorter data is relabelled so that each group or supergroup is one unit, and each step uses only the data of its units. Thus each supergroup has its own stratigraphic order. --- packages/map2loop/src/map2loop/project.py | 43 ++- packages/map2loop/src/map2loop/sorter.py | 349 +++++++++++++++++- .../tests/sorter/test_sorter_hierarchical.py | 181 +++++++++ 3 files changed, 556 insertions(+), 17 deletions(-) create mode 100644 packages/map2loop/tests/sorter/test_sorter_hierarchical.py 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