Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,7 @@ this guide then work unchanged. CI does not use Nix; this is a dev shell only.
crates/
ddx-core/ # the v1 engine — sqlparser only. Start here.
ddx-datafusion/ # DataFusion adapter: AnalyzerRule (bare grad) + ddx_sql
ddx-ad/ # v2 query-level reverse-mode AD over Substrait (M3/M4)
ddx-ad/ # v2: query-level reverse-mode AD over Substrait
python/ddxdb/ # PyO3/maturin wheel: rewrite_sql + a DataFusion Context
tests/ # cross-engine numeric-agreement suites (vs JAX)
docs/design.md # the design (source of truth)
Expand All @@ -86,7 +86,8 @@ docs/spikes/ # runnable evidence behind the design
```

`ddx-core` is where most work happens; everything else is a thin layer over it.
`ddx-ad` is under construction for M3/M4.
`ddx-ad` is v2. Its integration tests live in `ddx-datafusion/tests/`, where
there is an engine to run the plans it emits.

## Working on the Python code

Expand Down
51 changes: 42 additions & 9 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -62,10 +62,43 @@ On DataFusion, `ddx-datafusion` installs an `AnalyzerRule` so bare `grad()` work
in ordinary SQL *and* through the DataFrame API, with columns resolved by the
planner rather than syntactically.

## Training: gradients of whole queries

A model's gradient is a gradient of a whole query, and in ddx that is `grad`
too, in a `FROM` clause. `grad(loss, table.column)` is the gradient of the loss
a CTE computes, as a relation shaped like the table, so an SGD step,
`params - lr * grad(loss)(params)`, is a join:

```sql
WITH h AS (
SELECT x.sample, w.out, tanh(SUM(x.val * w.val)) AS val
FROM x JOIN w ON x.inp = w.inp GROUP BY x.sample, w.out),
loss AS (
SELECT SUM(power(h.val - y.val, 2)) AS l
FROM h JOIN y ON h.sample = y.sample AND h.out = y.out)
SELECT w.inp, w.out, w.val - 0.1 * g.val AS val
FROM w JOIN grad(loss, w.val) g ON w.inp = g.inp AND w.out = g.out
```

That runs as written through `ddxdb.Context` in Python and
`ddx_datafusion::ad::sql` in Rust. The loss is ordinary SQL: ddx differentiates
it the way `jax.grad` differentiates a function, reverse mode, one transpose
rule per relational operator (joins, sums, averages, max and min, window
rankings), and nothing in it is labelled. The backward pass is a sequence of
plain Substrait plans the engine runs.
[`examples/nn`](crates/ddx-datafusion/examples/nn) trains a small classifier
this way: a two-layer MLP over 6×6 images, adapted from [a neural network written entirely in SQL](https://github.com/xqlsystems/xarray-sql/pull/196)
in xarray-sql, which computed its gradients with hand-written backward queries.
Here its loss query is used unchanged, and ddx's gradients equal those
hand-written ones to 1e-12. The spikes' MLP, attention and max-pool gradients
equal `jax.grad` to 1e-12.

## Status

**M2 landed and released.** The scalar engine, the DataFusion adapter and the
Python wheel are all published.
**M2 is released; M3 and M4 are built.** The scalar engine, the DataFusion
adapter and the Python wheel are published. Query-level AD (`ddx-ad`, with
`ddx_datafusion::ad` and `ddxdb.ad` on top) trains an MLP on DataFusion and
matches `jax.grad`, and ships with the next release.

**DataFusion is the engine with native support**: `ddx-datafusion` installs an
`AnalyzerRule`, so bare `grad()` works in ordinary SQL and through the DataFrame
Expand All @@ -76,13 +109,13 @@ extension, with `grad()` understood in-database, comes eventually (M5).
| | | |
|---|---|---|
| [`ddx-core`](crates/ddx-core) | the v1 engine | [crates.io](https://crates.io/crates/ddx-core) |
| [`ddxdb`](python/ddxdb) | Python wheel — `rewrite_sql` + a DataFusion `Context` | [PyPI](https://pypi.org/project/ddxdb/) |
| [`ddx-datafusion`](crates/ddx-datafusion) | DataFusion adapter: `AnalyzerRule` + `ddx_sql` | [crates.io](https://crates.io/crates/ddx-datafusion) |
| [`ddx-ad`](crates/ddx-ad) | v2 — query-level reverse-mode AD over Substrait | M3/M4 |
| [`ddxdb`](python/ddxdb) | Python wheel — `rewrite_sql`, a DataFusion `Context`, and `ad` | [PyPI](https://pypi.org/project/ddxdb/) |
| [`ddx-datafusion`](crates/ddx-datafusion) | DataFusion adapter: `AnalyzerRule` + `ddx_sql` + `ad` | [crates.io](https://crates.io/crates/ddx-datafusion) |
| [`ddx-ad`](crates/ddx-ad) | v2 — query-level reverse-mode AD over Substrait | next release |
| `ddx-duckdb` | DuckDB community extension | M5 |

Next is **M3/M4** — reverse-mode AD over whole *queries*, where a gradient step
becomes a query rather than a column. See [docs/design.md](docs/design.md) §8.
Next is **M5**, DuckDB: the `ddx('<sql>')` extension, and v2's backward program
run through DuckDB's Substrait consumer. See [docs/design.md](docs/design.md) §8.

## Correctness

Expand All @@ -108,8 +141,8 @@ a plausible number. What backs that up:
crates/
ddx-core/ # v1 engine — differentiate sqlparser::ast::Expr + rewrite_sql
ddx-ad/ # v2 engine — query-level reverse-mode AD over Substrait
ddx-datafusion/ # DataFusion adapter: AnalyzerRule (bare grad) + ddx_sql
python/ddxdb/ # PyO3/maturin wheel: rewrite_sql + a DataFusion Context
ddx-datafusion/ # DataFusion adapter: AnalyzerRule (bare grad) + ddx_sql + ad (v2)
python/ddxdb/ # PyO3/maturin wheel: rewrite_sql + a DataFusion Context + ad (v2)
tests/ # cross-engine numeric-agreement suites (vs JAX)
docs/spikes/ # runnable evidence for every design claim
docs/design.md # the design
Expand Down
41 changes: 36 additions & 5 deletions crates/ddx-datafusion/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,35 @@ SELECT i, grad(x * y, x) AS dfdx, grad(x * y, y) AS dfdy FROM g
request through planning and are always rewritten away *before* execution, so
what DataFusion runs is an ordinary expression it already knows how to evaluate.

All the calculus lives in [`ddx-core`](https://crates.io/crates/ddx-core); this
crate only connects it to an engine.
All the calculus lives in [`ddx-core`](https://crates.io/crates/ddx-core) and
[`ddx-ad`](https://crates.io/crates/ddx-ad); this crate only connects them to an
engine.

## Gradients of whole queries: `ad`

`ad` differentiates a query rather than an expression (ddx v2). In SQL,
`grad(loss, table.column)` in a `FROM` clause is the gradient of the loss a CTE
computes, as a relation shaped like the table:

```rust
use ddx_datafusion::ad;

let step = ad::sql(&ctx, "
WITH loss AS (SELECT SUM(val * val) AS l FROM w)
SELECT w.i, w.val - 0.1 * g.val AS val
FROM w JOIN grad(loss, w.val) g ON w.i = g.i").await?;
```

Nothing in the loss is labelled; the one function ddx claims is
`ddx_stop_gradient` (`register_stop_gradient`). `ad::sql_all` runs several
statements that take `grad` of one loss for one backward pass, and
`ad::grad`/`ad::vjp` with `ad::run` give the program underneath. A program's
tables are named under a prefix of its own (`__ddx_{id}_`, a reserved
prefix), so programs never touch each other's tables or yours; `ad::sql`
leaves none behind, and `ad::release` drops what `ad::run` keeps.
[`examples/nn`](examples/nn) trains a small MLP written entirely in SQL
(adapted from [a neural network written entirely in SQL](https://github.com/xqlsystems/xarray-sql/pull/196) in xarray-sql) with one SQL statement per
parameter table: `cargo run -p ddx-datafusion --example nn`.

## Two routes to the same rewrite

Expand Down Expand Up @@ -81,9 +108,13 @@ bridge does not compile, so the dependency is pinned exactly and a test asserts
the resolved tree still contains exactly one `sqlparser` — a future bump fails at
the pin, with an explanation, rather than confusingly at the bridge.

| `ddx-datafusion` | `datafusion` | `sqlparser` |
|---|---|---|
| 0.1 | 54 | 0.62 |
v2 has the same shape one layer down: `ddx-ad` reads and writes the plans
`datafusion-substrait` produces and consumes, so the two must resolve the same
`substrait`, and `tests/substrait_pin.rs` asserts it.

| `ddx-datafusion` | `datafusion` | `sqlparser` | `substrait` |
|---|---|---|---|
| 0.1 | 54 | 0.62 | 0.63 |

## Status

Expand Down
10 changes: 6 additions & 4 deletions crates/ddx-datafusion/examples/nn/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,14 @@
//
// SPDX-License-Identifier: Apache-2.0

//! Train nn.py's MLP (xarray-sql#196) with `grad` in SQL.
//! Train a small neural network written entirely in SQL, with `grad` in SQL.
//!
//! nn.py trains a small network entirely in SQL, but writes its backward pass
//! by hand: a query per layer for the error, and one per weight and bias
//! The network is a two-layer MLP classifying 6×6 images, adapted from a
//! pure-SQL demo in xarray-sql (<https://github.com/xqlsystems/xarray-sql/pull/196>,
//! called nn.py in the comments here). That demo writes its backward pass by
//! hand: a query per layer for the error, and one per weight and bias
//! gradient. Here the backward pass is one expression per parameter table,
//! `grad(loss, weight.val)`, where `loss` is nn.py's own loss query, and a
//! `grad(loss, weight.val)`, where `loss` is the demo's own loss query, and a
//! training step is the SGD update written as a join:
//!
//! ```sql
Expand Down
7 changes: 5 additions & 2 deletions crates/ddx-datafusion/examples/nn/model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,11 @@
//
// SPDX-License-Identifier: Apache-2.0

//! nn.py (xarray-sql#196) with ddx: nn.py's forward pass and loss as SQL, the
//! data, and an SGD step that takes `grad` in SQL.
//! A small neural network written entirely in SQL, trained with ddx: its
//! forward pass and loss as SQL, the data, and an SGD step that takes `grad`
//! in SQL. It is adapted from a pure-SQL demo in xarray-sql
//! (<https://github.com/xqlsystems/xarray-sql/pull/196>), nn.py below, which
//! wrote its backward pass by hand.
//!
//! Shared by the `nn` example, which trains it, and by `tests/nn.rs`, which
//! checks its gradients against nn.py's hand-written backward pass.
Expand Down
Loading
Loading