| from __future__ import annotations |
|
|
| import random |
| from collections.abc import Iterator, Sequence |
| from pathlib import Path |
| from typing import Any |
|
|
| from .common import write_jsonl |
|
|
|
|
| class DataObject(Sequence[dict[str, str]]): |
| def __init__( |
| self, |
| records: Sequence[dict[str, str]] | None = None, |
| *, |
| name: str = "dataset", |
| warnings: Sequence[str] | None = None, |
| ) -> None: |
| self.records = list(records or []) |
| self.name = name |
| self.warnings = list(warnings or []) |
|
|
| def __len__(self) -> int: |
| return len(self.records) |
|
|
| def __iter__(self) -> Iterator[dict[str, str]]: |
| return iter(self.records) |
|
|
| def __getitem__(self, index: int | slice) -> dict[str, str] | list[dict[str, str]]: |
| return self.records[index] |
|
|
| def take(self, amount: int) -> list[dict[str, str]]: |
| if amount < 0: |
| raise ValueError("amount must be non-negative") |
|
|
| return self.records[:amount] |
|
|
| def shuffle(self, *, seed: int | None = 42) -> "DataObject": |
| shuffled = list(self.records) |
| rng = random.Random(seed) |
| rng.shuffle(shuffled) |
|
|
| return DataObject( |
| shuffled, |
| name=self.name, |
| warnings=self.warnings, |
| ) |
|
|
| def to_jsonl(self, path: str | Path) -> None: |
| write_jsonl(self.records, path) |
|
|
| def extend(self, other: "DataObject") -> None: |
| self.records.extend(other.records) |
| self.warnings.extend(other.warnings) |
|
|
| @classmethod |
| def concat(cls, datasets: Sequence["DataObject"], *, name: str = "all") -> "DataObject": |
| records: list[dict[str, str]] = [] |
| warnings: list[str] = [] |
|
|
| for dataset in datasets: |
| records.extend(dataset.records) |
| warnings.extend(dataset.warnings) |
|
|
| return cls(records, name=name, warnings=warnings) |
|
|
| def summary(self) -> dict[str, Any]: |
| return { |
| "name": self.name, |
| "rows": len(self.records), |
| "warnings": self.warnings, |
| } |
|
|