google-deepmind / google-deepmind/weathernext
Question Regarding Model Sharding
- Dominant language
- Python
- Stars
- 7.7k
- Forks
- 986
- PR merge metrics
- No merged PRs in 30d
Description
Hi, I'm working on reimplementing FGN within the publicly available GenCast codebase. I was wondering if model sharding across ensembles is implemented similarly to AIFS-CRPS? And if so, could you provide some sample functions (eg all-gather) of how to handle the sharding of gradients across TPUs in JAX? Thank you!
Contributor guide
Research direction
No files, tests, or entry points are named. Start by locating the GenCast model-sharding implementation and comparing it with AIFS-CRPS, then inspect the TPU gradient and all-gather paths in JAX. Done means the repository documents whether ensemble sharding is supported and provides the requested sample functions.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- distributed-systems, machine-learning
- Issue type
- Documentation
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 25/100