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
11 changes: 11 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,17 @@ remove; each such change is listed here with the migration in one line.

## Unreleased

**`Latu.to_nx/2`, `to_nx!/2` and `stream_nx/2`** turn a result into `Nx` tensors. A numeric
column with no nulls becomes a 1-D tensor whose binary *is* the Arrow buffer — no copy for a
single batch — and a column of equal-length numeric lists, or of dense MLlib `Vector`s, becomes
one `{rows, width}` tensor. Everything else is refused by name.

This is the only way to read a `Vector` column into Elixir: Spark describes one as a UDT with
no SQL type, so `collect/2` and `to_explorer/2` both refuse it, while the Arrow stream carries
its own schema and says exactly what it is. `Latu.Result.Arrow` is the reader — the IPC
streaming format, no dependency — and `Latu.Result.Nx` the mapping, behind the now-optional
`:nx`. Adding `{:nx, "~> 0.13"}` is what turns them on; without it `to_nx/2` says so.

**`Latu.disconnect/2` closes the socket within a second.** Gun waited its default 15 s for a
close the Spark server never sends, so a client that connects per unit of work leaked sockets
for 15 s each — enough to hit an open-files limit at a few connections a second. No migration.
Expand Down
19 changes: 19 additions & 0 deletions dev/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -193,3 +193,22 @@ offline that it still covers the exported surface, so a Spark bump that adds or
function fails in a second. **Re-run the probe at any Spark bump or registry regeneration, and
diff**: an unchanged file means no wrapper's shape moved. `dev/probe_dtypes.exs` is the same
instrument for the Arrow types the schema guard refuses.

## `make_arrow_fixtures.py`

python dev/make_arrow_fixtures.py # rewrite test/arrow/*.arrow
python dev/make_arrow_fixtures.py --check # exit 1 if any file is out of date

The Arrow IPC streams `Latu.Result.Arrow` and `Latu.Result.Nx` are tested against. Each file is
one complete stream — schema, one record batch, end marker — which is what Spark Connect sends
per batch and what `Latu.to_arrow/2` hands back.

Written by **pyarrow**, so the reader is checked against Arrow's own encoder rather than against
itself; the expected values live in the tests, where they can be read. It reaches no server.
`dev/.venv` already has pyarrow, as a PySpark dependency.

Two fixtures are shaped by something other than the type they carry. `vector_dense` and
`vector_sparse` are Spark's `VectorUDT` sqlType — `struct<type, size, indices, values>` — which
is the whole reason `to_nx/2` can read a column `collect/2` refuses. And `big_two` is over
64 bytes on purpose: below that the BEAM copies a binary into the process heap however it was
made, so a small batch cannot show whether a column's buffer still points into the whole one.
131 changes: 131 additions & 0 deletions dev/make_arrow_fixtures.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
"""Arrow IPC stream fixtures for `Latu.Result.Arrow` and `Latu.Result.Nx`.

python dev/make_arrow_fixtures.py # rewrite test/arrow/*.arrow
python dev/make_arrow_fixtures.py --check # exit 1 if any file is out of date

Each file is one complete IPC *stream* — schema, one record batch, end marker — which is the
shape Spark Connect sends per batch, and the shape `Latu.to_arrow/2` hands back. Written by
pyarrow so the reader is checked against Arrow's own encoder rather than against itself; the
expected values live in the tests, where they can be read.

Needs pyarrow, which `dev/.venv` already has as a PySpark dependency. It reaches no server:
these are bytes, not results.
"""

from __future__ import annotations

import argparse
import pathlib
import sys

import pyarrow as pa

OUT = pathlib.Path("test/arrow")


def stream(table: pa.Table) -> bytes:
sink = pa.BufferOutputStream()
with pa.ipc.new_stream(sink, table.schema) as writer:
writer.write_table(table)
return sink.getvalue().to_pybytes()


def vector_udt(vectors, dense=True):
"""Spark's VectorUDT sqlType, which is what a Vector column is in Arrow.

struct<type:int8, size:int32, indices:list<int32>, values:list<double>>, with size and
indices null on a dense row.
"""
fields = [
pa.field("type", pa.int8()),
pa.field("size", pa.int32()),
pa.field("indices", pa.list_(pa.int32())),
pa.field("values", pa.list_(pa.float64())),
]
rows = [
{
"type": 1 if dense else 0,
"size": None if dense else len(v),
"indices": None if dense else list(range(len(v))),
"values": list(v),
}
for v in vectors
]
return pa.array(rows, type=pa.struct(fields))


def cases() -> dict[str, pa.Table]:
return {
# The arms that decode.
"doubles": pa.table({"v": pa.array([1.5, 2.5, 0.0, -3.25], pa.float64())}),
"int64s": pa.table({"v": pa.array([1, -2, 3, 4], pa.int64())}),
"int32s": pa.table({"v": pa.array([1, -2, 3], pa.int32())}),
"float32s": pa.table({"v": pa.array([1.5, 2.5], pa.float32())}),
"two_columns": pa.table(
{
"a": pa.array([1.0, 2.0, 3.0], pa.float64()),
"b": pa.array([10, 20, 30], pa.int64()),
}
),
# What `vector_to_array` gives, and what a fitted `features` column is.
"list_uniform": pa.table(
{"v": pa.array([[1.5, 2.5], [0.5, 3.5], [0.0, 0.0]], pa.list_(pa.float64()))}
),
"vector_dense": pa.table(
{"features": vector_udt([[1.5, 2.5], [0.5, 3.5], [0.0, 0.0], [4.0, 4.5]])}
),
# What must be refused.
"list_ragged": pa.table(
{"v": pa.array([[1.0, 2.0], [3.0]], pa.list_(pa.float64()))}
),
"vector_sparse": pa.table(
{"features": vector_udt([[1.5, 2.5], [0.5, 3.5]], dense=False)}
),
"doubles_with_null": pa.table({"v": pa.array([1.0, None, 3.0], pa.float64())}),
"strings": pa.table({"v": pa.array(["a", "b"], pa.string())}),
"booleans": pa.table({"v": pa.array([True, False, True], pa.bool_())}),
# A schema and no batch at all, which pyarrow writes for an empty table and Spark does
# not — the reader has to answer for both.
"empty": pa.table({"v": pa.array([], pa.float64())}),
# Big enough that its buffers are refc binaries. Under 64 bytes the BEAM copies into
# the process heap whatever the reader does, so a small batch cannot show whether a
# column's buffer still points into the whole one.
"big_two": pa.table(
{
"a": pa.array([float(i) for i in range(5000)], pa.float64()),
"b": pa.array(list(range(5000)), pa.int64()),
}
),
}


def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--check", action="store_true", help="exit 1 if a file is stale")
args = parser.parse_args()

OUT.mkdir(parents=True, exist_ok=True)
stale = []

for name, table in cases().items():
path = OUT / f"{name}.arrow"
wanted = stream(table)

if args.check:
if not path.exists() or path.read_bytes() != wanted:
stale.append(path)
else:
path.write_bytes(wanted)
print(f"{name:20s} {len(wanted):7d} bytes {table.num_rows} rows")

if args.check:
for path in stale:
print(f"stale: {path}", file=sys.stderr)
print("up to date" if not stale else f"{len(stale)} stale", file=sys.stderr)
return 1 if stale else 0

return 0


if __name__ == "__main__":
raise SystemExit(main())
2 changes: 2 additions & 0 deletions docs/cheatsheet.cheatmd
Original file line number Diff line number Diff line change
Expand Up @@ -151,6 +151,7 @@ Called qualified, the way `Enum` is.
| `show/2` | Print the table Spark renders, and return `:ok`. **+ !** |
| `storage_level/1` | How the server is storing this frame, if at all. **+ !** |
| `stream/2` | The result as a lazy stream of `Explorer.DataFrame`s, one per Arrow batch. |
| `stream_nx/2` | A lazy stream of `to_nx/2`'s tensors, one map per Arrow batch. |
| `summary/2` | Summary statistics: one row per statistic, one column per column Spark can… |
| `tail/3` | The last `count` rows, as maps. **+ !** |
| `take/3` | The first `count` rows, as maps — `limit/2` then `collect/2`, as in PySpark. **+… |
Expand All @@ -159,6 +160,7 @@ Called qualified, the way `Enum` is.
| `to_explorer/2` | The result as one `Explorer.DataFrame`. **+ !** |
| `to_explorer_with_metrics/2` | `to_explorer/2`, and the metrics `observe/3` asked for. **+ !** |
| `to_html/2` | The table `show/2` prints, as an HTML string. **+ !** |
| `to_nx/2` | The result as `Nx` tensors, one per column. **+ !** |
| `tree_string/2` | The schema tree as a string, where `print_schema/2` prints it. **+ !** |
| `unpersist/2` | Drop the server's cache of this frame, and hand it back. **+ !** |
| `write/2` | Write to a path. **+ !** |
Expand Down
24 changes: 24 additions & 0 deletions docs/decisions.md
Original file line number Diff line number Diff line change
Expand Up @@ -1264,3 +1264,27 @@ server answers `UNRESOLVED_ROUTINE`. Measured on 4.2.0 by `latu_ml`'s `dev/probe
unset resolves, `true` resolves, `false` does not. So a tidying pass that populated every proto
field explicitly would silently take the ML functions away from `latu_ml`, and nothing in this
repo would go red. Setting it to `true` is never needed, so there is no `internal:` option.

## 2026-09-07 — `to_explorer/2` keeps refusing a Vector column

Spark 4.2.0 describes `features`, `rawPrediction` and `probability` as a UDT with `sql_type`
unset. `Latu.Result.Schema` refuses what the server declines to describe, because the guard
exists so an unsupported dtype names its column instead of panicking inside Polars' NIF.

**Polars can in fact read those bytes** — measured 2026-09-07. `to_arrow/2` on an
`array_to_vector` column, through `Explorer.DataFrame.load_ipc_stream/1`, gives
`v struct[4] [%{"type" => 1, "size" => nil, "indices" => nil, "values" => [...]}]`. So the
refusal is a choice rather than a limit, and this entry exists so the next person to find
`sql_type: nil` does not have to rediscover that by experiment.

It stands, because what comes back is Spark's **internal** `VectorUDT` layout and not a vector:
a caller would unpack the struct, take `values`, and rebuild. Both routes that give a useful
shape already exist — `vector_to_array` server-side, for an `array<double>` that Explorer reads
as `list[f64]`, and `Latu.to_nx/2` for an `{n, d}` tensor off the same bytes with no unpacking
at all.

Reversing it means checking the **Arrow** schema where the proto schema says nothing.
`Latu.Result.Arrow.schema/1` makes that possible and the batches are already in hand, so it
would cost no extra round trip; the price is a second decodability table, in Arrow's vocabulary
rather than Spark's, and a guard with two sources of truth. Not worth it for a struct nobody
wants.
55 changes: 55 additions & 0 deletions lib/latu.ex
Original file line number Diff line number Diff line change
Expand Up @@ -2960,4 +2960,59 @@ defmodule Latu do
@doc "Like `to_arrow/2`, raising on failure."
@spec to_arrow!(DataFrame.t(), keyword()) :: [binary()]
defdelegate to_arrow!(df, opts \\ []), to: DataFrame

@doc """
The result as `Nx` tensors, one per column.

Bypasses the Explorer decoder and the schema guard, exactly as `to_arrow/2` does — which is
what makes a `Vector` column readable here when `collect/2` and `to_explorer/2` refuse it.

A numeric column with no nulls becomes a 1-D tensor; a column of equal-length numeric lists,
or of dense `Vector`s, becomes one `{rows, width}` tensor. Nulls, strings, booleans, ragged
lists and sparse vectors are refused by name.

Unbounded, like `collect/2`. Bound the plan, or use `stream_nx/2` for a result too large to
hold.

## Options

* `:columns` — keep only these columns, by name. Pruning copies, because an Arrow buffer is
a slice of the whole batch and would otherwise hold the rest alive. Defaults to `nil`,
every column.
* `:progress` — as `collect/2` describes. Defaults to `nil`.

## Examples

{:ok, tensors} = Latu.to_nx(df)
{:ok, %{"features" => t}} = Latu.to_nx(scored, columns: ["features"])

Needs the optional `:nx` dependency. See `Latu.DataFrame.to_nx/2`.
"""
@spec to_nx(DataFrame.t(), keyword()) ::
{:ok, %{String.t() => term()}} | {:error, Error.t()}
defdelegate to_nx(df, opts \\ []), to: DataFrame

@doc "Like `to_nx/2`, raising on failure."
@spec to_nx!(DataFrame.t(), keyword()) :: %{String.t() => term()}
defdelegate to_nx!(df, opts \\ []), to: DataFrame

@doc """
A lazy stream of `to_nx/2`'s tensors, one map per Arrow batch.

What `stream/2` is for Explorer. The tensors are per batch, so stacking them is the caller's
business — `to_nx/2` is the one that concatenates. Raises `Latu.Error` on failure.

## Options

* `:columns` — as `to_nx/2` describes. Defaults to `nil`.
* `:progress` — as `collect/2` describes. Defaults to `nil`.

## Examples

df |> Latu.stream_nx(columns: ["features"]) |> Enum.map(&Nx.sum(&1["features"]))

See `Latu.DataFrame.stream_nx/2`.
"""
@spec stream_nx(DataFrame.t(), keyword()) :: Enumerable.t()
defdelegate stream_nx(df, opts \\ []), to: DataFrame
end
81 changes: 81 additions & 0 deletions lib/latu/data_frame.ex
Original file line number Diff line number Diff line change
Expand Up @@ -1849,6 +1849,71 @@ defmodule Latu.DataFrame do
@spec to_arrow!(t(), keyword()) :: [binary()]
def to_arrow!(%__MODULE__{} = df, opts \\ []), do: unwrap!(to_arrow(df, opts))

@doc """
The result as `Nx` tensors, one per column.

Bypasses the Explorer decoder and the schema guard, as `to_arrow/2` does, and for the same
reason: these bytes are read by `Latu.Result.Arrow` rather than by Polars, and what Polars
cannot take is not this path's concern. That is what makes a `Vector` column readable here
when `collect/2` and `to_explorer/2` both refuse it.

Two shapes decode. A numeric column with no nulls becomes a 1-D tensor, and a column of
equal-length numeric lists — or of dense `Vector`s — becomes one `{rows, width}` tensor.
Anything else is refused by name: nulls, strings, booleans, ragged lists, sparse vectors.

**Unbounded**, like `collect/2` and `to_arrow/2`: bound the plan, or use `stream_nx/2`.

{:ok, %{"features" => t}} = Latu.to_nx(scored, columns: ["features"])

Needs the optional `:nx` dependency; without it this says so rather than failing obscurely.
"""
@spec to_nx(t(), keyword()) :: {:ok, %{String.t() => term()}} | {:error, Error.t()}
def to_nx(%__MODULE__{} = df, opts \\ []) do
opts = Keyword.validate!(opts, progress: nil, columns: nil)

with {:ok, batches, _execution} <-
Client.execute(df.session, Plan.new(df.plan), watch(opts)) do
tensors(Enum.map(batches, & &1.data), opts)
end
end

@doc "Like `to_nx/2`, raising on failure."
@spec to_nx!(t(), keyword()) :: %{String.t() => term()}
def to_nx!(%__MODULE__{} = df, opts \\ []), do: unwrap!(to_nx(df, opts))

@doc """
A lazy stream of `to_nx/2`'s tensors, one map per Arrow batch.

Backpressure for results too large to hold, as `stream/2` is for Explorer. Each batch decodes
on its own, so the tensors are per batch and stacking them is the caller's business — that is
the difference from `to_nx/2`, which concatenates. Raises `Latu.Error` on failure, since an
enumeration has no way to return one.

df |> Latu.stream_nx(columns: ["features"]) |> Enum.map(&Nx.sum(&1["features"]))
"""
@spec stream_nx(t(), keyword()) :: Enumerable.t()
def stream_nx(%__MODULE__{} = df, opts \\ []) do
opts = Keyword.validate!(opts, progress: nil, columns: nil)

df.session
|> Client.responses(Plan.new(df.plan))
|> Client.watched(watch(opts))
|> Stream.flat_map(fn
{:ok, batch} ->
case tensors([batch.data], opts) do
{:ok, decoded} -> [decoded]
{:error, error} -> raise error
end

{:error, error} ->
raise error

# The schema guard is deliberately not run here; see `to_nx/2`.
_progress_or_schema_or_done ->
[]
end)
end

# =============================================
# Observed metrics
# =============================================
Expand Down Expand Up @@ -2063,6 +2128,22 @@ defmodule Latu.DataFrame do
end
end

# `:nx` is optional, so `Latu.Result.Nx` may not exist at all, and refusing by name beats a
# `NoSuchModule` from three frames down. The call goes through `apply/3` on purpose: a direct
# one is a compile-time warning in a project that did not take the dependency, and this
# module is compiled by every one of them.
defp tensors(binaries, opts) do
if Code.ensure_loaded?(Result.Nx) do
case apply(Result.Nx, :decode, [binaries, [columns: opts[:columns]]]) do
{:ok, decoded} -> {:ok, decoded}
{:error, message} -> {:error, Error.new(:decode, message)}
end
else
{:error,
Error.new(:decode, "to_nx/2 needs the optional :nx dependency; add it to your deps")}
end
end

defp keys!(keys) when keys in [:atoms, :strings], do: keys

defp keys!(other) do
Expand Down
Loading