google/grain: a declarative data pipeline for JAX training
Library for reading and processing ML training data.
At a glance
- What is it?
- Grain is a Python library for reading and transforming training data for JAX models, with a MapDataset API that chains shuffle, map and batch steps. It installs from PyPI as grain and runs its transformations on the CPU.
- Who is it for?
- Adopt Grain if you are training JAX models and want data loading expressed as a chain of shuffle, map and batch calls rather than a hand-written loader loop, and if your pipeline can live on the CPU. Do not adopt it if you need GPU or TPU processing inside the transformation steps, or if you are on Windows ARM, which the supported-platform table does not list.
- Can I use it commercially?
- Yes. Apache-2.0 is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
- Is it still maintained?
- Yes. The repository last received commits 1 day ago.
- What is it written in?
- Mainly Python, according to GitHub's language statistics.
Answers come from the project's GitHub data, last synced on September 30, 2026, and from our analysis. They are not legal advice.
Editorial analysis
What Grain replaces in a JAX training loop
Training a JAX model needs a stream of batches, and that stream usually starts as a hand-written Python generator: open files, decode records, shuffle a buffer, group into batches, hand the result to the training step. The README positions Grain as the library that takes over that job. Its stated purpose is reading and processing data for training and evaluating JAX models, and it describes itself as flexible, fast and deterministic. The audience is narrow and identifiable: people already working in JAX. The README lists MaxText, Gemma, kauldron and maxdiffusion as existing users, all of which are JAX-adjacent projects, so the design assumptions come from that world. The library does not require JAX to run, according to the README, which means the dataset code can be exercised on its own before a model exists. That separation is the practical reason to care: pipeline bugs and model bugs stop being the same debugging session.
How MapDataset pipelines are actually built
The mechanism is a chain of transformations on a dataset object. The README's example starts with grain.MapDataset.source([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10]), then calls .shuffle(seed=42), .map(lambda x: x+1) and .batch(batch_size=2). Each call returns something that the next call consumes, and the final object is iterable, so a for loop over it yields batches. Two details in that chain carry weight. First, the shuffle takes an explicit seed, and the README claims determinism as a property of the library, so the same seed should produce the same order. Second, the comment on shuffle says it shuffles elements globally, which is a different guarantee from the bounded shuffle buffer common in streaming loaders: a global shuffle implies the whole dataset is available to the shuffle step, which is fine for an in-memory list and a real constraint for a source that does not fit. The batch step groups consecutive elements, so ordering after the shuffle is what determines batch composition. Nothing in the README's example shows sharding across devices or prefetching, and the README does not document a checkpoint or resume mechanism for a partially consumed pipeline.
Installing Grain and running a first pipeline
The README gives one installation route: the package is on PyPI and installs with pip install grain. Python 3.11 or newer is required according to pyproject.toml, which sets requires-python to >=3.11 and lists classifiers for 3.11 through 3.14. Run the install in an environment that already has a compatible interpreter, otherwise pip will refuse the package rather than resolve an older one.
pip install grainAfter that, the README's own example is the shortest thing that proves the install worked. It builds a ten-element source, shuffles with a fixed seed, increments each element, and batches in pairs, then iterates. If the loop prints or processes batches of two, the pipeline is running.
import grain
dataset = (
grain.MapDataset.source([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10])
.shuffle(seed=42)
.map(lambda x: x+1)
.batch(batch_size=2)
)
for batch in dataset:
passThe README points to a Basic Dataset tutorial in the reference docs for the next step, and to a changelog page for release history. Parquet input is not in the default dependency set: pyproject.toml defines an optional parquet extra that pulls in pyarrow, so a Parquet source needs that extra installed. The README does not show the extra's install command, so confirm the spelling against pyproject.toml before assuming it.
CPU-only transformations and the platform table
The README is direct about where the work happens: Grain does not directly use GPU or TPU in its transformations, and processing within Grain is done on the CPU by default. For a training setup where the accelerator is saturated, that is usually the right division of labour, since the loader feeds the device rather than competing with it. It becomes a limitation when a transformation is itself expensive and would benefit from accelerator hardware, because Grain's declarative steps will not place it there. The platform matrix is equally specific and has gaps worth reading carefully. Linux x86_64 and Windows x86_64 are listed as supported; macOS x86_64 is not; Linux aarch64 and macOS aarch64 are supported; Windows aarch64 is marked n/a. An Intel Mac is therefore outside the documented support set, which is an unusual omission for a Python data library and worth checking before a team standardises on it. The dependency list in pyproject.toml reinforces the split: array-record is required only when sys_platform is not win32, so the Windows path is built on a different input stack than the Linux and macOS paths. Behaviour that depends on ArrayRecord files should be expected to differ on Windows.
Building from source and the GRAIN_SKIP_EXTRA_BUILD switch
Installing a published wheel is not the same as installing from the repository. The build-system table in pyproject.toml requires setuptools, grpcio-tools and pybind11, and setup.py explains why: it generates protobuf bindings and compiles C++ extensions through pybind11. The proto step runs grpc_tools.protoc against grain/proto/execution_summary.proto, so a source build needs a working protobuf toolchain, not just a compiler. setup.py documents an escape hatch for the case where binaries are already built, for example when packaging wheels: setting GRAIN_SKIP_EXTRA_BUILD=1 skips both the C++ compilation and the proto generation. That variable is the difference between a fast packaging step and a full toolchain build, and it is only safe when the prebuilt artifacts are already in place. Anyone building from source on a machine without grpcio-tools or pybind11 will hit the build step before they hit any Grain code.
Where Grain sits next to a tf.data pipeline
The closest comparison is tf.data, which solves the same problem for TensorFlow models and, in practice, gets used for JAX training too. The difference is the shape of the API. A tf.data pipeline is built from tf.data.Dataset methods and typically ends in a prefetch and a device placement call, and its execution model is a graph the runtime schedules. Grain's pipeline in the README is plain Python objects chained with method calls, with an explicit seed on the shuffle and no prefetch step shown. That makes Grain easier to read and to unit test in isolation, and it makes the ordering semantics explicit rather than runtime-dependent. tf.data has the larger input-format surface and a longer track record outside Google. If your team already runs tf.data loaders and they are not the bottleneck, switching buys readability and not throughput, and the README makes no throughput claim that would justify the migration on its own.
Licence, release cadence and upgrade cost
Grain is Apache-2.0, and pyproject.toml declares the licence as a file reference to LICENSE rather than an SPDX string. Apache-2.0 permits commercial and internal use and includes a patent grant, and it requires that the licence text and notices be preserved when the code is redistributed. That is the general shape of the licence; whether a specific redistribution complies is a question for your own counsel, not something the repository answers. On cadence, the releases listed run v0.2.16 on 2026-02-25, v0.2.17 on 2026-06-01, and v0.2.18 on 2026-06-17, with pyproject.toml on the main branch already at version 0.2.19. The last push to the repository was on 2026-09-10, so the project is current rather than dormant. The practical upgrade cost is the version pin: pyproject.toml requires protobuf>=5.28.3 and array-record>=0.8.1, and those floors move with the project, so a shared environment that pins protobuf for another library can block a Grain upgrade. Pin Grain explicitly and read CHANGELOG.md before bumping, since the README does not document a compatibility policy across minor versions.
Editorial conclusion
Adopt Grain if you are training JAX models and want data loading expressed as a chain of shuffle, map and batch calls rather than a hand-written loader loop, and if your pipeline can live on the CPU. Do not adopt it if you need GPU or TPU processing inside the transformation steps, or if you are on Windows ARM, which the supported-platform table does not list. Before committing, verify the Python version your environment offers against the project's requires-python of >=3.11, and check whether the array-record dependency matters on your platform, since pyproject.toml excludes it on win32.
Frequently asked questions
How do I install google/grain?
The README states that Grain is available on PyPI and installs with pip install grain. Python 3.11 or newer is required according to pyproject.toml, which sets requires-python to >=3.11.
Does google/grain require JAX to run?
No. The README says Grain is designed to work with JAX models but does not require JAX to run, and that it can be used with other frameworks as well.
Does google/grain use the GPU or TPU for its transformations?
The README states that Grain does not directly use GPU or TPU in its transformations, and that processing within Grain is done on the CPU by default.
Which platforms does google/grain support?
The README's platform table lists Linux x86_64, Linux aarch64, macOS aarch64 and Windows x86_64 as supported, marks macOS x86_64 as not supported, and marks Windows aarch64 as n/a.
Official sources
Add this badge to your README
If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.
[](https://hysenlabs.com/projects/google-grain)