File size: 2,055 Bytes
61449ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
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,
        }