diff --git a/pytomoatt/src_rec.py b/pytomoatt/src_rec.py index 36ee128..4aad51d 100644 --- a/pytomoatt/src_rec.py +++ b/pytomoatt/src_rec.py @@ -476,82 +476,181 @@ def write(self, fname="src_rec_file"): rec_points_cs = self.rec_points_cs rec_points_cr = self.rec_points_cr - for src in tqdm.tqdm( - src_points.itertuples(), + # Pre-format receiver records by source so the source loop only + # performs dictionary lookups and writes complete string blocks. + rec_lines_by_src = {} + for row in zip( + rec_points["src_index"].to_numpy(), + rec_points["rec_index"].to_numpy(), + rec_points["staname"].to_numpy(), + rec_points["stla"].to_numpy(), + rec_points["stlo"].to_numpy(), + rec_points["stel"].to_numpy(), + rec_points["phase"].to_numpy(), + rec_points["tt"].to_numpy(), + rec_points["weight"].to_numpy(), + ): + ( + src_index, + rec_index, + staname, + stla, + stlo, + stel, + phase, + tt, + weight, + ) = row + rec_lines_by_src.setdefault(src_index, []).append( + ( + f"{src_index:7d} {rec_index:7d} {staname!s:>6} " + f"{stla:9.4f} {stlo:9.4f} {stel:9.4f} " + f"{phase!s} {tt:8.4f} {weight:7.4f}\n" + ) + ) + rec_lines_by_src = { + src_index: "".join(lines) + for src_index, lines in rec_lines_by_src.items() + } + + rec_cs_lines_by_src = {} + if not rec_points_cs.empty: + for row in zip( + rec_points_cs["src_index"].to_numpy(), + rec_points_cs["rec_index1"].to_numpy(), + rec_points_cs["staname1"].to_numpy(), + rec_points_cs["stla1"].to_numpy(), + rec_points_cs["stlo1"].to_numpy(), + rec_points_cs["stel1"].to_numpy(), + rec_points_cs["rec_index2"].to_numpy(), + rec_points_cs["staname2"].to_numpy(), + rec_points_cs["stla2"].to_numpy(), + rec_points_cs["stlo2"].to_numpy(), + rec_points_cs["stel2"].to_numpy(), + rec_points_cs["phase"].to_numpy(), + rec_points_cs["tt"].to_numpy(), + rec_points_cs["weight"].to_numpy(), + ): + ( + src_index, + rec_index1, + staname1, + stla1, + stlo1, + stel1, + rec_index2, + staname2, + stla2, + stlo2, + stel2, + phase, + tt, + weight, + ) = row + rec_cs_lines_by_src.setdefault(src_index, []).append( + ( + f"{src_index:7d} {rec_index1:7d} " + f"{staname1!s:>6} {stla1:9.4f} {stlo1:9.4f} " + f"{stel1:9.4f} {rec_index2:7d} {staname2!s:>6} " + f"{stla2:9.4f} {stlo2:9.4f} {stel2:9.4f} " + f"{phase!s} {tt:8.4f} {weight:7.4f}\n" + ) + ) + rec_cs_lines_by_src = { + src_index: "".join(lines) + for src_index, lines in rec_cs_lines_by_src.items() + } + + rec_cr_lines_by_src = {} + if not rec_points_cr.empty: + for row in zip( + rec_points_cr["src_index"].to_numpy(), + rec_points_cr["rec_index"].to_numpy(), + rec_points_cr["staname"].to_numpy(), + rec_points_cr["stla"].to_numpy(), + rec_points_cr["stlo"].to_numpy(), + rec_points_cr["stel"].to_numpy(), + rec_points_cr["src_index2"].to_numpy(), + rec_points_cr["event_id2"].to_numpy(), + rec_points_cr["evla2"].to_numpy(), + rec_points_cr["evlo2"].to_numpy(), + rec_points_cr["evdp2"].to_numpy(), + rec_points_cr["phase"].to_numpy(), + rec_points_cr["tt"].to_numpy(), + rec_points_cr["weight"].to_numpy(), + ): + ( + src_index, + rec_index, + staname, + stla, + stlo, + stel, + src_index2, + event_id2, + evla2, + evlo2, + evdp2, + phase, + tt, + weight, + ) = row + rec_cr_lines_by_src.setdefault(src_index, []).append( + ( + f"{src_index:7d} {rec_index:7d} {staname!s:>6} " + f"{stla:9.4f} {stlo:9.4f} {stel:9.4f} " + f"{src_index2:7d} {event_id2!s:>6} " + f"{evla2:9.4f} {evlo2:9.4f} {evdp2:9.4f} " + f"{phase!s} {tt:8.4f} {weight:7.4f}\n" + ) + ) + rec_cr_lines_by_src = { + src_index: "".join(lines) + for src_index, lines in rec_cr_lines_by_src.items() + } + + source_columns = [ + "origin_time", + "evla", + "evlo", + "evdp", + "mag", + "num_rec", + "event_id", + "weight", + ] + source_rows = src_points[source_columns].itertuples(name=None) + for row in tqdm.tqdm( + source_rows, total=src_points.shape[0], desc="Writing src_rec file", ): - idx = src.Index - time_lst = ( - src.origin_time.strftime("%Y_%m_%d_%H_%M_%S.%f").split("_") + ( + idx, + origin_time, + evla, + evlo, + evdp, + mag, + num_rec, + event_id, + weight, + ) = row + time_fields = " ".join( + origin_time.strftime("%Y_%m_%d_%H_%M_%S.%f").split("_") + ) + output.write( + f"{idx:d} {time_fields} {evla:.4f} {evlo:.4f} " + f"{evdp:.4f} {mag:.4f} {num_rec} {event_id} " + f"{weight:.4f}\n" ) - output.write("{:d} {} {} {} {} {} {} {:.4f} {:.4f} {:.4f} {:.4f} {} {} {:.4f}\n".format( - idx, - *time_lst, - src.evla, - src.evlo, - src.evdp, - src.mag, - src.num_rec, - src.event_id, - src.weight, - )) if self.src_only: continue - rec_data = rec_points[rec_points["src_index"] == idx] - for rec in rec_data.itertuples(): - output.write(" {:d} {:d} {} {:6.4f} {:6.4f} {:6.4f} {} {:6.4f} {:6.4f}\n".format( - idx, - rec.rec_index, - rec.staname, - rec.stla, - rec.stlo, - rec.stel, - rec.phase, - rec.tt, - rec.weight, - )) - - if not rec_points_cs.empty: - rec_data = rec_points_cs[rec_points_cs["src_index"] == idx] - for rec in rec_data.itertuples(): - output.write(" {:d} {:d} {} {:6.4f} {:6.4f} {:6.4f} {:d} {} {:6.4f} {:6.4f} {:6.4f} {} {:.4f} {:6.4f}\n".format( - idx, - rec.rec_index1, - rec.staname1, - rec.stla1, - rec.stlo1, - rec.stel1, - rec.rec_index2, - rec.staname2, - rec.stla2, - rec.stlo2, - rec.stel2, - rec.phase, - rec.tt, - rec.weight, - )) - if not rec_points_cr.empty: - rec_data = rec_points_cr[rec_points_cr["src_index"] == idx] - for rec in rec_data.itertuples(): - output.write(" {:d} {:d} {} {:6.4f} {:6.4f} {:6.4f} {:d} {} {:6.4f} {:6.4f} {:6.4f} {} {:.4f} {:6.4f}\n".format( - idx, - rec.rec_index, - rec.staname, - rec.stla, - rec.stlo, - rec.stel, - rec.src_index2, - rec.event_id2, - rec.evla2, - rec.evlo2, - rec.evdp2, - rec.phase, - rec.tt, - rec.weight, - )) + output.write(rec_lines_by_src.get(idx, "")) + output.write(rec_cs_lines_by_src.get(idx, "")) + output.write(rec_cr_lines_by_src.get(idx, "")) with open(fname, "w") as f: f.write(output.getvalue()) diff --git a/test/test_src_rec.py b/test/test_src_rec.py index a51e666..419effc 100644 --- a/test/test_src_rec.py +++ b/test/test_src_rec.py @@ -625,7 +625,66 @@ def test_write_sources_and_receivers_format_weights(self): 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): + def test_write_roundtrip_preserves_microsecond_and_weight_precision(self): + """Write→read roundtrip preserves microsecond origin time and 4-decimal weight precision.""" + sr = SrcRec("unused") + + origin_time_with_us = pd.Timestamp("2013-10-06 09:20:53.123456") + src_weight = 1.0 / 3.0 # 0.3333... + rec_weight = 2.0 / 3.0 # 0.6666... + + src_df = pd.DataFrame({ + "origin_time": [origin_time_with_us], + "evla": [-1.7673], + "evlo": [-0.6619], + "evdp": [9.55], + "mag": [2.66], + "num_rec": [1], + "event_id": ["EVT_0001"], + "weight": [src_weight], + }) + src_df.index = pd.Index([0], name="src_index") + sr.src_points = src_df + + rec_df = pd.DataFrame({ + "src_index": [0], + "rec_index": [0], + "staname": ["STA0"], + "stla": [-1.0351], + "stlo": [-0.3383], + "stel": [219.0], + "phase": ["P"], + "tt": [14.786], + "weight": [rec_weight], + }) + sr.rec_points = rec_df + + with TemporaryDirectory() as directory: + output_file = join(directory, "src_rec.dat") + sr.write(output_file) + reread = SrcRec.read(output_file) + + # Verify microsecond precision is preserved in origin time + self.assertEqual( + reread.src_points["origin_time"].iloc[0], + origin_time_with_us, + ) + + # Verify source weight is preserved to 4 decimal places + self.assertAlmostEqual( + reread.src_points["weight"].iloc[0], + src_weight, + places=4, + ) + + # Verify receiver weight is preserved to 4 decimal places + self.assertAlmostEqual( + reread.rec_points["weight"].iloc[0], + rec_weight, + places=4, + ) + + sr = SrcRec.read(self.fname, src_only=True) figure = sr.plot(color_by="weight")