google-deepmind / google-deepmind/weathernext

Question Regarding Model Sharding

Open
#186 1 comment 0 reactions 0 assignees View on GitHub
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

Open the contributing 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.