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/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,11 @@
1
+ from .loader import load_dataset, Dataset
2
+ from .fingerprint import fingerprint_path
3
+ from .inspector import inspect_path, InspectionReport
4
+
5
+ __all__ = [
6
+ "load_dataset",
7
+ "Dataset",
8
+ "fingerprint_path",
9
+ "inspect_path",
10
+ "InspectionReport",
11
+ ]
@@ -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
+ )