Tutorial: train a model on Ray, end to end
A worked example rather than a reference. It builds one real model on a public dataset, from raw files to a trained artifact with measured metrics.
For what each field means, see Train Models on Ray. This page is the path through it.
What you will build. A model that predicts which products a shopper will buy again in their next basket, from the Instacart public dataset — 3.4M orders from 200,000 shoppers. Base rate 9.77%, so a useless model scores 0.5 and there is real room above it.
1. Get the data into a format a batch engine likes
The raw dataset is four CSVs. Convert them before you point a pipeline at them. CSV has no schema, no column statistics and no predicate pushdown, so Spark reparses every byte of every file on every run — including columns your query never selects.
The conversion is a one-off, and it does not need to be a staging query. Athena will do it:
CREATE TABLE instacart.order_products_prior
WITH (table_type = 'ICEBERG', format = 'PARQUET',
location = 's3://<your-bucket>/kaggle_instacart/iceberg/order_products_prior/')
AS SELECT order_id, product_id, add_to_cart_order, reordered
FROM instacart_raw.order_products_prior_csv;The full script, including the external tables over the raw CSVs, is in docs/instacart-athena-conversion.sql.
It is worth the step. On this dataset:
| file | CSV | Parquet |
|---|---|---|
orders | 109 MB | 19 MB |
order_products_prior | 578 MB | 77 MB |
Check the row counts afterwards. A CSV SerDe will silently mis-split a quoted field
containing a comma — one Instacart product name has both an escaped quote and commas, and
it lands a string in an integer column. Count rows against wc -l on the raw file before
trusting the conversion.
2. Staging queries: shape the raw data
python/test/canary/staging_queries/aws_databricks/instacart.py
Four of them, each reading the Parquet by path:
_PARQUET = "s3://<your-bucket>/kaggle_instacart/iceberg"
purchases = StagingQuery(
query=f"""
SELECT order_id, product_id, add_to_cart_order, reordered, ...
FROM parquet.`{_PARQUET}/order_products_prior/data/`
""",
output_namespace="workspace_iceberg.poc",
engine_type=EngineType.SPARK,
version=2,
)Read by path rather than by catalog name when the converted tables live in one catalog and your team's Spark session points at another.
The one that matters conceptually is candidates, which produces the spine and the
label: one row per (shopper, product they have bought before), with in_next_basket as 1 or 0.
Keep the thing you are predicting out of the thing you predict from. purchases reads only eval_set = 'prior'; candidates reads 'train'. That separation is what
makes point-in-time correctness real rather than declared — no aggregation can see the
basket it is predicting, because that basket is not in the table it reads.
3. GroupBys: the features
python/test/canary/group_bys/aws_databricks/instacart.py
Three GroupBys over the same event stream, keyed three different ways, because the signal lives at three levels:
user_product = GroupBy( # the pair
sources=[_source()],
keys=["user_id", "product_id"],
aggregations=[
Aggregation(input_column="order_id", operation=Operation.COUNT,
windows=[Window(7, TimeUnit.DAYS), Window(45, TimeUnit.DAYS)]),
...
],
)- pair
(user_id, product_id)— how often this shopper buys this product - shopper
(user_id)— how often they buy anything, to normalise the pair counts - product
(product_id)— how reliably it returns for anyone
Every aggregation needs a window. Unwindowed means "since the beginning," which needs a start partition and will scan far more than you expect.
4. The Join: assemble the training set
python/test/canary/joins/aws_databricks/instacart.py
v1 = Join(
left=_left, # the spine, with timestamps and labels
row_ids=["candidate_id"],
right_parts=[JoinPart(group_by=history.user_product), ...],
derivations=[
Derivation(name="*", expression="*"),
Derivation(name="share_of_orders_with_product",
expression="user_id_product_id_order_id_count_45d / "
"NULLIF(user_id_order_id_approx_unique_count_45d, 0)"),
],
)Two things that will bite you.
Adding any derivation replaces the output rather than extending it. Without Derivation(name="*", expression="*") the aggregation columns are pruned and never
written, and training then fails looking for a column the table does not have.
And derivations are where cross-key ratios live. share_of_orders_with_product divides a
pair-keyed count by a shopper-keyed one — no single GroupBy can span both. It is also the
single most valuable feature here: the raw count scores AUC 0.710, the same number
divided by the shopper's order count scores 0.771.
5. The trainer
python/test/canary/model_examples/listing_ctr_model/trainer/
Takes its table, features, label column and split from arguments, so a new modelling problem usually costs configs and no Python. It writes three things:
model.joblibmetadata.jsonexperiment_report.json— this is what Hub reads metrics from
Hold out by date, never at random:
train, train_y, evaluate, eval_y, split = split_by_date(features, labels, eval_days)Rows are time-ordered. A random split trains on rows that come after the ones it is scored against, and reports optimistic metrics.
6. Package it
The archive is content-addressed — its hash is the model version:
python model_examples/listing_ctr_model/build_archive.pyTwo consequences worth internalising:
- Adding a runtime module changes the hash, so it cuts a new model version.
- Two model variants that run identical trainer bytes share a version. Give each variant its own model name, or both write their reports to the same key and a comparison reads whichever finished last, twice.
7. Declare the model
python/test/canary/models/aws_databricks/instacart.py
Model(
version="afce3c0e00af", # the archive's content hash
inference_spec=InferenceSpec(model_backend=ModelBackend.RAY, ...),
model_artifact_base_uri=_artifact_base,
training_conf=TrainingSpec(
training_data_source=_training_source(features),
python_module="trainer.train",
resource_config=ResourceConfig(min_replica_count=0, max_replica_count=0),
job_configs={
"features": "user_id_product_id_order_id_count_45d:pair_count_45d,...",
"label-column": "in_next_basket",
"eval-days": "1",
},
),
)job_configs is passed straight to the trainer as CLI arguments.
8. Run it
zipline compile --chronon-root .
# features first — the join needs its GroupBys, which need their staging queries
zipline hub backfill compiled/joins/aws_databricks/instacart.v1__2
--chronon-root . --start-ds 2026-08-18 --end-ds 2026-08-24
# then training
zipline hub backfill compiled/models/aws_databricks/instacart.instacart_basket_model
--chronon-root . --start-ds 2026-08-24 --end-ds 2026-08-24Watch it in the UI: the workflow page shows each node, and a failed step links to the engine's own console — Spark on EMR for the feature work, Ray on EKS for training.
9. Check the result
Open experiment_report.json under the model's artifact URI. evaluation_status tells
you whether it was scored at all, which is not the same as scoring badly:
{
"metrics": {"validation_auc": 0.771, "validation_logloss": 0.283, ...},
"feature_importance": [...],
"evaluation_status": "evaluated"
}A perfect score means a bug, not a discovery. An early version of this example used
the dataset's reordered column as the label — but that column is true exactly when the
shopper has bought the item before, which is the definition of being on the spine. It
scored AUC 1.0000, which is how the mistake was found.
Things that cost time here, so they do not cost you any
- Jobs that die just past an hour. A Spark job against a Unity Catalog-backed Iceberg table authenticates with a token that lasts one hour; the client's renewal is broken, so the job dies on wall-clock rather than on workload. Converting to Parquet is the first thing to try, because it usually brings the run comfortably under the limit.
- A long single staging query. If one is slow, split it into chained staging queries rather than making it faster.
- Every aggregation needs a window. Compile will reject unwindowed ones.
Derivation(name="*", expression="*"). See §4. It is a silent, confusing failure.