From b2f8b6dd7e2ae1c36b1516ef03c5240dd5322bed Mon Sep 17 00:00:00 2001 From: xumi1993 Date: Tue, 21 Jul 2026 20:54:08 +0800 Subject: [PATCH 01/12] Replace yamlium with ruamel.yaml for YAML parsing across the project --- pyproject.toml | 2 +- pytomoatt/para.py | 20 +++++++++----------- test/test_para.py | 12 ++++++++---- test/test_script.py | 7 +++++-- test/vis_model_file/model_vis.ipynb | 13 +++++++------ 5 files changed, 30 insertions(+), 24 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index d1e9e51..c09f9a1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,7 +27,7 @@ dependencies = [ "pandas>=1.4.0", "scipy", "h5py", - "yamlium>=0.3.0", + "ruamel.yaml", "xarray", "tqdm", "pyproj", diff --git a/pytomoatt/para.py b/pytomoatt/para.py index 8d55c6d..e9a859e 100644 --- a/pytomoatt/para.py +++ b/pytomoatt/para.py @@ -1,9 +1,8 @@ -from pathlib import Path - -from yamlium import parse - +from ruamel.yaml import YAML from .utils.common import init_axis, str2val +yaml = YAML() +yaml.default_flow_style = True class ATTPara: """Class for read and write parameter file with ``yaml`` format @@ -14,7 +13,9 @@ def __init__(self, fname: str) -> None: :type fname: str """ self.fname = fname - self.input_params = parse(Path(fname)) + with open(fname, encoding='utf-8') as f: + file_data = f.read() + self.input_params = yaml.load(file_data) def init_axis(self): dep, lat, lon, dd, dt, dp = init_axis( @@ -34,11 +35,7 @@ def update_param(self, key: str, value) -> None: keys = key.split('.') param = self.input_params for k in keys[:-1]: - if k not in param: - # Assignment lets yamlium wrap the dict in its Mapping node. - # dict.setdefault() bypasses yamlium's conversion logic. - param[k] = {} - param = param[k] + param = param.setdefault(k, {}) param[keys[-1]] = str2val(value) def write(self, fname=None): @@ -49,4 +46,5 @@ def write(self, fname=None): """ if fname is None: fname = self.fname - self.input_params.yaml_dump(fname) + with open(fname, 'w') as f: + yaml.dump(self.input_params, f) diff --git a/test/test_para.py b/test/test_para.py index 18c7419..415ea8c 100644 --- a/test/test_para.py +++ b/test/test_para.py @@ -1,10 +1,11 @@ import unittest import os import shutil -from yamlium import from_dict, parse +from ruamel.yaml import YAML from pytomoatt.para import ATTPara import numpy as np +yaml = YAML() class TestATTPara(unittest.TestCase): def setUp(self): @@ -27,7 +28,8 @@ def setUp(self): } } self.fname = 'test_params.yml' - from_dict(self.yaml_content).yaml_dump(self.fname) + with open(self.fname, 'w') as f: + yaml.dump(self.yaml_content, f) def tearDown(self): os.chdir(self.cwd) @@ -82,7 +84,8 @@ def test_write(self): self.assertTrue(os.path.exists(out_fname)) # Verify content - new_params = parse(out_fname) + with open(out_fname, 'r') as f: + new_params = yaml.load(f) self.assertEqual(new_params['domain']['n_rtp'], [30, 30, 30]) def test_write_overwrite(self): @@ -90,7 +93,8 @@ def test_write_overwrite(self): para.update_param('domain.n_rtp', '40,40,40') para.write() # Should overwrite self.fname - new_params = parse(self.fname) + with open(self.fname, 'r') as f: + new_params = yaml.load(f) self.assertEqual(new_params['domain']['n_rtp'], [40, 40, 40]) if __name__ == '__main__': diff --git a/test/test_script.py b/test/test_script.py index 8e423f1..bef4a22 100644 --- a/test/test_script.py +++ b/test/test_script.py @@ -3,11 +3,13 @@ import os import shutil import sys -from yamlium import parse +from ruamel.yaml import YAML import h5py import numpy as np from pytomoatt.script import PTA +yaml = YAML() + class TestScripts(unittest.TestCase): def setUp(self): @@ -40,7 +42,8 @@ def test_setpar(self): with patch.object(sys, 'argv', ['pta', 'setpar', 'input_params.yml', 'domain.n_rtp', '10,10,10']): PTA() - params = parse('input_params.yml') + with open('input_params.yml', 'r') as f: + params = yaml.load(f) self.assertEqual(params['domain']['n_rtp'], [10, 10, 10]) def test_model2vtk(self): diff --git a/test/vis_model_file/model_vis.ipynb b/test/vis_model_file/model_vis.ipynb index 9bf2210..ce750e3 100644 --- a/test/vis_model_file/model_vis.ipynb +++ b/test/vis_model_file/model_vis.ipynb @@ -7,16 +7,17 @@ "outputs": [], "source": [ "# read grid information from yaml file\n", - "from yamlium import parse\n", + "import yaml\n", "\n", "fname = 'input_file_example.yml'\n", "\n", - "data = parse(fname)\n", + "with open(fname, 'r') as f:\n", + " data = yaml.load(f, Loader=yaml.FullLoader)\n", "\n", - "min_max_dep = data['domain']['min_max_dep']\n", - "min_max_lat = data['domain']['min_max_lat']\n", - "min_max_lon = data['domain']['min_max_lon']\n", - "n_rtp = data['domain']['n_rtp']" + " min_max_dep = data['domain']['min_max_dep']\n", + " min_max_lat = data['domain']['min_max_lat']\n", + " min_max_lon = data['domain']['min_max_lon']\n", + " n_rtp = data['domain']['n_rtp']" ] }, { From a2a3ce9670a2ce394370f10bd194220aeb615689 Mon Sep 17 00:00:00 2001 From: xumi1993 Date: Tue, 21 Jul 2026 21:34:41 +0800 Subject: [PATCH 02/12] Add linear regression functionality and tests for SrcRec class --- pytomoatt/src_rec.py | 121 ++++++++++++++++++++++++++++++- pytomoatt/utils/src_rec_utils.py | 33 +++++++++ test/test_src_rec.py | 50 ++++++++++++- 3 files changed, 200 insertions(+), 4 deletions(-) diff --git a/pytomoatt/src_rec.py b/pytomoatt/src_rec.py index 19c6212..af5c68e 100644 --- a/pytomoatt/src_rec.py +++ b/pytomoatt/src_rec.py @@ -4,7 +4,8 @@ from .distaz import DistAZ from .setuplog import SetupLog from .utils.src_rec_utils import define_rec_cols, setup_rec_points_dd, \ - get_rec_points_types, update_position + get_rec_points_types, update_position, \ + linear_regression from sklearn.metrics.pairwise import haversine_distances import copy from io import StringIO @@ -1203,6 +1204,124 @@ def select_by_distance(self, dist_min_max, recalc_dist=False, **kwargs): "rec_points after selection: {}".format(self._count_records()) ) + @staticmethod + def _regression_keep_mask(records, std_multiplier): + """Return a mask for finite records within the residual limit.""" + finite = np.isfinite(records["dist_deg"].to_numpy(dtype=float)) & \ + np.isfinite(records["tt"].to_numpy(dtype=float)) + valid = records.loc[finite] + keep = pd.Series(False, index=records.index, dtype=bool) + + if len(valid) < 2 or valid["dist_deg"].nunique() < 2: + keep.loc[valid.index] = True + return keep + + slope, intercept, residual_std = linear_regression( + valid["dist_deg"], valid["tt"] + ) + residual = valid["tt"] - (slope * valid["dist_deg"] + intercept) + keep.loc[valid.index] = ( + np.isclose(residual, 0.0) if residual_std == 0 + else np.abs(residual) <= std_multiplier * residual_std + ) + return keep + + def _filter_double_difference_by_arrivals(self): + """Remove double differences whose absolute arrivals were rejected.""" + arrivals = set( + self.rec_points[["src_index", "staname", "phase"]] + .itertuples(index=False, name=None) + ) + specs = ( + ("rec_points_cs", ",cs", + (("src_index", "staname1"), ("src_index", "staname2")), + "common-source"), + ("rec_points_cr", ",cr", + (("src_index", "staname"), ("src_index2", "staname")), + "common-receiver"), + ) + + for attr, suffix, endpoints, label in specs: + records = getattr(self, attr) + if records.empty: + continue + + phases = records["phase"].map( + lambda phase: phase[:-len(suffix)] + if isinstance(phase, str) and phase.endswith(suffix) + else phase + ) + keep = np.ones(len(records), dtype=bool) + for src_col, sta_col in endpoints: + keys = zip(records[src_col], records[sta_col], phases) + keep &= np.fromiter( + (key in arrivals for key in keys), bool, len(records) + ) + + setattr(self, attr, records.loc[keep]) + self.log.SrcReclog.info( + "Removed {} corresponding {} records".format( + len(records) - np.count_nonzero(keep), label + ) + ) + + def select_by_linear_regression(self, std_multiplier=3.0, + recalc_dist=False, separate_phase=True, + **kwargs): + """Select absolute travel times by linear-regression residual. + + A straight line is fitted between epicentral distance and travel time. + Records whose absolute residual is greater than ``std_multiplier`` + times the residual standard deviation are removed. By default each + phase is fitted separately so that phases with different apparent + velocities are not mixed. + + .. note:: + This criterion only applies to absolute travel-time data in + :attr:`rec_points`. A double-difference record is removed when + either of its corresponding absolute travel times is rejected. + + :param std_multiplier: Multiplier applied to the residual standard + deviation, defaults to 3. + :type std_multiplier: float + :param recalc_dist: Recalculate epicentral distance even when + ``dist_deg`` exists, defaults to False. + :type recalc_dist: bool + :param separate_phase: Fit each phase separately, defaults to True. + :type separate_phase: bool + """ + if (not np.isscalar(std_multiplier) + or not np.isfinite(std_multiplier) + or std_multiplier <= 0): + raise ValueError("std_multiplier must be a positive finite number") + + self.log.SrcReclog.info( + "rec_points before travel-time selection: {}".format( + self.rec_points.shape[0] + ) + ) + if ("dist_deg" not in self.rec_points) or recalc_dist: + self.log.SrcReclog.info("Calculating epicentral distance...") + self.calc_distaz() + + keep = pd.Series(False, index=self.rec_points.index, dtype=bool) + groups = (self.rec_points.groupby("phase", dropna=False) + if separate_phase else [("all", self.rec_points)]) + + for _, records in groups: + keep.loc[records.index] = self._regression_keep_mask( + records, std_multiplier + ) + + self.rec_points = self.rec_points.loc[keep] + self._filter_double_difference_by_arrivals() + self.update(**kwargs) + self.log.SrcReclog.info( + "rec_points after travel-time selection: {}".format( + self.rec_points.shape[0] + ) + ) + def select_by_azi_gap(self, max_azi_gap: float, **kwargs): """Select sources with azimuthal gap greater and equal than a number diff --git a/pytomoatt/utils/src_rec_utils.py b/pytomoatt/utils/src_rec_utils.py index 32c37ed..6f5a5b5 100644 --- a/pytomoatt/utils/src_rec_utils.py +++ b/pytomoatt/utils/src_rec_utils.py @@ -1,5 +1,6 @@ import io import tqdm +import numpy as np def define_rec_cols(dist_in_data, name_net_and_sta): @@ -244,3 +245,35 @@ def download_src_rec_file(url): else: response.release_conn() return None + + +def linear_regression(x, y): + """Fit a line and return its slope, intercept and residual standard deviation. + + Non-finite pairs are ignored. At least two samples with different + x-coordinates are required. + + :param x: Independent variable. + :type x: array-like + :param y: Dependent variable. + :type y: array-like + :return: Slope, intercept and standard deviation of the residuals. + :rtype: tuple of float + """ + x = np.asarray(x, dtype=float) + y = np.asarray(y, dtype=float) + if x.ndim != 1 or y.ndim != 1 or x.shape != y.shape: + raise ValueError("x and y must be one-dimensional arrays of equal length") + + finite = np.isfinite(x) & np.isfinite(y) + x = x[finite] + y = y[finite] + if x.size < 2: + raise ValueError("at least two finite samples are required") + if np.ptp(x) == 0: + raise ValueError("x must contain at least two distinct values") + + slope, intercept = np.polyfit(x, y, deg=1) + residual = y - (slope * x + intercept) + std = np.std(residual) + return float(slope), float(intercept), float(std) diff --git a/test/test_src_rec.py b/test/test_src_rec.py index 811a126..b5d375e 100644 --- a/test/test_src_rec.py +++ b/test/test_src_rec.py @@ -5,11 +5,13 @@ get_rec_points_types, setup_rec_points_dd, update_position, - download_src_rec_file + download_src_rec_file, + linear_regression, ) from os.path import dirname, join from unittest.mock import MagicMock, patch import pandas as pd +import numpy as np import io @@ -70,8 +72,52 @@ def test_subcase_10(self): sr = SrcRec.read(self.fname) sr.box_weighting(0.4, 10, obj='both') + def test_select_by_linear_regression(self): + sr = SrcRec('unused') + distance = np.concatenate((np.arange(21, dtype=float), [10.0, 0.0])) + travel_time = 2.0 * distance + 5.0 + travel_time[10] += 100.0 + sr.rec_points = pd.DataFrame({ + 'src_index': [0] * 21 + [1, 1], + 'staname': [f'STA{i:02d}' for i in range(21)] + ['STA10', 'STA00'], + 'dist_deg': distance, + 'tt': travel_time, + 'phase': 'P', + }) + sr.rec_points_cs = pd.DataFrame({ + 'src_index': [0, 0], + 'staname1': ['STA10', 'STA00'], + 'staname2': ['STA00', 'STA01'], + 'phase': ['P,cs', 'P,cs'], + }) + sr.rec_points_cr = pd.DataFrame({ + 'src_index': [0, 0], + 'src_index2': [1, 1], + 'staname': ['STA10', 'STA00'], + 'phase': ['P,cr', 'P,cr'], + }) + + with patch.object(sr, 'update') as update: + sr.select_by_linear_regression(std_multiplier=3.0) + + self.assertEqual(sr.rec_points.shape[0], 22) + self.assertNotIn(10, sr.rec_points.index) + self.assertEqual(sr.rec_points_cs.shape[0], 1) + self.assertEqual(sr.rec_points_cs.iloc[0]['staname1'], 'STA00') + self.assertEqual(sr.rec_points_cr.shape[0], 1) + self.assertEqual(sr.rec_points_cr.iloc[0]['staname'], 'STA00') + update.assert_called_once_with() + class TestSrcRecUtils(unittest.TestCase): + def test_linear_regression(self): + slope, intercept, std = linear_regression( + [0.0, 1.0, 2.0], [1.0, 3.0, 5.0] + ) + self.assertAlmostEqual(slope, 2.0) + self.assertAlmostEqual(intercept, 1.0) + self.assertAlmostEqual(std, 0.0) + def test_define_rec_cols(self): # Case 1: dist_in_data=False, name_net_and_sta=False cols, last_col = define_rec_cols(False, False) @@ -198,5 +244,3 @@ def test_download_src_rec_file(self): if __name__ == '__main__': unittest.main() - - From 463a86fd06ba5ea46b3169fc591394ed57378709 Mon Sep 17 00:00:00 2001 From: xumi1993 Date: Tue, 21 Jul 2026 21:47:50 +0800 Subject: [PATCH 03/12] Update select_by_linear_regression to return regression parameters and add tests for slope and intercept --- pytomoatt/src_rec.py | 16 ++++++++++++---- test/test_src_rec.py | 11 ++++++++++- 2 files changed, 22 insertions(+), 5 deletions(-) diff --git a/pytomoatt/src_rec.py b/pytomoatt/src_rec.py index af5c68e..368ceb8 100644 --- a/pytomoatt/src_rec.py +++ b/pytomoatt/src_rec.py @@ -1214,7 +1214,7 @@ def _regression_keep_mask(records, std_multiplier): if len(valid) < 2 or valid["dist_deg"].nunique() < 2: keep.loc[valid.index] = True - return keep + return keep, None slope, intercept, residual_std = linear_regression( valid["dist_deg"], valid["tt"] @@ -1224,7 +1224,7 @@ def _regression_keep_mask(records, std_multiplier): np.isclose(residual, 0.0) if residual_std == 0 else np.abs(residual) <= std_multiplier * residual_std ) - return keep + return keep, (slope, intercept) def _filter_double_difference_by_arrivals(self): """Remove double differences whose absolute arrivals were rejected.""" @@ -1289,6 +1289,9 @@ def select_by_linear_regression(self, std_multiplier=3.0, :type recalc_dist: bool :param separate_phase: Fit each phase separately, defaults to True. :type separate_phase: bool + :return: Mapping from phase name to ``(slope, intercept)``. When + ``separate_phase=False``, the key is ``"all"``. + :rtype: dict """ if (not np.isscalar(std_multiplier) or not np.isfinite(std_multiplier) @@ -1307,11 +1310,15 @@ def select_by_linear_regression(self, std_multiplier=3.0, keep = pd.Series(False, index=self.rec_points.index, dtype=bool) groups = (self.rec_points.groupby("phase", dropna=False) if separate_phase else [("all", self.rec_points)]) + regression_params = {} - for _, records in groups: - keep.loc[records.index] = self._regression_keep_mask( + for phase, records in groups: + group_keep, params = self._regression_keep_mask( records, std_multiplier ) + keep.loc[records.index] = group_keep + if params is not None: + regression_params[phase] = params self.rec_points = self.rec_points.loc[keep] self._filter_double_difference_by_arrivals() @@ -1321,6 +1328,7 @@ def select_by_linear_regression(self, std_multiplier=3.0, self.rec_points.shape[0] ) ) + return regression_params def select_by_azi_gap(self, max_azi_gap: float, **kwargs): """Select sources with azimuthal gap greater and equal than a number diff --git a/test/test_src_rec.py b/test/test_src_rec.py index b5d375e..3df0f37 100644 --- a/test/test_src_rec.py +++ b/test/test_src_rec.py @@ -98,8 +98,17 @@ def test_select_by_linear_regression(self): }) with patch.object(sr, 'update') as update: - sr.select_by_linear_regression(std_multiplier=3.0) + regression_params = sr.select_by_linear_regression( + std_multiplier=3.0 + ) + expected_slope, expected_intercept = np.polyfit( + distance, travel_time, deg=1 + ) + self.assertIn('P', regression_params) + slope, intercept = regression_params['P'] + self.assertAlmostEqual(slope, expected_slope) + self.assertAlmostEqual(intercept, expected_intercept) self.assertEqual(sr.rec_points.shape[0], 22) self.assertNotIn(10, sr.rec_points.index) self.assertEqual(sr.rec_points_cs.shape[0], 1) From 6c49eaa1e68fc2ed11b321168609d3a53b64a80e Mon Sep 17 00:00:00 2001 From: JingChen-Thu Date: Wed, 22 Jul 2026 11:54:55 +0800 Subject: [PATCH 04/12] add rotation in crust1.0 model --- pytomoatt/io/crustmodel.py | 32 +++++++++++++++++++++++--------- pytomoatt/model.py | 6 ++++-- 2 files changed, 27 insertions(+), 11 deletions(-) diff --git a/pytomoatt/io/crustmodel.py b/pytomoatt/io/crustmodel.py index f7fa937..cb5aa1e 100644 --- a/pytomoatt/io/crustmodel.py +++ b/pytomoatt/io/crustmodel.py @@ -2,6 +2,7 @@ from os.path import dirname, abspath, join from ..utils.common import init_axis from ..setuplog import SetupLog +from ..utils.rotate import rtp_rotation_reverse import pickle import sys from tqdm import tqdm @@ -40,7 +41,7 @@ def __init__(self, fname=join(dirname(dirname(abspath(__file__))), 'data', 'crus self.points_dict = pickle.load(f) self.log = SetupLog() - def griddata(self, min_max_dep, min_max_lat, min_max_lon, n_rtp, type='vp'): + def griddata(self, min_max_dep, min_max_lat, min_max_lon, n_rtp, type='vp', rotate=None): """Linearly interpolate velocity into regular grids :param min_max_dep: min and max depth, ``[min_dep, max_dep]`` @@ -63,25 +64,38 @@ def griddata(self, min_max_dep, min_max_lat, min_max_lon, n_rtp, type='vp'): else: self.log.Modellog.error(f"Velocity type {type} not supported in CRUST1.0 model") sys.exit(1) + self.dd, self.tt, self.pp, _, _, _, = init_axis( min_max_dep, min_max_lat, min_max_lon, n_rtp ) + tt_2d, pp_2d = np.meshgrid(self.tt, self.pp, indexing='ij') + + # rotate reversely, from computational grid to physical grid + if rotate is not None: + central_lat = rotate[0] + central_lon = rotate[1] + rotation_angle = rotate[2] + tt_2d, pp_2d = rtp_rotation_reverse(tt_2d, pp_2d, central_lat, central_lon, rotation_angle) + # Grid data self.log.Modellog.info('Grid data, please wait for a few minutes') vel = np.zeros(n_rtp) with tqdm(total=self.n_rtp[1] * self.n_rtp[2], desc='Gridding') as pbar: for ilat in range(self.n_rtp[1]): - new_lat = self.tt[ilat] - idx_lat_left, ratio_lat = degree_to_idx_and_ratio(new_lat) - idx_lat_right = idx_lat_left + 1 - if idx_lat_left == -1: - idx_lat_left = 0 - idx_lat_right = 1 - for ilon in range(self.n_rtp[2]): pbar.update(1) - new_lon = self.pp[ilon] + + # latitude index and ratio + new_lat = tt_2d[ilat, ilon] + idx_lat_left, ratio_lat = degree_to_idx_and_ratio(new_lat) + idx_lat_right = idx_lat_left + 1 + if idx_lat_left == -1: + idx_lat_left = 0 + idx_lat_right = 1 + + # longitude index and ratio + new_lon = pp_2d[ilat, ilon] idx_lon_left, ratio_lon = degree_to_idx_and_ratio(new_lon) idx_lon_right = idx_lon_left + 1 if idx_lon_left == -1: # between -179.5 and +179.5 diff --git a/pytomoatt/model.py b/pytomoatt/model.py index 14c5102..b41d082 100644 --- a/pytomoatt/model.py +++ b/pytomoatt/model.py @@ -102,18 +102,20 @@ def to_xarray(self): ) return dataset - def grid_data_crust1(self, type='vp'): + def grid_data_crust1(self, type='vp', rotate=None): """Grid data from CRUST1.0 model :param type: Specify velocity type of ``vp`` or ``vs``, defaults to 'vp' :type type: str, optional + :rotate: Rotation parameters [theta0, phi0, psi], defaults to None """ cm = CrustModel() self.vel = cm.griddata( self.min_max_dep, self.min_max_lat, self.min_max_lon, - self.n_rtp, type=type + self.n_rtp, type=type, + rotate=rotate ) def grid_data_ascii(self, model_fname:str, **kwargs): From d6ac9f3683b711d3060976f1a23e9ace86ef87cf Mon Sep 17 00:00:00 2001 From: JingChen-Thu Date: Wed, 22 Jul 2026 12:05:04 +0800 Subject: [PATCH 05/12] exclude zeta in Checker --- pytomoatt/checkerboard.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/pytomoatt/checkerboard.py b/pytomoatt/checkerboard.py index 589c577..309a73b 100644 --- a/pytomoatt/checkerboard.py +++ b/pytomoatt/checkerboard.py @@ -22,7 +22,10 @@ def __init__(self, model_fname:str, para_fname='input_params.yml') -> None: self.vel = f['vel'][:] self.eta = f['eta'][:] self.xi = f['xi'][:] - self.zeta = f['zeta'][:] + try: # some model may not have zeta + self.zeta = f['zeta'][:] + except: + pass self._init_axis() def _init_axis(self): From 8079307153946f4ea50c80ba7d7376ac3161b146 Mon Sep 17 00:00:00 2001 From: JingChen-Thu Date: Wed, 22 Jul 2026 12:06:09 +0800 Subject: [PATCH 06/12] exclude zeta in Checker --- pytomoatt/checkerboard.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pytomoatt/checkerboard.py b/pytomoatt/checkerboard.py index 309a73b..6724d51 100644 --- a/pytomoatt/checkerboard.py +++ b/pytomoatt/checkerboard.py @@ -25,7 +25,7 @@ def __init__(self, model_fname:str, para_fname='input_params.yml') -> None: try: # some model may not have zeta self.zeta = f['zeta'][:] except: - pass + self.zeta = np.zeros_like(self.vel) self._init_axis() def _init_axis(self): From 3b216809084d5958970182d9eb96ce46eb95eff0 Mon Sep 17 00:00:00 2001 From: xumi1993 Date: Wed, 22 Jul 2026 14:12:52 +0800 Subject: [PATCH 07/12] Update numpy dependency to version 1.26.4 in pyproject.toml --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index c09f9a1..c9be658 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,7 +23,7 @@ classifiers = [ "Programming Language :: Python :: 3.13", ] dependencies = [ - "numpy>=1.19.0", + "numpy>=1.26.4", "pandas>=1.4.0", "scipy", "h5py", From 46dadf484285737b2f6963a48531194ba73c63de Mon Sep 17 00:00:00 2001 From: Mijian Xu Date: Wed, 22 Jul 2026 14:17:03 +0800 Subject: [PATCH 08/12] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- pytomoatt/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pytomoatt/model.py b/pytomoatt/model.py index b41d082..ef1adc8 100644 --- a/pytomoatt/model.py +++ b/pytomoatt/model.py @@ -107,7 +107,7 @@ def grid_data_crust1(self, type='vp', rotate=None): :param type: Specify velocity type of ``vp`` or ``vs``, defaults to 'vp' :type type: str, optional - :rotate: Rotation parameters [theta0, phi0, psi], defaults to None + :param rotate: Rotation parameters [central_lat, central_lon, rotation_angle] in degrees, defaults to None """ cm = CrustModel() self.vel = cm.griddata( From e4c6effa4a4fbe66195cf7b4f9fc30f5b513aabd Mon Sep 17 00:00:00 2001 From: Mijian Xu Date: Wed, 22 Jul 2026 14:17:54 +0800 Subject: [PATCH 09/12] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- pytomoatt/io/crustmodel.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/pytomoatt/io/crustmodel.py b/pytomoatt/io/crustmodel.py index cb5aa1e..15bdc66 100644 --- a/pytomoatt/io/crustmodel.py +++ b/pytomoatt/io/crustmodel.py @@ -73,9 +73,11 @@ def griddata(self, min_max_dep, min_max_lat, min_max_lon, n_rtp, type='vp', rota # rotate reversely, from computational grid to physical grid if rotate is not None: - central_lat = rotate[0] - central_lon = rotate[1] - rotation_angle = rotate[2] + try: + central_lat, central_lon, rotation_angle = rotate + except (TypeError, ValueError): + self.log.Modellog.error("rotate must be a 3-item sequence: [central_lat, central_lon, rotation_angle]") + sys.exit(1) tt_2d, pp_2d = rtp_rotation_reverse(tt_2d, pp_2d, central_lat, central_lon, rotation_angle) # Grid data From a577245eb5dd3eaf9bcbfe02935c89051bc061e6 Mon Sep 17 00:00:00 2001 From: Mijian Xu Date: Wed, 22 Jul 2026 15:18:23 +0800 Subject: [PATCH 10/12] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- pytomoatt/checkerboard.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pytomoatt/checkerboard.py b/pytomoatt/checkerboard.py index 6724d51..a65db5c 100644 --- a/pytomoatt/checkerboard.py +++ b/pytomoatt/checkerboard.py @@ -22,9 +22,9 @@ def __init__(self, model_fname:str, para_fname='input_params.yml') -> None: self.vel = f['vel'][:] self.eta = f['eta'][:] self.xi = f['xi'][:] - try: # some model may not have zeta + if 'zeta' in f: # some model may not have zeta self.zeta = f['zeta'][:] - except: + else: self.zeta = np.zeros_like(self.vel) self._init_axis() From 8ab2001b978b358883f1aece64e63ffee6964b04 Mon Sep 17 00:00:00 2001 From: xumi1993 Date: Sun, 2 Aug 2026 08:39:38 +0800 Subject: [PATCH 11/12] Enhance SrcRec class for improved file handling and new selection methods - Updated `SrcRec.read` method to handle remote file downloads and improved error handling for missing files. - Implemented reassignment of source indices to ensure uniqueness when reading source-receiver data. - Added `select_by_constant_velocity` method to filter arrivals based on a constant-velocity travel-time curve. - Introduced validation for input parameters in `select_by_constant_velocity`. - Refactored plotting functions into a new `vis.py` module for better organization and maintainability. - Added comprehensive unit tests for new functionality, including handling of duplicate source indices and validation of selection methods. - Created a new test data file to validate behavior with duplicate source indices. --- pyproject.toml | 1 + pytomoatt/distaz.py | 333 +++++++++++++----------------- pytomoatt/src_rec.py | 227 ++++++++++++++++++-- pytomoatt/utils/vis.py | 344 +++++++++++++++++++++++++++++++ test/src_rec_duplicate_index.dat | 6 + test/test_src_rec.py | 212 +++++++++++++++++++ 6 files changed, 915 insertions(+), 208 deletions(-) create mode 100644 pytomoatt/utils/vis.py create mode 100644 test/src_rec_duplicate_index.dat diff --git a/pyproject.toml b/pyproject.toml index c9be658..ecb3950 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,6 +25,7 @@ classifiers = [ dependencies = [ "numpy>=1.26.4", "pandas>=1.4.0", + "matplotlib", "scipy", "h5py", "ruamel.yaml", diff --git a/pytomoatt/distaz.py b/pytomoatt/distaz.py index 98d4f8d..c7d5e6f 100644 --- a/pytomoatt/distaz.py +++ b/pytomoatt/distaz.py @@ -1,209 +1,154 @@ import numpy as np +# Geographic latitude is converted to geocentric latitude before calculating +# the spherical arc. This is the factor used by the reference MATLAB +# implementation (distaz.m). +_GEOCENTRIC_FACTOR = 0.993270 +_AZIMUTH_ZERO_TOL = 1.0e-5 + + +def _geocentric_latitude(latitude): + """Return geocentric latitude in radians.""" + latitude = np.deg2rad(latitude) + + # atan2 is equivalent to atan(factor * tan(latitude)) for valid + # geographic latitudes, but remains well defined at both poles. + return np.arctan2( + _GEOCENTRIC_FACTOR * np.sin(latitude), + np.cos(latitude), + ) + + +def _normalize_azimuth(angle): + """Normalize an angle to [0, 360), mapping round-off near 360 to zero.""" + angle = np.mod(np.rad2deg(angle), 360.0) + near_zero = ( + (np.abs(angle) < _AZIMUTH_ZERO_TOL) + | (np.abs(angle - 360.0) < _AZIMUTH_ZERO_TOL) + ) + return np.where(near_zero, 0.0, angle) + + +def _scalar_or_array(value): + """Return a float for scalar input and an ndarray for broadcast input.""" + value = np.asarray(value, dtype=float) + if value.ndim == 0: + return value.item() + return value + + class DistAZ: """ - DistAZ class - - Calculate the distance, azimuth and back-azimuth between two points on the - Earth's surface. - - :param lat1: Latitude of point 1 - :type lat1: float - :param lon1: Longitude of point 1 - :type lon1: float - :param lat2: Latitude of point 2 - :type lat2: float or np.ndarray - :param lon2: Longitude of point 2 - :type lon2: float or np.ndarray - :return: An instance of DistAZ - :rtype: DistAZ - - Subroutine to calculate the Great Circle Arc distance - between two sets of geographic coordinates - - Equations take from Bullen, pages 154, 155 - - T. Owens, September 19, 1991 - Sept. 25 -- fixed az and baz calculations - - P. Crotwell, Setember 27, 1995 - Converted to c to fix annoying problem of fortran giving wrong - answers if the input doesn't contain a decimal point. - - H. P. Crotwell, September 18, 1997 - Java version for direct use in java programs. - * - * C. Groves, May 4, 2004 - * Added enough convenience constructors to choke a horse and made public double - * values use accessors so we can use this class as an immutable - H.P. Crotwell, May 31, 2006 - Port to python, thus adding to the great list of languages to which - distaz has been ported from the origin fortran: C, Tcl, Java and now python - and I vaguely remember a perl port. Long live distaz! - Mijian Xu, Jan 01, 2016: Add compatibility of np.ndarray for input variables. - Mijian XU, Jun 12, 2025: Update to latest numpy standards. - + Calculate distance, azimuth and back-azimuth between two surface points. + + ``lat1``/``lon1`` describe point 1 (normally the station), and + ``lat2``/``lon2`` describe point 2 (normally the event). Inputs may be + scalars or any mutually broadcastable array-like objects. + + The public angle names retain seispy's historical convention: + + * ``baz`` is the bearing from point 1 to point 2 (MATLAB ``daze``). + * ``az`` is the bearing from point 2 to point 1 (MATLAB ``dazs``). + * ``delta`` is the geocentric great-circle arc in degrees (MATLAB ``dd``). + + The calculation follows the supplied MATLAB ``distaz.m`` implementation, + while using division-free ``atan2`` expressions so that poles and + arbitrary-dimensional NumPy broadcasting are handled safely. + + Parameters + ---------- + lat1, lon1, lat2, lon2 : float or array-like + Geographic coordinates in degrees. + + Notes + ----- + Equations are from Bullen, sections 10.2, pages 154--155. The original + routine was written by T. Owens (1991) and subsequently ported through + Fortran, C, Tcl, Java, and Python. NumPy array support was added to seispy + by Mijian Xu. + + ObsPy's :func:`obspy.geodetics.locations2degrees` uses geographic + latitudes on a sphere, while this routine first converts them to + geocentric latitudes to match ``distaz.m``. ObsPy's + :func:`obspy.geodetics.gps2dist_azimuth` uses a WGS84 ellipsoid by + default. Results agree when the same geocentric latitudes and a spherical + Earth (``f=0``) are used. """ def __init__(self, lat1, lon1, lat2, lon2): - self.stalat = lat1 self.stalon = lon1 self.evtlat = lat2 self.evtlon = lon2 - ''' - if (lat1 == lat2) and (lon1 == lon2): - self.delta = 0.0 - self.az = 0.0 - self.baz = 0.0 - return - ''' - - rad = 2. * np.pi / 360.0 - ''' - scolat and ecolat are the geocentric colatitudes - as defined by Richter (pg. 318) - - Earth Flattening of 1/298.257 take from Bott (pg. 3) - - ''' - sph = 1.0 / 298.257 - - scolat = np.pi / 2.0 - np.arctan((1. - sph) * (1. - sph) * np.tan(lat1 * rad)) - ecolat = np.pi / 2.0 - np.arctan((1. - sph) * (1. - sph) * np.tan(lat2 * rad)) - slon = lon1 * rad - elon = lon2 * rad - """ - - a - e are as defined by Bullen (pg. 154, Sec 10.2) - These are defined for the pt. 1 - - """ - a = np.sin(scolat) * np.cos(slon) - b = np.sin(scolat) * np.sin(slon) - c = np.cos(scolat) - d = np.sin(slon) - e = -np.cos(slon) - g = -c * e - h = c * d - k = -np.sin(scolat) - """ - c - c aa - ee are the same as a - e, except for pt. 2 - c - """ - aa = np.sin(ecolat) * np.cos(elon) - bb = np.sin(ecolat) * np.sin(elon) - cc = np.cos(ecolat) - dd = np.sin(elon) - ee = -np.cos(elon) - gg = -cc * ee - hh = cc * dd - kk = -np.sin(ecolat) - """ - c - c Bullen, Sec 10.2, eqn. 4 - c - """ - clamped_value = np.clip(a * aa + b * bb + c * cc, -1.0, 1.0) - delrad = np.arccos(clamped_value) - self.delta = delrad / rad - """ - c - c Bullen, Sec 10.2, eqn 7 / eqn 8 - c - c pt. 1 is unprimed, so this is technically the baz - c - c Calculate baz this way to avoid quadrant problems - c - """ - rhs1 = (aa - d) * (aa - d) + (bb - e) * (bb - e) + cc * cc - 2. - rhs2 = (aa - g) * (aa - g) + (bb - h) * (bb - h) + (cc - k) * (cc - k) - 2. - dbaz = np.arctan2(rhs1, rhs2) - - # dbaz_idx = np.where(dbaz < 0.0)[0] - dbaz_idx = np.atleast_1d(dbaz < 0.0).nonzero()[0] - if len(dbaz_idx) != 0: - if np.isscalar(dbaz): - dbaz += 2 * np.pi - else: - dbaz[dbaz_idx] += 2 * np.pi - - self.baz = dbaz / rad - """ - c - c Bullen, Sec 10.2, eqn 7 / eqn 8 - c - c pt. 2 is unprimed, so this is technically the az - c - """ - rhs1 = (a - dd) * (a - dd) + (b - ee) * (b - ee) + c * c - 2. - rhs2 = (a - gg) * (a - gg) + (b - hh) * (b - hh) + (c - kk) * (c - kk) - 2. - daz = np.arctan2(rhs1, rhs2) - - # daz_idx = np.where(daz < 0.0)[0] - daz_idx = np.atleast_1d(dbaz < 0.0).nonzero()[0] - if len(daz_idx) != 0: - if np.isscalar(daz): - daz += 2 * np.pi - else: - daz[daz_idx] += 2 * np.pi - - self.az = daz / rad - """ - c - c Make sure 0.0 is always 0.0, not 360. - c - """ - # idx = np.where(np.abs(self.baz - 360.) < .00001)[0] - idx = np.atleast_1d(np.abs(self.baz - 360.) < .00001).nonzero()[0] - if len(idx) != 0: - if np.isscalar(self.baz): - self.baz = 0.0 - else: - self.baz[idx] = 0.0 - # idx = np.where(np.abs(self.baz) < .00001)[0] - idx = np.atleast_1d(np.abs(self.baz) < .00001).nonzero()[0] - if len(idx) != 0: - if np.isscalar(self.baz): - self.baz = 0.0 - else: - self.baz[idx] = 0.0 - - # idx = np.where(np.abs(self.az - 360.) < .00001)[0] - idx = np.atleast_1d(np.abs(self.az - 360.) < .00001).nonzero()[0] - if len(idx) != 0: - if isinstance(self.az, float): - self.az = 0.0 - else: - self.az[idx] = 0.0 - # idx = np.where(np.abs(self.az) < .00001)[0] - idx = np.atleast_1d(np.abs(self.az) < .00001).nonzero()[0] - if len(idx) != 0: - if isinstance(self.az, float): - self.az = 0.0 - else: - self.az[idx] = 0.0 - - # la_idx = np.where(lat1 == lat2)[0] - # lo_idx = np.where(lon1 == lon2)[0] - la_idx = np.atleast_1d(lat1 == lat2).nonzero()[0] - lo_idx = np.atleast_1d(lon1 == lon2).nonzero()[0] - idx = np.intersect1d(la_idx, lo_idx) - if len(idx) != 0: - if isinstance(self.delta, float): - self.delta = 0. - else: - self.delta[idx] = 0. - if isinstance(self.az, float): - self.az = 0. - else: - self.az[idx] = 0. - if isinstance(self.baz, float): - self.baz = 0. - else: - self.baz[idx] = 0. + + lat1_array, lon1_array, lat2_array, lon2_array = np.broadcast_arrays( + np.asarray(lat1, dtype=float), + np.asarray(lon1, dtype=float), + np.asarray(lat2, dtype=float), + np.asarray(lon2, dtype=float), + ) + + geocentric_lat1 = _geocentric_latitude(lat1_array) + geocentric_lat2 = _geocentric_latitude(lat2_array) + + sin_lat1 = np.sin(geocentric_lat1) + cos_lat1 = np.cos(geocentric_lat1) + sin_lat2 = np.sin(geocentric_lat2) + cos_lat2 = np.cos(geocentric_lat2) + + # Reducing the longitude difference before converting to radians + # avoids avoidable precision loss for longitudes outside [-180, 180]. + longitude_difference = lon2_array - lon1_array + wrapped_longitude_difference = ( + np.remainder(longitude_difference + 180.0, 360.0) - 180.0 + ) + longitude_difference_rad = np.deg2rad(wrapped_longitude_difference) + sin_dlon = np.sin(longitude_difference_rad) + cos_dlon = np.cos(longitude_difference_rad) + + # The two components below are also the numerator and denominator of + # the point-1-to-point-2 bearing. Together with the dot product they + # form atan2(|cross product|, dot product), a stable equivalent of the + # MATLAB acos/atan distance calculation. + east = cos_lat2 * sin_dlon + north = ( + cos_lat1 * sin_lat2 + - sin_lat1 * cos_lat2 * cos_dlon + ) + dot_product = ( + sin_lat1 * sin_lat2 + + cos_lat1 * cos_lat2 * cos_dlon + ) + cross_product_norm = np.hypot(east, north) + + delta = np.rad2deg(np.arctan2(cross_product_norm, dot_product)) + baz = _normalize_azimuth(np.arctan2(east, north)) + + # Reverse bearing: point 2 to point 1. This is MATLAB's ``dazs`` and + # seispy's historical ``az`` attribute. + reverse_east = -cos_lat1 * sin_dlon + reverse_north = ( + cos_lat2 * sin_lat1 + - sin_lat2 * cos_lat1 * cos_dlon + ) + az = _normalize_azimuth(np.arctan2(reverse_east, reverse_north)) + + # Longitudes differing by full rotations identify the same point. + # At either pole, longitude is immaterial. Bearings for coincident + # points are undefined, so retain distaz.m's established zero value. + same_latitude = lat1_array == lat2_array + same_longitude = np.remainder(longitude_difference, 360.0) == 0.0 + same_pole = same_latitude & (np.abs(lat1_array) == 90.0) + coincident = same_latitude & (same_longitude | same_pole) + + delta = np.where(coincident, 0.0, delta) + az = np.where(coincident, 0.0, az) + baz = np.where(coincident, 0.0, baz) + + self.delta = _scalar_or_array(delta) + self.az = _scalar_or_array(az) + self.baz = _scalar_or_array(baz) def getDelta(self): return self.delta diff --git a/pytomoatt/src_rec.py b/pytomoatt/src_rec.py index 368ceb8..7d3d2c1 100644 --- a/pytomoatt/src_rec.py +++ b/pytomoatt/src_rec.py @@ -3,13 +3,16 @@ import pandas as pd from .distaz import DistAZ from .setuplog import SetupLog +from .utils import _EARTH_RADIUS_KM from .utils.src_rec_utils import define_rec_cols, setup_rec_points_dd, \ get_rec_points_types, update_position, \ linear_regression from sklearn.metrics.pairwise import haversine_distances import copy from io import StringIO +from numbers import Real import os +from urllib.parse import urlparse pd.options.mode.chained_assignment = None # default='warn' @@ -220,6 +223,11 @@ def read(cls, fname: str, dist_in_data=False, name_net_and_sta=False, **kwargs): """ Read source <--> receiver file to pandas.DataFrame + Source indices in the file are treated as local labels and may be + duplicated. Sources are reassigned consecutive ``src_index`` values + from zero in file order, and all receiver records are remapped to the + new indices using source blocks and unique event IDs. + :param fname: Path to src_rec file :type fname: str :param dist_in_data: Whether distance is included in the src_rec file @@ -230,7 +238,14 @@ def read(cls, fname: str, dist_in_data=False, name_net_and_sta=False, **kwargs): :rtype: SrcRec """ sr = cls(fname=fname, **kwargs) - if not os.path.exists(fname): + parsed_url = urlparse(str(fname)) + is_remote = ( + parsed_url.scheme in {"http", "https"} + and bool(parsed_url.netloc) + ) + if os.path.exists(fname): + src_rec_data = fname + elif is_remote: sr.log.SrcReclog.info("Downloading src_rec file from {}".format(fname)) try: from .utils.src_rec_utils import download_src_rec_file @@ -242,18 +257,47 @@ def read(cls, fname: str, dist_in_data=False, name_net_and_sta=False, **kwargs): sr.log.SrcReclog.error("Failed to download src_rec file from {}".format(fname)) return sr else: - src_rec_data = fname + raise FileNotFoundError(f"src_rec file not found: {fname}") + alldf = pd.read_csv( src_rec_data, sep=r"\s+", header=None, comment="#", low_memory=False, dtype={12: str} ) last_col_src = 12 dd_col = 11 - # this is a source line if the last column is not NaN - # sr.src_points = alldf[pd.notna(alldf[last_col_src])] - sr.src_points = alldf[~(alldf[dd_col].astype(str).str.contains("cs")| \ - alldf[dd_col].astype(str).str.contains("cr")| \ - pd.isna(alldf[last_col_src]))] + source_mask = ~( + alldf[dd_col].astype(str).str.contains("cs") + | alldf[dd_col].astype(str).str.contains("cr") + | pd.isna(alldf[last_col_src]) + ) + source_event_ids = alldf.loc[source_mask, last_col_src].astype(str) + duplicated_event_ids = source_event_ids[ + source_event_ids.duplicated(keep=False) + ].unique() + if duplicated_event_ids.size: + raise ValueError( + "event_id must be unique; duplicated values: {}".format( + ", ".join(duplicated_event_ids) + ) + ) + + # File src_index values are not guaranteed to be unique. Assign each + # source block a new consecutive index and propagate it to all records + # belonging to that block before parsing the individual record types. + block_src_index = pd.Series(np.nan, index=alldf.index) + block_src_index.loc[source_mask] = np.arange(source_mask.sum()) + block_src_index = block_src_index.ffill() + if block_src_index.isna().any(): + raise ValueError( + "Receiver data found before the first source record" + ) + alldf.loc[:, 0] = block_src_index.astype(int) + + event_id_to_src_index = dict(zip( + source_event_ids, + np.arange(source_mask.sum()), + )) + sr.src_points = alldf[source_mask] # add weight column if not included if sr.src_points.shape[1] == last_col_src + 1: # add another column for weight @@ -378,6 +422,20 @@ def read(cls, fname: str, dist_in_data=False, name_net_and_sta=False, **kwargs): cols, data_type = setup_rec_points_dd(type='cr') sr.rec_points_cr.columns = cols sr.rec_points_cr = sr.rec_points_cr.astype(data_type) + if not sr.rec_points_cr.empty: + mapped_src_index2 = sr.rec_points_cr["event_id2"].map( + event_id_to_src_index + ) + if mapped_src_index2.isna().any(): + missing_event_ids = sr.rec_points_cr.loc[ + mapped_src_index2.isna(), "event_id2" + ].unique() + raise ValueError( + "Unknown event_id2 in common-receiver data: {}".format( + ", ".join(missing_event_ids) + ) + ) + sr.rec_points_cr["src_index2"] = mapped_src_index2.astype(int) # read common source data sr.rec_points_cs = alldf[ @@ -1330,6 +1388,94 @@ def select_by_linear_regression(self, std_multiplier=3.0, ) return regression_params + def select_by_constant_velocity( + self, + velocity, + tt_res_range, + recalc_dist=False, + **kwargs, + ): + """Select arrivals around a constant-velocity travel-time curve. + + An arrival is retained when its travel-time residual satisfies + + ``tt_res_range[0] <= tt - distance_km / velocity <= tt_res_range[1]``. + + ``dist_deg`` is converted to epicentral arc distance in kilometres + using the package Earth radius. The residual bounds are inclusive and + may be asymmetric. Non-finite distances or travel times are removed. + + .. note:: + This criterion only applies to absolute travel-time data in + :attr:`rec_points`. A double-difference record is removed when + either corresponding absolute arrival is rejected. + + :param velocity: Constant reference velocity in kilometres per second. + :type velocity: float + :param tt_res_range: Inclusive travel-time residual range in seconds, + ``[min_residual, max_residual]``. + :type tt_res_range: list or tuple + :param recalc_dist: Recalculate epicentral distance even when + ``dist_deg`` exists, defaults to False. + :type recalc_dist: bool + """ + if ( + not isinstance(velocity, Real) + or isinstance(velocity, (bool, np.bool_)) + or not np.isfinite(velocity) + or velocity <= 0 + ): + raise ValueError("velocity must be a positive finite number") + + try: + min_residual, max_residual = tt_res_range + except (TypeError, ValueError): + raise ValueError( + "tt_res_range must contain exactly two finite numbers" + ) from None + if not all( + isinstance(value, Real) + and not isinstance(value, (bool, np.bool_)) + and np.isfinite(value) + for value in (min_residual, max_residual) + ): + raise ValueError( + "tt_res_range must contain exactly two finite numbers" + ) + if min_residual > max_residual: + raise ValueError( + "tt_res_range minimum must not exceed its maximum" + ) + + self.log.SrcReclog.info( + "rec_points before constant-velocity selection: {}".format( + self.rec_points.shape[0] + ) + ) + if ("dist_deg" not in self.rec_points) or recalc_dist: + self.log.SrcReclog.info("Calculating epicentral distance...") + self.calc_distaz() + + distances_deg = self.rec_points["dist_deg"].to_numpy(dtype=float) + travel_times = self.rec_points["tt"].to_numpy(dtype=float) + distances_km = np.deg2rad(distances_deg) * _EARTH_RADIUS_KM + residuals = travel_times - distances_km / velocity + keep = ( + np.isfinite(distances_deg) + & np.isfinite(travel_times) + & (residuals >= min_residual) + & (residuals <= max_residual) + ) + + self.rec_points = self.rec_points.loc[keep] + self._filter_double_difference_by_arrivals() + self.update(**kwargs) + self.log.SrcReclog.info( + "rec_points after constant-velocity selection: {}".format( + self.rec_points.shape[0] + ) + ) + def select_by_azi_gap(self, max_azi_gap: float, **kwargs): """Select sources with azimuthal gap greater and equal than a number @@ -1935,20 +2081,73 @@ def from_seispy(cls, rf_path: str): return sr - # implemented in vis.py - def plot(self, weight=False, fname=None): - """Plot source and receivers for preview + # implemented in utils/vis.py + def plot(self, color_by="depth", fname=None, **kwargs): + """Plot sources and receivers with source-depth sections. - :param weight: Draw colors of weights, defaults to False - :type weight: bool, optional + :param color_by: Source attribute used for color mapping; either + ``"depth"`` or ``"weight"``, defaults to ``"depth"`` + :type color_by: str, optional :param fname: Path to output file, defaults to None :type fname: str, optional + :param kwargs: Additional keyword arguments passed to Matplotlib's + ``Axes.scatter`` for source points, such as ``cmap``, + ``s``, ``alpha``, ``marker``, ``vmin`` and ``vmax`` :return: matplotlib figure :rtype: matplotlib.figure.Figure """ - from .vis import plot_srcrec + from .utils.vis import plot_src_rec + + return plot_src_rec( + self, color_by=color_by, fname=fname, **kwargs + ) - return plot_srcrec(self, weight=weight, fname=fname) + def plot_travel_time( + self, + color="tab:blue", + fname=None, + fig=None, + ylim="adaptive", + **kwargs, + ): + """Plot absolute travel time against epicentral distance. + + If ``dist_deg`` is unavailable, it is calculated before plotting. + The returned Matplotlib figure remains editable; use + ``figure.axes[0]`` to add lines, annotations, or other content. + + :param color: Matplotlib-compatible point color, defaults to + ``"tab:blue"``. + :param fname: Path to output file, defaults to None. + :type fname: str, optional + :param fig: Existing Matplotlib figure on which to draw, defaults to + None. Its current axis is used, or one is created when + necessary. + :type fig: matplotlib.figure.Figure, optional + :param ylim: Y-axis scaling strategy. Use ``"adaptive"`` to derive + limits from travel times at the minimum and maximum + epicentral distances, ``"auto"`` for Matplotlib + autoscaling, ``"inherit"`` to preserve the limits of an + existing figure, or pass ``(min, max)`` explicitly. + :param kwargs: Additional keyword arguments passed to Matplotlib's + ``Axes.scatter``. + :return: Matplotlib figure. + :rtype: matplotlib.figure.Figure + """ + if "dist_deg" not in self.rec_points: + self.log.SrcReclog.info("Calculating epicentral distance...") + self.calc_distaz() + + from .utils.vis import plot_travel_time + + return plot_travel_time( + self, + color=color, + fname=fname, + fig=fig, + ylim=ylim, + **kwargs, + ) if __name__ == "__main__": diff --git a/pytomoatt/utils/vis.py b/pytomoatt/utils/vis.py new file mode 100644 index 0000000..280823d --- /dev/null +++ b/pytomoatt/utils/vis.py @@ -0,0 +1,344 @@ +"""Visualization helpers for source--receiver data.""" + +from __future__ import annotations + +from os import PathLike +from typing import TYPE_CHECKING, Literal + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.figure import Figure + +if TYPE_CHECKING: + from pytomoatt.src_rec import SrcRec + + +def _axis_limits(values: np.ndarray, padding: float = 0.1) -> tuple[float, float]: + """Calculate padded limits, including for a constant-valued coordinate.""" + lower = float(np.min(values)) + upper = float(np.max(values)) + span = upper - lower + pad = span * padding if span else max(abs(lower) * padding, 0.5) + return lower - pad, upper + pad + + +def _endpoint_travel_time_limits( + distances: np.ndarray, + travel_times: np.ndarray, + padding: float = 0.05, +) -> tuple[float, float]: + """Calculate travel-time limits from the minimum/maximum distances.""" + minimum_distance = np.min(distances) + maximum_distance = np.max(distances) + if np.isclose(minimum_distance, maximum_distance): + return _axis_limits(travel_times, padding=padding) + + minimum_distance_times = travel_times[ + np.isclose(distances, minimum_distance) + ] + maximum_distance_times = travel_times[ + np.isclose(distances, maximum_distance) + ] + endpoint_times = np.array([ + np.min(minimum_distance_times), + np.max(maximum_distance_times), + ]) + return _axis_limits(endpoint_times, padding=padding) + + +def _check_columns(data, columns: set[str], name: str) -> None: + missing = columns.difference(data.columns) + if missing: + missing_names = ", ".join(sorted(missing)) + raise ValueError(f"{name} is missing required columns: {missing_names}") + + +def plot_src_rec( + src_rec: "SrcRec", + *, + color_by: Literal["depth", "weight"] = "depth", + fname: str | PathLike[str] | None = None, + **kwargs, +) -> Figure: + """Plot source and receiver locations with two source-depth sections. + + Parameters + ---------- + src_rec + A :class:`~pytomoatt.src_rec.SrcRec` instance. + color_by + Source attribute used for color mapping: ``"depth"`` or + ``"weight"``. Receiver symbols remain red so they can be + distinguished from sources. + fname + Optional output path. The format is inferred by Matplotlib from the + filename extension. + **kwargs + Additional keyword arguments passed to Matplotlib's + :meth:`~matplotlib.axes.Axes.scatter` for source points. Common + options include ``cmap``, ``s``, ``alpha``, ``marker``, ``vmin``, + ``vmax``, ``norm``, ``edgecolors`` and ``linewidths``. + + Returns + ------- + matplotlib.figure.Figure + The created figure. + + Notes + ----- + This function does not add temporary columns to ``src_rec`` or otherwise + mutate it. + """ + sources = src_rec.src_points + _check_columns(sources, {"evlo", "evla", "evdp"}, "src_points") + if sources.empty: + raise ValueError("Cannot plot an SrcRec object without sources") + + source_values = sources[["evlo", "evla", "evdp"]].to_numpy(dtype=float) + finite_sources = np.isfinite(source_values).all(axis=1) + if not finite_sources.any(): + raise ValueError("src_points contains no finite source coordinates") + + source_lon, source_lat, source_depth = source_values[finite_sources].T + + receivers = src_rec.receivers + if receivers is None or receivers.empty: + receiver_records = src_rec.rec_points + if receiver_records is not None and not receiver_records.empty: + receivers = receiver_records.drop_duplicates(subset="staname") + + if receivers is None or receivers.empty: + receiver_lon = np.empty(0) + receiver_lat = np.empty(0) + else: + _check_columns(receivers, {"stlo", "stla"}, "receivers") + receiver_values = receivers[["stlo", "stla"]].to_numpy(dtype=float) + finite_receivers = np.isfinite(receiver_values).all(axis=1) + receiver_lon, receiver_lat = receiver_values[finite_receivers].T + + all_lon = np.concatenate((source_lon, receiver_lon)) + all_lat = np.concatenate((source_lat, receiver_lat)) + lon_limits = _axis_limits(all_lon) + lat_limits = _axis_limits(all_lat) + depth_limits = _axis_limits(np.concatenate((source_depth, [0.0])), padding=0.05) + + if color_by not in {"depth", "weight"}: + raise ValueError("color_by must be either 'depth' or 'weight'") + + if color_by == "weight": + _check_columns(sources, {"weight"}, "src_points") + color_values = np.asarray(sources.loc[finite_sources, "weight"], dtype=float) + if not np.isfinite(color_values).all(): + raise ValueError("src_points contains non-finite source weights") + colorbar_label = "Source weight" + cmap = "plasma" + else: + color_values = source_depth + colorbar_label = "Source depth (km)" + cmap = "viridis" + + figure = plt.figure(figsize=(8, 8), layout="constrained") + grid = figure.add_gridspec( + 2, + 2, + width_ratios=(4, 1.35), + height_ratios=(4, 1.35), + ) + map_axis = figure.add_subplot(grid[0, 0]) + latitude_depth_axis = figure.add_subplot(grid[0, 1], sharey=map_axis) + longitude_depth_axis = figure.add_subplot(grid[1, 0], sharex=map_axis) + colorbar_host = figure.add_subplot(grid[1, 1]) + colorbar_host.set_axis_off() + colorbar_axis = colorbar_host.inset_axes((0.05, 0.52, 0.9, 0.12)) + + scatter_options = { + "c": color_values, + "cmap": cmap, + "s": 8, + "label": "Sources", + } + scatter_options.update(kwargs) + source_scatter = map_axis.scatter(source_lon, source_lat, **scatter_options) + latitude_depth_axis.scatter(source_depth, source_lat, **scatter_options) + longitude_depth_axis.scatter(source_lon, source_depth, **scatter_options) + + if receiver_lon.size: + map_axis.scatter( + receiver_lon, + receiver_lat, + c="tab:red", + edgecolors="white", + linewidths=0.5, + label="Receivers", + marker="v", + s=55, + ) + + map_axis.set( + xlabel="Longitude", + ylabel="Latitude", + xlim=lon_limits, + ylim=lat_limits, + ) + map_axis.legend() + + latitude_depth_axis.set( + xlabel="Depth (km)", + ylabel="Latitude", + xlim=depth_limits, + ylim=lat_limits, + ) + longitude_depth_axis.set( + xlabel="Longitude", + ylabel="Depth (km)", + xlim=lon_limits, + ylim=depth_limits, + ) + longitude_depth_axis.invert_yaxis() + + colorbar = figure.colorbar( + source_scatter, + cax=colorbar_axis, + orientation="horizontal", + ) + colorbar.set_label(colorbar_label) + + if fname is not None: + figure.savefig(fname, dpi=300, bbox_inches="tight") + + return figure + + +def fig_ev_st_distribution_dep( + src_rec: "SrcRec", + fname: str | PathLike[str] | None = None, + *, + color_by: Literal["depth", "weight"] = "depth", + **kwargs, +) -> Figure: + """Backward-compatible name for :func:`plot_src_rec`.""" + return plot_src_rec(src_rec, color_by=color_by, fname=fname, **kwargs) + + +def plot_travel_time( + src_rec: "SrcRec", + *, + color="tab:blue", + fname: str | PathLike[str] | None = None, + fig: Figure | None = None, + ylim="adaptive", + **kwargs, +) -> Figure: + """Plot absolute travel time against epicentral distance. + + Parameters + ---------- + src_rec + A :class:`~pytomoatt.src_rec.SrcRec` instance whose ``rec_points`` + contains ``dist_deg`` and ``tt`` columns. + color + Any Matplotlib-compatible color specification for the points. + fname + Optional output path. The format is inferred by Matplotlib from the + filename extension. + fig + Optional existing Matplotlib figure. Points are added to its current + axis; an axis is created when the figure has none. + ylim + Y-axis scaling strategy. ``"adaptive"`` (default) uses travel times + at the minimum and maximum epicentral distances, ``"auto"`` or + ``None`` uses Matplotlib autoscaling, ``"inherit"`` preserves the + current limits of an existing figure, and a ``(min, max)`` pair sets + explicit limits. + **kwargs + Additional keyword arguments passed to Matplotlib's + :meth:`~matplotlib.axes.Axes.scatter`. + + Returns + ------- + matplotlib.figure.Figure + The created figure. Its axis is available as ``figure.axes[0]`` for + adding lines, annotations, or other content. + """ + records = src_rec.rec_points + _check_columns(records, {"dist_deg", "tt"}, "rec_points") + if records.empty: + raise ValueError("Cannot plot travel times without receiver records") + + values = records[["dist_deg", "tt"]].to_numpy(dtype=float) + finite = np.isfinite(values).all(axis=1) + if not finite.any(): + raise ValueError("rec_points contains no finite distance--time pairs") + + inherited_limits = None + inherit_requested = isinstance(ylim, str) and ylim == "inherit" + if fig is None: + if inherit_requested: + raise ValueError( + "ylim='inherit' requires an existing figure and axis" + ) + figure, axis = plt.subplots(figsize=(6, 4.5), layout="constrained") + elif isinstance(fig, Figure): + figure = fig + if inherit_requested and not figure.axes: + raise ValueError( + "ylim='inherit' requires a figure with an existing axis" + ) + axis = figure.gca() + if inherit_requested: + inherited_limits = axis.get_ylim() + else: + raise TypeError("fig must be a matplotlib.figure.Figure or None") + + scatter_options = {"color": color, "s": 4} + scatter_options.update(kwargs) + axis.scatter(values[finite, 0], values[finite, 1], **scatter_options) + axis.set( + xlabel="Epicentral distance (degree)", + ylabel="Travel time (s)", + ) + + travel_times = values[finite, 1] + distances = values[finite, 0] + if isinstance(ylim, str): + if ylim == "adaptive": + axis.set_ylim( + _endpoint_travel_time_limits(distances, travel_times) + ) + elif ylim == "inherit": + if inherited_limits is None: + raise ValueError( + "ylim='inherit' requires an existing figure and axis" + ) + axis.set_ylim(inherited_limits) + elif ylim != "auto": + raise ValueError( + "ylim must be 'adaptive', 'auto', 'inherit', None, " + "or a (min, max) pair" + ) + elif ylim is not None: + try: + limits = np.asarray(ylim, dtype=float) + except (TypeError, ValueError): + raise ValueError( + "ylim must be 'adaptive', 'auto', 'inherit', None, " + "or a (min, max) pair" + ) from None + if limits.shape != (2,): + raise ValueError( + "ylim must be 'adaptive', 'auto', 'inherit', None, " + "or a (min, max) pair" + ) + lower_limit, upper_limit = limits + if ( + not np.isfinite(lower_limit) + or not np.isfinite(upper_limit) + or lower_limit >= upper_limit + ): + raise ValueError("ylim must contain two increasing finite values") + axis.set_ylim(lower_limit, upper_limit) + + if fname is not None: + figure.savefig(fname, dpi=300, bbox_inches="tight", facecolor="white") + + return figure diff --git a/test/src_rec_duplicate_index.dat b/test/src_rec_duplicate_index.dat new file mode 100644 index 0000000..256516f --- /dev/null +++ b/test/src_rec_duplicate_index.dat @@ -0,0 +1,6 @@ +5 2020 1 1 0 0 0.0 10.0000 20.0000 5.0000 3.0000 3 EVT_A 1.0000 +5 0 STA1 11.0000 21.0000 0.0000 P 2.0000 1.0000 +5 0 STA1 11.0000 21.0000 0.0000 5 EVT_B 12.0000 22.0000 6.0000 P,cr -1.0000 1.0000 +5 0 STA1 11.0000 21.0000 0.0000 1 STA2 13.0000 23.0000 0.0000 P,cs -1.0000 1.0000 +5 2020 1 2 0 0 0.0 12.0000 22.0000 6.0000 3.0000 1 EVT_B 1.0000 +5 0 STA2 13.0000 23.0000 0.0000 P 3.0000 1.0000 diff --git a/test/test_src_rec.py b/test/test_src_rec.py index 3df0f37..5c7f6f8 100644 --- a/test/test_src_rec.py +++ b/test/test_src_rec.py @@ -9,15 +9,48 @@ linear_regression, ) from os.path import dirname, join +from tempfile import TemporaryDirectory from unittest.mock import MagicMock, patch import pandas as pd import numpy as np import io +import matplotlib.pyplot as plt +from matplotlib.colors import to_rgba class TestSrcRec(unittest.TestCase): fname: str = join(dirname(dirname(__file__)), 'examples', 'src_rec_file_eg') fname1: str = join(dirname(__file__), 'test_srcrec_a.dat') + duplicate_index_fname: str = join( + dirname(__file__), 'src_rec_duplicate_index.dat' + ) + + def test_read_missing_local_file(self): + missing_file = join(dirname(__file__), "missing_src_rec.dat") + + with self.assertRaisesRegex( + FileNotFoundError, "src_rec file not found" + ): + SrcRec.read(missing_file) + + def test_read_reindexes_duplicate_file_src_indices(self): + sr = SrcRec.read(self.duplicate_index_fname) + + self.assertEqual(sr.src_points.index.tolist(), [0, 1]) + self.assertEqual(sr.src_points['event_id'].tolist(), ['EVT_A', 'EVT_B']) + self.assertEqual(sr.rec_points['src_index'].tolist(), [0, 1]) + self.assertEqual(sr.rec_points_cs['src_index'].tolist(), [0]) + self.assertEqual(sr.rec_points_cr['src_index'].tolist(), [0]) + self.assertEqual(sr.rec_points_cr['src_index2'].tolist(), [1]) + + with TemporaryDirectory() as directory: + output_file = join(directory, 'src_rec.dat') + sr.write(output_file) + reread = SrcRec.read(output_file) + + self.assertEqual(reread.src_points.index.tolist(), [0, 1]) + self.assertEqual(reread.rec_points['src_index'].tolist(), [0, 1]) + self.assertEqual(reread.rec_points_cr['src_index2'].tolist(), [1]) def test_subcase_01(self): sr = SrcRec.read(self.fname) @@ -117,6 +150,185 @@ def test_select_by_linear_regression(self): self.assertEqual(sr.rec_points_cr.iloc[0]['staname'], 'STA00') update.assert_called_once_with() + def test_select_by_constant_velocity(self): + sr = SrcRec('unused') + distance = np.array([0.0, 1.0, 2.0, 3.0, 0.0]) + reference_tt = np.deg2rad(distance) * 6371.0 / 10.0 + sr.rec_points = pd.DataFrame({ + 'src_index': [0, 0, 0, 0, 1], + 'staname': ['STA0', 'STA1', 'STA2', 'STA3', 'STA0'], + 'dist_deg': distance, + 'tt': reference_tt + np.array([-1.0, 0.0, 2.0, 2.1, 0.0]), + 'phase': ['P'] * 5, + }) + sr.rec_points_cs = pd.DataFrame({ + 'src_index': [0, 0], + 'staname1': ['STA0', 'STA0'], + 'staname2': ['STA1', 'STA3'], + 'phase': ['P,cs', 'P,cs'], + }) + sr.rec_points_cr = pd.DataFrame({ + 'src_index': [0, 0], + 'src_index2': [1, 1], + 'staname': ['STA0', 'STA3'], + 'phase': ['P,cr', 'P,cr'], + }) + + with patch.object(sr, 'update') as update: + sr.select_by_constant_velocity( + velocity=10.0, + tt_res_range=(-1.0, 2.0), + ) + + self.assertEqual(sr.rec_points.shape[0], 4) + self.assertNotIn('STA3', sr.rec_points['staname'].values) + self.assertEqual(sr.rec_points_cs.shape[0], 1) + self.assertEqual(sr.rec_points_cr.shape[0], 1) + update.assert_called_once_with() + + def test_select_by_constant_velocity_validates_parameters(self): + sr = SrcRec('unused') + + with self.assertRaisesRegex(ValueError, 'velocity'): + sr.select_by_constant_velocity(0.0, (-1.0, 1.0)) + with self.assertRaisesRegex(ValueError, 'tt_res_range'): + sr.select_by_constant_velocity(1.0, (2.0, 1.0)) + + def test_plot(self): + sr = SrcRec.read(self.fname) + original_columns = sr.src_points.columns.copy() + + figure = sr.plot() + figure.canvas.draw() + + self.assertIsNotNone(figure) + self.assertTrue(np.allclose(figure.get_size_inches(), (8.0, 8.0))) + self.assertTrue(original_columns.equals(sr.src_points.columns)) + map_position = figure.axes[0].get_position() + latitude_depth_position = figure.axes[1].get_position() + longitude_depth_position = figure.axes[2].get_position() + self.assertAlmostEqual(map_position.y0, latitude_depth_position.y0) + self.assertAlmostEqual(map_position.y1, latitude_depth_position.y1) + self.assertAlmostEqual(map_position.x0, longitude_depth_position.x0) + self.assertAlmostEqual(map_position.x1, longitude_depth_position.x1) + plt.close(figure) + + def test_plot_source_only(self): + sr = SrcRec.read(self.fname, src_only=True) + + figure = sr.plot(color_by="weight") + + self.assertIsNotNone(figure) + plt.close(figure) + + def test_plot_rejects_invalid_color_by(self): + sr = SrcRec.read(self.fname, src_only=True) + + with self.assertRaisesRegex(ValueError, "color_by"): + sr.plot(color_by="magnitude") + + def test_plot_accepts_matplotlib_scatter_options(self): + sr = SrcRec.read(self.fname, src_only=True) + + figure = sr.plot(cmap="jet", s=12, alpha=0.5, marker="x") + source_collection = figure.axes[0].collections[0] + + self.assertEqual(source_collection.get_cmap().name, "jet") + self.assertEqual(source_collection.get_sizes()[0], 12) + self.assertEqual(source_collection.get_alpha(), 0.5) + plt.close(figure) + + def test_plot_travel_time_returns_editable_figure(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_deg': [0.0, 1.0, np.nan], + 'tt': [1.0, 3.0, 5.0], + }) + + figure = sr.plot_travel_time(color='red', s=12, alpha=0.5) + axis = figure.axes[0] + collection = axis.collections[0] + line = axis.plot([0.0, 1.0], [1.0, 3.0])[0] + + self.assertEqual(collection.get_offsets().shape[0], 2) + self.assertTrue(np.allclose(figure.get_size_inches(), (6.0, 4.5))) + self.assertTrue( + np.allclose(collection.get_facecolors()[0], to_rgba('red', 0.5)) + ) + self.assertIn(line, axis.lines) + plt.close(figure) + + def test_plot_travel_time_calculates_missing_distance(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({'tt': [1.0, 2.0]}) + + def add_distance(): + sr.rec_points['dist_deg'] = [0.0, 1.0] + + with patch.object(sr, 'calc_distaz', side_effect=add_distance) as calc: + figure = sr.plot_travel_time() + + calc.assert_called_once_with() + plt.close(figure) + + def test_plot_travel_time_uses_existing_figure(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_deg': [0.0, 1.0], + 'tt': [1.0, 3.0], + }) + existing_figure, axis = plt.subplots() + existing_line = axis.plot([0.0, 1.0], [0.0, 2.0])[0] + + returned_figure = sr.plot_travel_time( + fig=existing_figure, + color='red', + ) + + self.assertIs(returned_figure, existing_figure) + self.assertIn(existing_line, axis.lines) + self.assertEqual(len(axis.collections), 1) + plt.close(existing_figure) + + def test_plot_travel_time_inherits_y_limits(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_deg': [0.0, 1.0], + 'tt': [1.0, 30.0], + }) + existing_figure, axis = plt.subplots() + axis.set_ylim(5.0, 20.0) + + returned_figure = sr.plot_travel_time( + fig=existing_figure, + ylim='inherit', + ) + + self.assertIs(returned_figure, existing_figure) + self.assertEqual(axis.get_ylim(), (5.0, 20.0)) + plt.close(existing_figure) + + def test_plot_travel_time_y_limits(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_deg': [0.0, 1.0, 2.0, 3.0, 4.0], + 'tt': [1.0, 2.0, 1000.0, 4.0, 5.0], + }) + + adaptive_figure = sr.plot_travel_time() + auto_figure = sr.plot_travel_time(ylim='auto') + explicit_figure = sr.plot_travel_time(ylim=(0.0, 10.0)) + + adaptive_limits = adaptive_figure.axes[0].get_ylim() + self.assertLess(adaptive_limits[0], 1.0) + self.assertGreater(adaptive_limits[1], 5.0) + self.assertLess(adaptive_limits[1], 1000.0) + self.assertGreater(auto_figure.axes[0].get_ylim()[1], 1000.0) + self.assertEqual(explicit_figure.axes[0].get_ylim(), (0.0, 10.0)) + plt.close(adaptive_figure) + plt.close(auto_figure) + plt.close(explicit_figure) + class TestSrcRecUtils(unittest.TestCase): def test_linear_regression(self): From 8de6ba67ec98d876f5f3590b21684afdd761ac61 Mon Sep 17 00:00:00 2001 From: xumi1993 Date: Mon, 3 Aug 2026 00:35:49 +0800 Subject: [PATCH 12/12] Enhance plotting functions and add receiver weighting features - Updated `plot_src_rec` to support coloring receivers by weight and added a colorbar for receiver weights. - Introduced shared normalization for scatter plots when using weights. - Modified `plot_travel_time` to allow selection of distance metrics and updated documentation accordingly. - Added tests for new features including receiver weighting and distance selection in travel time plotting. - Improved handling of receiver conflicts and ensured proper logging of warnings. - Enhanced the `geo_weighting` method to deduplicate receiver names and normalize weights. --- pytomoatt/distaz.py | 3 +- pytomoatt/src_rec.py | 592 ++++++++++++++++++++++++++++++++--------- pytomoatt/utils/vis.py | 164 ++++++++++-- test/test_src_rec.py | 538 ++++++++++++++++++++++++++++++++++++- 4 files changed, 1139 insertions(+), 158 deletions(-) diff --git a/pytomoatt/distaz.py b/pytomoatt/distaz.py index c7d5e6f..3588330 100644 --- a/pytomoatt/distaz.py +++ b/pytomoatt/distaz.py @@ -1,4 +1,5 @@ import numpy as np +from .utils.common import deg2km # Geographic latitude is converted to geocentric latitude before calculating @@ -160,4 +161,4 @@ def getBaz(self): return self.baz def degreesToKilometers(self): - return self.delta * 111.19 \ No newline at end of file + return deg2km(self.delta) \ No newline at end of file diff --git a/pytomoatt/src_rec.py b/pytomoatt/src_rec.py index 7d3d2c1..36ee128 100644 --- a/pytomoatt/src_rec.py +++ b/pytomoatt/src_rec.py @@ -6,7 +6,7 @@ from .utils import _EARTH_RADIUS_KM from .utils.src_rec_utils import define_rec_cols, setup_rec_points_dd, \ get_rec_points_types, update_position, \ - linear_regression + linear_regression as fit_linear_regression from sklearn.metrics.pairwise import haversine_distances import copy from io import StringIO @@ -130,6 +130,8 @@ def rec_points(self): ================ =========================================================================== ``netname`` Name of the network (when ``name_net_and_sta=True`` in ``SrcRec.read``) ``dist_deg`` Epicentral distance in deg (when ``dist_in_data=True`` in ``SrcRec.read``) + ``dist_km`` Epicentral distance in km + ``dist_3d_km`` Three-dimensional source--receiver distance in km ================ =========================================================================== """ @@ -606,9 +608,25 @@ def update_unique_src_rec(self): [receivers, self.rec_points_cr[ ["staname", "stla", "stlo", "stel"] ].values]) - self.receivers = pd.DataFrame( + receiver_rows = pd.DataFrame( receivers, columns=rec_col ).drop_duplicates(ignore_index=True) + conflicting_receiver_mask = receiver_rows.duplicated( + subset="staname", keep=False + ) + conflicting_receiver_names = receiver_rows.loc[ + conflicting_receiver_mask, "staname" + ].drop_duplicates().astype(str).tolist() + if conflicting_receiver_names: + self.log.SrcReclog.warning( + "Found conflicting coordinates/elevations for %d " + "receiver(s): %s.", + len(conflicting_receiver_names), + ", ".join(conflicting_receiver_names), + ) + self.receivers = receiver_rows.drop_duplicates( + subset="staname", keep="first", ignore_index=True + ) self.receivers = self.receivers.astype( { "stla": float, @@ -1213,76 +1231,183 @@ def select_by_depth(self, dep_min_max, **kwargs): ) def calc_distaz(self): - """Calculate epicentral distance and azimuth for each receiver""" - self.rec_points["dist_deg"] = 0.0 - self.rec_points["az"] = 0.0 - self.rec_points["baz"] = 0.0 - rec_group = self.rec_points.groupby("src_index") - for idx, rec in rec_group: - da = DistAZ( - self.src_points.loc[idx]["evla"], - self.src_points.loc[idx]["evlo"], - rec["stla"].values, - rec["stlo"].values, + """Calculate distance and azimuth for each source--receiver pair. + + ``dist_deg`` and ``dist_km`` are the epicentral distance in degrees + and kilometres, respectively. ``dist_3d_km`` is the three-dimensional + Euclidean distance calculated from ``dist_km`` and the vertical + separation. Source depth is in km and receiver elevation is converted + from m to km. + """ + output_columns = ( + "dist_deg", + "dist_km", + "dist_3d_km", + "az", + "baz", + ) + if self.rec_points.empty: + for column in output_columns: + self.rec_points[column] = pd.Series(dtype=float) + return + + source_indices = self.rec_points["src_index"].to_numpy() + missing_source_indices = pd.Index(np.unique(source_indices)).difference( + self.src_points.index + ) + if not missing_source_indices.empty: + missing = ", ".join(map(str, missing_source_indices.tolist())) + raise KeyError( + f"rec_points references missing src_index values: {missing}" ) - self.rec_points.loc[rec.index, "dist_deg"] = da.delta - self.rec_points.loc[rec.index, "az"] = da.az - self.rec_points.loc[rec.index, "baz"] = da.baz - def select_by_distance(self, dist_min_max, recalc_dist=False, **kwargs): - """Select stations in a range of distance + source_rows = self.src_points.reindex(source_indices) + distaz = DistAZ( + source_rows["evla"].to_numpy(dtype=float), + source_rows["evlo"].to_numpy(dtype=float), + self.rec_points["stla"].to_numpy(dtype=float), + self.rec_points["stlo"].to_numpy(dtype=float), + ) + epicentral_distance_km = np.asarray( + distaz.degreesToKilometers(), dtype=float + ) + vertical_distance_km = ( + source_rows["evdp"].to_numpy(dtype=float) + + self.rec_points["stel"].to_numpy(dtype=float) / 1000.0 + ) + + self.rec_points["dist_deg"] = distaz.delta + self.rec_points["dist_km"] = epicentral_distance_km + self.rec_points["dist_3d_km"] = np.hypot( + epicentral_distance_km, vertical_distance_km + ) + self.rec_points["az"] = distaz.az + self.rec_points["baz"] = distaz.baz + + def select_by_distance( + self, + dist_min_max, + recalc_dist=False, + distance="dist_deg", + **kwargs, + ): + """Select source--receiver pairs in a range of distance. .. note:: - This criteria only works for absolute travel time data. + Absolute arrivals are retained when their epicentral distance is + in range. Common-source and common-receiver records are retained + only when both of their source--receiver distances are in range. - :param dist_min_max: limit of distance in deg, ``[dist_min, dist_max]`` + :param dist_min_max: Distance limits, ``[dist_min, dist_max]``. Their + unit follows ``distance``. :type dist_min_max: list or tuple + :param recalc_dist: Recalculate distance fields even when the selected + field exists, defaults to ``False``. + :type recalc_dist: bool + :param distance: Distance field used for selection. Choose + ``"dist_deg"`` for degrees or ``"dist_km"`` for km, + defaults to ``"dist_deg"``. + :type distance: str """ + if distance not in {"dist_deg", "dist_km"}: + raise ValueError( + "distance must be either 'dist_deg' or 'dist_km'" + ) + self.log.SrcReclog.info( "rec_points before selection: {}".format(self._count_records()) ) - # rec_group = self.rec_points.groupby('src_index') - if ("dist_deg" not in self.rec_points) or recalc_dist: + if (distance not in self.rec_points) or recalc_dist: self.log.SrcReclog.info("Calculating epicentral distance...") self.calc_distaz() - elif not recalc_dist: - pass - else: - self.log.SrcReclog.error( - "No such field of dist, please set up recalc_dist to True" - ) - # for _, rec in rec_group: - mask = (self.rec_points["dist_deg"] < dist_min_max[0]) | ( - self.rec_points["dist_deg"] > dist_min_max[1] + + selected_distance = self.rec_points[distance] + mask = (selected_distance < dist_min_max[0]) | ( + selected_distance > dist_min_max[1] ) drop_idx = self.rec_points[mask].index self.rec_points = self.rec_points.drop(index=drop_idx) + self._filter_double_difference_by_distance( + dist_min_max, distance=distance + ) self.update(**kwargs) self.log.SrcReclog.info( "rec_points after selection: {}".format(self._count_records()) ) @staticmethod - def _regression_keep_mask(records, std_multiplier): + def _regression_keep_mask(records, std_multiplier, distance="dist_deg"): """Return a mask for finite records within the residual limit.""" - finite = np.isfinite(records["dist_deg"].to_numpy(dtype=float)) & \ + finite = np.isfinite(records[distance].to_numpy(dtype=float)) & \ np.isfinite(records["tt"].to_numpy(dtype=float)) valid = records.loc[finite] keep = pd.Series(False, index=records.index, dtype=bool) - if len(valid) < 2 or valid["dist_deg"].nunique() < 2: + if len(valid) < 2 or valid[distance].nunique() < 2: keep.loc[valid.index] = True return keep, None - slope, intercept, residual_std = linear_regression( - valid["dist_deg"], valid["tt"] + slope, intercept, residual_std = fit_linear_regression( + valid[distance], valid["tt"] ) - residual = valid["tt"] - (slope * valid["dist_deg"] + intercept) + residual = valid["tt"] - (slope * valid[distance] + intercept) keep.loc[valid.index] = ( np.isclose(residual, 0.0) if residual_std == 0 else np.abs(residual) <= std_multiplier * residual_std ) - return keep, (slope, intercept) + return keep, (slope, intercept, residual_std) + + def linear_regression( + self, + phase=None, + recalc_dist=False, + distance="dist_3d_km", + ): + """Fit travel time as a linear function of distance. + + This method only computes regression parameters; it does not filter + travel-time records. Non-finite distance/travel-time pairs are ignored + by the fit. + + :param phase: Fit only this phase. When ``None``, all phases are used, + defaults to ``None``. + :type phase: str or None + :param recalc_dist: Recalculate distance fields even when the selected + field exists, defaults to ``False``. + :type recalc_dist: bool + :param distance: Independent variable used by the fit. Choose + ``"dist_deg"`` for epicentral distance in degrees or + ``"dist_km"`` for epicentral distance in km, or + ``"dist_3d_km"`` for three-dimensional + source--receiver distance in km, defaults to + ``"dist_3d_km"``. + :type distance: str + :return: ``(slope, intercept, residual_std)``. Slope is in s/degree + for ``dist_deg`` or s/km for the kilometre fields; the other + values are in s. + :rtype: tuple of float + """ + if phase is not None and not isinstance(phase, str): + raise TypeError("phase must be a string or None") + if distance not in {"dist_deg", "dist_km", "dist_3d_km"}: + raise ValueError( + "distance must be 'dist_deg', 'dist_km', or 'dist_3d_km'" + ) + + if (distance not in self.rec_points) or recalc_dist: + self.log.SrcReclog.info("Calculating source--receiver distance...") + self.calc_distaz() + + records = self.rec_points + if phase is not None: + records = records.loc[records["phase"] == phase] + if records.empty: + raise ValueError( + "No absolute travel-time records found for phase " + f"{phase!r}" + ) + + return fit_linear_regression(records[distance], records["tt"]) def _filter_double_difference_by_arrivals(self): """Remove double differences whose absolute arrivals were rejected.""" @@ -1323,13 +1448,92 @@ def _filter_double_difference_by_arrivals(self): ) ) + def _filter_double_difference_by_distance( + self, + dist_min_max, + distance="dist_deg", + ): + """Keep double differences whose two endpoint distances are in range.""" + min_distance, max_distance = dist_min_max + + def calculate_distance(lat1, lon1, lat2, lon2): + distaz = DistAZ(lat1, lon1, lat2, lon2) + if distance == "dist_km": + return distaz.degreesToKilometers() + return distaz.delta + + def in_range(values): + values = np.asarray(values, dtype=float) + return ( + np.isfinite(values) + & (values >= min_distance) + & (values <= max_distance) + ) + + if not self.rec_points_cs.empty: + records = self.rec_points_cs + source_latitudes = self.src_points["evla"].reindex( + records["src_index"] + ).to_numpy() + source_longitudes = self.src_points["evlo"].reindex( + records["src_index"] + ).to_numpy() + distance1 = calculate_distance( + source_latitudes, + source_longitudes, + records["stla1"].to_numpy(), + records["stlo1"].to_numpy(), + ) + distance2 = calculate_distance( + source_latitudes, + source_longitudes, + records["stla2"].to_numpy(), + records["stlo2"].to_numpy(), + ) + keep = in_range(distance1) & in_range(distance2) + self.rec_points_cs = records.loc[keep] + self.log.SrcReclog.info( + "Removed {} common-source records outside the distance " + "range".format(len(records) - np.count_nonzero(keep)) + ) + + if not self.rec_points_cr.empty: + records = self.rec_points_cr + source_latitudes = self.src_points["evla"].reindex( + records["src_index"] + ).to_numpy() + source_longitudes = self.src_points["evlo"].reindex( + records["src_index"] + ).to_numpy() + station_latitudes = records["stla"].to_numpy() + station_longitudes = records["stlo"].to_numpy() + distance1 = calculate_distance( + source_latitudes, + source_longitudes, + station_latitudes, + station_longitudes, + ) + distance2 = calculate_distance( + records["evla2"].to_numpy(), + records["evlo2"].to_numpy(), + station_latitudes, + station_longitudes, + ) + keep = in_range(distance1) & in_range(distance2) + self.rec_points_cr = records.loc[keep] + self.log.SrcReclog.info( + "Removed {} common-receiver records outside the distance " + "range".format(len(records) - np.count_nonzero(keep)) + ) + def select_by_linear_regression(self, std_multiplier=3.0, recalc_dist=False, separate_phase=True, + distance="dist_3d_km", **kwargs): """Select absolute travel times by linear-regression residual. - A straight line is fitted between epicentral distance and travel time. - Records whose absolute residual is greater than ``std_multiplier`` + A straight line is fitted between the selected distance and travel + time. Records whose absolute residual is greater than ``std_multiplier`` times the residual standard deviation are removed. By default each phase is fitted separately so that phases with different apparent velocities are not mixed. @@ -1342,12 +1546,20 @@ def select_by_linear_regression(self, std_multiplier=3.0, :param std_multiplier: Multiplier applied to the residual standard deviation, defaults to 3. :type std_multiplier: float - :param recalc_dist: Recalculate epicentral distance even when - ``dist_deg`` exists, defaults to False. + :param recalc_dist: Recalculate distance fields even when the selected + field exists, defaults to False. :type recalc_dist: bool :param separate_phase: Fit each phase separately, defaults to True. :type separate_phase: bool - :return: Mapping from phase name to ``(slope, intercept)``. When + :param distance: Independent variable used by the fit. Choose + ``"dist_deg"`` for epicentral distance in degrees or + ``"dist_km"`` for epicentral distance in km, or + ``"dist_3d_km"`` for three-dimensional + source--receiver distance in km, defaults to + ``"dist_3d_km"``. + :type distance: str + :return: Mapping from phase name to + ``(slope, intercept, residual_std)``. When ``separate_phase=False``, the key is ``"all"``. :rtype: dict """ @@ -1355,14 +1567,18 @@ def select_by_linear_regression(self, std_multiplier=3.0, or not np.isfinite(std_multiplier) or std_multiplier <= 0): raise ValueError("std_multiplier must be a positive finite number") + if distance not in {"dist_deg", "dist_km", "dist_3d_km"}: + raise ValueError( + "distance must be 'dist_deg', 'dist_km', or 'dist_3d_km'" + ) self.log.SrcReclog.info( "rec_points before travel-time selection: {}".format( self.rec_points.shape[0] ) ) - if ("dist_deg" not in self.rec_points) or recalc_dist: - self.log.SrcReclog.info("Calculating epicentral distance...") + if (distance not in self.rec_points) or recalc_dist: + self.log.SrcReclog.info("Calculating source--receiver distance...") self.calc_distaz() keep = pd.Series(False, index=self.rec_points.index, dtype=bool) @@ -1372,7 +1588,7 @@ def select_by_linear_regression(self, std_multiplier=3.0, for phase, records in groups: group_keep, params = self._regression_keep_mask( - records, std_multiplier + records, std_multiplier, distance=distance ) keep.loc[records.index] = group_keep if params is not None: @@ -1447,6 +1663,11 @@ def select_by_constant_velocity( "tt_res_range minimum must not exceed its maximum" ) + self.log.SrcReclog.info( + "src_points before constant-velocity selection: {}".format( + self.src_points.shape[0] + ) + ) self.log.SrcReclog.info( "rec_points before constant-velocity selection: {}".format( self.rec_points.shape[0] @@ -1470,6 +1691,11 @@ def select_by_constant_velocity( self.rec_points = self.rec_points.loc[keep] self._filter_double_difference_by_arrivals() self.update(**kwargs) + self.log.SrcReclog.info( + "src_points after constant-velocity selection: {}".format( + self.src_points.shape[0] + ) + ) self.log.SrcReclog.info( "rec_points after constant-velocity selection: {}".format( self.rec_points.shape[0] @@ -1598,17 +1824,29 @@ def select_one_event_in_each_subgrid(self, d_deg: float, d_km: float, **kwargs): # self.remove_rec_by_new_src() self.update(**kwargs) - def box_weighting(self, d_deg: float, d_km: float, obj="both", dd_weight='average'): + def box_weighting( + self, + d_deg: float, + d_km: float | None = None, + obj="both", + dd_weight="average", + ): """Weighting sources and receivers by number in each subgrid :param d_deg: grid size along lat and lon in degree :type d_deg: float - :param d_km: grid size along depth axis in km, (only used when obj=``src`` or ``both``) - :type d_km: float + :param d_km: Grid size along the depth axis in km. Required only when + ``obj="src"`` or ``obj="both"``, defaults to ``None``. + :type d_km: float, optional :param obj: Object to be weighted, options: ``src``, ``rec`` or ``both``, defaults to ``both`` :type obj: str, optional :param dd_weight: Weighting method for double difference, options: ``average``, `multiply`, defaults to ``average`` """ + if obj in {"src", "both"} and d_km is None: + raise ValueError( + "d_km is required when obj is 'src' or 'both'" + ) + if obj == "src": self._box_weighting_ev(d_deg, d_km) elif obj == "rec": @@ -1658,72 +1896,81 @@ def _box_weighting_ev(self, d_deg: float, d_km: float): ) def _box_weighting_st(self, d_deg: float, dd_weight='average'): - """Weighting receivers by number of sources in each subgrid + """Weight receivers by density in two-dimensional horizontal cells. + + Receiver elevation is intentionally ignored. Stations are grouped + only by latitude and longitude. :param d_deg: grid size along lat and lon in degree :type d_deg: float """ + if not isinstance(d_deg, Real) or not np.isfinite(d_deg) or d_deg <= 0: + raise ValueError("d_deg must be a positive finite number") + self.log.SrcReclog.info( "Box weighting for receivers: d_deg={}".format(d_deg) ) - # group events by grid size - self.receivers["lat_group"] = self.receivers["stla"].apply( - lambda x: int(x / d_deg) - ) - self.receivers["lon_group"] = self.receivers["stlo"].apply( - lambda x: int(x / d_deg) + duplicate_receiver_count = self.receivers.duplicated( + subset="staname" + ).sum() + if duplicate_receiver_count: + self.log.SrcReclog.warning( + "Found %d duplicate receiver rows by staname; keeping the " + "first occurrence for box_weighting", + duplicate_receiver_count, + ) + self.receivers = self.receivers.drop_duplicates( + subset="staname", keep="first", ignore_index=True + ) + + horizontal_coordinates = self.receivers[["stla", "stlo"]].to_numpy( + dtype=float ) + if not np.isfinite(horizontal_coordinates).all(): + raise ValueError("receiver latitude and longitude must be finite") + horizontal_groups = np.trunc( + horizontal_coordinates / d_deg + ).astype(np.int64) + self.receivers[["lat_group", "lon_group"]] = horizontal_groups - # count num of sources in the same lat_group and lon_group self.receivers["num_receivers"] = self.receivers.groupby( ["lat_group", "lon_group"] )["lat_group"].transform("count") + self.receivers["weight"] = np.reciprocal( + np.sqrt(self.receivers["num_receivers"].to_numpy(dtype=float)) + ) - # calculate weight for each event - self.receivers["weight"] = 1 / np.sqrt(self.receivers["num_receivers"]) - - # assign weight to rec_points - self.rec_points["weight"] = self.rec_points.apply( - lambda x: self.receivers[ - (self.receivers["staname"] == x["staname"]) - ]["weight"].values[0], - axis=1, + receiver_weights = dict(zip( + self.receivers["staname"], + self.receivers["weight"], + )) + self.rec_points["weight"] = self.rec_points["staname"].map( + receiver_weights ) - # assign weight to rec_points_cs - # the weight is the average of the two receivers if not self.rec_points_cs.empty: - self.rec_points_cs["weight"] = self.rec_points_cs.apply( - lambda x: self._cal_dd_weight( - self.receivers[ - (self.receivers["staname"] == x["staname1"]) - ]["weight"].values[0], - self.receivers[ - (self.receivers["staname"] == x["staname2"]) - ]["weight"].values[0], - dd_weight - ), - axis=1, + weight1 = self.rec_points_cs["staname1"].map(receiver_weights) + weight2 = self.rec_points_cs["staname2"].map(receiver_weights) + self.rec_points_cs["weight"] = self._cal_dd_weight( + weight1, weight2, dd_weight ) - - # assign weight to rec_points_cr - # the weight is the average of the one receiver and the other source + if not self.rec_points_cr.empty: - self.rec_points_cr["weight"] = self.rec_points_cr.apply( - lambda x: self._cal_dd_weight( - self.receivers[ - (self.receivers["staname"] == x["staname"]) - ]["weight"].values[0], - self.src_points[ - (self.src_points["event_id"] == x["event_id2"]) - ]["weight"].values[0], - dd_weight - ), - axis=1, + source_weights = dict(zip( + self.src_points["event_id"], + self.src_points["weight"], + )) + receiver_weight = self.rec_points_cr["staname"].map( + receiver_weights + ) + source_weight = self.rec_points_cr["event_id2"].map( + source_weights + ) + self.rec_points_cr["weight"] = self._cal_dd_weight( + receiver_weight, source_weight, dd_weight ) - # drop 'lat_group' and 'lon_group' self.receivers = self.receivers.drop( columns=["lat_group", "lon_group", "num_receivers"] ) @@ -1888,12 +2135,31 @@ def _count_records(self): return count def _calc_weights(self, lat, lon, scale): - points = pd.concat([lon, lat], axis=1) - points_rad = points * (np.pi / 180) - dist = haversine_distances(points_rad) * 6371.0 / 111.19 - dist_ref = scale * np.mean(dist) - om = np.exp(-((dist / dist_ref) ** 2)) * points.shape[0] - return 1 / np.mean(om, axis=0) + """Calculate inverse-density weights normalized to a maximum of one.""" + if not isinstance(scale, Real) or not np.isfinite(scale) or scale <= 0: + raise ValueError("scale must be a positive finite number") + + points_rad = np.column_stack(( + np.asarray(lat, dtype=float), + np.asarray(lon, dtype=float), + )) + if len(points_rad) == 0: + return np.empty(0, dtype=float) + if not np.isfinite(points_rad).all(): + raise ValueError("latitude and longitude must be finite") + + np.deg2rad(points_rad, out=points_rad) + distances = haversine_distances(points_rad) + mean_distance = distances.mean() + if np.isclose(mean_distance, 0.0): + return np.ones(len(points_rad), dtype=float) + + distances /= scale * mean_distance + np.square(distances, out=distances) + distances *= -1.0 + np.exp(distances, out=distances) + weights = np.reciprocal(distances.sum(axis=0)) + return weights / weights.max() def _cal_dd_weight(self, w1, w2, dd_weight='average'): if dd_weight == "average": @@ -1904,7 +2170,12 @@ def _cal_dd_weight(self, w1, w2, dd_weight='average'): raise ValueError("Only 'average' or 'multiply' are supported for dd_weight") def geo_weighting(self, scale=0.5, obj="both", dd_weight="average"): - """Calculating geographical weights for sources + """Calculate and assign normalized geographical weights. + + Source and receiver weights are normalized so that the maximum of + each calculated population is one before weights are propagated to + absolute and double-difference records. Consequently, all generated + weights are no greater than one. :param scale: Scale of reference distance parameter. See equation 22 in Ruan et al., (2019). The reference distance is given by ``scale* dis_average``, defaults to 0.5 @@ -1914,39 +2185,78 @@ def geo_weighting(self, scale=0.5, obj="both", dd_weight="average"): :param dd_weight: Weighting method for double difference data, options: ``average`` or ``multiply``, defaults to ``average`` """ - if obj == "src" or obj == "both": + if obj not in {"src", "rec", "both"}: + raise ValueError("obj must be 'src', 'rec', or 'both'") + if dd_weight not in {"average", "multiply"}: + raise ValueError( + "Only 'average' or 'multiply' are supported for dd_weight" + ) + + if obj in {"src", "both"}: self.src_points["weight"] = self._calc_weights( self.src_points["evla"], self.src_points["evlo"], scale ) - # assign weight to sources - self.sources["weight"] = self.sources.apply( - lambda x: self.src_points[ - (self.src_points["event_id"] == x["event_id"]) - ]["weight"].values[0], - axis=1, + source_weights = dict(zip( + self.src_points["event_id"], + self.src_points["weight"], + )) + self.sources["weight"] = self.sources["event_id"].map( + source_weights ) - if obj == "rec" or obj == "both": + + if obj in {"rec", "both"}: + duplicate_receiver_count = self.receivers.duplicated( + subset="staname" + ).sum() + if duplicate_receiver_count: + self.log.SrcReclog.warning( + "Found %d duplicate receiver rows by staname; keeping " + "the first occurrence for geo_weighting", + duplicate_receiver_count, + ) + self.receivers = self.receivers.drop_duplicates( + subset="staname", keep="first", ignore_index=True + ) + weights = self._calc_weights( self.receivers['stla'], self.receivers['stlo'], scale ) - # apply weights to rec_points self.receivers['weight'] = weights - for row in self.receivers.itertuples(index=False): - self.rec_points.loc[self.rec_points['staname'] == row.staname, 'weight'] = row.weight + receiver_weights = dict(zip( + self.receivers["staname"], + self.receivers["weight"], + )) + self.rec_points["weight"] = self.rec_points["staname"].map( + receiver_weights + ) if not self.rec_points_cs.empty: - for row in self.rec_points_cs.itertuples(index=True): - w1 = self.receivers.loc[self.receivers['staname'] == row.staname1, 'weight'].values[0] - w2 = self.receivers.loc[self.receivers['staname'] == row.staname2, 'weight'].values[0] - self.rec_points_cs.loc[row.Index, 'weight'] = self._cal_dd_weight(w1, w2, dd_weight) + weight1 = self.rec_points_cs["staname1"].map( + receiver_weights + ) + weight2 = self.rec_points_cs["staname2"].map( + receiver_weights + ) + self.rec_points_cs["weight"] = self._cal_dd_weight( + weight1, weight2, dd_weight + ) if not self.rec_points_cr.empty: - for row in self.rec_points_cr.itertuples(index=True): - w1 = self.receivers.loc[self.receivers['staname'] == row.staname, 'weight'].values[0] - w2 = self.src_points.loc[self.src_points['event_id'] == row.event_id2, 'weight'].values[0] - self.rec_points_cr.loc[row.Index, 'weight'] = self._cal_dd_weight(w1, w2, dd_weight) + source_weights = dict(zip( + self.src_points["event_id"], + self.src_points["weight"], + )) + receiver_weight = self.rec_points_cr["staname"].map( + receiver_weights + ) + source_weight = self.rec_points_cr["event_id2"].map( + source_weights + ) + self.rec_points_cr["weight"] = self._cal_dd_weight( + receiver_weight, source_weight, dd_weight + ) def add_noise(self, range_in_sec=0.1, mean_in_sec=0.0, shape="gaussian"): """Add random noise on travel time @@ -2041,7 +2351,12 @@ def write_receivers(self, fname: str): :param fname: Path to output txt file of receivers """ - self.receivers.to_csv(fname, sep=" ", header=False, index=False) + receivers = self.receivers.copy() + if "weight" in receivers: + receivers["weight"] = receivers["weight"].map( + lambda value: "" if pd.isna(value) else f"{value:.4f}" + ) + receivers.to_csv(fname, sep=" ", header=False, index=False) def write_sources(self, fname: str): """ @@ -2049,7 +2364,12 @@ def write_sources(self, fname: str): :param fname: Path to output txt file of sources """ - self.sources.to_csv(fname, sep=" ", header=False, index=False) + sources = self.sources.copy() + if "weight" in sources: + sources["weight"] = sources["weight"].map( + lambda value: "" if pd.isna(value) else f"{value:.4f}" + ) + sources.to_csv(fname, sep=" ", header=False, index=False) @classmethod def from_seispy(cls, rf_path: str): @@ -2104,20 +2424,28 @@ def plot(self, color_by="depth", fname=None, **kwargs): def plot_travel_time( self, - color="tab:blue", + color=None, fname=None, fig=None, ylim="adaptive", + distance="dist_3d_km", **kwargs, ): - """Plot absolute travel time against epicentral distance. + """Plot absolute travel time against source--receiver distance. - If ``dist_deg`` is unavailable, it is calculated before plotting. + If the selected distance field is unavailable, epicentral distance is + calculated before plotting. The returned Matplotlib figure remains editable; use ``figure.axes[0]`` to add lines, annotations, or other content. - :param color: Matplotlib-compatible point color, defaults to - ``"tab:blue"``. + :param distance: Distance field used for the x-axis. Choose + ``"dist_3d_km"`` for three-dimensional distance, + ``"dist_deg"`` for epicentral distance in degrees, + or ``"dist_km"`` for epicentral distance in km, + defaults to ``"dist_3d_km"``. + :type distance: str, optional + :param color: Matplotlib-compatible point color. When ``None``, the + next color from the current axis color cycle is used. :param fname: Path to output file, defaults to None. :type fname: str, optional :param fig: Existing Matplotlib figure on which to draw, defaults to @@ -2134,14 +2462,20 @@ def plot_travel_time( :return: Matplotlib figure. :rtype: matplotlib.figure.Figure """ - if "dist_deg" not in self.rec_points: - self.log.SrcReclog.info("Calculating epicentral distance...") + if distance not in {"dist_deg", "dist_km", "dist_3d_km"}: + raise ValueError( + "distance must be 'dist_deg', 'dist_km', or 'dist_3d_km'" + ) + + if distance not in self.rec_points: + self.log.SrcReclog.info("Calculating source--receiver distance...") self.calc_distaz() from .utils.vis import plot_travel_time return plot_travel_time( self, + distance=distance, color=color, fname=fname, fig=fig, diff --git a/pytomoatt/utils/vis.py b/pytomoatt/utils/vis.py index 280823d..c56383c 100644 --- a/pytomoatt/utils/vis.py +++ b/pytomoatt/utils/vis.py @@ -7,7 +7,9 @@ import matplotlib.pyplot as plt import numpy as np +from matplotlib.colors import Normalize from matplotlib.figure import Figure +from matplotlib.transforms import Bbox if TYPE_CHECKING: from pytomoatt.src_rec import SrcRec @@ -110,11 +112,19 @@ def plot_src_rec( if receivers is None or receivers.empty: receiver_lon = np.empty(0) receiver_lat = np.empty(0) + receiver_color_values = None else: _check_columns(receivers, {"stlo", "stla"}, "receivers") receiver_values = receivers[["stlo", "stla"]].to_numpy(dtype=float) finite_receivers = np.isfinite(receiver_values).all(axis=1) receiver_lon, receiver_lat = receiver_values[finite_receivers].T + receiver_color_values = None + if color_by == "weight" and "weight" in receivers: + receiver_color_values = np.asarray( + receivers.loc[finite_receivers, "weight"], dtype=float + ) + if not np.isfinite(receiver_color_values).all(): + raise ValueError("receivers contains non-finite weights") all_lon = np.concatenate((source_lon, receiver_lon)) all_lat = np.concatenate((source_lat, receiver_lat)) @@ -149,7 +159,10 @@ def plot_src_rec( longitude_depth_axis = figure.add_subplot(grid[1, 0], sharex=map_axis) colorbar_host = figure.add_subplot(grid[1, 1]) colorbar_host.set_axis_off() - colorbar_axis = colorbar_host.inset_axes((0.05, 0.52, 0.9, 0.12)) + source_colorbar_y = 0.68 if receiver_color_values is not None else 0.52 + colorbar_axis = colorbar_host.inset_axes( + (0.05, source_colorbar_y, 0.9, 0.12) + ) scatter_options = { "c": color_values, @@ -158,20 +171,48 @@ def plot_src_rec( "label": "Sources", } scatter_options.update(kwargs) + if color_by == "weight" and "norm" not in scatter_options: + vmin = scatter_options.pop("vmin", None) + vmax = scatter_options.pop("vmax", None) + shared_norm = Normalize(vmin=vmin, vmax=vmax) + shared_norm.autoscale_None(color_values) + scatter_options["norm"] = shared_norm + source_scatter = map_axis.scatter(source_lon, source_lat, **scatter_options) - latitude_depth_axis.scatter(source_depth, source_lat, **scatter_options) - longitude_depth_axis.scatter(source_lon, source_depth, **scatter_options) + section_scatter_options = scatter_options.copy() + section_scatter_options["norm"] = source_scatter.norm + section_scatter_options.pop("vmin", None) + section_scatter_options.pop("vmax", None) + latitude_depth_axis.scatter( + source_depth, source_lat, **section_scatter_options + ) + longitude_depth_axis.scatter( + source_lon, source_depth, **section_scatter_options + ) + receiver_scatter = None if receiver_lon.size: - map_axis.scatter( + receiver_scatter_options = { + "edgecolors": "white", + "linewidths": 0.5, + "label": "Receivers", + "marker": "v", + "s": 55, + } + if receiver_color_values is None: + receiver_scatter_options["c"] = "tab:red" + else: + receiver_norm = Normalize() + receiver_norm.autoscale_None(receiver_color_values) + receiver_scatter_options.update({ + "c": receiver_color_values, + "cmap": source_scatter.cmap, + "norm": receiver_norm, + }) + receiver_scatter = map_axis.scatter( receiver_lon, receiver_lat, - c="tab:red", - edgecolors="white", - linewidths=0.5, - label="Receivers", - marker="v", - s=55, + **receiver_scatter_options, ) map_axis.set( @@ -180,6 +221,58 @@ def plot_src_rec( xlim=lon_limits, ylim=lat_limits, ) + map_axis.set_aspect("equal", adjustable="box") + + def _align_latitude_depth_axis(axis, renderer): + lower_section_position = longitude_depth_axis.get_position( + original=True + ) + map_position = map_axis.get_position() + # Reuse the lower section's automatically calculated padding so the + # right and lower gaps are equal in physical units for any figure size. + section_gap_inches = ( + map_position.y0 - lower_section_position.y1 + ) * figure.get_figheight() + horizontal_gap = section_gap_inches / figure.get_figwidth() + depth_length_inches = ( + lower_section_position.height * figure.get_figheight() + ) + section_width = depth_length_inches / figure.get_figwidth() + section_x0 = map_position.x1 + horizontal_gap + return Bbox.from_extents( + section_x0, + map_position.y0, + section_x0 + section_width, + map_position.y1, + ) + + def _align_longitude_depth_axis(axis, renderer): + section_position = axis.get_position(original=True) + map_position = map_axis.get_position() + return Bbox.from_extents( + map_position.x0, + section_position.y0, + map_position.x1, + section_position.y1, + ) + + def _align_colorbar_host(axis, renderer): + right_position = _align_latitude_depth_axis( + latitude_depth_axis, renderer + ) + lower_position = _align_longitude_depth_axis( + longitude_depth_axis, renderer + ) + return Bbox.from_extents( + right_position.x0, + lower_position.y0, + right_position.x1, + lower_position.y1, + ) + + latitude_depth_axis.set_axes_locator(_align_latitude_depth_axis) + longitude_depth_axis.set_axes_locator(_align_longitude_depth_axis) + colorbar_host.set_axes_locator(_align_colorbar_host) map_axis.legend() latitude_depth_axis.set( @@ -188,6 +281,8 @@ def plot_src_rec( xlim=depth_limits, ylim=lat_limits, ) + latitude_depth_axis.yaxis.tick_right() + latitude_depth_axis.yaxis.set_label_position("right") longitude_depth_axis.set( xlabel="Longitude", ylabel="Depth (km)", @@ -203,6 +298,17 @@ def plot_src_rec( ) colorbar.set_label(colorbar_label) + if receiver_color_values is not None and receiver_scatter is not None: + receiver_colorbar_axis = colorbar_host.inset_axes( + (0.05, 0.22, 0.9, 0.12) + ) + receiver_colorbar = figure.colorbar( + receiver_scatter, + cax=receiver_colorbar_axis, + orientation="horizontal", + ) + receiver_colorbar.set_label("Receiver weight") + if fname is not None: figure.savefig(fname, dpi=300, bbox_inches="tight") @@ -223,21 +329,29 @@ def fig_ev_st_distribution_dep( def plot_travel_time( src_rec: "SrcRec", *, - color="tab:blue", + distance: Literal["dist_deg", "dist_km", "dist_3d_km"] = "dist_3d_km", + color=None, fname: str | PathLike[str] | None = None, fig: Figure | None = None, ylim="adaptive", **kwargs, ) -> Figure: - """Plot absolute travel time against epicentral distance. + """Plot absolute travel time against source--receiver distance. Parameters ---------- src_rec A :class:`~pytomoatt.src_rec.SrcRec` instance whose ``rec_points`` - contains ``dist_deg`` and ``tt`` columns. + contains the selected distance column and ``tt``. + distance + Distance column used for the x-axis: ``"dist_3d_km"`` (default) for + three-dimensional source--receiver distance in kilometres, + ``"dist_deg"`` for epicentral distance in degrees, or ``"dist_km"`` + for epicentral distance in kilometres. color - Any Matplotlib-compatible color specification for the points. + Any Matplotlib-compatible color specification for the points. When + ``None`` (default), Matplotlib selects the next color from the current + axis color cycle. fname Optional output path. The format is inferred by Matplotlib from the filename extension. @@ -246,7 +360,7 @@ def plot_travel_time( axis; an axis is created when the figure has none. ylim Y-axis scaling strategy. ``"adaptive"`` (default) uses travel times - at the minimum and maximum epicentral distances, ``"auto"`` or + at the minimum and maximum selected distances, ``"auto"`` or ``None`` uses Matplotlib autoscaling, ``"inherit"`` preserves the current limits of an existing figure, and a ``(min, max)`` pair sets explicit limits. @@ -260,12 +374,17 @@ def plot_travel_time( The created figure. Its axis is available as ``figure.axes[0]`` for adding lines, annotations, or other content. """ + if distance not in {"dist_deg", "dist_km", "dist_3d_km"}: + raise ValueError( + "distance must be 'dist_deg', 'dist_km', or 'dist_3d_km'" + ) + records = src_rec.rec_points - _check_columns(records, {"dist_deg", "tt"}, "rec_points") + _check_columns(records, {distance, "tt"}, "rec_points") if records.empty: raise ValueError("Cannot plot travel times without receiver records") - values = records[["dist_deg", "tt"]].to_numpy(dtype=float) + values = records[[distance, "tt"]].to_numpy(dtype=float) finite = np.isfinite(values).all(axis=1) if not finite.any(): raise ValueError("rec_points contains no finite distance--time pairs") @@ -290,11 +409,18 @@ def plot_travel_time( else: raise TypeError("fig must be a matplotlib.figure.Figure or None") - scatter_options = {"color": color, "s": 4} + scatter_options = {"s": 4} + if color is not None: + scatter_options["color"] = color scatter_options.update(kwargs) axis.scatter(values[finite, 0], values[finite, 1], **scatter_options) + distance_label = { + "dist_deg": "Epicentral distance (degree)", + "dist_km": "Epicentral distance (km)", + "dist_3d_km": "3-D source-receiver distance (km)", + }[distance] axis.set( - xlabel="Epicentral distance (degree)", + xlabel=distance_label, ylabel="Travel time (s)", ) diff --git a/test/test_src_rec.py b/test/test_src_rec.py index 5c7f6f8..a51e666 100644 --- a/test/test_src_rec.py +++ b/test/test_src_rec.py @@ -16,6 +16,7 @@ import io import matplotlib.pyplot as plt from matplotlib.colors import to_rgba +from sklearn.metrics.pairwise import haversine_distances class TestSrcRec(unittest.TestCase): @@ -33,6 +34,32 @@ def test_read_missing_local_file(self): ): SrcRec.read(missing_file) + def test_conflicting_receiver_warning_lists_station_names(self): + sr = SrcRec("unused") + sr.src_points = pd.DataFrame({ + "event_id": ["EVENT_0"], + "evla": [1.0], + "evlo": [2.0], + "evdp": [3.0], + }) + sr.rec_points = pd.DataFrame({ + "staname": ["STA_CONFLICT", "STA_CONFLICT", "STA_OK"], + "stla": [10.0, 10.1, 20.0], + "stlo": [30.0, 30.0, 40.0], + "stel": [0.0, 0.0, 0.0], + }) + + with self.assertLogs("SrcRec", level="WARNING") as captured_logs: + sr.update_unique_src_rec() + + warning = captured_logs.output[0] + self.assertIn("1 receiver(s): STA_CONFLICT", warning) + self.assertNotIn("Keeping", warning) + self.assertEqual( + sr.receivers["staname"].tolist(), + ["STA_CONFLICT", "STA_OK"], + ) + def test_read_reindexes_duplicate_file_src_indices(self): sr = SrcRec.read(self.duplicate_index_fname) @@ -105,15 +132,82 @@ def test_subcase_10(self): sr = SrcRec.read(self.fname) sr.box_weighting(0.4, 10, obj='both') + def test_box_weighting_receiver_does_not_require_depth_size(self): + sr = SrcRec('unused') + + with ( + patch.object(sr, '_box_weighting_ev') as weight_sources, + patch.object(sr, '_box_weighting_st') as weight_receivers, + ): + sr.box_weighting(d_deg=0.4, obj='rec') + + weight_sources.assert_not_called() + weight_receivers.assert_called_once_with(0.4, 'average') + + def test_box_weighting_requires_depth_size_for_sources(self): + sr = SrcRec('unused') + + with self.assertRaisesRegex(ValueError, 'd_km'): + sr.box_weighting(d_deg=0.4, obj='src') + with self.assertRaisesRegex(ValueError, 'd_km'): + sr.box_weighting(d_deg=0.4, obj='both') + + def test_box_weighting_receiver_uses_horizontal_cells_only(self): + sr = SrcRec('unused') + sr.receivers = pd.DataFrame({ + 'staname': ['STA0', 'STA1', 'STA2'], + 'stla': [0.1, 0.2, 2.1], + 'stlo': [0.1, 0.2, 2.1], + 'stel': [0.0, 5000.0, 100.0], + }) + sr.rec_points = pd.DataFrame({ + 'staname': ['STA0', 'STA1', 'STA2'], + 'weight': [1.0, 1.0, 1.0], + }) + sr.src_points = pd.DataFrame({ + 'event_id': ['E0'], + 'weight': [0.5], + }) + sr.rec_points_cs = pd.DataFrame({ + 'staname1': ['STA0'], + 'staname2': ['STA1'], + 'weight': [1.0], + }) + sr.rec_points_cr = pd.DataFrame({ + 'staname': ['STA2'], + 'event_id2': ['E0'], + 'weight': [1.0], + }) + + sr.box_weighting(d_deg=1.0, obj='rec') + + expected_dense_weight = 1.0 / np.sqrt(2.0) + receiver_weights = sr.receivers.set_index('staname')['weight'] + self.assertAlmostEqual(receiver_weights['STA0'], expected_dense_weight) + self.assertAlmostEqual(receiver_weights['STA1'], expected_dense_weight) + self.assertAlmostEqual(receiver_weights['STA2'], 1.0) + self.assertTrue(np.allclose( + sr.rec_points['weight'], + sr.rec_points['staname'].map(receiver_weights), + )) + self.assertAlmostEqual( + sr.rec_points_cs.iloc[0]['weight'], expected_dense_weight + ) + self.assertAlmostEqual(sr.rec_points_cr.iloc[0]['weight'], 0.75) + def test_select_by_linear_regression(self): sr = SrcRec('unused') distance = np.concatenate((np.arange(21, dtype=float), [10.0, 0.0])) + distance_km = distance * 100.0 + distance_3d_km = distance_km + 10.0 travel_time = 2.0 * distance + 5.0 travel_time[10] += 100.0 sr.rec_points = pd.DataFrame({ 'src_index': [0] * 21 + [1, 1], 'staname': [f'STA{i:02d}' for i in range(21)] + ['STA10', 'STA00'], 'dist_deg': distance, + 'dist_km': distance_km, + 'dist_3d_km': distance_3d_km, 'tt': travel_time, 'phase': 'P', }) @@ -132,16 +226,18 @@ def test_select_by_linear_regression(self): with patch.object(sr, 'update') as update: regression_params = sr.select_by_linear_regression( - std_multiplier=3.0 + std_multiplier=3.0, + distance='dist_3d_km', ) - expected_slope, expected_intercept = np.polyfit( - distance, travel_time, deg=1 + expected_slope, expected_intercept, expected_std = linear_regression( + distance_3d_km, travel_time ) self.assertIn('P', regression_params) - slope, intercept = regression_params['P'] + slope, intercept, residual_std = regression_params['P'] self.assertAlmostEqual(slope, expected_slope) self.assertAlmostEqual(intercept, expected_intercept) + self.assertAlmostEqual(residual_std, expected_std) self.assertEqual(sr.rec_points.shape[0], 22) self.assertNotIn(10, sr.rec_points.index) self.assertEqual(sr.rec_points_cs.shape[0], 1) @@ -150,6 +246,245 @@ def test_select_by_linear_regression(self): self.assertEqual(sr.rec_points_cr.iloc[0]['staname'], 'STA00') update.assert_called_once_with() + def test_linear_regression_method(self): + sr = SrcRec('unused') + distance = np.arange(5, dtype=float) + p_travel_time = np.array([4.0, 7.0, 10.0, 13.0, 17.0]) + sr.rec_points = pd.DataFrame({ + 'dist_deg': np.concatenate((distance, distance)), + 'dist_km': np.concatenate((distance * 10.0, distance * 10.0)), + 'dist_3d_km': np.concatenate( + (distance * 10.0 + 2.0, distance * 10.0 + 2.0) + ), + 'tt': np.concatenate((p_travel_time, 5.0 * distance + 2.0)), + 'phase': ['P'] * 5 + ['S'] * 5, + }) + original = sr.rec_points.copy(deep=True) + + result = sr.linear_regression(phase='P') + expected = linear_regression(distance * 10.0 + 2.0, p_travel_time) + + for actual_value, expected_value in zip(result, expected): + self.assertAlmostEqual(actual_value, expected_value) + + result_deg = sr.linear_regression(phase='P', distance='dist_deg') + expected_deg = linear_regression(distance, p_travel_time) + for actual_value, expected_value in zip(result_deg, expected_deg): + self.assertAlmostEqual(actual_value, expected_value) + pd.testing.assert_frame_equal(sr.rec_points, original) + + def test_linear_regression_method_calculates_distance(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'tt': [2.0, 5.0, 8.0], + 'phase': ['P', 'P', 'P'], + }) + + def add_distance(): + sr.rec_points['dist_3d_km'] = [0.0, 1.0, 2.0] + + with patch.object(sr, 'calc_distaz', side_effect=add_distance) as calc: + slope, intercept, residual_std = sr.linear_regression() + + self.assertAlmostEqual(slope, 3.0) + self.assertAlmostEqual(intercept, 2.0) + self.assertAlmostEqual(residual_std, 0.0) + calc.assert_called_once_with() + + def test_calc_distaz_calculates_source_receiver_distance(self): + sr = SrcRec('unused') + sr.src_points = pd.DataFrame({ + 'evla': [10.0], + 'evlo': [20.0], + 'evdp': [10.0], + }) + sr.rec_points = pd.DataFrame({ + 'src_index': [0], + 'stla': [10.0], + 'stlo': [20.0], + 'stel': [1000.0], + }) + + sr.calc_distaz() + + self.assertAlmostEqual(sr.rec_points.loc[0, 'dist_deg'], 0.0) + self.assertAlmostEqual(sr.rec_points.loc[0, 'dist_km'], 0.0) + self.assertAlmostEqual(sr.rec_points.loc[0, 'dist_3d_km'], 11.0) + + def test_calc_weights_uses_lat_lon_order_and_normalizes(self): + sr = SrcRec('unused') + latitude = np.array([0.0, 60.0, 10.0]) + longitude = np.array([0.0, 10.0, 170.0]) + scale = 0.5 + + weights = sr._calc_weights(latitude, longitude, scale) + + points_rad = np.deg2rad(np.column_stack((latitude, longitude))) + distances = haversine_distances(points_rad) + reference_distance = scale * distances.mean() + expected = np.reciprocal( + np.exp(-((distances / reference_distance) ** 2)).sum(axis=0) + ) + expected /= expected.max() + self.assertTrue(np.allclose(weights, expected)) + self.assertAlmostEqual(weights.max(), 1.0) + + def test_geo_weighting_maps_and_normalizes_all_weights(self): + sr = SrcRec('unused') + sr.src_points = pd.DataFrame({ + 'event_id': ['E0', 'E1', 'E2'], + 'evla': [0.0, 1.0, 3.0], + 'evlo': [0.0, 2.0, 1.0], + 'evdp': [5.0, 10.0, 15.0], + 'weight': [1.0, 1.0, 1.0], + }) + sr.rec_points = pd.DataFrame({ + 'src_index': [0, 1, 2], + 'staname': ['STA0', 'STA1', 'STA2'], + 'stla': [0.0, 2.0, 1.0], + 'stlo': [0.0, 1.0, 4.0], + 'stel': [0.0, 100.0, 200.0], + 'phase': ['P', 'P', 'P'], + 'tt': [1.0, 2.0, 3.0], + 'weight': [1.0, 1.0, 1.0], + }) + sr.rec_points_cs = pd.DataFrame({ + 'src_index': [0], + 'staname1': ['STA0'], + 'stla1': [0.0], + 'stlo1': [0.0], + 'stel1': [0.0], + 'staname2': ['STA1'], + 'stla2': [2.0], + 'stlo2': [1.0], + 'stel2': [100.0], + 'phase': ['P,cs'], + 'weight': [1.0], + }) + sr.rec_points_cr = pd.DataFrame({ + 'src_index': [0], + 'src_index2': [1], + 'event_id2': ['E1'], + 'evla2': [1.0], + 'evlo2': [2.0], + 'evdp2': [10.0], + 'staname': ['STA0'], + 'stla': [0.0], + 'stlo': [0.0], + 'stel': [0.0], + 'phase': ['P,cr'], + 'weight': [1.0], + }) + sr.update_unique_src_rec() + + sr.geo_weighting(scale=0.5, obj='both', dd_weight='multiply') + + source_weights = sr.src_points.set_index('event_id')['weight'] + receiver_weights = sr.receivers.set_index('staname')['weight'] + self.assertTrue(np.allclose( + sr.sources['weight'], sr.sources['event_id'].map(source_weights) + )) + self.assertTrue(np.allclose( + sr.rec_points['weight'], + sr.rec_points['staname'].map(receiver_weights), + )) + expected_cs_weight = ( + receiver_weights['STA0'] * receiver_weights['STA1'] + ) + expected_cr_weight = receiver_weights['STA0'] * source_weights['E1'] + self.assertAlmostEqual( + sr.rec_points_cs.iloc[0]['weight'], expected_cs_weight + ) + self.assertAlmostEqual( + sr.rec_points_cr.iloc[0]['weight'], expected_cr_weight + ) + + all_weights = np.concatenate(( + sr.src_points['weight'].to_numpy(), + sr.sources['weight'].to_numpy(), + sr.receivers['weight'].to_numpy(), + sr.rec_points['weight'].to_numpy(), + sr.rec_points_cs['weight'].to_numpy(), + sr.rec_points_cr['weight'].to_numpy(), + )) + self.assertAlmostEqual(all_weights.max(), 1.0) + self.assertTrue(np.all(all_weights <= 1.0)) + + def test_geo_weighting_deduplicates_receiver_names(self): + sr = SrcRec('unused') + sr.receivers = pd.DataFrame({ + 'staname': ['STA0', 'STA0', 'STA1'], + 'stla': [0.0, 0.1, 1.0], + 'stlo': [0.0, 0.1, 2.0], + 'stel': [0.0, 10.0, 20.0], + }) + sr.rec_points = pd.DataFrame({ + 'staname': ['STA0', 'STA1'], + 'weight': [1.0, 1.0], + }) + + sr.geo_weighting(scale=0.5, obj='rec') + + self.assertEqual(sr.receivers['staname'].tolist(), ['STA0', 'STA1']) + receiver_weights = sr.receivers.set_index('staname')['weight'] + self.assertTrue(np.allclose( + sr.rec_points['weight'], + sr.rec_points['staname'].map(receiver_weights), + )) + self.assertAlmostEqual(sr.receivers['weight'].max(), 1.0) + + def test_select_by_distance_filters_double_differences(self): + sr = SrcRec('unused') + sr.src_points = pd.DataFrame({ + 'evla': [0.0, 0.0], + 'evlo': [0.0, 0.5], + }, index=[0, 1]) + sr.rec_points = pd.DataFrame({ + 'src_index': [0, 0], + 'staname': ['ABS_IN', 'ABS_OUT'], + 'dist_deg': [0.5, 2.0], + 'dist_km': [55.6, 222.4], + 'phase': ['P', 'P'], + }) + sr.rec_points_cs = pd.DataFrame({ + 'src_index': [0], + 'staname1': ['STA1'], + 'stla1': [0.0], + 'stlo1': [0.5], + 'staname2': ['STA2'], + 'stla2': [0.0], + 'stlo2': [2.0], + 'phase': ['P,cs'], + }) + sr.rec_points_cr = pd.DataFrame({ + 'src_index': [0, 0], + 'src_index2': [1, 1], + 'event_id2': ['EVT1', 'EVT1'], + 'staname': ['STA1', 'STA2'], + 'stla': [0.0, 0.0], + 'stlo': [0.5, 0.5], + 'evla2': [0.0, 0.0], + 'evlo2': [0.2, 2.0], + 'phase': ['P,cr', 'P,cr'], + }) + + with patch.object(sr, 'update') as update: + sr.select_by_distance( + [0.0, 112.0], + distance='dist_km', + ) + + self.assertEqual(sr.rec_points['staname'].tolist(), ['ABS_IN']) + self.assertTrue(sr.rec_points_cs.empty) + self.assertEqual(sr.rec_points_cr['staname'].tolist(), ['STA1']) + update.assert_called_once_with() + + def test_select_by_distance_rejects_invalid_distance(self): + sr = SrcRec('unused') + + with self.assertRaisesRegex(ValueError, 'distance'): + sr.select_by_distance([0.0, 1.0], distance='dist_3d_km') + def test_select_by_constant_velocity(self): sr = SrcRec('unused') distance = np.array([0.0, 1.0, 2.0, 3.0, 0.0]) @@ -211,8 +546,85 @@ def test_plot(self): self.assertAlmostEqual(map_position.y1, latitude_depth_position.y1) self.assertAlmostEqual(map_position.x0, longitude_depth_position.x0) self.assertAlmostEqual(map_position.x1, longitude_depth_position.x1) + right_gap = ( + latitude_depth_position.x0 - map_position.x1 + ) * figure.get_figwidth() + lower_gap = ( + map_position.y0 - longitude_depth_position.y1 + ) * figure.get_figheight() + self.assertAlmostEqual(right_gap, lower_gap) + right_depth_length = ( + latitude_depth_position.width * figure.get_figwidth() + ) + lower_depth_length = ( + longitude_depth_position.height * figure.get_figheight() + ) + self.assertAlmostEqual(right_depth_length, lower_depth_length) + colorbar_position = figure.axes[3].get_position() + self.assertAlmostEqual( + colorbar_position.x0, latitude_depth_position.x0 + ) + self.assertAlmostEqual( + colorbar_position.x1, latitude_depth_position.x1 + ) + self.assertAlmostEqual( + colorbar_position.y0, longitude_depth_position.y0 + ) + self.assertAlmostEqual( + colorbar_position.y1, longitude_depth_position.y1 + ) + self.assertEqual( + figure.axes[1].yaxis.get_ticks_position(), "right" + ) + self.assertEqual( + figure.axes[1].yaxis.get_label_position(), "right" + ) + self.assertEqual(figure.axes[0].get_aspect(), 1.0) + self.assertEqual(figure.axes[0].get_adjustable(), "box") + map_xlim = figure.axes[0].get_xlim() + map_ylim = figure.axes[0].get_ylim() + longitude_scale = map_position.width / (map_xlim[1] - map_xlim[0]) + latitude_scale = map_position.height / (map_ylim[1] - map_ylim[0]) + self.assertAlmostEqual(longitude_scale, latitude_scale) plt.close(figure) + def test_write_sources_and_receivers_format_weights(self): + sr = SrcRec("unused") + sr.sources = pd.DataFrame({ + "event_id": ["EVENT_0", "EVENT_1"], + "evla": [1.0, 2.0], + "evlo": [3.0, 4.0], + "evdp": [5.0, 6.0], + "weight": [1.0 / 3.0, 1.0], + }) + sr.receivers = pd.DataFrame({ + "staname": ["STA0", "STA1"], + "stla": [1.0, 2.0], + "stlo": [3.0, 4.0], + "stel": [5.0, 6.0], + "weight": [2.0 / 3.0, 1.0], + }) + + with TemporaryDirectory() as output_directory: + source_file = join(output_directory, "sources.txt") + receiver_file = join(output_directory, "receivers.txt") + sr.write_sources(source_file) + sr.write_receivers(receiver_file) + + with open(source_file) as output: + source_weights = [ + line.split()[-1] for line in output if line.strip() + ] + with open(receiver_file) as output: + receiver_weights = [ + line.split()[-1] for line in output if line.strip() + ] + + self.assertEqual(source_weights, ["0.3333", "1.0000"]) + self.assertEqual(receiver_weights, ["0.6667", "1.0000"]) + self.assertEqual(sr.sources.loc[0, "weight"], 1.0 / 3.0) + self.assertEqual(sr.receivers.loc[0, "weight"], 2.0 / 3.0) + def test_plot_source_only(self): sr = SrcRec.read(self.fname, src_only=True) @@ -221,6 +633,55 @@ def test_plot_source_only(self): self.assertIsNotNone(figure) plt.close(figure) + def test_plot_uses_shared_norm_for_constant_weights(self): + sr = SrcRec.read(self.fname, src_only=True) + sr.src_points['weight'] = 1.0 + + figure = sr.plot(color_by='weight') + figure.canvas.draw() + source_collections = [axis.collections[0] for axis in figure.axes[:3]] + + self.assertIs( + source_collections[0].norm, source_collections[1].norm + ) + self.assertIs( + source_collections[0].norm, source_collections[2].norm + ) + reference_colors = source_collections[0].get_facecolors() + for collection in source_collections[1:]: + self.assertTrue(np.allclose( + collection.get_facecolors(), reference_colors + )) + plt.close(figure) + + def test_plot_colors_receivers_by_weight(self): + sr = SrcRec.read(self.fname) + sr.geo_weighting(obj='both') + + figure = sr.plot(color_by='weight') + figure.canvas.draw() + source_collection = figure.axes[0].collections[0] + receiver_collection = figure.axes[0].collections[1] + source_index = sr.src_points['weight'].to_numpy().argmax() + receiver_index = sr.receivers['weight'].to_numpy().argmax() + + self.assertIsNot(source_collection.norm, receiver_collection.norm) + self.assertIs(source_collection.cmap, receiver_collection.cmap) + self.assertTrue(np.allclose( + source_collection.get_facecolors()[source_index], + receiver_collection.get_facecolors()[receiver_index], + )) + self.assertEqual( + source_collection.norm(sr.src_points['weight'].max()), + receiver_collection.norm(sr.receivers['weight'].max()), + ) + colorbar_labels = { + axis.get_xlabel() for axis in figure.axes[3].child_axes + } + self.assertIn('Source weight', colorbar_labels) + self.assertIn('Receiver weight', colorbar_labels) + plt.close(figure) + def test_plot_rejects_invalid_color_by(self): sr = SrcRec.read(self.fname, src_only=True) @@ -241,7 +702,7 @@ def test_plot_accepts_matplotlib_scatter_options(self): def test_plot_travel_time_returns_editable_figure(self): sr = SrcRec('unused') sr.rec_points = pd.DataFrame({ - 'dist_deg': [0.0, 1.0, np.nan], + 'dist_3d_km': [0.0, 1.0, np.nan], 'tt': [1.0, 3.0, 5.0], }) @@ -263,7 +724,7 @@ def test_plot_travel_time_calculates_missing_distance(self): sr.rec_points = pd.DataFrame({'tt': [1.0, 2.0]}) def add_distance(): - sr.rec_points['dist_deg'] = [0.0, 1.0] + sr.rec_points['dist_3d_km'] = [0.0, 1.0] with patch.object(sr, 'calc_distaz', side_effect=add_distance) as calc: figure = sr.plot_travel_time() @@ -271,10 +732,54 @@ def add_distance(): calc.assert_called_once_with() plt.close(figure) - def test_plot_travel_time_uses_existing_figure(self): + def test_plot_travel_time_uses_kilometres(self): sr = SrcRec('unused') sr.rec_points = pd.DataFrame({ 'dist_deg': [0.0, 1.0], + 'dist_km': [0.0, 111.19], + 'tt': [1.0, 3.0], + }) + + figure = sr.plot_travel_time(distance='dist_km') + axis = figure.axes[0] + offsets = axis.collections[0].get_offsets() + + self.assertTrue(np.allclose(offsets[:, 0], [0.0, 111.19])) + self.assertEqual(axis.get_xlabel(), 'Epicentral distance (km)') + plt.close(figure) + + def test_plot_travel_time_defaults_to_3d_distance(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_deg': [1.0, 2.0], + 'dist_3d_km': [120.0, 230.0], + 'tt': [10.0, 20.0], + }) + + figure = sr.plot_travel_time() + axis = figure.axes[0] + offsets = axis.collections[0].get_offsets() + + self.assertTrue(np.allclose(offsets[:, 0], [120.0, 230.0])) + self.assertEqual( + axis.get_xlabel(), '3-D source-receiver distance (km)' + ) + plt.close(figure) + + def test_plot_travel_time_rejects_invalid_distance(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_3d_km': [0.0, 1.0], + 'tt': [1.0, 3.0], + }) + + with self.assertRaisesRegex(ValueError, 'distance'): + sr.plot_travel_time(distance='miles') + + def test_plot_travel_time_uses_existing_figure(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_3d_km': [0.0, 1.0], 'tt': [1.0, 3.0], }) existing_figure, axis = plt.subplots() @@ -293,7 +798,7 @@ def test_plot_travel_time_uses_existing_figure(self): def test_plot_travel_time_inherits_y_limits(self): sr = SrcRec('unused') sr.rec_points = pd.DataFrame({ - 'dist_deg': [0.0, 1.0], + 'dist_3d_km': [0.0, 1.0], 'tt': [1.0, 30.0], }) existing_figure, axis = plt.subplots() @@ -308,10 +813,25 @@ def test_plot_travel_time_inherits_y_limits(self): self.assertEqual(axis.get_ylim(), (5.0, 20.0)) plt.close(existing_figure) + def test_plot_travel_time_uses_matplotlib_color_cycle(self): + sr = SrcRec('unused') + sr.rec_points = pd.DataFrame({ + 'dist_3d_km': [0.0, 1.0], + 'tt': [1.0, 2.0], + }) + + figure = sr.plot_travel_time() + first_color = figure.axes[0].collections[-1].get_facecolors()[0] + sr.plot_travel_time(fig=figure, ylim='inherit') + second_color = figure.axes[0].collections[-1].get_facecolors()[0] + + self.assertFalse(np.allclose(first_color, second_color)) + plt.close(figure) + def test_plot_travel_time_y_limits(self): sr = SrcRec('unused') sr.rec_points = pd.DataFrame({ - 'dist_deg': [0.0, 1.0, 2.0, 3.0, 4.0], + 'dist_3d_km': [0.0, 1.0, 2.0, 3.0, 4.0], 'tt': [1.0, 2.0, 1000.0, 4.0, 5.0], })