taktiny 0.0.1__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.
- taktiny/__init__.py +34 -0
- taktiny/data/__init__.py +54 -0
- taktiny/data/loader.py +435 -0
- taktiny/data/transforms.py +750 -0
- taktiny/nn/__init__.py +22 -0
- taktiny/nn/base.py +683 -0
- taktiny/nn/block.py +689 -0
- taktiny/nn/flatten.py +250 -0
- taktiny/nn/modules/__init__.py +23 -0
- taktiny/nn/modules/activation.py +282 -0
- taktiny/nn/modules/convolution.py +2386 -0
- taktiny/nn/modules/embedding.py +203 -0
- taktiny/nn/modules/linear.py +567 -0
- taktiny/nn/modules/normalization.py +1001 -0
- taktiny/nn/modules/peft.py +951 -0
- taktiny/nn/modules/recurrent.py +1717 -0
- taktiny/nn/modules/transformer.py +1471 -0
- taktiny/nn/regularization.py +498 -0
- taktiny/nn/resampling.py +351 -0
- taktiny/nn/rng.py +208 -0
- taktiny/nn/utils.py +650 -0
- taktiny/py.typed +0 -0
- taktiny/takt/__init__.py +37 -0
- taktiny/takt/adapter/__init__.py +34 -0
- taktiny/takt/adapter/adalora.py +51 -0
- taktiny/takt/adapter/base.py +117 -0
- taktiny/takt/adapter/dora.py +51 -0
- taktiny/takt/adapter/loha.py +51 -0
- taktiny/takt/adapter/lokr.py +51 -0
- taktiny/takt/adapter/lora.py +47 -0
- taktiny/takt/adapter/vera.py +113 -0
- taktiny/takt/base.py +135 -0
- taktiny/takt/optimizer.py +143 -0
- taktiny/trainer/__init__.py +19 -0
- taktiny/trainer/callbacks.py +165 -0
- taktiny/trainer/checkpoint.py +314 -0
- taktiny/trainer/config.py +258 -0
- taktiny/trainer/evaluate.py +254 -0
- taktiny/trainer/trainer.py +1447 -0
- taktiny/utils/__init__.py +21 -0
- taktiny/utils/format.py +113 -0
- taktiny/utils/ops.py +554 -0
- taktiny/utils/quantization.py +480 -0
- taktiny/utils/sharding.py +88 -0
- taktiny/utils/spmd.py +175 -0
- taktiny/utils/trainer.py +352 -0
- taktiny/utils/transforms.py +286 -0
- taktiny/utils/typing.py +146 -0
- taktiny/utils/weights.py +101 -0
- taktiny-0.0.1.dist-info/METADATA +366 -0
- taktiny-0.0.1.dist-info/RECORD +53 -0
- taktiny-0.0.1.dist-info/WHEEL +4 -0
- taktiny-0.0.1.dist-info/licenses/LICENSE.md +202 -0
taktiny/__init__.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
# Copyright 2026 Shinapri
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
from importlib.metadata import version
|
|
16
|
+
|
|
17
|
+
__author__ = "Shinapri"
|
|
18
|
+
__version__ = version('taktiny')
|
|
19
|
+
__description__ = (
|
|
20
|
+
"Build, train, and scale neural networks with JAX."
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
from taktiny.takt import Takt
|
|
24
|
+
from taktiny.takt.optimizer import Optimizer
|
|
25
|
+
from taktiny.utils import typing
|
|
26
|
+
from taktiny.utils.transforms import scan, vmap
|
|
27
|
+
|
|
28
|
+
__all__ = [
|
|
29
|
+
'Takt',
|
|
30
|
+
'vmap',
|
|
31
|
+
'scan',
|
|
32
|
+
'Optimizer',
|
|
33
|
+
'typing'
|
|
34
|
+
]
|
taktiny/data/__init__.py
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
# Copyright 2026 Shinapri
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
"""Generic preprocessing for caller-provided data; no automatic downloads.
|
|
15
|
+
|
|
16
|
+
Use DataLoader with composable operations, or call Map/Compose/MapFields on
|
|
17
|
+
individual examples. Text-specific helpers also remain available from .text.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from taktiny.data.loader import DataLoader, RandomAccessSource, train_validation_split
|
|
21
|
+
from taktiny.data.transforms import (
|
|
22
|
+
ApplyTemplate,
|
|
23
|
+
Batch,
|
|
24
|
+
BatchMap,
|
|
25
|
+
Compose,
|
|
26
|
+
Filter,
|
|
27
|
+
FlatMap,
|
|
28
|
+
IndexMap,
|
|
29
|
+
Map,
|
|
30
|
+
MapFields,
|
|
31
|
+
Pack,
|
|
32
|
+
RandomMap,
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
__all__ = [
|
|
36
|
+
'ApplyTemplate',
|
|
37
|
+
'Batch',
|
|
38
|
+
'BatchMap',
|
|
39
|
+
'CausalLMBatch',
|
|
40
|
+
'Compose',
|
|
41
|
+
'DataLoader',
|
|
42
|
+
'DatasetUtils',
|
|
43
|
+
'Filter',
|
|
44
|
+
'FlatMap',
|
|
45
|
+
'IndexMap',
|
|
46
|
+
'Map',
|
|
47
|
+
'MapFields',
|
|
48
|
+
'Pack',
|
|
49
|
+
'PackSequences',
|
|
50
|
+
'RandomAccessSource',
|
|
51
|
+
'RandomMap',
|
|
52
|
+
'tokenize',
|
|
53
|
+
'train_validation_split'
|
|
54
|
+
]
|
taktiny/data/loader.py
ADDED
|
@@ -0,0 +1,435 @@
|
|
|
1
|
+
# Copyright 2026 Shinapri
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
"""Load already-available, random-access data without prescribing a modality."""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import sys
|
|
19
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
20
|
+
from dataclasses import replace
|
|
21
|
+
from typing import Any, Protocol
|
|
22
|
+
|
|
23
|
+
import grain.python as grain
|
|
24
|
+
import jax
|
|
25
|
+
import numpy as np
|
|
26
|
+
from absl import flags
|
|
27
|
+
from absl.flags import UnparsedFlagAccessError
|
|
28
|
+
from grain._src.python.dataset import base as dataset_base
|
|
29
|
+
from jax.sharding import NamedSharding, PartitionSpec
|
|
30
|
+
|
|
31
|
+
from taktiny.data.transforms import Batch, _expand_operations
|
|
32
|
+
from taktiny.utils.spmd import logical_to_mesh_axes
|
|
33
|
+
from taktiny.utils.typing import AxisNames
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class _PlacedIterator:
|
|
37
|
+
def __init__(self, parent: Any, axis_names: Any, specs: Any, mesh: Any, default_spec: Any):
|
|
38
|
+
self._parent = parent
|
|
39
|
+
self._axis_names = axis_names
|
|
40
|
+
self._specs = specs
|
|
41
|
+
self._mesh = mesh
|
|
42
|
+
self._default_spec = default_spec
|
|
43
|
+
|
|
44
|
+
def __iter__(self):
|
|
45
|
+
return self
|
|
46
|
+
|
|
47
|
+
def __getattr__(self, name: str) -> Any:
|
|
48
|
+
return getattr(self._parent, name)
|
|
49
|
+
|
|
50
|
+
def __next__(self):
|
|
51
|
+
def place(value, names, spec):
|
|
52
|
+
if isinstance(names, Mapping) or isinstance(spec, Mapping):
|
|
53
|
+
if not isinstance(value, Mapping):
|
|
54
|
+
raise TypeError('Field sharding specifications require a mapping batch')
|
|
55
|
+
for config in (names, spec):
|
|
56
|
+
if isinstance(config, Mapping) and config.keys() - value.keys():
|
|
57
|
+
raise ValueError('Sharding configuration contains unknown batch fields')
|
|
58
|
+
return {key: place(item,
|
|
59
|
+
names.get(key) if isinstance(names, Mapping) else names,
|
|
60
|
+
spec.get(key, self._default_spec) if isinstance(spec, Mapping) else spec)
|
|
61
|
+
for key, item in value.items()}
|
|
62
|
+
if spec is None:
|
|
63
|
+
return value
|
|
64
|
+
|
|
65
|
+
def put(array):
|
|
66
|
+
array = array if isinstance(array, jax.Array) else np.asarray(array)
|
|
67
|
+
if names is not None and len(names) != array.ndim:
|
|
68
|
+
raise ValueError('axis_names length must match batch array ndim')
|
|
69
|
+
return jax.device_put(array, NamedSharding(self._mesh, spec))
|
|
70
|
+
|
|
71
|
+
return jax.tree.map(put, value)
|
|
72
|
+
|
|
73
|
+
return place(next(self._parent), self._axis_names, self._specs)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class RandomAccessSource(Protocol):
|
|
77
|
+
"""Structural interface for finite sources; no framework inheritance needed."""
|
|
78
|
+
|
|
79
|
+
def __len__(self) -> int: ...
|
|
80
|
+
|
|
81
|
+
def __getitem__(self, index: int, /) -> Any: ...
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
class _ShardSampler:
|
|
85
|
+
"""Present a shard as a local sampler, including uneven and empty tails.
|
|
86
|
+
|
|
87
|
+
Grain 0.2.18 DataLoader floors sampler length / shard_count even when
|
|
88
|
+
drop_remainder=False. Do the partitioning here and disable its second
|
|
89
|
+
partitioning step. Keep global RNGs but expose local traversal indices.
|
|
90
|
+
"""
|
|
91
|
+
|
|
92
|
+
def __init__(
|
|
93
|
+
self, sampler: grain.IndexSampler | None,
|
|
94
|
+
shard_index: int, shard_count: int,
|
|
95
|
+
) -> None:
|
|
96
|
+
self.sampler = sampler
|
|
97
|
+
self.shard_index = shard_index
|
|
98
|
+
self.shard_count = shard_count
|
|
99
|
+
self._shard_options = grain.NoSharding()
|
|
100
|
+
|
|
101
|
+
def __len__(self) -> int:
|
|
102
|
+
size = len(self.sampler) if self.sampler is not None else 0
|
|
103
|
+
if size == sys.maxsize:
|
|
104
|
+
return sys.maxsize
|
|
105
|
+
return max(0, (size - self.shard_index + self.shard_count - 1) // self.shard_count)
|
|
106
|
+
|
|
107
|
+
def __getitem__(self, index: int) -> grain.RecordMetadata:
|
|
108
|
+
if index < 0 or index >= len(self) or self.sampler is None:
|
|
109
|
+
raise IndexError(index)
|
|
110
|
+
record = self.sampler[index * self.shard_count + self.shard_index]
|
|
111
|
+
return replace(record, index=index)
|
|
112
|
+
|
|
113
|
+
def __repr__(self) -> str:
|
|
114
|
+
return (f'_ShardSampler({self.sampler!r}, shard_index={self.shard_index}, '
|
|
115
|
+
f'shard_count={self.shard_count})')
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def _prepare_grain_workers() -> None:
|
|
119
|
+
"""Allow Grain workers to start when Abseil flags are still unparsed."""
|
|
120
|
+
if flags.FLAGS.is_parsed():
|
|
121
|
+
return
|
|
122
|
+
|
|
123
|
+
# Grain 0.2.18 reads this FlagHolder while constructing its worker pool.
|
|
124
|
+
# Reading a holder before absl.app.run() raises in notebooks and regular
|
|
125
|
+
# Python programs. The underlying Flag exposes the same live value without
|
|
126
|
+
# requiring TakTiny to parse the application's complete flag registry.
|
|
127
|
+
from grain._src.core import profiler
|
|
128
|
+
|
|
129
|
+
name = '_GRAIN_ENABLE_MULTIPROCESS_WORKER_PROFILING'
|
|
130
|
+
holder = getattr(profiler, name, None)
|
|
131
|
+
if holder is None:
|
|
132
|
+
return
|
|
133
|
+
try:
|
|
134
|
+
holder.value # noqa: B018
|
|
135
|
+
except UnparsedFlagAccessError:
|
|
136
|
+
setattr(profiler, name, flags.FLAGS[holder.name])
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
class DataLoader(grain.DataLoader):
|
|
140
|
+
"""Preprocess and iterate over caller-provided, random-access records.
|
|
141
|
+
|
|
142
|
+
Records may be arrays, mappings, tuples, images, audio, strings, or custom
|
|
143
|
+
objects. The source is never downloaded, decoded, or copied into memory.
|
|
144
|
+
Supply a list, array, already-loaded dataset, or an object implementing
|
|
145
|
+
__len__ and integer __getitem__. This loader does not accept streaming
|
|
146
|
+
generators; use a streaming backend directly, or explicitly materialize a
|
|
147
|
+
finite stream with list(source) if it fits in memory.
|
|
148
|
+
|
|
149
|
+
Args:
|
|
150
|
+
source: Caller-owned random-access data, not a repository ID or path.
|
|
151
|
+
operations: Ordered Taktiny or native Grain operations. Use Map for a
|
|
152
|
+
record callable, RandomMap for augmentation, and Filter to drop
|
|
153
|
+
records. Operations are lazy and run when the loader is iterated.
|
|
154
|
+
batch_size: Optional final batch size, applied after all operations.
|
|
155
|
+
None emits records unchanged. For operations after batching, put
|
|
156
|
+
Batch directly in operations and leave this argument unset.
|
|
157
|
+
drop_remainder: Drop an incomplete final batch; requires batch_size.
|
|
158
|
+
collate_fn: Optional function(rows) for final batching. None uses Grain
|
|
159
|
+
stacking; list preserves ragged or custom objects without padding.
|
|
160
|
+
Requires batch_size. Collation runs independently in each worker.
|
|
161
|
+
sampler: Optional native Grain sampler. When provided, it owns sampling,
|
|
162
|
+
epochs, seed, and sharding; the corresponding convenience arguments
|
|
163
|
+
are ignored after validation.
|
|
164
|
+
shuffle: Shuffle indices using seed; defaults to False.
|
|
165
|
+
seed: Unsigned 32-bit integer seed for sampling and RandomMap augmentation.
|
|
166
|
+
num_epochs: Positive epoch count (default 1); None repeats indefinitely.
|
|
167
|
+
Each new iterator starts from the beginning unless state is restored.
|
|
168
|
+
shard_index: This process's data shard, in [0, shard_count).
|
|
169
|
+
shard_count: Number of data shards. Shards may have unequal lengths;
|
|
170
|
+
this is input partitioning, not JAX device-array sharding.
|
|
171
|
+
worker_count: Child workers; 0 runs locally, None lets Grain choose.
|
|
172
|
+
Sources and transforms must be serializable when workers are used.
|
|
173
|
+
worker_buffer_size: Positive per-worker prefetch buffer size.
|
|
174
|
+
read_options: Grain reader settings. None uses Grain's defaults. For an
|
|
175
|
+
in-memory source, ``grain.ReadOptions(num_threads=0,
|
|
176
|
+
prefetch_buffer_size=0)`` avoids threaded record prefetching.
|
|
177
|
+
axis_names: Logical axis names for output arrays, or a mapping from
|
|
178
|
+
batch fields to axis names. Names describe the final batched rank
|
|
179
|
+
and override partition_spec for that field, using logical rules
|
|
180
|
+
active when the loader is constructed.
|
|
181
|
+
partition_spec: Explicit output PartitionSpec, or a mapping from batch
|
|
182
|
+
fields to specs. Placement uses the mesh active when iter(loader)
|
|
183
|
+
is called and occurs after Grain produces each batch. An active
|
|
184
|
+
mesh is required when a spec is resolved. Fields without names or
|
|
185
|
+
a spec are unchanged. Neither argument changes record sampling.
|
|
186
|
+
|
|
187
|
+
Iterators retain Grain's get_state()/set_state() checkpoint API. Restore
|
|
188
|
+
against the same source and pipeline. Custom iterator operations retain
|
|
189
|
+
their own Grain checkpoint limitations. Transforms should be deterministic
|
|
190
|
+
apart from RandomMap's supplied RNG and should not mutate source records.
|
|
191
|
+
|
|
192
|
+
Example:
|
|
193
|
+
>>> from taktiny.data import DataLoader, MapFields
|
|
194
|
+
>>> rows = [{'value': 255, 'label': 0}, {'value': 0, 'label': 1}]
|
|
195
|
+
>>> loader = DataLoader(rows, operations=[
|
|
196
|
+
... MapFields({'value': lambda x: x / 255})], batch_size=2)
|
|
197
|
+
>>> next(iter(loader))['value'].tolist()
|
|
198
|
+
[1.0, 0.0]
|
|
199
|
+
"""
|
|
200
|
+
def __init__(
|
|
201
|
+
self,
|
|
202
|
+
source: dataset_base.RandomAccessDataSource | Any,
|
|
203
|
+
*,
|
|
204
|
+
operations: Sequence[Any] = (),
|
|
205
|
+
batch_size: int | None = None,
|
|
206
|
+
drop_remainder: bool = False,
|
|
207
|
+
collate_fn: Callable[[Sequence[Any]], Any] | None = None,
|
|
208
|
+
sampler: grain.Sampler | None = None,
|
|
209
|
+
shuffle: bool = False,
|
|
210
|
+
seed: int = 0,
|
|
211
|
+
num_epochs: int | None = 1,
|
|
212
|
+
shard_index: int = 0,
|
|
213
|
+
shard_count: int = 1,
|
|
214
|
+
worker_count: int | None = 0,
|
|
215
|
+
worker_buffer_size: int = 1,
|
|
216
|
+
read_options: grain.ReadOptions | None = None,
|
|
217
|
+
axis_names: AxisNames | Mapping[str, AxisNames | None] | None = None,
|
|
218
|
+
partition_spec: PartitionSpec | Mapping[str, PartitionSpec | None] | None = None,
|
|
219
|
+
) -> None:
|
|
220
|
+
"""Create a Grain loader from a random-access dataset.
|
|
221
|
+
|
|
222
|
+
``operations`` are applied exactly in the supplied order. Mapping,
|
|
223
|
+
filtering, packing, batching, and collation therefore remain separate
|
|
224
|
+
concerns and can be composed using Grain transformations or custom
|
|
225
|
+
Grain operations.
|
|
226
|
+
|
|
227
|
+
When ``sampler`` is omitted, an :class:`grain.IndexSampler` is created
|
|
228
|
+
from the remaining sampling arguments. Supplying ``sampler`` transfers
|
|
229
|
+
sampling and sharding responsibility entirely to that object. The
|
|
230
|
+
default ``num_epochs=1`` creates a single-epoch (finite) loader; pass
|
|
231
|
+
``None`` for an unbounded loader.
|
|
232
|
+
"""
|
|
233
|
+
_validate_source(source)
|
|
234
|
+
self.axis_names = axis_names
|
|
235
|
+
|
|
236
|
+
def resolve(names, spec):
|
|
237
|
+
if isinstance(names, Mapping) or isinstance(spec, Mapping):
|
|
238
|
+
keys = set(names if isinstance(names, Mapping) else ())
|
|
239
|
+
keys.update(spec if isinstance(spec, Mapping) else ())
|
|
240
|
+
return {key: resolve(names.get(key) if isinstance(names, Mapping) else names,
|
|
241
|
+
spec.get(key) if isinstance(spec, Mapping) else spec)
|
|
242
|
+
for key in keys}
|
|
243
|
+
if spec is not None and not isinstance(spec, PartitionSpec):
|
|
244
|
+
raise TypeError('partition_spec must contain PartitionSpec values')
|
|
245
|
+
if names is not None:
|
|
246
|
+
if isinstance(names, (str, PartitionSpec)):
|
|
247
|
+
raise TypeError('axis_names must contain logical axis-name tuples')
|
|
248
|
+
return logical_to_mesh_axes(names)
|
|
249
|
+
return spec
|
|
250
|
+
|
|
251
|
+
self.partition_spec = partition_spec
|
|
252
|
+
self._resolved_specs = resolve(axis_names, partition_spec)
|
|
253
|
+
|
|
254
|
+
if operations is None or isinstance(operations, (str, bytes)):
|
|
255
|
+
raise TypeError('operations must be a sequence')
|
|
256
|
+
|
|
257
|
+
try:
|
|
258
|
+
operations = tuple(operations)
|
|
259
|
+
except TypeError as error:
|
|
260
|
+
raise TypeError('operations must be a sequence') from error
|
|
261
|
+
|
|
262
|
+
operations = _expand_operations(operations)
|
|
263
|
+
if not isinstance(drop_remainder, bool):
|
|
264
|
+
raise TypeError('drop_remainder must be a boolean')
|
|
265
|
+
|
|
266
|
+
if batch_size is None:
|
|
267
|
+
if drop_remainder or collate_fn is not None:
|
|
268
|
+
raise ValueError('drop_remainder and collate_fn require batch_size')
|
|
269
|
+
else:
|
|
270
|
+
operations += (Batch(batch_size, drop_remainder=drop_remainder,
|
|
271
|
+
collate_fn=collate_fn),)
|
|
272
|
+
|
|
273
|
+
if not isinstance(shuffle, bool):
|
|
274
|
+
raise TypeError('shuffle must be a boolean')
|
|
275
|
+
|
|
276
|
+
if isinstance(seed, bool) or not isinstance(seed, int):
|
|
277
|
+
raise TypeError('seed must be an integer')
|
|
278
|
+
|
|
279
|
+
if not 0 <= seed < 2**32:
|
|
280
|
+
raise ValueError('seed must be an unsigned 32-bit integer')
|
|
281
|
+
|
|
282
|
+
if (
|
|
283
|
+
num_epochs is not None
|
|
284
|
+
and (
|
|
285
|
+
isinstance(num_epochs, bool)
|
|
286
|
+
or not isinstance(num_epochs, int)
|
|
287
|
+
or num_epochs < 1
|
|
288
|
+
)
|
|
289
|
+
):
|
|
290
|
+
raise ValueError('num_epochs must be a positive integer or None')
|
|
291
|
+
|
|
292
|
+
if (
|
|
293
|
+
isinstance(shard_count, bool)
|
|
294
|
+
or not isinstance(shard_count, int)
|
|
295
|
+
or shard_count < 1
|
|
296
|
+
):
|
|
297
|
+
raise ValueError('shard_count must be a positive integer')
|
|
298
|
+
|
|
299
|
+
if (
|
|
300
|
+
isinstance(shard_index, bool)
|
|
301
|
+
or not isinstance(shard_index, int)
|
|
302
|
+
or not 0 <= shard_index < shard_count
|
|
303
|
+
):
|
|
304
|
+
raise ValueError(
|
|
305
|
+
'shard_index must be between zero and shard_count - 1'
|
|
306
|
+
)
|
|
307
|
+
|
|
308
|
+
if (
|
|
309
|
+
worker_count is not None
|
|
310
|
+
and (
|
|
311
|
+
isinstance(worker_count, bool)
|
|
312
|
+
or not isinstance(worker_count, int)
|
|
313
|
+
or worker_count < 0
|
|
314
|
+
)
|
|
315
|
+
):
|
|
316
|
+
raise ValueError('worker_count must be non-negative or None')
|
|
317
|
+
|
|
318
|
+
if (
|
|
319
|
+
isinstance(worker_buffer_size, bool)
|
|
320
|
+
or not isinstance(worker_buffer_size, int)
|
|
321
|
+
or worker_buffer_size < 1
|
|
322
|
+
):
|
|
323
|
+
raise ValueError('worker_buffer_size must be a positive integer')
|
|
324
|
+
|
|
325
|
+
if sampler is None:
|
|
326
|
+
try:
|
|
327
|
+
num_records = len(source)
|
|
328
|
+
except TypeError as error:
|
|
329
|
+
raise TypeError(
|
|
330
|
+
'source must have a finite length when sampler is omitted'
|
|
331
|
+
) from error
|
|
332
|
+
|
|
333
|
+
base_sampler = grain.IndexSampler(
|
|
334
|
+
num_records=num_records,
|
|
335
|
+
num_epochs=num_epochs,
|
|
336
|
+
shard_options=grain.NoSharding(),
|
|
337
|
+
shuffle=shuffle,
|
|
338
|
+
seed=seed,
|
|
339
|
+
) if num_records else None
|
|
340
|
+
sampler = _ShardSampler(base_sampler, shard_index, shard_count)
|
|
341
|
+
|
|
342
|
+
if worker_count is None or worker_count > 0:
|
|
343
|
+
_prepare_grain_workers()
|
|
344
|
+
|
|
345
|
+
super().__init__(
|
|
346
|
+
data_source=source,
|
|
347
|
+
sampler=sampler,
|
|
348
|
+
operations=operations,
|
|
349
|
+
worker_count=worker_count,
|
|
350
|
+
worker_buffer_size=worker_buffer_size,
|
|
351
|
+
read_options=read_options,
|
|
352
|
+
)
|
|
353
|
+
|
|
354
|
+
def __iter__(self):
|
|
355
|
+
if self.axis_names is None and self.partition_spec is None:
|
|
356
|
+
return super().__iter__()
|
|
357
|
+
mesh = jax.sharding.get_mesh()
|
|
358
|
+
if mesh.empty:
|
|
359
|
+
raise ValueError('DataLoader output sharding requires an active JAX mesh')
|
|
360
|
+
default_spec = self.partition_spec if isinstance(self.partition_spec, PartitionSpec) else None
|
|
361
|
+
return _PlacedIterator(super().__iter__(), self.axis_names, self._resolved_specs, mesh, default_spec)
|
|
362
|
+
|
|
363
|
+
|
|
364
|
+
|
|
365
|
+
def _validate_source(source: Any) -> None:
|
|
366
|
+
if isinstance(source, (str, bytes, Mapping)):
|
|
367
|
+
raise TypeError('source must contain records, not a path, repository ID, or column mapping')
|
|
368
|
+
if not hasattr(source, '__getitem__'):
|
|
369
|
+
raise TypeError('source must support random access; materialize finite iterables explicitly')
|
|
370
|
+
|
|
371
|
+
|
|
372
|
+
def train_validation_split(
|
|
373
|
+
source: RandomAccessSource,
|
|
374
|
+
validation_size: float,
|
|
375
|
+
*,
|
|
376
|
+
shuffle: bool = True,
|
|
377
|
+
seed: int = 0,
|
|
378
|
+
) -> tuple[Any, Any]:
|
|
379
|
+
"""Split a random-access source into ``(train, validation)`` views.
|
|
380
|
+
|
|
381
|
+
``validation_size`` may be a count (``int``) or a fraction (``float``
|
|
382
|
+
in ``(0, 1)``). The returned views are random-access and can be passed
|
|
383
|
+
directly to DataLoader. Fractions are rounded to the nearest integer;
|
|
384
|
+
both splits must be nonempty. The source itself is not copied or shuffled.
|
|
385
|
+
"""
|
|
386
|
+
_validate_source(source)
|
|
387
|
+
if not isinstance(shuffle, bool):
|
|
388
|
+
raise TypeError('shuffle must be a boolean')
|
|
389
|
+
if isinstance(seed, bool) or not isinstance(seed, int):
|
|
390
|
+
raise TypeError('seed must be an integer')
|
|
391
|
+
|
|
392
|
+
count = len(source)
|
|
393
|
+
if isinstance(validation_size, bool):
|
|
394
|
+
raise TypeError('validation_size must be an int or float')
|
|
395
|
+
if isinstance(validation_size, float):
|
|
396
|
+
if not 0.0 < validation_size < 1.0:
|
|
397
|
+
raise ValueError('float validation_size must be in (0, 1)')
|
|
398
|
+
validation_count = round(count * validation_size)
|
|
399
|
+
elif isinstance(validation_size, int):
|
|
400
|
+
if not 0 < validation_size < count:
|
|
401
|
+
raise ValueError(
|
|
402
|
+
f'int validation_size must be in (0, {count})'
|
|
403
|
+
)
|
|
404
|
+
validation_count = validation_size
|
|
405
|
+
else:
|
|
406
|
+
raise TypeError('validation_size must be an int or float')
|
|
407
|
+
|
|
408
|
+
if not 0 < validation_count < count:
|
|
409
|
+
raise ValueError('validation_size must leave both splits nonempty')
|
|
410
|
+
|
|
411
|
+
indices = np.arange(count)
|
|
412
|
+
if shuffle:
|
|
413
|
+
indices = np.random.default_rng(seed).permutation(count)
|
|
414
|
+
train = _IndexedView(source, indices[validation_count:])
|
|
415
|
+
validation = _IndexedView(source, indices[:validation_count])
|
|
416
|
+
return train, validation
|
|
417
|
+
|
|
418
|
+
|
|
419
|
+
class _IndexedView:
|
|
420
|
+
"""Random-access view over a subset of a source's indices."""
|
|
421
|
+
|
|
422
|
+
def __init__(self, source: RandomAccessSource, indices: np.ndarray) -> None:
|
|
423
|
+
self._source = source
|
|
424
|
+
self._indices = indices
|
|
425
|
+
|
|
426
|
+
def __len__(self) -> int:
|
|
427
|
+
return len(self._indices)
|
|
428
|
+
|
|
429
|
+
def __getitem__(self, index: int | slice) -> Any:
|
|
430
|
+
if isinstance(index, slice):
|
|
431
|
+
return _IndexedView(self._source, self._indices[index])
|
|
432
|
+
return self._source[int(self._indices[index])]
|
|
433
|
+
|
|
434
|
+
|
|
435
|
+
__all__ = ['DataLoader', 'RandomAccessSource', 'train_validation_split']
|