tensorless 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- tensorless/__init__.py +40 -0
- tensorless/_version.py +6 -0
- tensorless/api.py +230 -0
- tensorless/auto/__init__.py +4 -0
- tensorless/auto/config.py +113 -0
- tensorless/auto/detector.py +95 -0
- tensorless/checkpoint/__init__.py +3 -0
- tensorless/checkpoint/manager.py +66 -0
- tensorless/cli/__init__.py +3 -0
- tensorless/cli/main.py +121 -0
- tensorless/config.py +121 -0
- tensorless/data/__init__.py +11 -0
- tensorless/data/fingerprint.py +71 -0
- tensorless/data/inspector.py +161 -0
- tensorless/data/loader.py +255 -0
- tensorless/data/tabular.py +179 -0
- tensorless/devices/__init__.py +3 -0
- tensorless/devices/device.py +107 -0
- tensorless/errors.py +36 -0
- tensorless/models/__init__.py +5 -0
- tensorless/models/mlp.py +64 -0
- tensorless/models/registry.py +53 -0
- tensorless/models/transformer.py +175 -0
- tensorless/runtime.py +158 -0
- tensorless/serialization/__init__.py +3 -0
- tensorless/serialization/tl_format.py +89 -0
- tensorless/tokenization/__init__.py +3 -0
- tensorless/tokenization/char_tokenizer.py +84 -0
- tensorless/training/__init__.py +4 -0
- tensorless/training/data_prep.py +215 -0
- tensorless/training/early_stopping.py +31 -0
- tensorless/training/trainer.py +234 -0
- tensorless-0.1.0.dist-info/METADATA +111 -0
- tensorless-0.1.0.dist-info/RECORD +38 -0
- tensorless-0.1.0.dist-info/WHEEL +5 -0
- tensorless-0.1.0.dist-info/entry_points.txt +2 -0
- tensorless-0.1.0.dist-info/licenses/LICENSE +21 -0
- tensorless-0.1.0.dist-info/top_level.txt +1 -0
tensorless/config.py
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
"""Configuration objects.
|
|
2
|
+
|
|
3
|
+
`TrainConfig` holds every knob a user can override in `tl.train(...)`.
|
|
4
|
+
Any field left as `None` means "let Tensorless decide automatically".
|
|
5
|
+
`ResolvedConfig` is what the auto-configuration system produces after
|
|
6
|
+
filling in every `None` with a concrete value.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from dataclasses import dataclass, field, asdict
|
|
12
|
+
from typing import Optional, Any, Dict
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass
|
|
16
|
+
class TrainConfig:
|
|
17
|
+
"""User-facing training configuration. Every field is optional;
|
|
18
|
+
unset fields are chosen automatically by the auto-config system.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
# --- output / lifecycle ---
|
|
22
|
+
out: Optional[str] = None # output .tl path, default "model.tl"
|
|
23
|
+
force: bool = False # force retraining even if unchanged
|
|
24
|
+
resume: Optional[bool] = None # force/forbid resume (None = auto)
|
|
25
|
+
ask_on_data_change: bool = False # raise instead of auto-retrain on data change
|
|
26
|
+
|
|
27
|
+
# --- task / architecture ---
|
|
28
|
+
task: Optional[str] = None # "text-generation", "classification", "regression"
|
|
29
|
+
model_type: Optional[str] = None # "transformer", "mlp"
|
|
30
|
+
d_model: Optional[int] = None
|
|
31
|
+
layers: Optional[int] = None
|
|
32
|
+
heads: Optional[int] = None
|
|
33
|
+
ff_mult: Optional[int] = None
|
|
34
|
+
dropout: Optional[float] = None
|
|
35
|
+
max_seq_len: Optional[int] = None
|
|
36
|
+
|
|
37
|
+
# --- optimization ---
|
|
38
|
+
optimizer: Optional[str] = None # "adamw", "adam", "sgd"
|
|
39
|
+
learning_rate: Optional[float] = None
|
|
40
|
+
weight_decay: Optional[float] = None
|
|
41
|
+
batch_size: Optional[int] = None
|
|
42
|
+
epochs: Optional[int] = None
|
|
43
|
+
max_steps: Optional[int] = None
|
|
44
|
+
grad_clip: Optional[float] = None
|
|
45
|
+
warmup_steps: Optional[int] = None
|
|
46
|
+
|
|
47
|
+
# --- validation / early stopping ---
|
|
48
|
+
val_split: Optional[float] = None
|
|
49
|
+
patience: Optional[int] = None
|
|
50
|
+
min_delta: Optional[float] = None
|
|
51
|
+
|
|
52
|
+
# --- hardware ---
|
|
53
|
+
device: Optional[str] = None # "cpu", "cuda", "tpu", or None = auto
|
|
54
|
+
precision: Optional[str] = None # "fp32", "fp16", "bf16"
|
|
55
|
+
|
|
56
|
+
# --- checkpointing ---
|
|
57
|
+
checkpoint_every: Optional[int] = None # steps between checkpoints
|
|
58
|
+
checkpoint_dir: Optional[str] = None
|
|
59
|
+
|
|
60
|
+
# --- misc ---
|
|
61
|
+
seed: int = 42
|
|
62
|
+
verbose: bool = True
|
|
63
|
+
extra: Dict[str, Any] = field(default_factory=dict)
|
|
64
|
+
|
|
65
|
+
def overrides(self) -> Dict[str, Any]:
|
|
66
|
+
"""Return only the fields the user explicitly set (non-None)."""
|
|
67
|
+
d = asdict(self)
|
|
68
|
+
d.pop("extra", None)
|
|
69
|
+
return {k: v for k, v in d.items() if v is not None and v is not False}
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
@dataclass
|
|
73
|
+
class ResolvedConfig:
|
|
74
|
+
"""Fully resolved configuration -- every field has a concrete value.
|
|
75
|
+
This is what actually gets used for training and what gets embedded
|
|
76
|
+
in checkpoints and the final `.tl` file.
|
|
77
|
+
"""
|
|
78
|
+
|
|
79
|
+
out: str
|
|
80
|
+
force: bool
|
|
81
|
+
resume: Optional[bool]
|
|
82
|
+
ask_on_data_change: bool
|
|
83
|
+
|
|
84
|
+
task: str
|
|
85
|
+
model_type: str
|
|
86
|
+
d_model: int
|
|
87
|
+
layers: int
|
|
88
|
+
heads: int
|
|
89
|
+
ff_mult: int
|
|
90
|
+
dropout: float
|
|
91
|
+
max_seq_len: int
|
|
92
|
+
|
|
93
|
+
optimizer: str
|
|
94
|
+
learning_rate: float
|
|
95
|
+
weight_decay: float
|
|
96
|
+
batch_size: int
|
|
97
|
+
epochs: int
|
|
98
|
+
max_steps: Optional[int]
|
|
99
|
+
grad_clip: float
|
|
100
|
+
warmup_steps: int
|
|
101
|
+
|
|
102
|
+
val_split: float
|
|
103
|
+
patience: int
|
|
104
|
+
min_delta: float
|
|
105
|
+
|
|
106
|
+
device: str
|
|
107
|
+
precision: str
|
|
108
|
+
|
|
109
|
+
checkpoint_every: int
|
|
110
|
+
checkpoint_dir: str
|
|
111
|
+
|
|
112
|
+
seed: int
|
|
113
|
+
verbose: bool
|
|
114
|
+
|
|
115
|
+
def to_dict(self) -> Dict[str, Any]:
|
|
116
|
+
return asdict(self)
|
|
117
|
+
|
|
118
|
+
@classmethod
|
|
119
|
+
def from_dict(cls, d: Dict[str, Any]) -> "ResolvedConfig":
|
|
120
|
+
known = {f for f in cls.__dataclass_fields__.keys()}
|
|
121
|
+
return cls(**{k: v for k, v in d.items() if k in known})
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""Dataset fingerprinting.
|
|
2
|
+
|
|
3
|
+
We hash the *content* of every data file under a path (plus filenames and
|
|
4
|
+
sizes) so that Tensorless can cheaply detect whether a dataset has changed
|
|
5
|
+
between runs -- this is what powers the "don't retrain if nothing changed"
|
|
6
|
+
and "resume if interrupted" behaviour.
|
|
7
|
+
|
|
8
|
+
The fingerprint is intentionally content-based (not mtime-based) so that
|
|
9
|
+
copying a dataset to a new machine, or touching a file without changing it,
|
|
10
|
+
does not trigger an unnecessary retrain.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import hashlib
|
|
16
|
+
import os
|
|
17
|
+
from typing import List
|
|
18
|
+
|
|
19
|
+
# Cap how many bytes of large files we hash, to keep fingerprinting fast on
|
|
20
|
+
# huge datasets. We hash the size + a content sample (head/tail) rather than
|
|
21
|
+
# the whole file when it exceeds this threshold.
|
|
22
|
+
_MAX_FULL_HASH_BYTES = 25 * 1024 * 1024 # 25 MB
|
|
23
|
+
_SAMPLE_BYTES = 1 * 1024 * 1024 # 1 MB head + 1 MB tail for big files
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _iter_files(path: str) -> List[str]:
|
|
27
|
+
if os.path.isfile(path):
|
|
28
|
+
return [path]
|
|
29
|
+
files = []
|
|
30
|
+
for root, dirs, names in os.walk(path):
|
|
31
|
+
dirs.sort()
|
|
32
|
+
for name in sorted(names):
|
|
33
|
+
if name.startswith("."):
|
|
34
|
+
continue
|
|
35
|
+
files.append(os.path.join(root, name))
|
|
36
|
+
return files
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _hash_file(path: str, hasher: "hashlib._Hash") -> None:
|
|
40
|
+
size = os.path.getsize(path)
|
|
41
|
+
hasher.update(str(size).encode("utf-8"))
|
|
42
|
+
with open(path, "rb") as f:
|
|
43
|
+
if size <= _MAX_FULL_HASH_BYTES:
|
|
44
|
+
while True:
|
|
45
|
+
chunk = f.read(1024 * 1024)
|
|
46
|
+
if not chunk:
|
|
47
|
+
break
|
|
48
|
+
hasher.update(chunk)
|
|
49
|
+
else:
|
|
50
|
+
hasher.update(f.read(_SAMPLE_BYTES))
|
|
51
|
+
f.seek(max(0, size - _SAMPLE_BYTES))
|
|
52
|
+
hasher.update(f.read(_SAMPLE_BYTES))
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def fingerprint_path(path: str) -> str:
|
|
56
|
+
"""Return a stable hex digest fingerprinting the dataset at `path`.
|
|
57
|
+
|
|
58
|
+
The fingerprint changes if any file's content, size, name, or the set
|
|
59
|
+
of files itself changes.
|
|
60
|
+
"""
|
|
61
|
+
if not os.path.exists(path):
|
|
62
|
+
raise FileNotFoundError(f"Path '{path}' does not exist.")
|
|
63
|
+
path = os.path.abspath(path)
|
|
64
|
+
hasher = hashlib.sha256()
|
|
65
|
+
files = _iter_files(path)
|
|
66
|
+
for f in files:
|
|
67
|
+
rel = os.path.relpath(f, path if os.path.isdir(path) else os.path.dirname(path))
|
|
68
|
+
hasher.update(rel.encode("utf-8"))
|
|
69
|
+
_hash_file(f, hasher)
|
|
70
|
+
hasher.update(str(len(files)).encode("utf-8"))
|
|
71
|
+
return hasher.hexdigest()
|
|
@@ -0,0 +1,161 @@
|
|
|
1
|
+
"""Dataset inspection: `tl.inspect("./data")`.
|
|
2
|
+
|
|
3
|
+
Loads the dataset, detects its task type, and reports size, samples,
|
|
4
|
+
detected problems, and recommendations -- without training anything.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from dataclasses import dataclass, field
|
|
10
|
+
from typing import Any, Dict, List
|
|
11
|
+
|
|
12
|
+
from .loader import load_dataset, Dataset
|
|
13
|
+
from .fingerprint import fingerprint_path
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@dataclass
|
|
17
|
+
class InspectionReport:
|
|
18
|
+
path: str
|
|
19
|
+
fingerprint: str
|
|
20
|
+
kind: str
|
|
21
|
+
task: str
|
|
22
|
+
n_examples: int
|
|
23
|
+
n_files: int
|
|
24
|
+
columns: List[str] = field(default_factory=list)
|
|
25
|
+
sample: Any = None
|
|
26
|
+
warnings: List[str] = field(default_factory=list)
|
|
27
|
+
recommendations: List[str] = field(default_factory=list)
|
|
28
|
+
stats: Dict[str, Any] = field(default_factory=dict)
|
|
29
|
+
|
|
30
|
+
def __str__(self) -> str:
|
|
31
|
+
lines = [
|
|
32
|
+
f"Dataset: {self.path}",
|
|
33
|
+
f" fingerprint : {self.fingerprint[:16]}...",
|
|
34
|
+
f" detected kind : {self.kind}",
|
|
35
|
+
f" detected task : {self.task}",
|
|
36
|
+
f" examples : {self.n_examples}",
|
|
37
|
+
f" files : {self.n_files}",
|
|
38
|
+
]
|
|
39
|
+
if self.columns:
|
|
40
|
+
lines.append(f" columns : {', '.join(self.columns)}")
|
|
41
|
+
for k, v in self.stats.items():
|
|
42
|
+
lines.append(f" {k:14s}: {v}")
|
|
43
|
+
if self.warnings:
|
|
44
|
+
lines.append(" Warnings:")
|
|
45
|
+
lines += [f" - {w}" for w in self.warnings]
|
|
46
|
+
if self.recommendations:
|
|
47
|
+
lines.append(" Recommendations:")
|
|
48
|
+
lines += [f" - {r}" for r in self.recommendations]
|
|
49
|
+
return "\n".join(lines)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _compute_stats(ds: Dataset) -> Dict[str, Any]:
|
|
53
|
+
stats: Dict[str, Any] = {}
|
|
54
|
+
if ds.kind in ("text", "text_labeled"):
|
|
55
|
+
lengths = [len(t) for t in ds.texts]
|
|
56
|
+
if lengths:
|
|
57
|
+
stats["avg_chars"] = round(sum(lengths) / len(lengths), 1)
|
|
58
|
+
stats["min_chars"] = min(lengths)
|
|
59
|
+
stats["max_chars"] = max(lengths)
|
|
60
|
+
if ds.kind == "text_labeled":
|
|
61
|
+
classes = sorted(set(ds.labels))
|
|
62
|
+
stats["n_classes"] = len(classes)
|
|
63
|
+
stats["classes"] = classes[:20]
|
|
64
|
+
else:
|
|
65
|
+
stats["n_rows"] = len(ds.records)
|
|
66
|
+
stats["n_columns"] = len(ds.columns)
|
|
67
|
+
return stats
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _warnings_and_recommendations(ds: Dataset, task: str) -> (List[str], List[str]):
|
|
71
|
+
warnings: List[str] = []
|
|
72
|
+
recs: List[str] = []
|
|
73
|
+
n = len(ds)
|
|
74
|
+
|
|
75
|
+
if n == 0:
|
|
76
|
+
warnings.append("Dataset is empty.")
|
|
77
|
+
return warnings, recs
|
|
78
|
+
|
|
79
|
+
if n < 20:
|
|
80
|
+
warnings.append(
|
|
81
|
+
f"Only {n} example(s) found. This is very small for training a "
|
|
82
|
+
f"useful model; results may be poor / mostly memorization."
|
|
83
|
+
)
|
|
84
|
+
recs.append("Collect more data if possible (aim for hundreds+ examples).")
|
|
85
|
+
|
|
86
|
+
if ds.kind == "text_labeled":
|
|
87
|
+
from collections import Counter
|
|
88
|
+
|
|
89
|
+
counts = Counter(ds.labels)
|
|
90
|
+
if len(counts) < 2:
|
|
91
|
+
warnings.append("Only one class detected; classification needs 2+ classes.")
|
|
92
|
+
else:
|
|
93
|
+
majority = max(counts.values())
|
|
94
|
+
minority = min(counts.values())
|
|
95
|
+
if majority > 3 * max(minority, 1):
|
|
96
|
+
warnings.append(
|
|
97
|
+
f"Class imbalance detected (largest class {majority} vs "
|
|
98
|
+
f"smallest {minority}). Consider balancing or using "
|
|
99
|
+
f"class weights."
|
|
100
|
+
)
|
|
101
|
+
recs.append(
|
|
102
|
+
"Tensorless will still train, but consider collecting more "
|
|
103
|
+
"examples for underrepresented classes."
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
if ds.kind == "tabular":
|
|
107
|
+
missing_cols = set()
|
|
108
|
+
for r in ds.records[:200]:
|
|
109
|
+
for c in ds.columns:
|
|
110
|
+
v = r.get(c, "")
|
|
111
|
+
if v is None or (isinstance(v, str) and v.strip() == ""):
|
|
112
|
+
missing_cols.add(c)
|
|
113
|
+
if missing_cols:
|
|
114
|
+
warnings.append(
|
|
115
|
+
f"Missing values detected in column(s): {', '.join(sorted(missing_cols))}."
|
|
116
|
+
)
|
|
117
|
+
recs.append(
|
|
118
|
+
"Tensorless will impute missing numeric values with the column "
|
|
119
|
+
"mean and missing categorical values with a placeholder token."
|
|
120
|
+
)
|
|
121
|
+
|
|
122
|
+
if ds.kind in ("text", "text_labeled"):
|
|
123
|
+
avg_len = sum(len(t) for t in ds.texts) / max(1, len(ds.texts))
|
|
124
|
+
if avg_len > 20000:
|
|
125
|
+
recs.append(
|
|
126
|
+
"Texts are long; Tensorless will truncate to the configured "
|
|
127
|
+
"max_seq_len. Pass max_seq_len=... to change this."
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
return warnings, recs
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def inspect_path(path: str) -> InspectionReport:
|
|
134
|
+
# Local import to avoid a circular import between `data` and `auto`.
|
|
135
|
+
from ..auto.detector import detect_task
|
|
136
|
+
|
|
137
|
+
ds = load_dataset(path)
|
|
138
|
+
task = detect_task(ds)
|
|
139
|
+
fp = fingerprint_path(path)
|
|
140
|
+
stats = _compute_stats(ds)
|
|
141
|
+
warnings, recs = _warnings_and_recommendations(ds, task)
|
|
142
|
+
|
|
143
|
+
sample: Any = None
|
|
144
|
+
if ds.kind in ("text", "text_labeled") and ds.texts:
|
|
145
|
+
sample = ds.texts[0][:300]
|
|
146
|
+
elif ds.records:
|
|
147
|
+
sample = ds.records[0]
|
|
148
|
+
|
|
149
|
+
return InspectionReport(
|
|
150
|
+
path=path,
|
|
151
|
+
fingerprint=fp,
|
|
152
|
+
kind=ds.kind,
|
|
153
|
+
task=task,
|
|
154
|
+
n_examples=len(ds),
|
|
155
|
+
n_files=ds.n_files,
|
|
156
|
+
columns=list(ds.columns),
|
|
157
|
+
sample=sample,
|
|
158
|
+
warnings=warnings,
|
|
159
|
+
recommendations=recs,
|
|
160
|
+
stats=stats,
|
|
161
|
+
)
|
|
@@ -0,0 +1,255 @@
|
|
|
1
|
+
"""Dataset loading.
|
|
2
|
+
|
|
3
|
+
Turns a path (file or directory) into a normalized `Dataset` object that
|
|
4
|
+
the rest of Tensorless can reason about, regardless of whether the data
|
|
5
|
+
started life as .txt, .json, .jsonl, .csv, or a directory of any of those
|
|
6
|
+
(optionally organized into class subfolders).
|
|
7
|
+
|
|
8
|
+
Design goal: never silently drop or mutate user data. If something looks
|
|
9
|
+
wrong (empty dataset, unreadable file, inconsistent columns) we raise a
|
|
10
|
+
`DataError` with an actionable message rather than guessing.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import csv
|
|
16
|
+
import json
|
|
17
|
+
import os
|
|
18
|
+
from dataclasses import dataclass, field
|
|
19
|
+
from typing import Any, Dict, List, Optional
|
|
20
|
+
|
|
21
|
+
from ..errors import DataError
|
|
22
|
+
|
|
23
|
+
TEXT_EXTENSIONS = {".txt", ".md"}
|
|
24
|
+
JSON_EXTENSIONS = {".json"}
|
|
25
|
+
JSONL_EXTENSIONS = {".jsonl", ".ndjson"}
|
|
26
|
+
CSV_EXTENSIONS = {".csv", ".tsv"}
|
|
27
|
+
|
|
28
|
+
# Common field names we look for when a JSON/JSONL record represents a
|
|
29
|
+
# single piece of free text (e.g. for language modeling).
|
|
30
|
+
_TEXT_FIELD_CANDIDATES = ("text", "content", "body", "document", "sentence")
|
|
31
|
+
_LABEL_FIELD_CANDIDATES = ("label", "target", "class", "category", "y")
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@dataclass
|
|
35
|
+
class Dataset:
|
|
36
|
+
"""Normalized in-memory representation of a loaded dataset.
|
|
37
|
+
|
|
38
|
+
kind is one of:
|
|
39
|
+
- "text" : a corpus of raw text (language modeling)
|
|
40
|
+
- "text_labeled": (text, label) pairs (text classification)
|
|
41
|
+
- "tabular" : rows of named columns (classification/regression)
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
kind: str
|
|
45
|
+
source: str
|
|
46
|
+
texts: List[str] = field(default_factory=list)
|
|
47
|
+
labels: List[Any] = field(default_factory=list)
|
|
48
|
+
records: List[Dict[str, Any]] = field(default_factory=list)
|
|
49
|
+
columns: List[str] = field(default_factory=list)
|
|
50
|
+
n_files: int = 0
|
|
51
|
+
|
|
52
|
+
def __len__(self) -> int:
|
|
53
|
+
if self.kind in ("text", "text_labeled"):
|
|
54
|
+
return len(self.texts)
|
|
55
|
+
return len(self.records)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _read_text_file(path: str) -> str:
|
|
59
|
+
try:
|
|
60
|
+
with open(path, "r", encoding="utf-8", errors="strict") as f:
|
|
61
|
+
return f.read()
|
|
62
|
+
except UnicodeDecodeError as e:
|
|
63
|
+
raise DataError(
|
|
64
|
+
f"File '{path}' is not valid UTF-8 text. Tensorless expects "
|
|
65
|
+
f"text datasets to be UTF-8 encoded."
|
|
66
|
+
) from e
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _read_json_records(path: str) -> List[Dict[str, Any]]:
|
|
70
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
71
|
+
data = json.load(f)
|
|
72
|
+
if isinstance(data, dict):
|
|
73
|
+
# Either a single record, or {"data": [...]}
|
|
74
|
+
for key in ("data", "records", "items", "examples"):
|
|
75
|
+
if key in data and isinstance(data[key], list):
|
|
76
|
+
data = data[key]
|
|
77
|
+
break
|
|
78
|
+
else:
|
|
79
|
+
data = [data]
|
|
80
|
+
if not isinstance(data, list):
|
|
81
|
+
raise DataError(f"Unsupported JSON structure in '{path}': expected a list of records.")
|
|
82
|
+
return data
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def _read_jsonl_records(path: str) -> List[Dict[str, Any]]:
|
|
86
|
+
records = []
|
|
87
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
88
|
+
for i, line in enumerate(f):
|
|
89
|
+
line = line.strip()
|
|
90
|
+
if not line:
|
|
91
|
+
continue
|
|
92
|
+
try:
|
|
93
|
+
records.append(json.loads(line))
|
|
94
|
+
except json.JSONDecodeError as e:
|
|
95
|
+
raise DataError(f"Malformed JSON on line {i + 1} of '{path}': {e}") from e
|
|
96
|
+
return records
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def _read_csv_records(path: str) -> List[Dict[str, Any]]:
|
|
100
|
+
delimiter = "\t" if path.lower().endswith(".tsv") else ","
|
|
101
|
+
records = []
|
|
102
|
+
with open(path, "r", encoding="utf-8", newline="") as f:
|
|
103
|
+
reader = csv.DictReader(f, delimiter=delimiter)
|
|
104
|
+
if reader.fieldnames is None:
|
|
105
|
+
raise DataError(f"'{path}' has no header row / is empty.")
|
|
106
|
+
for row in reader:
|
|
107
|
+
records.append(dict(row))
|
|
108
|
+
return records
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def _extract_text_field(record: Dict[str, Any]) -> Optional[str]:
|
|
112
|
+
for key in _TEXT_FIELD_CANDIDATES:
|
|
113
|
+
if key in record and isinstance(record[key], str):
|
|
114
|
+
return record[key]
|
|
115
|
+
return None
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def _extract_label_field(record: Dict[str, Any]) -> Optional[Any]:
|
|
119
|
+
for key in _LABEL_FIELD_CANDIDATES:
|
|
120
|
+
if key in record:
|
|
121
|
+
return record[key]
|
|
122
|
+
return None
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def _records_to_dataset(records: List[Dict[str, Any]], source: str) -> Dataset:
|
|
126
|
+
if not records:
|
|
127
|
+
raise DataError(f"No records found in '{source}'.")
|
|
128
|
+
|
|
129
|
+
text_field_present = all(_extract_text_field(r) is not None for r in records[:50])
|
|
130
|
+
if text_field_present:
|
|
131
|
+
texts = [_extract_text_field(r) or "" for r in records]
|
|
132
|
+
labels = [_extract_label_field(r) for r in records]
|
|
133
|
+
if any(l is not None for l in labels):
|
|
134
|
+
return Dataset(kind="text_labeled", source=source, texts=texts, labels=labels)
|
|
135
|
+
return Dataset(kind="text", source=source, texts=texts)
|
|
136
|
+
|
|
137
|
+
# Otherwise treat as generic tabular data.
|
|
138
|
+
# Preserve column order as it appears in the data (important: the
|
|
139
|
+
# "use the last column as the target" heuristic in auto/detector.py
|
|
140
|
+
# depends on this being the *original* column order, not alphabetical).
|
|
141
|
+
columns: List[str] = []
|
|
142
|
+
seen = set()
|
|
143
|
+
for r in records:
|
|
144
|
+
for k in r.keys():
|
|
145
|
+
if k not in seen:
|
|
146
|
+
seen.add(k)
|
|
147
|
+
columns.append(k)
|
|
148
|
+
return Dataset(kind="tabular", source=source, records=records, columns=columns)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def _load_directory(path: str) -> Dataset:
|
|
152
|
+
entries = [e for e in sorted(os.listdir(path)) if not e.startswith(".")]
|
|
153
|
+
if not entries:
|
|
154
|
+
raise DataError(f"Directory '{path}' is empty.")
|
|
155
|
+
|
|
156
|
+
subdirs = [e for e in entries if os.path.isdir(os.path.join(path, e))]
|
|
157
|
+
files = [e for e in entries if os.path.isfile(os.path.join(path, e))]
|
|
158
|
+
|
|
159
|
+
# Case 1: class subfolders of text files -> text classification.
|
|
160
|
+
if subdirs and not files:
|
|
161
|
+
texts, labels = [], []
|
|
162
|
+
n_files = 0
|
|
163
|
+
for label in subdirs:
|
|
164
|
+
sub = os.path.join(path, label)
|
|
165
|
+
for fname in sorted(os.listdir(sub)):
|
|
166
|
+
fpath = os.path.join(sub, fname)
|
|
167
|
+
if not os.path.isfile(fpath):
|
|
168
|
+
continue
|
|
169
|
+
ext = os.path.splitext(fname)[1].lower()
|
|
170
|
+
if ext in TEXT_EXTENSIONS:
|
|
171
|
+
texts.append(_read_text_file(fpath))
|
|
172
|
+
labels.append(label)
|
|
173
|
+
n_files += 1
|
|
174
|
+
if not texts:
|
|
175
|
+
raise DataError(
|
|
176
|
+
f"Directory '{path}' contains subfolders but no readable .txt files inside them."
|
|
177
|
+
)
|
|
178
|
+
ds = Dataset(kind="text_labeled", source=path, texts=texts, labels=labels, n_files=n_files)
|
|
179
|
+
return ds
|
|
180
|
+
|
|
181
|
+
# Case 2: flat directory of files -> merge by type.
|
|
182
|
+
all_records: List[Dict[str, Any]] = []
|
|
183
|
+
all_texts: List[str] = []
|
|
184
|
+
n_files = 0
|
|
185
|
+
for fname in files:
|
|
186
|
+
fpath = os.path.join(path, fname)
|
|
187
|
+
ext = os.path.splitext(fname)[1].lower()
|
|
188
|
+
if ext in TEXT_EXTENSIONS:
|
|
189
|
+
all_texts.append(_read_text_file(fpath))
|
|
190
|
+
n_files += 1
|
|
191
|
+
elif ext in JSON_EXTENSIONS:
|
|
192
|
+
all_records.extend(_read_json_records(fpath))
|
|
193
|
+
n_files += 1
|
|
194
|
+
elif ext in JSONL_EXTENSIONS:
|
|
195
|
+
all_records.extend(_read_jsonl_records(fpath))
|
|
196
|
+
n_files += 1
|
|
197
|
+
elif ext in CSV_EXTENSIONS:
|
|
198
|
+
all_records.extend(_read_csv_records(fpath))
|
|
199
|
+
n_files += 1
|
|
200
|
+
# silently skip unknown extensions (e.g. README, .gitkeep) -- but
|
|
201
|
+
# never silently skip *data*-looking files; this is only for
|
|
202
|
+
# incidental non-data files.
|
|
203
|
+
|
|
204
|
+
if all_records and all_texts:
|
|
205
|
+
raise DataError(
|
|
206
|
+
f"Directory '{path}' mixes plain text files with structured "
|
|
207
|
+
f"(json/csv) files. Please keep one data format per directory."
|
|
208
|
+
)
|
|
209
|
+
if all_records:
|
|
210
|
+
ds = _records_to_dataset(all_records, path)
|
|
211
|
+
ds.n_files = n_files
|
|
212
|
+
return ds
|
|
213
|
+
if all_texts:
|
|
214
|
+
return Dataset(kind="text", source=path, texts=all_texts, n_files=n_files)
|
|
215
|
+
|
|
216
|
+
raise DataError(
|
|
217
|
+
f"No supported data files found in '{path}'. Supported: "
|
|
218
|
+
f".txt, .md, .json, .jsonl, .csv, .tsv"
|
|
219
|
+
)
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def load_dataset(path: str) -> Dataset:
|
|
223
|
+
"""Load a dataset from `path` (file or directory) into a `Dataset`."""
|
|
224
|
+
if not os.path.exists(path):
|
|
225
|
+
raise DataError(f"Path '{path}' does not exist.")
|
|
226
|
+
|
|
227
|
+
if os.path.isdir(path):
|
|
228
|
+
return _load_directory(path)
|
|
229
|
+
|
|
230
|
+
ext = os.path.splitext(path)[1].lower()
|
|
231
|
+
if ext in TEXT_EXTENSIONS:
|
|
232
|
+
text = _read_text_file(path)
|
|
233
|
+
if not text.strip():
|
|
234
|
+
raise DataError(f"'{path}' is empty.")
|
|
235
|
+
return Dataset(kind="text", source=path, texts=[text], n_files=1)
|
|
236
|
+
if ext in JSON_EXTENSIONS:
|
|
237
|
+
records = _read_json_records(path)
|
|
238
|
+
ds = _records_to_dataset(records, path)
|
|
239
|
+
ds.n_files = 1
|
|
240
|
+
return ds
|
|
241
|
+
if ext in JSONL_EXTENSIONS:
|
|
242
|
+
records = _read_jsonl_records(path)
|
|
243
|
+
ds = _records_to_dataset(records, path)
|
|
244
|
+
ds.n_files = 1
|
|
245
|
+
return ds
|
|
246
|
+
if ext in CSV_EXTENSIONS:
|
|
247
|
+
records = _read_csv_records(path)
|
|
248
|
+
ds = _records_to_dataset(records, path)
|
|
249
|
+
ds.n_files = 1
|
|
250
|
+
return ds
|
|
251
|
+
|
|
252
|
+
raise DataError(
|
|
253
|
+
f"Unsupported file type '{ext}' for '{path}'. Supported: "
|
|
254
|
+
f".txt, .md, .json, .jsonl, .csv, .tsv, or a directory containing these."
|
|
255
|
+
)
|