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
21 changes: 8 additions & 13 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,8 @@ import xarray_sql as qr
ds = xr.tutorial.open_dataset('air_temperature')

# The same as a dask-sql Context; i.e. an Apache DataFusion Context.
c = qr.Context(ds)
c.create_table('air', ds, chunks=dict(time=24))
c = qr.XarrayContext()
c.from_dataset('air', ds, chunks=dict(time=24))

df = c.sql('''
SELECT
Expand All @@ -33,10 +33,10 @@ df = c.sql('''
''')

# A table of the average temperature for each location across time.
df.compute()
df.to_pandas()

# Alternatively, you can just create the DataFrame from the Dataset:
df = qr.read_xarray(ds)
df = qr.read_xarray(ds).to_pandas()
df.head()
```

Expand Down Expand Up @@ -68,6 +68,9 @@ All chunks in an Xarray Dataset are transformed into a Dask DataFrame via
`from_map()` and `to_dataframe()`. For SQL support, we just use `dask-sql`.
That's it!

_2025 update_: This library now implements a dask-like `from_map` interface in
pure `datafusion` and `pyarrow`, but works with the same principle!

## Why does this work?

Underneath Xarray, Dask, and Pandas, there are NumPy arrays. These are paged in
Expand All @@ -81,15 +84,7 @@ worth the convenience of DataFrames.

## What are the current limitations?

Dask doesn't support
`MultiIndex`s ([dask/dask#1493](https://github.com/dask/dask/issues/1493)). If
it did, I suspect performance for many types of queries would greatly improve.

Further, while this does play well with `dask-geopandas` (for geospatial query
support), certain types of operations don't quite match standard geopandas.
Spatial joins come to mind as a killer feature, but only inner joins are
supported ([geopandas/dask-geopandas#72](https://github.com/geopandas/dask-geopandas/issues/72))
.
_2025 update_: TBD, `datafusion` provides a whole new world!

## What would a deeper integration look like?

Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,8 @@ classifiers = [
"Topic :: Database :: Front-Ends",
]
dependencies = [
"dask-sql>=2024.5.0",
"dask>=2024.8.0",
"datafusion>=47.0.0",
"xarray>=2024.7.0",
]

Expand Down
538 changes: 16 additions & 522 deletions uv.lock

Large diffs are not rendered by default.

4 changes: 2 additions & 2 deletions xarray_sql/__init__.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,2 @@
from .df import read_xarray
from .sql import Context
from .df import read_xarray, from_map
from .sql import XarrayContext
6 changes: 3 additions & 3 deletions xarray_sql/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,12 @@

# deprecated
def get_columns(ds: xr.Dataset) -> t.List[str]:
return list(ds.dims.keys()) + list(ds.data_vars.keys())
return list(ds.sizes.keys()) + list(ds.data_vars.keys())


# Deprecated
def unravel(ds: xr.Dataset) -> t.Iterator[Row]:
dim_keys, dim_vals = zip(*ds.dims.items())
dim_keys, dim_vals = zip(*ds.sizes.items())

for idx in itertools.product(*(range(d) for d in dim_vals)):
coord_idx = dict(zip(dim_keys, idx))
Expand All @@ -27,7 +27,7 @@ def unravel(ds: xr.Dataset) -> t.Iterator[Row]:
# Deprecated
def unbounded_unravel(ds: xr.Dataset) -> np.ndarray:
"""Unravel with unbounded memory (as a NumPy Array)."""
dim_keys, dim_vals = zip(*ds.dims.items())
dim_keys, dim_vals = zip(*ds.sizes.items())
columns = get_columns(ds)

N = np.prod([d for d in dim_vals])
Expand Down
99 changes: 59 additions & 40 deletions xarray_sql/df.py
Original file line number Diff line number Diff line change
@@ -1,24 +1,14 @@
import itertools
import typing as t

import dask
import dask.dataframe as dd
import numpy as np
import pandas as pd
import pyarrow as pa
import xarray as xr
from dask.dataframe.io import from_map

from . import core

Block = t.Dict[str, slice]
Chunks = t.Optional[t.Dict[str, int]]

# Turn on Dask-Expr
dask.config.set({'dataframe.query-planning-warning': False})
dask.config.set({'dataframe.query-planning': True})
# Turn on Copy-On-Write (needs Pandas 2.0).
pd.options.mode.copy_on_write = True


# Borrowed from Xarray
def _get_chunk_slicer(
Expand All @@ -42,7 +32,7 @@ def block_slices(ds: xr.Dataset, chunks: Chunks = None) -> t.Iterator[Block]:
else:
chunks = ds.chunks

assert chunks, 'Dataset `ds` must be chunked or `chunks` must be provided.'
assert chunks, "Dataset `ds` must be chunked or `chunks` must be provided."

chunk_bounds = {dim: np.cumsum((0,) + c) for dim, c in chunks.items()}
ichunk = {dim: range(len(c)) for dim, c in chunks.items()}
Expand All @@ -67,8 +57,60 @@ def _block_len(block: Block) -> int:
return np.prod([v.stop - v.start for v in block.values()])


def read_xarray(ds: xr.Dataset, chunks: Chunks = None) -> dd.DataFrame:
"""Pivots an Xarray Dataset into a Dask Dataframe, partitioned by chunks.
def from_map(
func: t.Callable, *iterables, args: t.Optional[t.Tuple] = None, **kwargs
) -> pa.Table:
"""Create a PyArrow Table by mapping a function over iterables.

This is equivalent to dask's from_map but returns a PyArrow Table
that can be used with DataFusion instead of a Dask DataFrame.

Args:
func: Function to apply to each element of the iterables.
*iterables: Iterable objects to map the function over.
args: Additional positional arguments to pass to func.
**kwargs: Additional keyword arguments to pass to func.

Returns:
A PyArrow Table containing the concatenated results.
"""
if args is None:
args = ()

# Apply the function to each combination of iterable elements
results = []
for items in zip(*iterables) if len(iterables) > 1 else iterables[0]:
if isinstance(items, tuple):
result = func(*items, *args, **kwargs)
else:
result = func(items, *args, **kwargs)

# Convert result to PyArrow Table
if isinstance(result, pd.DataFrame):
pa_table = pa.Table.from_pandas(result)
elif isinstance(result, pa.Table):
pa_table = result
else:
# Try to convert to pandas first, then to PyArrow
try:
df = pd.DataFrame(result)
pa_table = pa.Table.from_pandas(df)
except Exception as e:
raise ValueError(
f"Cannot convert function result to PyArrow Table: {e}"
)

results.append(pa_table)

# Concatenate all results
if not results:
raise ValueError("No results to concatenate")

return pa.concat_tables(results)


def read_xarray(ds: xr.Dataset, chunks: Chunks = None) -> pa.Table:
"""Pivots an Xarray Dataset into a PyArrow Table, partitioned by chunks.

Args:
ds: An Xarray Dataset. All `data_vars` mush share the same dimensions.
Expand All @@ -77,39 +119,16 @@ def read_xarray(ds: xr.Dataset, chunks: Chunks = None) -> dd.DataFrame:
dataframe partition.

Returns:
A Dask Dataframe, which is a table representation of the input Dataset.
A PyArrow Table, which is a table representation of the input Dataset.
"""
fst = next(iter(ds.values())).dims
assert all(
da.dims == fst for da in ds.values()
), 'All dimensions must be equal. Please filter data_vars in the Dataset.'
), "All dimensions must be equal. Please filter data_vars in the Dataset."

blocks = list(block_slices(ds, chunks))

block_lengths = [_block_len(b) for b in blocks]
divisions = tuple(np.cumsum([0] + block_lengths)) # 0 ==> start partition.

def pivot(b: Block) -> pd.DataFrame:
return ds.isel(b).to_dataframe().reset_index()

# Token is needed to prevent Dask from spending too many cycles calculating
# it's own token from the constituent parts.
token = (
'xarray-Dataset-'
f'{"_".join(list(ds.dims.keys()))}'
'__'
f'{"_".join(list(ds.data_vars.keys()))}'
)

columns = pivot(blocks[0]).columns

# TODO(#18): Is it possible to pass the length (known now) here?
meta = {c: ds[c].dtype for c in columns}

return from_map(
pivot,
blocks,
meta=meta,
divisions=divisions,
token=token,
)
return from_map(pivot, blocks)
114 changes: 85 additions & 29 deletions xarray_sql/df_test.py
Original file line number Diff line number Diff line change
@@ -1,19 +1,19 @@
import itertools
import unittest

import dask.dataframe as dd
import numpy as np
import pandas as pd
import pyarrow as pa
import xarray as xr

from .df import explode, read_xarray, block_slices
from .df import explode, read_xarray, block_slices, from_map


def rand_wx(start: str, end: str) -> xr.Dataset:
np.random.seed(42)
lat = np.linspace(-90, 90, num=720)
lon = np.linspace(-180, 180, num=1440)
time = pd.date_range(start, end, freq='H')
time = pd.date_range(start, end, freq='h')
level = np.array([1000, 500], dtype=np.int32)
reference_time = pd.Timestamp(start)
temperature = 15 + 8 * np.random.randn(720, 1440, len(time), len(level))
Expand Down Expand Up @@ -59,7 +59,7 @@ def test_dim_sizes__one(self):
ds = next(iter(explode(self.air)))
for k, v in self.chunks.items():
self.assertIn(k, ds.dims)
self.assertEqual(v, ds.dims[k])
self.assertEqual(v, ds.sizes[k])

def skip_test_dim_sizes__all(self):
# TODO(alxmrs): Why is this test slow?
Expand All @@ -71,55 +71,111 @@ def skip_test_dim_sizes__all(self):

def test_data_equal__one__first(self):
ds = next(iter(explode(self.air)))
iselection = {dim: slice(0, s) for dim, s in ds.dims.items()}
iselection = {dim: slice(0, s) for dim, s in ds.sizes.items()}
self.assertEqual(self.air.isel(iselection), ds)

def test_data_equal__one__last(self):
dss = list(explode(self.air))
ds = dss[-1]
iselection = {dim: slice(0, s) for dim, s in ds.dims.items()}
iselection = {dim: slice(0, s) for dim, s in ds.sizes.items()}
self.assertEqual(self.air.isel(iselection), ds)


class DaskDataframeTest(DaskTestCase):
class PyArrowTableTest(DaskTestCase):

def test_sanity(self):
df = read_xarray(self.air_small).compute()
self.assertIsNotNone(df)
self.assertEqual(len(df), np.prod(list(self.air_small.dims.values())))
table = read_xarray(self.air_small)
self.assertIsNotNone(table)
self.assertIsInstance(table, pa.Table)
self.assertEqual(len(table), np.prod(list(self.air_small.sizes.values())))

def test_columns(self):
df = read_xarray(self.air_small).compute()
cols = list(df.columns)
table = read_xarray(self.air_small)
cols = table.column_names
self.assertEqual(cols, ['lat', 'time', 'lon', 'air'])

def test_dtypes(self):
df: dd.DataFrame = read_xarray(self.air_small).compute()
table = read_xarray(self.air_small)
# Convert to pandas to check dtypes
df = table.to_pandas()
types = list(df.dtypes)
self.assertEqual([self.air_small[c].dtype for c in df.columns], types)

def test_partitions_dont_match_dataset_chunks(self):
standard_blocks = list(block_slices(self.air_small))
default: dd.DataFrame = read_xarray(self.air_small)
chunked: dd.DataFrame = read_xarray(self.air_small, dict(time=5))
def test_different_chunk_sizes(self):
default_table = read_xarray(self.air_small)
chunked_table = read_xarray(self.air_small, dict(time=5))

self.assertEqual(default.npartitions, len(standard_blocks))
self.assertNotEqual(chunked.npartitions, len(standard_blocks))
# Both should produce valid tables
self.assertIsInstance(default_table, pa.Table)
self.assertIsInstance(chunked_table, pa.Table)
# Should have same number of rows
self.assertEqual(len(default_table), len(chunked_table))

def test_chunk_perf(self):
df = read_xarray(self.air, chunks=dict(time=6)).compute()
self.assertIsNotNone(df)
self.assertEqual(len(df), np.prod(list(self.air.dims.values())))
table = read_xarray(self.air, chunks=dict(time=6))
self.assertIsNotNone(table)
self.assertEqual(len(table), np.prod(list(self.air.sizes.values())))

def test_column_metadata_preserved(self):
try:
_ = read_xarray(self.randwx, chunks=dict(time=24)).compute()
except ValueError as e:
if (
'The columns in the computed data do not match the columns in the'
' provided metadata' in str(e)
):
self.fail('Column metadata is incorrect.')
table = read_xarray(self.randwx, chunks=dict(time=24))
self.assertIsInstance(table, pa.Table)
except Exception as e:
self.fail(f'Unexpected error: {e}')


class FromMapTest(unittest.TestCase):

def test_basic_from_map(self):
"""Test basic from_map functionality with pandas DataFrames."""

def make_df(x):
return pd.DataFrame({'value': [x, x * 2], 'index': [0, 1]})

result = from_map(make_df, [1, 2, 3])
self.assertIsInstance(result, pa.Table)
self.assertEqual(len(result), 6) # 3 inputs * 2 rows each
self.assertEqual(result.column_names, ['value', 'index'])

def test_from_map_with_multiple_iterables(self):
"""Test from_map with multiple iterables."""

def add_values(x, y):
return pd.DataFrame({'sum': [x + y], 'x': [x], 'y': [y]})

result = from_map(add_values, [1, 2], [10, 20])
self.assertIsInstance(result, pa.Table)
self.assertEqual(len(result), 2)

# Convert to pandas to check values
df = result.to_pandas()
self.assertEqual(list(df['sum']), [11, 22])

def test_from_map_with_args(self):
"""Test from_map with additional arguments."""

def multiply_and_add(x, multiplier, add_value):
return pd.DataFrame({'result': [x * multiplier + add_value]})

result = from_map(multiply_and_add, [1, 2, 3], args=(2, 10))
self.assertIsInstance(result, pa.Table)
self.assertEqual(len(result), 3)

df = result.to_pandas()
self.assertEqual(
list(df['result']), [12, 14, 16]
) # (1*2+10, 2*2+10, 3*2+10)

def test_from_map_with_pyarrow_tables(self):
"""Test from_map when function returns PyArrow tables."""

def make_arrow_table(x):
df = pd.DataFrame({'value': [x]})
return pa.Table.from_pandas(df)

result = from_map(make_arrow_table, [1, 2, 3])
self.assertIsInstance(result, pa.Table)
self.assertEqual(len(result), 3)


if __name__ == '__main__':
Expand Down
Loading