Skip to content

Commit 10c0de8

Browse files
committed
Address Lance review readability feedback
Signed-off-by: Vibhu Jawa <vjawa@nvidia.com>
1 parent 3e09e3d commit 10c0de8

2 files changed

Lines changed: 67 additions & 47 deletions

File tree

nemo_curator/stages/text/io/reader/lance.py

Lines changed: 6 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -59,19 +59,14 @@ def _requested_blob_v2_columns(dataset: object, scanner_kwargs: dict[str, Any])
5959

6060

6161
def _restore_lance_blob_v2_columns(dataset: object, table: pa.Table, blob_columns: list[str]) -> pa.Table:
62-
if not blob_columns:
63-
return table
64-
6562
import lance
6663

6764
rowaddrs = [int(value) for value in table["_rowaddr"].combine_chunks().to_pylist()]
6865
for column in blob_columns:
69-
if column not in table.column_names:
70-
continue
71-
payloads_by_rowaddr = dict(
72-
dataset.read_blobs(column, addresses=rowaddrs, preserve_order=True) # type: ignore[attr-defined]
73-
)
74-
payloads = [payloads_by_rowaddr.get(rowaddr) for rowaddr in rowaddrs]
66+
payloads = [
67+
payload
68+
for _, payload in dataset.read_blobs(column, addresses=rowaddrs, preserve_order=True) # type: ignore[attr-defined]
69+
]
7570
table = table.set_column(table.schema.get_field_index(column), column, lance.blob_array(payloads))
7671
return table
7772

@@ -179,7 +174,8 @@ def process(self, task: LanceReadTask) -> DocumentBatch | None:
179174
if table.num_rows == 0:
180175
return None
181176
lance_schema = pa.schema([dataset.schema.field(name) for name in table.column_names if name in dataset.schema.names])
182-
table = _restore_lance_blob_v2_columns(dataset, table, blob_columns)
177+
if blob_columns:
178+
table = _restore_lance_blob_v2_columns(dataset, table, blob_columns)
183179
if self.include_lance_metadata:
184180
table = _add_lance_metadata(table)
185181
elif blob_columns and "_rowaddr" in table.column_names:

nemo_curator/stages/text/io/writer/lance.py

Lines changed: 61 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -51,8 +51,13 @@ def _metadata_lance_schema(task: DocumentBatch) -> pa.Schema | None:
5151
return json_to_schema(schema)
5252

5353

54-
def _schema_for_table(schema: pa.Schema, table: pa.Table) -> pa.Schema:
55-
fields = [schema.field(name) if name in schema.names else table.schema.field(name) for name in table.column_names]
54+
def _schema_for_table(lance_schema: pa.Schema, table: pa.Table) -> pa.Schema:
55+
fields = []
56+
for table_field in table.schema:
57+
if table_field.name in lance_schema.names:
58+
fields.append(lance_schema.field(table_field.name))
59+
else:
60+
fields.append(table_field)
5661
return pa.schema(fields)
5762

5863

@@ -220,8 +225,6 @@ def process(self, task: DocumentBatch) -> FileGroupTask:
220225
fragment_ids = sorted(int(value) for value in pc.unique(table[LANCE_FRAGID_COLUMN].combine_chunks()).to_pylist())
221226
for fragment_id in fragment_ids:
222227
update_table = self._update_table_for_fragment(table, fragment_id)
223-
if update_table.num_rows == 0:
224-
continue
225228
fragment = dataset.get_fragment(fragment_id)
226229
updated_fragment, fields_modified = fragment.update_columns(
227230
update_table,
@@ -256,6 +259,49 @@ def process(self, task: DocumentBatch) -> FileGroupTask:
256259
)
257260

258261

262+
def _validate_checkpoint_path(records: list[dict[str, Any]], path: str) -> None:
263+
dataset_paths = {record["dataset_path"] for record in records}
264+
if dataset_paths != {path}:
265+
msg = f"Checkpoint records are for {sorted(dataset_paths)}, not {path}"
266+
raise ValueError(msg)
267+
268+
269+
def _single_checkpoint_value(records: list[dict[str, Any]], key: str, label: str) -> object:
270+
values = {record[key] for record in records}
271+
if len(values) != 1:
272+
msg = f"Expected one {label}; got {sorted(values)}"
273+
raise ValueError(msg)
274+
return next(iter(values))
275+
276+
277+
def _decode_write_fragments(records: list[dict[str, Any]]) -> list[tuple[object, pa.Schema]]:
278+
from lance.schema import json_to_schema
279+
280+
return [
281+
(pickle.loads(base64.b64decode(record["fragment"])), json_to_schema(record["schema"])) # noqa: S301
282+
for record in records
283+
]
284+
285+
286+
def _annotation_records_by_fragment(records: list[dict[str, Any]]) -> dict[int, dict[str, Any]]:
287+
records_by_fragment = {int(record["fragment_id"]): record for record in records}
288+
if len(records_by_fragment) != len(records):
289+
msg = "Ensure each Lance fragment is updated by at most one writer task."
290+
raise ValueError(msg)
291+
return records_by_fragment
292+
293+
294+
def _decode_updated_fragments(records: list[dict[str, Any]]) -> list[object]:
295+
return [
296+
pickle.loads(base64.b64decode(record["updated_fragment"])) # noqa: S301
297+
for record in records
298+
]
299+
300+
301+
def _fields_modified(records: list[dict[str, Any]]) -> list[str]:
302+
return sorted({field for record in records for field in record["fields_modified"]})
303+
304+
259305
def commit_lance_checkpoint(
260306
path: str,
261307
commit_path: str,
@@ -265,32 +311,22 @@ def commit_lance_checkpoint(
265311
) -> int:
266312
"""Commit records written by ``LanceWriter`` and return the Lance version."""
267313
import lance
268-
from lance.schema import json_to_schema
269314
from lance_ray import LanceFragmentCommitter
270315

271316
records, committed_version = read_lance_checkpoint(commit_path, "lance_write", checkpoint_storage_options)
272317
if committed_version is not None:
273318
return committed_version
274319

275-
dataset_paths = {record["dataset_path"] for record in records}
276-
if dataset_paths != {path}:
277-
msg = f"Checkpoint records are for {sorted(dataset_paths)}, not {path}"
278-
raise ValueError(msg)
279-
modes = {record["mode"] for record in records}
280-
if len(modes) != 1:
281-
msg = f"Expected one write mode; got {sorted(modes)}"
282-
raise ValueError(msg)
283-
mode = str(next(iter(modes)))
284-
fragments = [
285-
(pickle.loads(base64.b64decode(record["fragment"])), json_to_schema(record["schema"])) # noqa: S301
286-
for record in records
287-
]
320+
_validate_checkpoint_path(records, path)
321+
mode = str(_single_checkpoint_value(records, "mode", "write mode"))
322+
fragments = _decode_write_fragments(records)
288323
schema = fragments[0][1]
289324

290325
committer = LanceFragmentCommitter(path, schema=schema, mode=mode, storage_options=storage_options)
291326
if mode == "append":
292327
committer.on_write_start(schema)
293-
committer.on_write_complete([[(pickle.dumps(fragment), pickle.dumps(schema)) for fragment, schema in fragments]])
328+
fragment_payloads = [(pickle.dumps(fragment), pickle.dumps(schema)) for fragment, schema in fragments]
329+
committer.on_write_complete([fragment_payloads])
294330
version = lance.dataset(path, storage_options=storage_options).version
295331
write_lance_checkpoint_marker(commit_path, version, checkpoint_storage_options)
296332
return version
@@ -312,24 +348,12 @@ def commit_lance_annotation_checkpoint(
312348
if committed_version is not None:
313349
return committed_version
314350

315-
dataset_paths = {record["dataset_path"] for record in records}
316-
if dataset_paths != {path}:
317-
msg = f"Checkpoint records are for {sorted(dataset_paths)}, not {path}"
318-
raise ValueError(msg)
319-
read_versions = {int(record["dataset_version"]) for record in records}
320-
if len(read_versions) != 1:
321-
msg = f"Expected one dataset version; got {sorted(read_versions)}"
322-
raise ValueError(msg)
323-
read_version = next(iter(read_versions))
324-
records_by_fragment = {int(record["fragment_id"]): record for record in records}
325-
if len(records_by_fragment) != len(records):
326-
msg = "Ensure each Lance fragment is updated by at most one writer task."
327-
raise ValueError(msg)
328-
updated_fragments = [
329-
pickle.loads(base64.b64decode(record["updated_fragment"])) # noqa: S301
330-
for record in records_by_fragment.values()
331-
]
332-
fields_modified = sorted({field for record in records_by_fragment.values() for field in record["fields_modified"]})
351+
_validate_checkpoint_path(records, path)
352+
read_version = int(_single_checkpoint_value(records, "dataset_version", "dataset version"))
353+
records_by_fragment = _annotation_records_by_fragment(records)
354+
fragment_records = list(records_by_fragment.values())
355+
updated_fragments = _decode_updated_fragments(fragment_records)
356+
fields_modified = _fields_modified(fragment_records)
333357
operation = lance.LanceOperation.Update(updated_fragments=updated_fragments, fields_modified=fields_modified)
334358
version = lance.LanceDataset.commit(
335359
path,

0 commit comments

Comments
 (0)