Skip to content
Open
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
35 changes: 28 additions & 7 deletions src/datasets/arrow_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,33 @@ def get_writer_batch_size_from_data_size(num_rows: int, num_bytes: int) -> int:
return max(1, num_rows * convert_file_size_to_int(config.MAX_ROW_GROUP_SIZE) // num_bytes) if num_bytes > 0 else 1


def get_parquet_column_options(features: Features, schema: pa.Schema) -> dict:
"""
Get the per-column `compression`, `use_dictionary` and `column_encoding` options for `pq.ParquetWriter`.
Parquet matches them against the leaf column paths (e.g. "col.list.element.field"), not the top-level column names.
"""

def leaf_paths(path: str, pa_type: pa.DataType) -> list[str]:
if isinstance(pa_type, pa.ExtensionType):
pa_type = pa_type.storage_type
if pa.types.is_struct(pa_type):
return [leaf for field in pa_type for leaf in leaf_paths(f"{path}.{field.name}", field.type)]
if pa.types.is_list(pa_type) or pa.types.is_large_list(pa_type) or pa.types.is_fixed_size_list(pa_type):
return leaf_paths(f"{path}.list.element", pa_type.value_type)
return [path]

compression, use_dictionary, column_encoding = {}, [], {}
for field in schema:
embed = require_storage_embed(features[field.name])
for leaf in leaf_paths(field.name, field.type):
compression[leaf] = "none" if embed else "snappy"
if embed:
column_encoding[leaf] = "PLAIN"
else:
use_dictionary.append(leaf)
return {"compression": compression, "use_dictionary": use_dictionary, "column_encoding": column_encoding}


class SchemaInferenceError(ValueError):
pass

Expand Down Expand Up @@ -836,13 +863,7 @@ def _build_writer(self, inferred_schema: pa.Schema):
self._schema,
use_content_defined_chunking=self.use_content_defined_chunking,
write_page_index=self.write_page_index,
compression={
col: "none" if require_storage_embed(feature) else "snappy" for col, feature in self._features.items()
},
use_dictionary=[col for col, feature in self._features.items() if not require_storage_embed(feature)],
column_encoding={
col: "PLAIN" for col, feature in self._features.items() if require_storage_embed(feature)
},
**get_parquet_column_options(self._features, self._schema),
)
if self.use_content_defined_chunking is not False:
self.pa_writer.add_key_value_metadata(
Expand Down
18 changes: 6 additions & 12 deletions src/datasets/io/parquet.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,11 @@
import pyarrow.parquet as pq

from .. import Dataset, Features, NamedSplit, config
from ..arrow_writer import get_writer_batch_size_from_data_size, get_writer_batch_size_from_features
from ..features.features import require_storage_embed
from ..arrow_writer import (
get_parquet_column_options,
get_writer_batch_size_from_data_size,
get_writer_batch_size_from_features,
)
from ..formatting import query_table
from ..packaged_modules import _PACKAGED_DATASETS_MODULES
from ..packaged_modules.parquet.parquet import Parquet
Expand Down Expand Up @@ -125,16 +128,7 @@ def _write(self, file_obj: BinaryIO, batch_size: int, **parquet_writer_kwargs) -
schema=schema,
use_content_defined_chunking=self.use_content_defined_chunking,
write_page_index=self.write_page_index,
compression={
col: "none" if require_storage_embed(feature) else "snappy"
for col, feature in self.dataset.features.items()
},
use_dictionary=[
col for col, feature in self.dataset.features.items() if not require_storage_embed(feature)
],
column_encoding={
col: "PLAIN" for col, feature in self.dataset.features.items() if require_storage_embed(feature)
},
**get_parquet_column_options(self.dataset.features, schema),
**parquet_writer_kwargs,
)

Expand Down
34 changes: 33 additions & 1 deletion tests/test_arrow_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@

from datasets import config
from datasets.arrow_writer import ArrowWriter, OptimizedTypedSequence, ParquetWriter, TypedSequence
from datasets.features import Array2D, ClassLabel, Features, Image, Value
from datasets.features import Array2D, ClassLabel, Features, Image, List, Value
from datasets.features.features import Array2DExtensionType, cast_to_python_objects

from .utils import require_pil
Expand Down Expand Up @@ -334,6 +334,38 @@ def test_parquet_writer_write():
assert pa_table.to_pydict() == {"col_1": ["foo", "bar"], "col_2": [1, 2]}


def test_parquet_writer_compresses_nested_columns():
features = Features(
{
"text": Value("string"),
"nested": List({"a": Value("int64"), "b": List(Value("string"))}),
"array": Array2D(shape=(2, 2), dtype="int32"),
"image": Image(),
}
)
output = pa.BufferOutputStream()
with ParquetWriter(stream=output, features=features) as writer:
writer.write(
{
"text": "foo",
"nested": [{"a": 1, "b": ["bar"]}],
"array": [[1, 2], [3, 4]],
"image": {"bytes": b"image_bytes", "path": None},
}
)
writer.finalize()
metadata = pq.ParquetFile(pa.BufferReader(output.getvalue())).metadata
columns = [metadata.row_group(0).column(i) for i in range(metadata.num_columns)]
assert {column.path_in_schema: column.compression for column in columns} == {
"text": "SNAPPY",
"nested.list.element.a": "SNAPPY",
"nested.list.element.b.list.element": "SNAPPY",
"array.list.element.list.element": "SNAPPY",
"image.bytes": "UNCOMPRESSED",
"image.path": "UNCOMPRESSED",
}


def test_parquet_writer_uses_content_defined_chunking():
def write_and_get_argument_and_metadata(**kwargs):
output = pa.BufferOutputStream()
Expand Down
Loading