Skip to content

Tracker module for ML based TC tracking algortihms - #207

Draft
surbhigoel77 wants to merge 45 commits into
mainfrom
ml-integration
Draft

surbhigoel77 wants to merge 45 commits into
mainfrom
ml-integration

Conversation

@surbhigoel77

@surbhigoel77 surbhigoel77 commented Jun 2, 2026 •

Copy link
Copy Markdown
Member

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 dataclass TCMLParameters to fetch ML Parameters and a base abstract class TCMLTracker that 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 in ml_tracker.py and 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 like TCTrackerMetadata and Trajectory are inherited from the original core/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)

@surbhigoel77
surbhigoel77 marked this pull request as draft June 2, 2026 17:50
This was referenced Jun 2, 2026
@surbhigoel77 surbhigoel77 self-assigned this Jun 3, 2026
@surbhigoel77 surbhigoel77 added enhancement New feature or request MLTracker labels Jun 24, 2026
@sjavis sjavis mentioned this pull request Aug 27, 2026
3 tasks done
@surbhigoel77
surbhigoel77 requested a review from sjavis August 28, 2026 08:30
@surbhigoel77

surbhigoel77 commented Aug 28, 2026 •

Copy link
Copy Markdown
Member Author

@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

  1. the hugging face repo for model weights, and
  2. the required ERA5 data and norm stats

Let me know if have issues accessing it.

The integration test is not to be reviewed atm, I am making changes to it

@sjavis sjavis left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread src/tctrack/machine_learning/cyclone_track_ml.py
Comment on lines +110 to +128
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.
"""

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I suggest making a separate class for the stitching parameters.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This will also need adding to the __init__.py file

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

And the parameters can probably loose the stitch_ prefix since they are now in the class.

Comment thread src/tctrack/machine_learning/cyclone_track_ml.py
Comment thread src/tctrack/machine_learning/cyclone_track_ml.py
Comment thread src/tctrack/machine_learning/cyclone_track_ml.py Outdated
Comment thread src/tctrack/machine_learning/cyclone_track_ml.py Outdated
self._trajectories: list[Trajectory] = []
self._scores: list[dict] = []
# TC locations found by detect(), consumed by stitch().
self._candidates: list[dict] = []

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@sjavis sjavis Sep 16, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Comment thread src/tctrack/machine_learning/cyclone_track_ml.py Outdated
Comment thread src/tctrack/machine_learning/cyclone_track_ml.py Outdated

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Just to clarify, is this file intended just for development and to be removed before merging?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@surbhigoel77
surbhigoel77 requested a review from sjavis September 15, 2026 09:24

@sjavis sjavis left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +110 to +128
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.
"""

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This will also need adding to the __init__.py file

Comment thread src/tctrack/machine_learning/cyclone_track_ml.py
Comment on lines +72 to +79
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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
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] = []

@sjavis sjavis Sep 16, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Comment on lines +8 to +9
import torch
from huggingface_hub import hf_hub_download

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

torch and huggingface_hub should be added as dependencies to the pyproject.toml

MLTracker,
)

RDS_DIR = "/home/sg2147/rds/rds-inspire-tc-TqEGHMWTn8A/sg2147"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
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

Comment on lines +86 to +87
model_path=MODEL_PATH,
normalisation_stats_path=stats_file,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can these arguments be ignored here or are the files different to the defaults?

Comment on lines +132 to +140
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}")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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")

Comment on lines +174 to +185
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)")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request MLTracker

Projects

None yet

2 participants