2.1 MiB
2.1 MiB
In [1]:
import pandas as pd
from obspy.core.event import read_events
import matplotlib.pyplot as plt
import seisbench.models as sbm
import torch
import torch.nn as nn
import seisbench.data as sbd
import seisbench.generate as sbg
import seisbench.models as sbm
from seisbench.util import worker_seeding
import numpy as np
from torch.utils.data import DataLoader
from pathlib import Path
import wandb
import os
import sys
from pathlib import Path
project_path = str(Path.cwd().parent)
sys.path.append(project_path)
from scripts import train[34m[1mwandb[0m: Currently logged in as: [33mkmilian[0m ([33mepos[0m). Use [1m`wandb login --relogin`[0m to force relogin [34m[1mwandb[0m: [33mWARNING[0m If you're specifying your api key in code, ensure this code is not shared publicly. [34m[1mwandb[0m: [33mWARNING[0m Consider setting the WANDB_API_KEY environment variable, or running `wandb login` from the command line. [34m[1mwandb[0m: Appending key for api.wandb.ai to your netrc file: /Users/krystynamilian/.netrc
In [2]:
model = train.load_model()
run = wandb.init()
artifact = run.use_artifact('epos/training_seisbench_models_on_igf_data/model:v113', type='model')
artifact_dir = artifact.download()
fname = artifact_dir + "/" + os.listdir(artifact_dir)[0]
model.load_state_dict(torch.load(fname))
model.eval()Out [2]:
wandb version 0.15.4 is available! To upgrade, please run:
$ pip install wandb --upgrade
Tracking run with wandb version 0.15.3
Run data is saved locally in
/Users/krystynamilian/Documents/praca/Cyfronet/epos/ai/repo/demo_scripts/notebooks/wandb/run-20230704_110541-ir19n1xv View project at https://wandb.ai/epos/demo_scripts-notebooks
[34m[1mwandb[0m: 1 of 1 files downloaded.
PhaseNet(
(inc): Conv1d(3, 8, kernel_size=(7,), stride=(1,), padding=same)
(in_bn): BatchNorm1d(8, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(down_branch): ModuleList(
(0): ModuleList(
(0): Conv1d(8, 8, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(1): BatchNorm1d(8, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(8, 8, kernel_size=(7,), stride=(4,), padding=(3,), bias=False)
(3): BatchNorm1d(8, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(1): ModuleList(
(0): Conv1d(8, 16, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(1): BatchNorm1d(16, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(16, 16, kernel_size=(7,), stride=(4,), bias=False)
(3): BatchNorm1d(16, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(2): ModuleList(
(0): Conv1d(16, 32, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(1): BatchNorm1d(32, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(32, 32, kernel_size=(7,), stride=(4,), bias=False)
(3): BatchNorm1d(32, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(3): ModuleList(
(0): Conv1d(32, 64, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(1): BatchNorm1d(64, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(64, 64, kernel_size=(7,), stride=(4,), bias=False)
(3): BatchNorm1d(64, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(4): ModuleList(
(0): Conv1d(64, 128, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(1): BatchNorm1d(128, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2-3): 2 x None
)
)
(up_branch): ModuleList(
(0): ModuleList(
(0): ConvTranspose1d(128, 64, kernel_size=(7,), stride=(4,), bias=False)
(1): BatchNorm1d(64, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(128, 64, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(3): BatchNorm1d(64, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(1): ModuleList(
(0): ConvTranspose1d(64, 32, kernel_size=(7,), stride=(4,), bias=False)
(1): BatchNorm1d(32, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(64, 32, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(3): BatchNorm1d(32, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(2): ModuleList(
(0): ConvTranspose1d(32, 16, kernel_size=(7,), stride=(4,), bias=False)
(1): BatchNorm1d(16, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(32, 16, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(3): BatchNorm1d(16, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
(3): ModuleList(
(0): ConvTranspose1d(16, 8, kernel_size=(7,), stride=(4,), bias=False)
(1): BatchNorm1d(8, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
(2): Conv1d(16, 8, kernel_size=(7,), stride=(1,), padding=same, bias=False)
(3): BatchNorm1d(8, eps=0.001, momentum=0.1, affine=True, track_running_stats=True)
)
)
(out): Conv1d(8, 2, kernel_size=(1,), stride=(1,), padding=same)
(softmax): Softmax(dim=1)
)In [3]:
data_path = '../../../data/igf/seisbench_format'
sampling_rate = 100
data = sbd.WaveformDataset(data_path, sampling_rate=sampling_rate)
data.filter(data.metadata.trace_Pg_arrival_sample.notna())
pick_mae = train.PickMAE(sampling_rate)
splits = ['train', 'dev', 'test']
In [4]:
for split in splits:
print(f"\n\nModel resutls for {split} set")
print("\nFixed window")
# the results depend on random selection of a window around a pick
gen = train.get_data_generator(split=split, station=None, sampling_rate=sampling_rate, path=data_path, window='fixed')
data_loader = DataLoader(gen, batch_size=256, shuffle=False, num_workers=0,
worker_init_fn=worker_seeding)
test_loss, test_mae = train.test_one_epoch(model, data_loader, pick_mae, wandb_log=False)
Model resutls for train set Fixed window train (12444, 17) 100 Test avg loss: 0.025157 Test avg mae: 0.047488 Model resutls for dev set Fixed window dev (2773, 17) 100 Test avg loss: 0.025309 Test avg mae: 0.051242 Model resutls for test set Fixed window test (2785, 17) 100 Test avg loss: 0.025555 Test avg mae: 0.047198
In [5]:
frames_per_station = []
for split in splits:
frames_per_station.append(data.get_split(split).metadata.groupby('station_code').count()['index'])
frames_per_station = pd.DataFrame(frames_per_station, index=splits).transpose()
frames_per_station.plot(kind='bar', figsize=(15,3), title='Frames per station')
Out [5]:
<Axes: title={'center': 'Frames per station'}, xlabel='station_code'>In [6]:
stations = data.metadata.station_code.unique()
results = []
# highest_dev_mae = 0
# highest_test_mae = 0
for split in splits:
split_results = {}
print(split)
for station in stations:
print(station)
split_results[station] = {}
gen = train.get_data_generator(split=split, station=station, sampling_rate=sampling_rate, path=data_path, window='fixed')
data_loader = DataLoader(gen, batch_size=256, shuffle=False, num_workers=0,
worker_init_fn=worker_seeding)
test_loss, test_mae = None, None
try:
test_loss, test_mae = train.test_one_epoch(model, data_loader, pick_mae, wandb_log=False)
test_mae = float(test_mae)
except Exception as e:
print(e)
split_results[station]['mae']=test_mae
split_results[station]['loss']=test_loss
results.append(split_results)train BRDW train (12444, 17) 100 Test avg loss: 0.029592 Test avg mae: 0.089062 GROD train (12444, 17) 100 Test avg loss: 0.025606 Test avg mae: 0.057617 GUZI train (12444, 17) 100 Test avg loss: 0.024516 Test avg mae: 0.040501 JEDR train (12444, 17) 100 Test avg loss: 0.024908 Test avg mae: 0.044932 MOSK2 train (12444, 17) 100 Test avg loss: 0.024886 Test avg mae: 0.040648 NWLU train (12444, 17) 100 Test avg loss: 0.024519 Test avg mae: 0.032965 PCHB train (12444, 17) 100 Test avg loss: 0.024587 Test avg mae: 0.042854 PPOL train (12444, 17) 100 Test avg loss: 0.026397 Test avg mae: 0.076008 RUDN train (12444, 17) 100 Test avg loss: 0.025373 Test avg mae: 0.047411 RYNR train (12444, 17) 100 Test avg loss: 0.025934 Test avg mae: 0.066802 RZEC train (12444, 17) 100 Test avg loss: 0.023816 Test avg mae: 0.029310 SGOR train (12444, 17) 100 Test avg loss: 0.024461 Test avg mae: 0.034385 TRBC2 train (12444, 17) 100 Test avg loss: 0.025840 Test avg mae: 0.046694 TRN2 train (12444, 17) 100 Test avg loss: 0.025408 Test avg mae: 0.051002 TRZS train (12444, 17) 100 Test avg loss: 0.025081 Test avg mae: 0.043014 ZMST train (12444, 17) 100 Test avg loss: 0.025119 Test avg mae: 0.049605 LUBW train (12444, 17) 100 Test avg loss: 0.029460 Test avg mae: 0.095455 DWOL train (12444, 17) 100 Test avg loss: 0.024191 Test avg mae: 0.026768 LUBZ train (12444, 17) 100 Test avg loss: 0.032006 Test avg mae: 0.180000 ZUKW2 train (12444, 17) 100 Test avg loss: 0.024766 Test avg mae: 0.032847 DABR train (12444, 17) 100 Test avg loss: 0.024297 Test avg mae: 0.030349 PEKW2 train (12444, 17) 100 Test avg loss: 0.025126 Test avg mae: 0.044390 KRZY train (12444, 17) 100 Test avg loss: 0.025487 Test avg mae: 0.060000 OBIS train (12444, 17) 100 Test avg loss: 0.024071 Test avg mae: 0.028000 KAZI train (12444, 17) 100 Test avg loss: 0.023955 Test avg mae: 0.025048 KWLC train (12444, 17) 100 division by zero dev BRDW dev (2773, 17) 100 Test avg loss: 0.029170 Test avg mae: 0.093000 GROD dev (2773, 17) 100 Test avg loss: 0.025056 Test avg mae: 0.045482 GUZI dev (2773, 17) 100 Test avg loss: 0.024413 Test avg mae: 0.039008 JEDR dev (2773, 17) 100 Test avg loss: 0.024724 Test avg mae: 0.017778 MOSK2 dev (2773, 17) 100 Test avg loss: 0.024798 Test avg mae: 0.039188 NWLU dev (2773, 17) 100 Test avg loss: 0.025681 Test avg mae: 0.037319 PCHB dev (2773, 17) 100 Test avg loss: 0.024584 Test avg mae: 0.044390 PPOL dev (2773, 17) 100 Test avg loss: 0.025954 Test avg mae: 0.074627 RUDN dev (2773, 17) 100 Test avg loss: 0.025528 Test avg mae: 0.049302 RYNR dev (2773, 17) 100 Test avg loss: 0.026649 Test avg mae: 0.071854 RZEC dev (2773, 17) 100 division by zero SGOR dev (2773, 17) 100 Test avg loss: 0.025600 Test avg mae: 0.083871 TRBC2 dev (2773, 17) 100 Test avg loss: 0.025220 Test avg mae: 0.032712 TRN2 dev (2773, 17) 100 Test avg loss: 0.024823 Test avg mae: 0.039834 TRZS dev (2773, 17) 100 Test avg loss: 0.026505 Test avg mae: 0.093793 ZMST dev (2773, 17) 100 Test avg loss: 0.024234 Test avg mae: 0.030628 LUBW dev (2773, 17) 100 Test avg loss: 0.029821 Test avg mae: 0.129583 DWOL dev (2773, 17) 100 Test avg loss: 0.024637 Test avg mae: 0.048252 LUBZ dev (2773, 17) 100 Test avg loss: 0.023594 Test avg mae: 0.010000 ZUKW2 dev (2773, 17) 100 Test avg loss: 0.024708 Test avg mae: 0.035825 DABR dev (2773, 17) 100 Test avg loss: 0.024622 Test avg mae: 0.036894 PEKW2 dev (2773, 17) 100 Test avg loss: 0.026771 Test avg mae: 0.046392 KRZY dev (2773, 17) 100 Test avg loss: 0.038841 Test avg mae: 0.099091 OBIS dev (2773, 17) 100 Test avg loss: 0.025738 Test avg mae: 0.112784 KAZI dev (2773, 17) 100 Test avg loss: 0.024498 Test avg mae: 0.029636 KWLC dev (2773, 17) 100 Test avg loss: 0.023677 Test avg mae: 0.030000 test BRDW test (2785, 17) 100 Test avg loss: 0.026271 Test avg mae: 0.062836 GROD test (2785, 17) 100 Test avg loss: 0.025284 Test avg mae: 0.057771 GUZI test (2785, 17) 100 Test avg loss: 0.025403 Test avg mae: 0.050159 JEDR test (2785, 17) 100 Test avg loss: 0.027117 Test avg mae: 0.039194 MOSK2 test (2785, 17) 100 Test avg loss: 0.024491 Test avg mae: 0.035829 NWLU test (2785, 17) 100 Test avg loss: 0.024571 Test avg mae: 0.030738 PCHB test (2785, 17) 100 Test avg loss: 0.024751 Test avg mae: 0.039241 PPOL test (2785, 17) 100 Test avg loss: 0.025289 Test avg mae: 0.047288 RUDN test (2785, 17) 100 Test avg loss: 0.025869 Test avg mae: 0.056053 RYNR test (2785, 17) 100 Test avg loss: 0.027814 Test avg mae: 0.101556 RZEC test (2785, 17) 100 division by zero SGOR test (2785, 17) 100 Test avg loss: 0.025836 Test avg mae: 0.043481 TRBC2 test (2785, 17) 100 Test avg loss: 0.024740 Test avg mae: 0.031455 TRN2 test (2785, 17) 100 Test avg loss: 0.027810 Test avg mae: 0.063457 TRZS test (2785, 17) 100 Test avg loss: 0.024783 Test avg mae: 0.036907 ZMST test (2785, 17) 100 Test avg loss: 0.025570 Test avg mae: 0.050635 LUBW test (2785, 17) 100 Test avg loss: 0.024866 Test avg mae: 0.038000 DWOL test (2785, 17) 100 Test avg loss: 0.024782 Test avg mae: 0.035500 LUBZ test (2785, 17) 100 Test avg loss: 0.023713 Test avg mae: 0.020000 ZUKW2 test (2785, 17) 100 Test avg loss: 0.024829 Test avg mae: 0.034684 DABR test (2785, 17) 100 Test avg loss: 0.025824 Test avg mae: 0.040625 PEKW2 test (2785, 17) 100 Test avg loss: 0.024243 Test avg mae: 0.036092 KRZY test (2785, 17) 100 Test avg loss: 0.025173 Test avg mae: 0.060000 OBIS test (2785, 17) 100 Test avg loss: 0.026633 Test avg mae: 0.056509 KAZI test (2785, 17) 100 Test avg loss: 0.025799 Test avg mae: 0.038860 KWLC test (2785, 17) 100 Test avg loss: 0.025239 Test avg mae: 0.052885
In [7]:
for i, split in enumerate(splits):
results_df = pd.DataFrame(results[i]).transpose()
# for station, values in results_df[['mae']].itertuples():
# results_df.loc[station, 'mae'] = np.mean([v for v in values if v is not None])
# for station, values in results_df[['loss']].itertuples():
# results_df.loc[station, 'loss'] = np.mean([v for v in values if v is not None])
results_df.plot(kind='bar', figsize=(15,4), title=f"Mean results per station in {split} set")
results[i] = results_dfIn [8]:
dev_res = results[1]
station_with_worst_res_dev_set = dev_res[dev_res.mae == dev_res.mae.max()].index[0]
highest_dev_mae = dev_res.loc[station_with_worst_res_dev_set, 'mae']
test_res = results[2]
station_with_worst_res_test_set = test_res[test_res.mae == test_res.mae.max()].index[0]
highest_test_mae = test_res.loc[station_with_worst_res_test_set, 'mae']
print(f"highest mean MAE in dev set: {highest_dev_mae:.2f} for station: {station_with_worst_res_dev_set}")
print(f"highest mean MAE in test set: {highest_test_mae:.2f} for station: {station_with_worst_res_test_set}")
highest mean MAE in dev set: 0.13 for station: LUBW highest mean MAE in test set: 0.10 for station: RYNR
In [9]:
def plot_sample(sample, model, i, desc=None):
fig = plt.figure(figsize=(15, 10))
axs = fig.subplots(2, 1, sharex=True, gridspec_kw={"hspace": 0, "height_ratios": [3, 2]})
axs[0].plot(sample["X"][0].T, label='x')
plt.legend()
axs[1].plot(sample["y"][0].T, label='y')
model.eval() # close the model for evaluation
with torch.no_grad():
pred = model(torch.tensor(sample["X"], device=model.device).unsqueeze(0)) # Add a fake batch dimension
pred = pred[0].cpu().numpy()
axs[1].plot(pred[0], label='pred', color='orange')
plt.legend()
pred_pick_idx = np.argmax(pred[0])
true_pick_idx = np.argmax(sample['y'][0])
mae_error = np.abs(pred_pick_idx - true_pick_idx) /100 #mae in seconds
fig.suptitle(f"Predictions for sample: {i} {desc}, mae: {mae_error}s")
plt.show()
In [10]:
##### dev setIn [11]:
mean_mae = 0
samples = []
split = 'dev'
station = station_with_worst_res_dev_set
window='fixed'
gen = train.get_data_generator(split=split, station=station , sampling_rate=sampling_rate, path=data_path,
window='fixed')
station_mae = []
with torch.no_grad():
for i in range(len(gen)):
# idx = np.random.randint(len(gen))
idx = i
sample = gen[idx]
samples.append(sample)
pred = model(torch.tensor(sample["X"], device=model.device).unsqueeze(0))
pred = pred[0].cpu().numpy()
pred_pick_idx = np.argmax(pred[0])
true_pick_idx = np.argmax(sample['y'][0])
mae_error = np.abs(pred_pick_idx - true_pick_idx) /100 #mae in seconds
station_mae.append(mae_error)
sorted = np.argsort(station_mae)[::-1]
mean_mae = np.mean(station_mae)
print(mean_mae, station_with_worst_res_dev_set)
print(np.array(station_mae)[sorted])
## plot samples with mae error at leas 0.2s
for idx in sorted:
if station_mae[idx] < 0.1:
break
print(idx, station_mae[idx])
plot_sample(samples[idx], model, idx, desc=f" from station {station} {split} set")
dev (2773, 17) 100
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
0.12958333333333336 LUBW [1.06 0.4 0.33 0.28 0.23 0.1 0.08 0.07 0.07 0.07 0.07 0.06 0.05 0.05 0.05 0.04 0.03 0.03 0.01 0.01 0.01 0.01 0. 0. ] 2 1.06
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
3 0.4
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
13 0.33
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
22 0.28
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
17 0.23
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
18 0.1
In [12]:
mean_mae = 0
samples = []
split = 'test'
station = station_with_worst_res_test_set
for i in range(3):
gen = train.get_data_generator(split=split, station=station , sampling_rate=sampling_rate, path=data_path,
window='fixed')
station_mae = []
with torch.no_grad():
for i in range(len(gen)):
# idx = np.random.randint(len(gen))
idx = i
sample = gen[idx]
samples.append(sample)
pred = model(torch.tensor(sample["X"], device=model.device).unsqueeze(0))
pred = pred[0].cpu().numpy()
pred_pick_idx = np.argmax(pred[0])
true_pick_idx = np.argmax(sample['y'][0])
mae_error = np.abs(pred_pick_idx - true_pick_idx) /100 #mae in seconds
station_mae.append(mae_error)
sorted = np.argsort(station_mae)[::-1]
mean_mae = np.mean(station_mae)
print(np.array(station_mae)[sorted])
## plot samples with mae error at leas 0.2s
for idx in sorted:
if station_mae[idx] < 0.2:
break
print(idx, station_mae[idx])
plot_sample(samples[idx], model, idx, desc=f" from station {station} {split} set")
test (2785, 17) 100 test (2785, 17) 100 test (2785, 17) 100
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
[3.11 2.33 1.23 1.21 0.59 0.53 0.39 0.23 0.22 0.12 0.11 0.11 0.1 0.1 0.1 0.09 0.09 0.08 0.08 0.08 0.07 0.07 0.07 0.06 0.06 0.06 0.06 0.06 0.05 0.05 0.05 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.04 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.03 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.02 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0.01 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. ] 91 3.11
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
46 2.33
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
63 1.23
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
40 1.21
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
28 0.59
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
106 0.53
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
24 0.39
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
100 0.23
No artists with labels found to put in legend. Note that artists whose label start with an underscore are ignored when legend() is called with no argument.
119 0.22
In [ ]: