diff --git a/.github/workflows/build-test-conda.yml b/.github/workflows/build-test-conda.yml index 70d5bdc..4d0288f 100644 --- a/.github/workflows/build-test-conda.yml +++ b/.github/workflows/build-test-conda.yml @@ -45,7 +45,7 @@ jobs: run: | conda install hatchling python -m pip install --upgrade pip - pip install pytest pytest-cov obspy + pip install pytest pytest-cov obspy pyvista pip install . - name: test run: | diff --git a/README.md b/README.md index e12bb94..20ff38e 100644 --- a/README.md +++ b/README.md @@ -5,7 +5,7 @@ [![Python Package using Conda](https://github.com/MIGG-NTU/PyTomoATT/actions/workflows/build-test-conda.yml/badge.svg?branch=devel)](https://github.com/MIGG-NTU/PyTomoATT/actions/workflows/build-test-conda.yml) [![Build documentations](https://github.com/MIGG-NTU/PyTomoATT/actions/workflows/build-docs.yml/badge.svg?branch=docs)](https://migg-ntu.github.io/PyTomoATT/) -[![codecov](https://codecov.io/gh/MIGG-NTU/PyTomoATT/branch/devel/graph/badge.svg?token=EYOV0WOA2Y)](https://codecov.io/gh/MIGG-NTU/PyTomoATT) +[![codecov](https://codecov.io/gh/TomoATT/PyTomoATT/graph/badge.svg?token=EYOV0WOA2Y)](https://codecov.io/gh/TomoATT/PyTomoATT) ![PyPI - License](https://img.shields.io/pypi/l/pytomoatt) ![PyPI](https://img.shields.io/pypi/v/pytomoatt) diff --git a/pyproject.toml b/pyproject.toml index aeb9e78..c08d425 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,9 +17,9 @@ dynamic = ["version"] classifiers = [ "Programming Language :: Python", - "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", ] dependencies = [ "numpy>=1.19.0", diff --git a/pytomoatt/_version.py b/pytomoatt/_version.py index 6c8aff6..f0ca935 100644 --- a/pytomoatt/_version.py +++ b/pytomoatt/_version.py @@ -1 +1 @@ -__version__ = '0.2.9' \ No newline at end of file +__version__ = '0.2.10' \ No newline at end of file diff --git a/pytomoatt/checkerboard.py b/pytomoatt/checkerboard.py index 1cc5e06..589c577 100644 --- a/pytomoatt/checkerboard.py +++ b/pytomoatt/checkerboard.py @@ -133,6 +133,20 @@ def copy(self): :rtype: Checker """ return copy.deepcopy(self) + + def to_attmodel(self): + """Convert to ATTModel object + + :return: ATTModel object + :rtype: ATTModel + """ + from .model import ATTModel + mod = ATTModel(self.para_fname) + mod.vel = self.vel_pert + mod.xi = self.xi + mod.eta = self.eta + mod.zeta = self.zeta + return mod def write(self, fname): """Write new model to h5 file diff --git a/pytomoatt/model.py b/pytomoatt/model.py index af0b714..e896f24 100644 --- a/pytomoatt/model.py +++ b/pytomoatt/model.py @@ -5,7 +5,7 @@ from .io.crustmodel import CrustModel from .io.asciimodel import ASCIIModel from .attarray import Dataset -from .utils.common import init_axis +from .utils.common import init_axis, km2deg import copy @@ -73,14 +73,7 @@ def to_ani(self): """Convert to anisotropic strength (epsilon) and azimuth (phi) """ self.epsilon = np.sqrt(self.eta**2+self.xi**2) - # self.phi = np.zeros_like(self.epsilon) self.phi = np.rad2deg(0.5*np.arctan2(self.eta, self.xi)) - # idx = np.where(self.xi <= 0) - # self.phi[idx] = 90 + 0.5*atand(self.eta[idx]/self.xi[idx]) - # idx = np.where((self.xi > 0) & (self.eta <= 0)) - # self.phi[idx] = 180 + 0.5*atand(self.eta[idx]/self.xi[idx]) - # idx = np.where((self.xi > 0) & (self.eta > 0)) - # self.phi[idx] = 0.5*atand(self.eta[idx]/self.xi[idx]) def to_xarray(self): """Convert to xarray @@ -140,18 +133,61 @@ def grid_data_ascii(self, model_fname:str, **kwargs): self.n_rtp, ) - def smooth(self, sigma=5.0): + def smooth(self, sigma=5.0, unit_deg=False, smooth_ani=False, **kwargs): """Gaussian smooth the 3D velocity model - :param sigma: Standard division of gaussian kernel in km, defaults to 10 - :type sigma: scalar or sequence of scalars , optional + :param sigma: Standard deviation for Gaussian kernel. + If scalar, apply to all dimensions. + If sequence of 3, apply to [depth, lat, lon]. + Depth is always in km. + Horizontal dimensions depend on unit_deg. + :type sigma: scalar or sequence of scalars + :param unit_deg: If True, horizontal sigma is in degrees. + If False, horizontal sigma is in km. + Defaults to False. + :type unit_deg: bool + :param smooth_ani: If True, also smooth anisotropic parameters (xi, eta, zeta). + Defaults to False. + :type smooth_ani: bool + :param kwargs: Additional arguments passed to scipy.ndimage.gaussian_filter + + Example + ------------------- + To smooth with 5 km in depth and 0.2 degrees in horizontal directions: + >>> model.smooth(sigma=[5.0, 0.2, 0.2], unit_deg=True) + + To smooth with 5 km in depth and 20 km in horizontal directions: + >>> model.smooth(sigma=[5.0, 20.0, 20.0], unit_deg=False) """ - if isinstance(sigma, (int, float)): - sigma_all = np.ones(3)*sigma/self.d_rtp/2/np.pi - elif len(sigma) == 3: - sigma_all = np.array(sigma)/self.d_rtp/2/np.pi - sigma_all[0:2] /= 111.19 - self.vel = gaussian_filter(self.vel, sigma) + if np.isscalar(sigma): + sigma = [sigma, sigma, sigma] + elif len(sigma) != 3: + raise ValueError('sigma should be a scalar or a sequence of three scalars') + + sigma = np.array(sigma, dtype=float) + sigma_pixel = np.zeros(3) + + # Depth direction (always km) + sigma_pixel[0] = sigma[0] / self.d_rtp[0] + + if unit_deg: + # Horizontal sigma is in degrees + sigma_pixel[1] = sigma[1] / self.d_rtp[1] + sigma_pixel[2] = sigma[2] / self.d_rtp[2] + else: + # Horizontal sigma is in km + # Latitude + sigma_pixel[1] = km2deg(sigma[1]) / self.d_rtp[1] + # Longitude + mean_lat = np.mean(self.latitudes) + # 1 deg lon = cos(lat) * 1 deg lat + # so X km = km2deg(X) deg lat = km2deg(X) / cos(lat) deg lon + sigma_pixel[2] = km2deg(sigma[2]) / np.cos(np.deg2rad(mean_lat)) / self.d_rtp[2] + self.vel = gaussian_filter(self.vel, sigma_pixel, **kwargs) + if smooth_ani: + self.xi = gaussian_filter(self.xi, sigma_pixel, **kwargs) + self.eta = gaussian_filter(self.eta, sigma_pixel, **kwargs) + self.zeta = gaussian_filter(self.zeta, sigma_pixel, **kwargs) def calc_dv_avg(self): """calculate anomalies relative to average velocity at each depth diff --git a/pytomoatt/utils/__init__.py b/pytomoatt/utils/__init__.py index e69de29..463c88f 100644 --- a/pytomoatt/utils/__init__.py +++ b/pytomoatt/utils/__init__.py @@ -0,0 +1 @@ +_EARTH_RADIUS_KM = 6371.0 \ No newline at end of file diff --git a/pytomoatt/utils/common.py b/pytomoatt/utils/common.py index 325de97..5f7699a 100644 --- a/pytomoatt/utils/common.py +++ b/pytomoatt/utils/common.py @@ -1,6 +1,6 @@ import numpy as np from scipy.interpolate import griddata -import pandas as pd +from . import _EARTH_RADIUS_KM def sind(deg): rad = np.radians(deg) @@ -37,6 +37,34 @@ def atand(x): return np.degrees(rad) +def km2deg(km): + """ Convert km to degree + + :param km: Distance in km + :type km: float + :return: Distance in degree + :rtype: float + """ + circum = 2*np.pi*_EARTH_RADIUS_KM + conv = circum / 360 + deg = km / conv + return deg + + +def deg2km(deg): + """ Convert degree to km + + :param deg: Distance in degree + :type deg: float + :return: Distance in km + :rtype: float + """ + circum = 2*np.pi*_EARTH_RADIUS_KM + conv = circum / 360 + km = deg * conv + return km + + def WGS84_to_cartesian(dep, lat, lon): """ Convert WGS84 coordinates to cartesian coordinates diff --git a/pytomoatt/vis.py b/pytomoatt/vis.py deleted file mode 100644 index 3f09a6d..0000000 --- a/pytomoatt/vis.py +++ /dev/null @@ -1,87 +0,0 @@ -# module for visualization functions -import geoviews as gv -import geoviews.feature as gf -from geoviews import opts -import numpy as np -from holoviews.operation.datashader import rasterize, spread - - -def plot_srcrec(SR, weight=False, fname=None): - """ - Plot source and receiver locations - SR: source and receiver object - weight: if plots weight of the source/receiver - fname: if not None, save the plot to a file - """ - - gv.extension('bokeh') - - #tiles = gv.tile_sources.Wikipedia() - #tiles = gv.tile_sources.StamenTerrain() - #tiles = gv.tile_sources.EsriReference() - #tiles = gv.tile_sources.EsriUSATopo() - #tiles = gv.tile_sources.OSM()*gv.feature.coastline() - tiles = gv.feature.coastline(scale='50m') - - if weight==False: - # source points - SR.src_points['_evdp'] = SR.src_points['evdp'] # for protting - ds = gv.Dataset(SR.src_points, kdims=['evlo','evla','evdp','_evdp'], vdims=['num_rec']) - - lola = rasterize(gv.Points(ds, kdims=['evlo','evla'], vdims=['num_rec'])).opts(cmap='viridis') - lod = rasterize(gv.Points(ds, kdims=['evlo','evdp'], vdims=['num_rec'])).opts(cmap='viridis') - lad = rasterize(gv.Points(ds, kdims=['_evdp','evla'], vdims=['num_rec'])).opts(cmap='viridis') #.opts(invert_yaxis=True) - - #lola = spread(rasterize(gv.Points(ds, kdims=['evlo','evla'], vdims=['num_rec']), aggregator="sum").opts(cmap='viridis')) - #lod = spread(rasterize(gv.Points(ds, kdims=['evlo','evdp'], vdims=['num_rec']), aggregator="sum").opts(cmap='viridis')) - #lad = spread(rasterize(gv.Points(ds, kdims=['_evdp','evla'], vdims=['num_rec']), aggregator="sum").opts(cmap='viridis')) #.opts(invert_yaxis=True) - - # station points - SR.count_events_per_station() - # add '_stdp' column for plotting - SR.rec_points['_stdp'] = -0.001*SR.rec_points['stel'] # converting elevation [m] to depth [km] - df_sta = SR.rec_points.groupby(['staname']).apply(lambda x: x.iloc[-1]) - r_lola = gv.Points(df_sta, kdims=['stlo', 'stla'], vdims=['num_events','staname']).opts(color='red', size=2, tools=['hover'],) - r_lod = gv.Points(df_sta, kdims=['stlo', '_stdp'], vdims=['num_events','staname']).opts(color='red', size=2, tools=['hover'],) - r_lad = gv.Points(df_sta, kdims=['_stdp','stla'], vdims=['num_events','staname']).opts( color='red', size=2, tools=['hover'],) - - # fix layout - layout=(lola.opts(width=500, height=500)*tiles*r_lola + - lad.opts(width=200,height=500)*r_lad + - lod.opts(width=500,height=200, invert_yaxis=True)*r_lod).cols(2).opts(title='Nevada+SC+NC') - - if fname is not None: - gv.save(layout, fname+'.html') - - return layout - - else: - # weight plot - - # source points - SR.src_points['_evdp'] = SR.src_points['evdp'] # for protting - ds = gv.Dataset(SR.src_points, kdims=['evlo','evla','evdp','_evdp'], vdims=['weight']) - - # station points - SR.count_events_per_station() - # add '_stdp' column for plotting - SR.rec_points['_stdp'] = -0.001*SR.rec_points['stel'] # converting elevation [m] to depth [km] - df_sta = SR.rec_points.groupby(['staname']).apply(lambda x: x.iloc[-1]) - - # plot - - # show min and max weight - print("min weight = ", np.min(SR.src_points['weight'])) - print("max weight = ", np.max(SR.src_points['weight'])) - - w_evs = spread(rasterize(gv.Points(ds, kdims=['evlo','evla'], vdims=['weight']), aggregator="mean").opts( cmap='plasma'), how='source').opts(width=500, height=500, colorbar=True, tools=['hover'], colorbar_position='bottom') - w_sta = spread(rasterize(gv.Points(df_sta, kdims=['stlo', 'stla'], vdims=['weight']),aggregator="mean").opts(cmap='viridis'), px=2, how='source').opts(width=500, height=500, colorbar=True, tools=['hover'], colorbar_position='bottom') - - # log scale for colorbar - w_evs.opts(cmap='plasma', colorbar=True, logz=True) - w_sta.opts(cmap='viridis', colorbar=True, logz=True) - - if fname is not None: - gv.save((w_evs*tiles + w_sta*tiles).cols(2), fname+'.html') - - return (w_evs*tiles + w_sta*tiles).cols(2) diff --git a/test/test_create_model.py b/test/test_create_model.py index d43293e..9c6d175 100644 --- a/test/test_create_model.py +++ b/test/test_create_model.py @@ -34,6 +34,8 @@ def test_checkerboard03(self): lim_y=[-0.5, 0.5], lim_z=[10, 120] ) + mod = cm.to_attmodel() + mod.to_ani() def test_read_model(self): mod = ATTModel.read(self.out_fname, para_fname=self.para_fname) @@ -46,6 +48,13 @@ def test_read_model(self): start_point=[mod.min_max_lon[0], mod.min_max_lat[1]], end_point=[mod.min_max_lon[1], mod.min_max_lat[1]], field='vel', flat_earth=True) + def test_smooth_model(self): + mod = ATTModel.read(self.out_fname, para_fname=self.para_fname) + mod_test1 = mod.copy() + mod_test1.smooth(sigma=[2.0, 0.1, 0.1], unit_deg=True, smooth_ani=False) + mod_test2 = mod.copy() + mod_test2.smooth(sigma=[5.0, 20, 20], smooth_ani=True) + mod_test1.write('smoothed_model1.h5') if __name__ == '__main__': test = TestATTModel() diff --git a/test/test_para.py b/test/test_para.py new file mode 100644 index 0000000..415ea8c --- /dev/null +++ b/test/test_para.py @@ -0,0 +1,101 @@ +import unittest +import os +import shutil +from ruamel.yaml import YAML +from pytomoatt.para import ATTPara +import numpy as np + +yaml = YAML() + +class TestATTPara(unittest.TestCase): + def setUp(self): + self.test_dir = 'test_para_output' + if os.path.exists(self.test_dir): + shutil.rmtree(self.test_dir) + os.makedirs(self.test_dir) + self.cwd = os.getcwd() + os.chdir(self.test_dir) + + self.yaml_content = { + 'domain': { + 'min_max_dep': [0, 100], + 'min_max_lat': [30, 40], + 'min_max_lon': [100, 110], + 'n_rtp': [11, 11, 11] + }, + 'test_section': { + 'key1': 'value1' + } + } + self.fname = 'test_params.yml' + with open(self.fname, 'w') as f: + yaml.dump(self.yaml_content, f) + + def tearDown(self): + os.chdir(self.cwd) + if os.path.exists(self.test_dir): + shutil.rmtree(self.test_dir) + + def test_init(self): + para = ATTPara(self.fname) + self.assertEqual(para.input_params['domain']['n_rtp'], [11, 11, 11]) + self.assertEqual(para.input_params['test_section']['key1'], 'value1') + + def test_init_axis(self): + para = ATTPara(self.fname) + dep, lat, lon, dd, dt, dp = para.init_axis() + + # Check shapes based on n_rtp [11, 11, 11] + self.assertEqual(len(dep), 11) + self.assertEqual(len(lat), 11) + self.assertEqual(len(lon), 11) + + # Check values + # Note: init_axis flips the depth array + self.assertEqual(dep[0], 100) + self.assertEqual(dep[-1], 0) + self.assertEqual(lat[0], 30) + self.assertEqual(lat[-1], 40) + self.assertEqual(lon[0], 100) + self.assertEqual(lon[-1], 110) + + def test_update_param(self): + para = ATTPara(self.fname) + + # Test updating existing nested key + para.update_param('domain.n_rtp', '20,20,20') + self.assertEqual(para.input_params['domain']['n_rtp'], [20, 20, 20]) + + # Test updating existing simple key + para.update_param('test_section.key1', 'new_value') + self.assertEqual(para.input_params['test_section']['key1'], 'new_value') + + # Test adding new key + para.update_param('new_section.new_key', '123.45') + self.assertEqual(para.input_params['new_section']['new_key'], 123.45) + + def test_write(self): + para = ATTPara(self.fname) + para.update_param('domain.n_rtp', '30,30,30') + + out_fname = 'out_params.yml' + para.write(out_fname) + + self.assertTrue(os.path.exists(out_fname)) + + # Verify content + 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): + para = ATTPara(self.fname) + para.update_param('domain.n_rtp', '40,40,40') + para.write() # Should overwrite 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__': + unittest.main() diff --git a/test/test_rotate.py b/test/test_rotate.py new file mode 100644 index 0000000..9ff9e37 --- /dev/null +++ b/test/test_rotate.py @@ -0,0 +1,124 @@ +import unittest +import numpy as np +from pytomoatt.utils.rotate import ( + rtp2xyz, xyz2rtp, rotate_x, rotate_y, rotate_z, + rtp_rotation, rtp_rotation_reverse +) + +class TestRotate(unittest.TestCase): + def test_rtp2xyz_xyz2rtp(self): + # Test point: r=1, lat=0, lon=0 -> x=1, y=0, z=0 + x, y, z = rtp2xyz(1, 0, 0) + self.assertAlmostEqual(x, 1.0) + self.assertAlmostEqual(y, 0.0) + self.assertAlmostEqual(z, 0.0) + + r, t, p = xyz2rtp(x, y, z) + self.assertAlmostEqual(r, 1.0) + self.assertAlmostEqual(t, 0.0) + self.assertAlmostEqual(p, 0.0) + + # Test point: r=1, lat=90, lon=0 -> x=0, y=0, z=1 + x, y, z = rtp2xyz(1, 90, 0) + self.assertAlmostEqual(x, 0.0) + self.assertAlmostEqual(y, 0.0) + self.assertAlmostEqual(z, 1.0) + + r, t, p = xyz2rtp(x, y, z) + self.assertAlmostEqual(r, 1.0) + self.assertAlmostEqual(t, 90.0) + self.assertAlmostEqual(p, 0.0) + + # Test point: r=1, lat=0, lon=90 -> x=0, y=1, z=0 + x, y, z = rtp2xyz(1, 0, 90) + self.assertAlmostEqual(x, 0.0) + self.assertAlmostEqual(y, 1.0) + self.assertAlmostEqual(z, 0.0) + + r, t, p = xyz2rtp(x, y, z) + self.assertAlmostEqual(r, 1.0) + self.assertAlmostEqual(t, 0.0) + self.assertAlmostEqual(p, 90.0) + + def test_rotate_x(self): + # Rotate (0, 1, 0) by 90 degrees around X -> (0, 0, 1) + # Formula: y' = y cos - z sin, z' = y sin + z cos + # y=1, z=0, theta=90: y' = 0, z' = 1 + x, y, z = 0, 1, 0 + nx, ny, nz = rotate_x(x, y, z, 90) + self.assertAlmostEqual(nx, 0.0) + self.assertAlmostEqual(ny, 0.0) + self.assertAlmostEqual(nz, 1.0) + + def test_rotate_y(self): + # Rotate (1, 0, 0) by 90 degrees around Y -> (0, 0, -1) + # Formula: x' = x cos + z sin, z' = -x sin + z cos + # x=1, z=0, theta=90: x' = 0, z' = -1 + x, y, z = 1, 0, 0 + nx, ny, nz = rotate_y(x, y, z, 90) + self.assertAlmostEqual(nx, 0.0) + self.assertAlmostEqual(ny, 0.0) + self.assertAlmostEqual(nz, -1.0) + + def test_rotate_z(self): + # Rotate (1, 0, 0) by 90 degrees around Z -> (0, 1, 0) + # Formula: x' = x cos - y sin, y' = x sin + y cos + # x=1, y=0, theta=90: x' = 0, y' = 1 + x, y, z = 1, 0, 0 + nx, ny, nz = rotate_z(x, y, z, 90) + self.assertAlmostEqual(nx, 0.0) + self.assertAlmostEqual(ny, 1.0) + self.assertAlmostEqual(nz, 0.0) + + def test_rtp_rotation_roundtrip(self): + # Test that rotation and reverse rotation return original coordinates + lat, lon = 30.0, 60.0 + theta0, phi0, psi = 10.0, 20.0, 45.0 + + new_lat, new_lon = rtp_rotation(lat, lon, theta0, phi0, psi) + orig_lat, orig_lon = rtp_rotation_reverse(new_lat, new_lon, theta0, phi0, psi) + + self.assertAlmostEqual(lat, orig_lat) + self.assertAlmostEqual(lon, orig_lon) + + def test_rtp_rotation_specific(self): + # Test specific rotation + # Center at lat=0, lon=0. Rotate point (0, 0) -> should be (0, 0) if psi=0 + lat, lon = 0.0, 0.0 + theta0, phi0, psi = 0.0, 0.0, 0.0 + new_lat, new_lon = rtp_rotation(lat, lon, theta0, phi0, psi) + self.assertAlmostEqual(new_lat, 0.0) + self.assertAlmostEqual(new_lon, 0.0) + + # Center at lat=0, lon=0. Point at lat=0, lon=90. + # Rotate so center moves to lat=0, lon=0 (it is already there). + # If we rotate coordinate system? + # rtp_rotation logic: + # 1. rtp2xyz + # 2. rotate_z(-phi0) -> brings center longitude to 0 + # 3. rotate_y(theta0) -> brings center latitude to 0? + # rotate_y: x' = x cos + z sin. + # If center is at (lat=theta0, lon=0), x=cos(theta0), z=sin(theta0). + # rotate_y(theta0): x' = cos^2 + sin^2 = 1. z' = -cos*sin + sin*cos = 0. + # So rotate_y(theta0) brings (theta0, 0) to (0, 0) (x-axis). + # Wait, rotate_y(theta) rotates vector by theta. + # If we want to bring P(theta0, 0) to X-axis (0, 0), we need to rotate by -theta0? + # Let's check rotate_y implementation. + # new_x = x cos + z sin. + # If x=cos(theta0), z=sin(theta0). + # new_x = cos(theta0)cos(theta) + sin(theta0)sin(theta) = cos(theta0-theta) + # new_z = -cos(theta0)sin(theta) + sin(theta0)cos(theta) = sin(theta0-theta) + # If we want new_z=0 (lat=0), we need theta0-theta = 0 => theta = theta0. + # So rotate_y(theta0) rotates (theta0, 0) to (0, 0). Correct. + + # So rtp_rotation transforms coordinates such that (theta0, phi0) becomes (0, 0). + + lat, lon = 10.0, 20.0 + theta0, phi0 = 10.0, 20.0 + psi = 0.0 + new_lat, new_lon = rtp_rotation(lat, lon, theta0, phi0, psi) + self.assertAlmostEqual(new_lat, 0.0) + self.assertAlmostEqual(new_lon, 0.0) + +if __name__ == '__main__': + unittest.main() diff --git a/test/test_script.py b/test/test_script.py new file mode 100644 index 0000000..bef4a22 --- /dev/null +++ b/test/test_script.py @@ -0,0 +1,71 @@ +import unittest +from unittest.mock import patch, MagicMock +import os +import shutil +import sys +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): + self.test_dir = 'test_script_output' + if os.path.exists(self.test_dir): + shutil.rmtree(self.test_dir) + os.makedirs(self.test_dir) + self.cwd = os.getcwd() + os.chdir(self.test_dir) + + def tearDown(self): + os.chdir(self.cwd) + if os.path.exists(self.test_dir): + shutil.rmtree(self.test_dir) + + def test_init_pjt(self): + project_name = 'my_project' + with patch('builtins.input', return_value='y'): + with patch.object(sys, 'argv', ['pta', 'init_pjt', project_name]): + PTA() + self.assertTrue(os.path.exists(project_name)) + self.assertTrue(os.path.exists(os.path.join(project_name, 'input_params.yml'))) + + def test_setpar(self): + project_name = 'my_project' + with patch.object(sys, 'argv', ['pta', 'init_pjt', project_name]): + PTA() + + os.chdir(project_name) + with patch.object(sys, 'argv', ['pta', 'setpar', 'input_params.yml', 'domain.n_rtp', '10,10,10']): + PTA() + + 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): + # Create a dummy model file + project_name = 'my_project' + with patch.object(sys, 'argv', ['pta', 'init_pjt', project_name]): + PTA() + os.chdir(project_name) + + # Create a dummy h5 model + # The shape must match n_rtp in input_params.yml which is [10, 50, 50] by default + with h5py.File('model.h5', 'w') as f: + f.create_dataset('vel', data=np.zeros((10, 50, 50))) + f.create_dataset('xi', data=np.zeros((10, 50, 50))) + f.create_dataset('eta', data=np.zeros((10, 50, 50))) + f.create_dataset('zeta', data=np.zeros((10, 50, 50))) + + with patch.dict(sys.modules, {'pyvista': MagicMock()}): + with patch.object(sys, 'argv', ['pta', 'model2vtk', 'input_params.yml', '-i', 'model.h5', '-o', 'model.vtk']): + PTA() + + # self.assertTrue(os.path.exists('model.vtk')) # Mocked pyvista won't write file + +if __name__ == '__main__': + unittest.main() diff --git a/test/test_src_rec.py b/test/test_src_rec.py index 31173e1..811a126 100644 --- a/test/test_src_rec.py +++ b/test/test_src_rec.py @@ -1,21 +1,32 @@ -import pytest +import unittest from pytomoatt.src_rec import SrcRec +from pytomoatt.utils.src_rec_utils import ( + define_rec_cols, + get_rec_points_types, + setup_rec_points_dd, + update_position, + download_src_rec_file +) from os.path import dirname, join +from unittest.mock import MagicMock, patch +import pandas as pd +import io -class TestSrcRec: +class TestSrcRec(unittest.TestCase): fname: str = join(dirname(dirname(__file__)), 'examples', 'src_rec_file_eg') fname1: str = join(dirname(__file__), 'test_srcrec_a.dat') def test_subcase_01(self): sr = SrcRec.read(self.fname) sr.select_by_distance([0, 1]) - assert sr.rec_points.shape[0] == 19378 + self.assertEqual(sr.rec_points.shape[0], 19378) def test_subcase_02(self): sr = SrcRec.read(self.fname) sr.select_by_box_region([-1, 0, -1, 0]) - assert sr.rec_points.shape[0] == 671 and sr.src_points.shape[0] == 85 + self.assertEqual(sr.rec_points.shape[0], 671) + self.assertEqual(sr.src_points.shape[0], 85) def test_subcase_03(self): sr = SrcRec.read(self.fname) @@ -30,18 +41,21 @@ def test_subcase_04(self): def test_subcase_05(self): sr = SrcRec.read(self.fname) sr.select_by_depth([0, 10]) - assert sr.src_points.shape[0] == 1413 and sr.rec_points.shape[0] == 11226 + self.assertEqual(sr.src_points.shape[0], 1413) + self.assertEqual(sr.rec_points.shape[0], 11226) def test_subcase_06(self): sr = SrcRec.read(self.fname) sr1 = SrcRec.read(self.fname1, dist_in_data=True) sr.append(sr1) - assert sr.src_points.shape[0] == 2506 and sr._count_records() == 19926 + self.assertEqual(sr.src_points.shape[0], 2506) + self.assertEqual(sr._count_records(), 19926) def test_subcase_07(self): sr = SrcRec.read(self.fname) sr.select_by_azi_gap(120) - assert sr.src_points.shape[0] == 329 and sr.rec_points.shape[0] == 3815 + self.assertEqual(sr.src_points.shape[0], 329) + self.assertEqual(sr.rec_points.shape[0], 3815) def test_subcase_08(self): sr = SrcRec.read(self.fname) @@ -56,17 +70,133 @@ def test_subcase_10(self): sr = SrcRec.read(self.fname) sr.box_weighting(0.4, 10, obj='both') + +class TestSrcRecUtils(unittest.TestCase): + 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) + self.assertNotIn("dist_deg", cols) + self.assertNotIn("netname", cols) + self.assertEqual(last_col, 7) + + # Case 2: dist_in_data=True, name_net_and_sta=False + cols, last_col = define_rec_cols(True, False) + self.assertIn("dist_deg", cols) + self.assertNotIn("netname", cols) + self.assertEqual(last_col, 8) + + # Case 3: dist_in_data=False, name_net_and_sta=True + cols, last_col = define_rec_cols(False, True) + self.assertNotIn("dist_deg", cols) + self.assertIn("netname", cols) + self.assertEqual(last_col, 8) + + # Case 4: dist_in_data=True, name_net_and_sta=True + cols, last_col = define_rec_cols(True, True) + self.assertIn("dist_deg", cols) + self.assertIn("netname", cols) + self.assertEqual(last_col, 9) + + def test_get_rec_points_types(self): + types = get_rec_points_types(False) + self.assertNotIn("dist_deg", types) + + types = get_rec_points_types(True) + self.assertIn("dist_deg", types) + self.assertEqual(types["dist_deg"], float) + + def test_setup_rec_points_dd(self): + cols, types = setup_rec_points_dd('cs') + self.assertIn("rec_index1", cols) + self.assertIn("rec_index2", cols) + + cols, types = setup_rec_points_dd('cr') + self.assertIn("src_index", cols) + self.assertIn("src_index2", cols) + + with self.assertRaises(ValueError): + setup_rec_points_dd('invalid') + + def test_update_position(self): + # Mock SrcRec object + sr = MagicMock() + + # Setup DataFrames + sr.sources = pd.DataFrame({ + 'event_id': [1, 2], + 'evlo': [10.0, 20.0], + 'evla': [30.0, 40.0] + }) + + sr.receivers = pd.DataFrame({ + 'staname': ['STA1', 'STA2'], + 'stlo': [100.0, 110.0], + 'stla': [50.0, 60.0] + }) + + sr.src_points = pd.DataFrame({ + 'event_id': [1, 2], + 'evlo': [0.0, 0.0], # Old values + 'evla': [0.0, 0.0] # Old values + }) + + sr.rec_points = pd.DataFrame({ + 'staname': ['STA1', 'STA2'], + 'stlo': [0.0, 0.0], # Old values + 'stla': [0.0, 0.0] # Old values + }) + + sr.rec_points_cs = pd.DataFrame({ + 'staname1': ['STA1'], + 'staname2': ['STA2'], + 'stlo1': [0.0], 'stla1': [0.0], + 'stlo2': [0.0], 'stla2': [0.0] + }) + + sr.rec_points_cr = pd.DataFrame({ + 'staname': ['STA1'], + 'event_id2': [2], + 'stlo': [0.0], 'stla': [0.0], + 'evlo2': [0.0], 'evla2': [0.0] + }) + + update_position(sr) + + # Check src_points updated + self.assertEqual(sr.src_points.iloc[0]['evlo'], 10.0) + self.assertEqual(sr.src_points.iloc[0]['evla'], 30.0) + + # Check rec_points updated + self.assertEqual(sr.rec_points.iloc[0]['stlo'], 100.0) + self.assertEqual(sr.rec_points.iloc[0]['stla'], 50.0) + + # Check rec_points_cs updated + self.assertEqual(sr.rec_points_cs.iloc[0]['stlo1'], 100.0) + self.assertEqual(sr.rec_points_cs.iloc[0]['stlo2'], 110.0) + + # Check rec_points_cr updated + self.assertEqual(sr.rec_points_cr.iloc[0]['stlo'], 100.0) + self.assertEqual(sr.rec_points_cr.iloc[0]['evlo2'], 20.0) + + def test_download_src_rec_file(self): + with patch('urllib3.PoolManager') as mock_pool: + mock_http = mock_pool.return_value + mock_response = MagicMock() + mock_response.status = 200 + mock_response.headers = {'Content-Length': '10'} + mock_response.read.side_effect = [b'test data', b''] + mock_http.request.return_value = mock_response + + data = download_src_rec_file('http://example.com/file') + self.assertEqual(data.getvalue(), 'test data') + + # Test failure case + mock_response.status = 404 + data = download_src_rec_file('http://example.com/file') + self.assertIsNone(data) + + if __name__ == '__main__': - tsr = TestSrcRec() - tsr.test_subcase_01() - tsr.test_subcase_02() - tsr.test_subcase_03() - tsr.test_subcase_04() - tsr.test_subcase_05() - tsr.test_subcase_06() - tsr.test_subcase_07() - tsr.test_subcase_08() - tsr.test_subcase_09() - tsr.test_subcase_10() + unittest.main()