Tracker module for ML based TC tracking algortihms - #207
surbhigoel77 wants to merge 45 commits into
Conversation
|
@sjavis you can try and run the ML Tracker using the 'run_tracker_on_real_data.py' script. You will need ERA5 data and the normalisation stats to run this. I have already provided the absolute paths/creds in the code for
Let me know if have issues accessing it. The integration test is not to be reviewed atm, I am making changes to it |
sjavis
left a comment
There was a problem hiding this comment.
I haven't reviewed the stitching or the integration test yet but I thought I would give my suggestions so far. Overall the structure seems good.
| stitch_max_distance_deg: float = 3.0 | ||
| """Maximum distance (degrees lat/lon) between a track's last point and a | ||
| candidate in the next timestep for them to be linked as the same storm. | ||
|
|
||
| Not part of the reference training pipeline, which only trains | ||
| per-pixel classification and has no timestep-to-timestep linking step. | ||
| The default assumes 6-hourly input (as used for training) and typical | ||
| cyclone translation speeds; reduce for higher-frequency input. | ||
| """ | ||
|
|
||
| stitch_max_gap: int = 1 | ||
| """Number of consecutive timesteps a track may go unmatched before it is | ||
| closed, tolerating brief drops below :attr:`~TCMLParameters.threshold`. | ||
| """ | ||
|
|
||
| stitch_min_length: int = 2 | ||
| """Minimum number of observations a track must have to be kept, filtering | ||
| out single-timestep detections likely to be noise. | ||
| """ |
There was a problem hiding this comment.
I suggest making a separate class for the stitching parameters.
There was a problem hiding this comment.
This will also need adding to the __init__.py file
There was a problem hiding this comment.
And the parameters can probably loose the stitch_ prefix since they are now in the class.
| self._trajectories: list[Trajectory] = [] | ||
| self._scores: list[dict] = [] | ||
| # TC locations found by detect(), consumed by stitch(). | ||
| self._candidates: list[dict] = [] |
There was a problem hiding this comment.
I suggest we make a candidate class (maybe using a TypedDict) since it is used in a few places and it would make the typing clearer than just using a dict.
There was a problem hiding this comment.
This is looking good and I realise it was more complicated than I anticipated as I thought extra_keys was available (rather than being new in python 3.15).
I suggest making the extra variables inside a data dictionary to resolve the need for casting, etc. I will make a separate PR to show you my suggestion as it requires changes in a few places. (Added in #259)
There was a problem hiding this comment.
Just to clarify, is this file intended just for development and to be removed before merging?
There was a problem hiding this comment.
I want to add it to the tutorials as we have tutorial scripts for all other models, so might as well add a script for the ml one.
sjavis
left a comment
There was a problem hiding this comment.
This is looking good. Thank you for the changes that you previously made. I've left some comments and will open a separate PR for the changes to Candidate that I suggested.
| stitch_max_distance_deg: float = 3.0 | ||
| """Maximum distance (degrees lat/lon) between a track's last point and a | ||
| candidate in the next timestep for them to be linked as the same storm. | ||
|
|
||
| Not part of the reference training pipeline, which only trains | ||
| per-pixel classification and has no timestep-to-timestep linking step. | ||
| The default assumes 6-hourly input (as used for training) and typical | ||
| cyclone translation speeds; reduce for higher-frequency input. | ||
| """ | ||
|
|
||
| stitch_max_gap: int = 1 | ||
| """Number of consecutive timesteps a track may go unmatched before it is | ||
| closed, tolerating brief drops below :attr:`~TCMLParameters.threshold`. | ||
| """ | ||
|
|
||
| stitch_min_length: int = 2 | ||
| """Minimum number of observations a track must have to be kept, filtering | ||
| out single-timestep detections likely to be noise. | ||
| """ |
There was a problem hiding this comment.
This will also need adding to the __init__.py file
| Plain ``hypot(dlat, dlon)`` on raw coordinate differences is wrong in two | ||
| ways: it overstates distances that cross the antimeridian (e.g. 179E to | ||
| 179W is 2 degrees apart, not 358), and it ignores that a degree of | ||
| longitude covers less physical distance away from the equator. This | ||
| wraps the longitude difference to ``(-180, 180]`` and scales it by the | ||
| cosine of the mean latitude to correct for both. Still an approximation | ||
| - the Haversine formula would be exact - but adequate at the latitudes | ||
| tropical cyclones occur at. |
There was a problem hiding this comment.
| Plain ``hypot(dlat, dlon)`` on raw coordinate differences is wrong in two | |
| ways: it overstates distances that cross the antimeridian (e.g. 179E to | |
| 179W is 2 degrees apart, not 358), and it ignores that a degree of | |
| longitude covers less physical distance away from the equator. This | |
| wraps the longitude difference to ``(-180, 180]`` and scales it by the | |
| cosine of the mean latitude to correct for both. Still an approximation | |
| - the Haversine formula would be exact - but adequate at the latitudes | |
| tropical cyclones occur at. | |
| An approximation - the Haversine formula would be exact - but adequate at small | |
| distances and the latitudes tropical cyclones occur at. |
A smaller description is probably sufficient.
| self._trajectories: list[Trajectory] = [] | ||
| self._scores: list[dict] = [] | ||
| # TC locations found by detect(), consumed by stitch(). | ||
| self._candidates: list[dict] = [] |
There was a problem hiding this comment.
This is looking good and I realise it was more complicated than I anticipated as I thought extra_keys was available (rather than being new in python 3.15).
I suggest making the extra variables inside a data dictionary to resolve the need for casting, etc. I will make a separate PR to show you my suggestion as it requires changes in a few places. (Added in #259)
| import torch | ||
| from huggingface_hub import hf_hub_download |
There was a problem hiding this comment.
torch and huggingface_hub should be added as dependencies to the pyproject.toml
| MLTracker, | ||
| ) | ||
|
|
||
| RDS_DIR = "/home/sg2147/rds/rds-inspire-tc-TqEGHMWTn8A/sg2147" |
There was a problem hiding this comment.
| RDS_DIR = "/home/sg2147/rds/rds-inspire-tc-TqEGHMWTn8A/sg2147" | |
| RDS_DIR = "/rds/project/rds-TqEGHMWTn8A/sg2147" |
Longer term it would be best if we could make this not need data on CSD3
| model_path=MODEL_PATH, | ||
| normalisation_stats_path=stats_file, |
There was a problem hiding this comment.
Can these arguments be ignored here or are the files different to the defaults?
| def fetch_norm_stats(stats_file: str) -> None: | ||
| """Pull the real normalisation statistics from the reference repo.""" | ||
| git = ["git", "-C", REFERENCE_REPO] | ||
| subprocess.run([*git, "fetch", "origin"], check=True) # noqa: S603 | ||
| with open(stats_file, "wb") as handle: | ||
| subprocess.run( # noqa: S603 | ||
| [*git, "show", REFERENCE_STATS_REF], stdout=handle, check=True | ||
| ) | ||
| print(f" fetched statistics from {REFERENCE_STATS_REF}") |
There was a problem hiding this comment.
Should this now be able to use the bundled stats file?
| print(f" fetched statistics from {REFERENCE_STATS_REF}") | ||
|
|
||
|
|
||
| def test_pipeline(input_file: str, stats_file: str) -> None: |
There was a problem hiding this comment.
I would suggest converting this and test_output_writers to be compatible with pytest. But make it skip them by default with something like:
@pytest.mark.skip(reason="Requires data on CSD3")
| trajectories = tracker.stitch() | ||
| assert tracker.read_trajectories() == trajectories, ( | ||
| "read_trajectories() must return what stitch() produced, else to_netcdf() " | ||
| "would write nothing" | ||
| ) | ||
| for trajectory in trajectories: | ||
| assert trajectory.observations >= tracker.stitch_parameters.stitch_min_length | ||
| for key in ("time", "lat", "lon"): | ||
| assert key in trajectory.data, f"trajectory missing '{key}'" | ||
| print(f" stitch(): {len(trajectories)} trajectories, all well-formed OK") | ||
| if not trajectories: | ||
| print(" (none found - fine, model quality is not under test)") |
There was a problem hiding this comment.
This should probably check that there is at least one trajectory detected (although I don't know if any would be detected from the input data). It doesn't really check anything otherwise.
Linked issue: #190
The PR aims to build the core infrastructure for running ML models within the TCTrack framework. It uses exisiting base classes and defines new ml-specific classes.
It includes the following files (based on #204):
ml_tracker.py: contains a base dataclassTCMLParametersto fetch ML Parameters and a base abstract classTCMLTrackerthat extends the existing Tracker by adding a load_model function (for ml models). These classes remain common for all ML models (basically it is a way to standardise processes for any incoming new model).cyclone_track_ml.py: inherits the base classes defined inml_tracker.pyand adds new model specific functions. Reads the model from hugging face hub repository and connects it to the run_tracker pipeline. For classes that are common for both ML and non-ML models likeTCTrackerMetadataandTrajectoryare inherited from the originalcore/tracker.py.'run_tracker_on_real_data.py' : script to run the ML Tracker on a subset of ERA5 data.
test_ml_tracker.py: unit, integration (Still in progress)Assumptions:
merge_distance_deg = 2.0 (How close (in degrees) two candidates in the same timestep can be before the weaker one is discarded as a duplicate)
stitch_max_distance_deg = 3.0 (How far (in degrees) a candidate can be from a track's last known position and still be claimed as that same storm continuing)
stitch_max_gap: int = 1 (How many timesteps in a row a track can go unmatched before it's retired)
stitch_min_length: int = 2 (The minimum number of points a track must have to be kept in the final output)