pytensorforge 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.
- cli.py +604 -0
- pytensorforge-0.1.0.dist-info/METADATA +103 -0
- pytensorforge-0.1.0.dist-info/RECORD +146 -0
- pytensorforge-0.1.0.dist-info/WHEEL +5 -0
- pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
- pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
- src/__init__.py +0 -0
- src/activations/Activation.py +4 -0
- src/activations/ELU.py +11 -0
- src/activations/GELU.py +6 -0
- src/activations/ReLU.py +27 -0
- src/activations/SELU.py +14 -0
- src/activations/Sigmoid.py +27 -0
- src/activations/Softmax.py +84 -0
- src/activations/Tanh.py +29 -0
- src/activations/__init__.py +17 -0
- src/config.py +120 -0
- src/core/Matrix.py +3 -0
- src/core/Scalar.py +18 -0
- src/core/Tensor.py +866 -0
- src/core/Vector.py +31 -0
- src/core/__init__.py +0 -0
- src/data/__init__.py +0 -0
- src/data/chat_dataset.py +188 -0
- src/data/corpus.py +104 -0
- src/data/document_stream.py +178 -0
- src/data/parallel_encode.py +86 -0
- src/data/prefetch.py +62 -0
- src/data/shard_builder.py +119 -0
- src/data/shard_writer.py +81 -0
- src/data/sharded_dataset.py +112 -0
- src/data/streaming_dataset.py +132 -0
- src/data/validation.py +212 -0
- src/inference/__init__.py +0 -0
- src/inference/chat_template.py +384 -0
- src/inference/config.py +48 -0
- src/inference/engine.py +241 -0
- src/inference/export.py +133 -0
- src/inference/kv_cache.py +65 -0
- src/inference/runtime.py +161 -0
- src/inference/sampling.py +42 -0
- src/inference/scheduler.py +473 -0
- src/inference/text.py +67 -0
- src/initializers/Constant.py +9 -0
- src/initializers/GlorotNormal.py +15 -0
- src/initializers/GlorotUniform.py +26 -0
- src/initializers/HeNormal.py +15 -0
- src/initializers/HeUniform.py +14 -0
- src/initializers/Initializer.py +4 -0
- src/initializers/LecunNormal.py +16 -0
- src/initializers/LecunUniform.py +14 -0
- src/initializers/Ones.py +6 -0
- src/initializers/Orthogonal.py +14 -0
- src/initializers/RandomNormal.py +14 -0
- src/initializers/RandomUniform.py +14 -0
- src/initializers/Zeros.py +8 -0
- src/initializers/__init__.py +17 -0
- src/loss/CategoricalCrossEntropy.py +9 -0
- src/loss/CrossEntropyLoss.py +34 -0
- src/loss/CrossEntropyWithLogitsLoss.py +59 -0
- src/loss/Hinge.py +5 -0
- src/loss/Huber.py +22 -0
- src/loss/Loss.py +6 -0
- src/loss/MSE.py +7 -0
- src/loss/MSELoss.py +10 -0
- src/loss/SparseCategoricalCrossEntropy.py +15 -0
- src/loss/__init__.py +18 -0
- src/loss/bce.py +34 -0
- src/loss/mae.py +16 -0
- src/math/__init__.py +0 -0
- src/math/clip.py +37 -0
- src/math/exp.py +27 -0
- src/math/log.py +25 -0
- src/math/sigmoid.py +5 -0
- src/models/__init__.py +0 -0
- src/models/embedding/Embedding.py +65 -0
- src/models/embedding/__init__.py +0 -0
- src/models/gpt/__init__.py +0 -0
- src/models/gpt/attention.py +158 -0
- src/models/gpt/block.py +74 -0
- src/models/gpt/config.py +103 -0
- src/models/gpt/context.py +44 -0
- src/models/gpt/model.py +165 -0
- src/models/gpt/recompute.py +35 -0
- src/models/gpt/rope.py +84 -0
- src/models/regression/Linear.py +51 -0
- src/models/regression/Logistic.py +36 -0
- src/models/regression/__init__.py +0 -0
- src/models/seq/Sequential.py +297 -0
- src/models/seq/__init__.py +0 -0
- src/models/svm/__init__.py +0 -0
- src/models/tokenizer/BPETokenizer.py +228 -0
- src/models/tokenizer/__init__.py +0 -0
- src/models/transformers/Dropout.py +35 -0
- src/models/transformers/LastToken.py +10 -0
- src/models/transformers/LayerNorm.py +54 -0
- src/models/transformers/Linear.py +18 -0
- src/models/transformers/MultiHeadAttention.py +130 -0
- src/models/transformers/TransformerBlock.py +79 -0
- src/models/transformers/__init__.py +0 -0
- src/neural/Dense.py +58 -0
- src/neural/LSTM.py +167 -0
- src/neural/Layer.py +72 -0
- src/neural/Parameter.py +30 -0
- src/neural/RNN.py +83 -0
- src/neural/__init__.py +0 -0
- src/ops/__init__.py +0 -0
- src/ops/stack.py +40 -0
- src/optimizers/Adagrad.py +31 -0
- src/optimizers/Adam.py +98 -0
- src/optimizers/AdamW.py +84 -0
- src/optimizers/Batch.py +11 -0
- src/optimizers/Nesterov.py +35 -0
- src/optimizers/Optimizer.py +18 -0
- src/optimizers/RMSProp.py +35 -0
- src/optimizers/SGD.py +30 -0
- src/optimizers/SGDMomentum.py +28 -0
- src/optimizers/__init__.py +9 -0
- src/scaling/StandardScaler.py +15 -0
- src/scaling/__init__.py +0 -0
- src/serialization/__init__.py +0 -0
- src/serialization/checkpoint.py +58 -0
- src/serialization/modelio.py +132 -0
- src/serving/__init__.py +0 -0
- src/serving/app.py +792 -0
- src/serving/config.py +216 -0
- src/serving/errors.py +51 -0
- src/serving/http.py +599 -0
- src/serving/metrics.py +293 -0
- src/serving/model_server.py +287 -0
- src/serving/protocol.py +377 -0
- src/serving/security.py +200 -0
- src/serving/server.py +121 -0
- src/tokenization/__init__.py +0 -0
- src/tokenization/base.py +75 -0
- src/tokenization/bpe.py +190 -0
- src/tokenization/bytebpe.py +476 -0
- src/tokenization/registry.py +28 -0
- src/training/__init__.py +0 -0
- src/training/checkpoint_manager.py +101 -0
- src/training/experiment.py +71 -0
- src/training/losses.py +42 -0
- src/training/precision.py +141 -0
- src/training/profiler.py +38 -0
- src/training/scheduler.py +50 -0
- src/training/trainer.py +594 -0
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import json
|
|
3
|
+
import os
|
|
4
|
+
import time
|
|
5
|
+
from multiprocessing import get_context
|
|
6
|
+
|
|
7
|
+
from src.data.document_stream import DocumentReader
|
|
8
|
+
from src.data.shard_writer import SHARD_FORMAT_VERSION, ShardWriter, dtype_for_vocab
|
|
9
|
+
from src.tokenization.registry import load_tokenizer
|
|
10
|
+
|
|
11
|
+
MANIFEST_NAME = "manifest.json"
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def corpus_version(files):
|
|
15
|
+
digest = hashlib.sha256()
|
|
16
|
+
|
|
17
|
+
for path in files:
|
|
18
|
+
stat = os.stat(path)
|
|
19
|
+
digest.update(f"{path}:{stat.st_size}:{int(stat.st_mtime)}".encode())
|
|
20
|
+
|
|
21
|
+
return digest.hexdigest()
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _partition(files, workers):
|
|
25
|
+
workers = max(1, min(workers, len(files)))
|
|
26
|
+
size, extra = divmod(len(files), workers)
|
|
27
|
+
parts = []
|
|
28
|
+
start = 0
|
|
29
|
+
|
|
30
|
+
for i in range(workers):
|
|
31
|
+
end = start + size + (1 if i < extra else 0)
|
|
32
|
+
parts.append(files[start:end])
|
|
33
|
+
start = end
|
|
34
|
+
|
|
35
|
+
return parts
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _worker(args):
|
|
39
|
+
worker_id, files, tokenizer_path, output_dir, shard_tokens, insert_eos, text_field, read_buffer_size = args
|
|
40
|
+
|
|
41
|
+
tokenizer = load_tokenizer(tokenizer_path)
|
|
42
|
+
writer = ShardWriter(
|
|
43
|
+
output_dir,
|
|
44
|
+
f"shard-w{worker_id:03d}",
|
|
45
|
+
shard_tokens,
|
|
46
|
+
dtype_for_vocab(tokenizer.vocab_size),
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
reader = DocumentReader(files, read_buffer_size=read_buffer_size, text_field=text_field)
|
|
50
|
+
eos = tokenizer.eos_id
|
|
51
|
+
|
|
52
|
+
for _, _, text, doc_end in reader.iter_documents():
|
|
53
|
+
if text:
|
|
54
|
+
writer.write(tokenizer.encode(text))
|
|
55
|
+
|
|
56
|
+
if doc_end and insert_eos and eos is not None:
|
|
57
|
+
writer.write([eos])
|
|
58
|
+
|
|
59
|
+
return writer.close()
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def build_shards(
|
|
63
|
+
corpus,
|
|
64
|
+
tokenizer_path,
|
|
65
|
+
output_dir,
|
|
66
|
+
shard_tokens=50_000_000,
|
|
67
|
+
workers=1,
|
|
68
|
+
insert_eos=True,
|
|
69
|
+
text_field="text",
|
|
70
|
+
read_buffer_size=1 << 20,
|
|
71
|
+
):
|
|
72
|
+
if not corpus.files:
|
|
73
|
+
raise ValueError("corpus contains no files")
|
|
74
|
+
|
|
75
|
+
tokenizer = load_tokenizer(tokenizer_path)
|
|
76
|
+
|
|
77
|
+
os.makedirs(output_dir, exist_ok=True)
|
|
78
|
+
|
|
79
|
+
parts = _partition(corpus.files, workers)
|
|
80
|
+
|
|
81
|
+
jobs = [
|
|
82
|
+
(i, part, tokenizer_path, output_dir, shard_tokens, insert_eos, text_field, read_buffer_size)
|
|
83
|
+
for i, part in enumerate(parts)
|
|
84
|
+
]
|
|
85
|
+
|
|
86
|
+
started = time.time()
|
|
87
|
+
|
|
88
|
+
if len(jobs) == 1:
|
|
89
|
+
results = [_worker(jobs[0])]
|
|
90
|
+
else:
|
|
91
|
+
with get_context("spawn").Pool(len(jobs)) as pool:
|
|
92
|
+
results = pool.map(_worker, jobs)
|
|
93
|
+
|
|
94
|
+
shards = [shard for result in results for shard in result]
|
|
95
|
+
|
|
96
|
+
manifest = {
|
|
97
|
+
"format_version": SHARD_FORMAT_VERSION,
|
|
98
|
+
"tokenizer_identity": tokenizer.identity,
|
|
99
|
+
"vocab_size": tokenizer.vocab_size,
|
|
100
|
+
"dtype": dtype_for_vocab(tokenizer.vocab_size).name,
|
|
101
|
+
"insert_eos": insert_eos,
|
|
102
|
+
"corpus_version": corpus_version(corpus.files),
|
|
103
|
+
"source_files": list(corpus.files),
|
|
104
|
+
"total_tokens": sum(s["token_count"] for s in shards),
|
|
105
|
+
"shards": shards,
|
|
106
|
+
"created_at": time.time(),
|
|
107
|
+
"build_seconds": time.time() - started,
|
|
108
|
+
}
|
|
109
|
+
|
|
110
|
+
tmp = os.path.join(output_dir, MANIFEST_NAME + ".tmp")
|
|
111
|
+
|
|
112
|
+
with open(tmp, "w") as f:
|
|
113
|
+
json.dump(manifest, f, indent=2)
|
|
114
|
+
f.flush()
|
|
115
|
+
os.fsync(f.fileno())
|
|
116
|
+
|
|
117
|
+
os.replace(tmp, os.path.join(output_dir, MANIFEST_NAME))
|
|
118
|
+
|
|
119
|
+
return manifest
|
src/data/shard_writer.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import os
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
SHARD_FORMAT_VERSION = 1
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def dtype_for_vocab(vocab_size):
|
|
10
|
+
return np.dtype(np.uint16) if vocab_size <= 65535 else np.dtype(np.uint32)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def file_sha256(path, block_size=1 << 20):
|
|
14
|
+
digest = hashlib.sha256()
|
|
15
|
+
|
|
16
|
+
with open(path, "rb") as f:
|
|
17
|
+
while True:
|
|
18
|
+
block = f.read(block_size)
|
|
19
|
+
if not block:
|
|
20
|
+
break
|
|
21
|
+
digest.update(block)
|
|
22
|
+
|
|
23
|
+
return digest.hexdigest()
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class ShardWriter:
|
|
27
|
+
|
|
28
|
+
def __init__(self, output_dir, prefix, shard_tokens, dtype):
|
|
29
|
+
self.output_dir = output_dir
|
|
30
|
+
self.prefix = prefix
|
|
31
|
+
self.shard_tokens = shard_tokens
|
|
32
|
+
self.dtype = np.dtype(dtype)
|
|
33
|
+
|
|
34
|
+
self._pending = []
|
|
35
|
+
self._pending_count = 0
|
|
36
|
+
self._shard_index = 0
|
|
37
|
+
self.shards = []
|
|
38
|
+
|
|
39
|
+
os.makedirs(output_dir, exist_ok=True)
|
|
40
|
+
|
|
41
|
+
def write(self, token_ids):
|
|
42
|
+
self._pending.append(np.asarray(token_ids, dtype=self.dtype))
|
|
43
|
+
self._pending_count += len(token_ids)
|
|
44
|
+
|
|
45
|
+
while self._pending_count >= self.shard_tokens:
|
|
46
|
+
self._flush(self.shard_tokens)
|
|
47
|
+
|
|
48
|
+
def _flush(self, count):
|
|
49
|
+
merged = np.concatenate(self._pending) if len(self._pending) > 1 else self._pending[0]
|
|
50
|
+
|
|
51
|
+
chunk = merged[:count]
|
|
52
|
+
rest = merged[count:]
|
|
53
|
+
|
|
54
|
+
self._pending = [rest] if rest.size else []
|
|
55
|
+
self._pending_count = int(rest.size)
|
|
56
|
+
|
|
57
|
+
name = f"{self.prefix}-{self._shard_index:05d}.bin"
|
|
58
|
+
final_path = os.path.join(self.output_dir, name)
|
|
59
|
+
tmp_path = final_path + ".tmp"
|
|
60
|
+
|
|
61
|
+
with open(tmp_path, "wb") as f:
|
|
62
|
+
chunk.tofile(f)
|
|
63
|
+
f.flush()
|
|
64
|
+
os.fsync(f.fileno())
|
|
65
|
+
|
|
66
|
+
os.replace(tmp_path, final_path)
|
|
67
|
+
|
|
68
|
+
self.shards.append({
|
|
69
|
+
"file": name,
|
|
70
|
+
"token_count": int(chunk.size),
|
|
71
|
+
"dtype": self.dtype.name,
|
|
72
|
+
"checksum": file_sha256(final_path),
|
|
73
|
+
})
|
|
74
|
+
|
|
75
|
+
self._shard_index += 1
|
|
76
|
+
|
|
77
|
+
def close(self):
|
|
78
|
+
if self._pending_count > 0:
|
|
79
|
+
self._flush(self._pending_count)
|
|
80
|
+
|
|
81
|
+
return self.shards
|
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from src.data.shard_builder import MANIFEST_NAME
|
|
7
|
+
from src.tokenization.base import identity_matches
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def is_shard_dir(path):
|
|
11
|
+
return isinstance(path, str) and os.path.isfile(os.path.join(path, MANIFEST_NAME))
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class ShardIndex:
|
|
15
|
+
|
|
16
|
+
def __init__(self, directory):
|
|
17
|
+
self.directory = directory
|
|
18
|
+
|
|
19
|
+
with open(os.path.join(directory, MANIFEST_NAME)) as f:
|
|
20
|
+
self.manifest = json.load(f)
|
|
21
|
+
|
|
22
|
+
self.files = [os.path.join(directory, s["file"]) for s in self.manifest["shards"]]
|
|
23
|
+
|
|
24
|
+
def __len__(self):
|
|
25
|
+
return len(self.files)
|
|
26
|
+
|
|
27
|
+
def state_dict(self):
|
|
28
|
+
return {
|
|
29
|
+
"corpus_version": self.manifest["corpus_version"],
|
|
30
|
+
"total_tokens": self.manifest["total_tokens"],
|
|
31
|
+
"shards": [s["file"] for s in self.manifest["shards"]],
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
def load_state_dict(self, state):
|
|
35
|
+
if state.get("shards") != [s["file"] for s in self.manifest["shards"]] or \
|
|
36
|
+
state.get("corpus_version") != self.manifest["corpus_version"]:
|
|
37
|
+
raise ValueError("token shard set does not match the checkpoint; resuming would corrupt training")
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class ShardedTokenDataset:
|
|
41
|
+
|
|
42
|
+
def __init__(self, directory, sequence_length, tokenizer=None):
|
|
43
|
+
self.corpus = ShardIndex(directory)
|
|
44
|
+
self.sequence_length = sequence_length
|
|
45
|
+
self.dtype = np.dtype(self.corpus.manifest["dtype"])
|
|
46
|
+
|
|
47
|
+
if tokenizer is not None and not identity_matches(self.corpus.manifest["tokenizer_identity"], tokenizer.identity):
|
|
48
|
+
raise ValueError("token shards were built with a different tokenizer than the one supplied")
|
|
49
|
+
|
|
50
|
+
self.shard_index = 0
|
|
51
|
+
self.token_offset = 0
|
|
52
|
+
self._carry = np.empty(0, dtype=np.int64)
|
|
53
|
+
|
|
54
|
+
def batches(self, batch_size):
|
|
55
|
+
needed = self.sequence_length + 1
|
|
56
|
+
meta = self.corpus.manifest["shards"]
|
|
57
|
+
read_tokens = max(needed * batch_size * 8, 1 << 16)
|
|
58
|
+
|
|
59
|
+
while True:
|
|
60
|
+
examples = []
|
|
61
|
+
|
|
62
|
+
while len(examples) < batch_size:
|
|
63
|
+
while self._carry.size < needed and self.shard_index < len(meta):
|
|
64
|
+
path = self.corpus.files[self.shard_index]
|
|
65
|
+
count = meta[self.shard_index]["token_count"]
|
|
66
|
+
|
|
67
|
+
if self.token_offset >= count:
|
|
68
|
+
self.shard_index += 1
|
|
69
|
+
self.token_offset = 0
|
|
70
|
+
continue
|
|
71
|
+
|
|
72
|
+
take = min(read_tokens, count - self.token_offset)
|
|
73
|
+
chunk = np.fromfile(
|
|
74
|
+
path,
|
|
75
|
+
dtype=self.dtype,
|
|
76
|
+
count=take,
|
|
77
|
+
offset=self.token_offset * self.dtype.itemsize,
|
|
78
|
+
).astype(np.int64)
|
|
79
|
+
|
|
80
|
+
self.token_offset += take
|
|
81
|
+
self._carry = np.concatenate([self._carry, chunk]) if self._carry.size else chunk
|
|
82
|
+
|
|
83
|
+
if self._carry.size < needed:
|
|
84
|
+
break
|
|
85
|
+
|
|
86
|
+
examples.append(self._carry[:needed])
|
|
87
|
+
self._carry = self._carry[needed:]
|
|
88
|
+
|
|
89
|
+
if not examples:
|
|
90
|
+
return
|
|
91
|
+
|
|
92
|
+
arr = np.stack(examples)
|
|
93
|
+
|
|
94
|
+
yield arr[:, :-1], arr[:, 1:]
|
|
95
|
+
|
|
96
|
+
def start_new_epoch(self):
|
|
97
|
+
self.shard_index = 0
|
|
98
|
+
self.token_offset = 0
|
|
99
|
+
self._carry = np.empty(0, dtype=np.int64)
|
|
100
|
+
|
|
101
|
+
def state_dict(self):
|
|
102
|
+
return {
|
|
103
|
+
"kind": "shards",
|
|
104
|
+
"shard_index": self.shard_index,
|
|
105
|
+
"token_offset": self.token_offset,
|
|
106
|
+
"carry": self._carry.tolist(),
|
|
107
|
+
}
|
|
108
|
+
|
|
109
|
+
def load_state_dict(self, data):
|
|
110
|
+
self.shard_index = data["shard_index"]
|
|
111
|
+
self.token_offset = data["token_offset"]
|
|
112
|
+
self._carry = np.asarray(data["carry"], dtype=np.int64)
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
from collections import deque
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from src.data.document_stream import DocumentReader
|
|
6
|
+
from src.data.parallel_encode import ParallelEncoder, encode_serial
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class DatasetState:
|
|
10
|
+
|
|
11
|
+
def __init__(self, file_index=0, byte_offset=0, token_buffer=None):
|
|
12
|
+
self.file_index = file_index
|
|
13
|
+
self.byte_offset = byte_offset
|
|
14
|
+
self.token_buffer = list(token_buffer) if token_buffer else []
|
|
15
|
+
|
|
16
|
+
def to_dict(self):
|
|
17
|
+
return {
|
|
18
|
+
"kind": "text",
|
|
19
|
+
"file_index": self.file_index,
|
|
20
|
+
"byte_offset": self.byte_offset,
|
|
21
|
+
"token_buffer": list(self.token_buffer),
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
@classmethod
|
|
25
|
+
def from_dict(cls, data):
|
|
26
|
+
return cls(
|
|
27
|
+
file_index=data.get("file_index", 0),
|
|
28
|
+
byte_offset=data.get("byte_offset", 0),
|
|
29
|
+
token_buffer=data.get("token_buffer", []),
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class StreamingTextDataset:
|
|
34
|
+
|
|
35
|
+
def __init__(
|
|
36
|
+
self,
|
|
37
|
+
corpus,
|
|
38
|
+
tokenizer,
|
|
39
|
+
sequence_length,
|
|
40
|
+
read_buffer_size=1 << 20,
|
|
41
|
+
insert_eos=True,
|
|
42
|
+
text_field="text",
|
|
43
|
+
state=None,
|
|
44
|
+
workers=1,
|
|
45
|
+
):
|
|
46
|
+
self.corpus = corpus
|
|
47
|
+
self.tokenizer = tokenizer
|
|
48
|
+
self.sequence_length = sequence_length
|
|
49
|
+
self.insert_eos = insert_eos
|
|
50
|
+
self.workers = max(1, int(workers))
|
|
51
|
+
|
|
52
|
+
self.reader = DocumentReader(
|
|
53
|
+
corpus.files,
|
|
54
|
+
read_buffer_size=read_buffer_size,
|
|
55
|
+
text_field=text_field,
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
self.state = state or DatasetState()
|
|
59
|
+
self.token_buffer = deque(self.state.token_buffer)
|
|
60
|
+
|
|
61
|
+
def _encoded_documents(self, doc_iter):
|
|
62
|
+
if self.workers > 1:
|
|
63
|
+
return ParallelEncoder(self.tokenizer, self.workers).encode(doc_iter)
|
|
64
|
+
|
|
65
|
+
return encode_serial(self.tokenizer, doc_iter)
|
|
66
|
+
|
|
67
|
+
def _fill_buffer(self, doc_iter, min_tokens):
|
|
68
|
+
while len(self.token_buffer) < min_tokens:
|
|
69
|
+
try:
|
|
70
|
+
file_index, offset, ids, doc_end = next(doc_iter)
|
|
71
|
+
except StopIteration:
|
|
72
|
+
return False
|
|
73
|
+
|
|
74
|
+
self.state.file_index = file_index
|
|
75
|
+
self.state.byte_offset = offset
|
|
76
|
+
|
|
77
|
+
if len(ids):
|
|
78
|
+
self.token_buffer.extend(ids.tolist() if isinstance(ids, np.ndarray) else ids)
|
|
79
|
+
|
|
80
|
+
if doc_end and self.insert_eos and self.tokenizer.eos_id is not None:
|
|
81
|
+
self.token_buffer.append(self.tokenizer.eos_id)
|
|
82
|
+
|
|
83
|
+
return True
|
|
84
|
+
|
|
85
|
+
def batches(self, batch_size):
|
|
86
|
+
needed = self.sequence_length + 1
|
|
87
|
+
doc_iter = self._encoded_documents(self.reader.iter_documents(
|
|
88
|
+
self.state.file_index,
|
|
89
|
+
self.state.byte_offset,
|
|
90
|
+
))
|
|
91
|
+
|
|
92
|
+
try:
|
|
93
|
+
yield from self._batches(doc_iter, batch_size, needed)
|
|
94
|
+
finally:
|
|
95
|
+
doc_iter.close()
|
|
96
|
+
|
|
97
|
+
def _batches(self, doc_iter, batch_size, needed):
|
|
98
|
+
while True:
|
|
99
|
+
examples = []
|
|
100
|
+
|
|
101
|
+
for _ in range(batch_size):
|
|
102
|
+
has_more = self._fill_buffer(doc_iter, needed)
|
|
103
|
+
|
|
104
|
+
if len(self.token_buffer) < needed:
|
|
105
|
+
break
|
|
106
|
+
|
|
107
|
+
block = [self.token_buffer.popleft() for _ in range(needed)]
|
|
108
|
+
examples.append(block)
|
|
109
|
+
|
|
110
|
+
if not has_more and len(self.token_buffer) < needed:
|
|
111
|
+
break
|
|
112
|
+
|
|
113
|
+
if not examples:
|
|
114
|
+
return
|
|
115
|
+
|
|
116
|
+
arr = np.array(examples, dtype=np.int64)
|
|
117
|
+
|
|
118
|
+
self.state.token_buffer = list(self.token_buffer)
|
|
119
|
+
|
|
120
|
+
yield arr[:, :-1], arr[:, 1:]
|
|
121
|
+
|
|
122
|
+
def start_new_epoch(self):
|
|
123
|
+
self.state = DatasetState()
|
|
124
|
+
self.token_buffer = deque()
|
|
125
|
+
|
|
126
|
+
def state_dict(self):
|
|
127
|
+
self.state.token_buffer = list(self.token_buffer)
|
|
128
|
+
return self.state.to_dict()
|
|
129
|
+
|
|
130
|
+
def load_state_dict(self, data):
|
|
131
|
+
self.state = DatasetState.from_dict(data)
|
|
132
|
+
self.token_buffer = deque(self.state.token_buffer)
|
src/data/validation.py
ADDED
|
@@ -0,0 +1,212 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
|
|
4
|
+
from src.data.document_stream import JSON_EXTENSIONS
|
|
5
|
+
from src.data.shard_builder import MANIFEST_NAME, corpus_version
|
|
6
|
+
from src.data.shard_writer import file_sha256
|
|
7
|
+
from src.tokenization.base import identity_matches
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class ValidationReport:
|
|
11
|
+
|
|
12
|
+
def __init__(self):
|
|
13
|
+
self.errors = []
|
|
14
|
+
self.warnings = []
|
|
15
|
+
self.stats = {}
|
|
16
|
+
|
|
17
|
+
@property
|
|
18
|
+
def ok(self):
|
|
19
|
+
return not self.errors
|
|
20
|
+
|
|
21
|
+
def to_dict(self):
|
|
22
|
+
return {"ok": self.ok, "errors": self.errors, "warnings": self.warnings, "stats": self.stats}
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def validate_corpus(corpus, text_field="text", sample_bytes=4 << 20, tokenizer=None, sample_docs=200):
|
|
26
|
+
report = ValidationReport()
|
|
27
|
+
|
|
28
|
+
if not corpus.files:
|
|
29
|
+
report.errors.append("corpus contains no supported files")
|
|
30
|
+
return report
|
|
31
|
+
|
|
32
|
+
total_bytes = 0
|
|
33
|
+
malformed = 0
|
|
34
|
+
bad_encoding = 0
|
|
35
|
+
empty_docs = 0
|
|
36
|
+
checked_docs = 0
|
|
37
|
+
unk_tokens = 0
|
|
38
|
+
sampled_tokens = 0
|
|
39
|
+
sampled_chars = 0
|
|
40
|
+
|
|
41
|
+
for path in corpus.files:
|
|
42
|
+
if not os.path.isfile(path):
|
|
43
|
+
report.errors.append(f"missing file: {path}")
|
|
44
|
+
continue
|
|
45
|
+
|
|
46
|
+
if not os.access(path, os.R_OK):
|
|
47
|
+
report.errors.append(f"unreadable file: {path}")
|
|
48
|
+
continue
|
|
49
|
+
|
|
50
|
+
total_bytes += os.path.getsize(path)
|
|
51
|
+
|
|
52
|
+
report.stats["files"] = len(corpus.files)
|
|
53
|
+
report.stats["total_bytes"] = total_bytes
|
|
54
|
+
|
|
55
|
+
remaining = sample_bytes
|
|
56
|
+
|
|
57
|
+
for path in corpus.files:
|
|
58
|
+
if remaining <= 0 or checked_docs >= sample_docs * 10:
|
|
59
|
+
break
|
|
60
|
+
|
|
61
|
+
if not os.path.isfile(path):
|
|
62
|
+
continue
|
|
63
|
+
|
|
64
|
+
ext = os.path.splitext(path)[1].lower()
|
|
65
|
+
|
|
66
|
+
with open(path, "rb") as f:
|
|
67
|
+
if ext in JSON_EXTENSIONS:
|
|
68
|
+
for line in f:
|
|
69
|
+
remaining -= len(line)
|
|
70
|
+
|
|
71
|
+
if not line.strip():
|
|
72
|
+
continue
|
|
73
|
+
|
|
74
|
+
checked_docs += 1
|
|
75
|
+
|
|
76
|
+
try:
|
|
77
|
+
text = line.decode("utf-8")
|
|
78
|
+
except UnicodeDecodeError:
|
|
79
|
+
bad_encoding += 1
|
|
80
|
+
continue
|
|
81
|
+
|
|
82
|
+
try:
|
|
83
|
+
obj = json.loads(text)
|
|
84
|
+
except json.JSONDecodeError:
|
|
85
|
+
malformed += 1
|
|
86
|
+
continue
|
|
87
|
+
|
|
88
|
+
value = obj.get(text_field) if isinstance(obj, dict) else None
|
|
89
|
+
|
|
90
|
+
if not isinstance(value, str):
|
|
91
|
+
malformed += 1
|
|
92
|
+
elif not value.strip():
|
|
93
|
+
empty_docs += 1
|
|
94
|
+
elif tokenizer is not None and sampled_tokens < 200000 and checked_docs <= sample_docs:
|
|
95
|
+
ids = tokenizer.encode(value)
|
|
96
|
+
sampled_tokens += len(ids)
|
|
97
|
+
sampled_chars += len(value)
|
|
98
|
+
unk = getattr(tokenizer, "unk_id", None)
|
|
99
|
+
if unk is not None:
|
|
100
|
+
unk_tokens += sum(1 for i in ids if i == unk)
|
|
101
|
+
|
|
102
|
+
if remaining <= 0:
|
|
103
|
+
break
|
|
104
|
+
else:
|
|
105
|
+
data = f.read(min(max(remaining, 0), sample_bytes))
|
|
106
|
+
remaining -= len(data)
|
|
107
|
+
checked_docs += 1
|
|
108
|
+
|
|
109
|
+
try:
|
|
110
|
+
data.decode("utf-8")
|
|
111
|
+
except UnicodeDecodeError as exc:
|
|
112
|
+
if exc.start < len(data) - 4:
|
|
113
|
+
bad_encoding += 1
|
|
114
|
+
|
|
115
|
+
if not data.strip():
|
|
116
|
+
empty_docs += 1
|
|
117
|
+
elif tokenizer is not None:
|
|
118
|
+
sample_text = data.decode("utf-8", errors="ignore")[:200000]
|
|
119
|
+
ids = tokenizer.encode(sample_text)
|
|
120
|
+
sampled_tokens += len(ids)
|
|
121
|
+
sampled_chars += len(sample_text)
|
|
122
|
+
unk = getattr(tokenizer, "unk_id", None)
|
|
123
|
+
if unk is not None:
|
|
124
|
+
unk_tokens += sum(1 for i in ids if i == unk)
|
|
125
|
+
|
|
126
|
+
report.stats.update({
|
|
127
|
+
"sampled_documents": checked_docs,
|
|
128
|
+
"malformed_documents": malformed,
|
|
129
|
+
"encoding_errors": bad_encoding,
|
|
130
|
+
"empty_documents": empty_docs,
|
|
131
|
+
})
|
|
132
|
+
|
|
133
|
+
if malformed:
|
|
134
|
+
report.warnings.append(f"{malformed} malformed JSON documents in sample")
|
|
135
|
+
|
|
136
|
+
if bad_encoding:
|
|
137
|
+
report.warnings.append(f"{bad_encoding} documents with invalid UTF-8 in sample")
|
|
138
|
+
|
|
139
|
+
if tokenizer is not None and sampled_tokens:
|
|
140
|
+
report.stats["chars_per_token"] = sampled_chars / sampled_tokens
|
|
141
|
+
|
|
142
|
+
if tokenizer is not None and sampled_tokens and getattr(tokenizer, "unk_id", None) is not None:
|
|
143
|
+
rate = unk_tokens / sampled_tokens
|
|
144
|
+
report.stats["unk_rate"] = rate
|
|
145
|
+
|
|
146
|
+
if rate > 0.05:
|
|
147
|
+
report.warnings.append(f"unknown-token rate {rate:.1%} in sample; tokenizer may not fit this corpus")
|
|
148
|
+
|
|
149
|
+
return report
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def validate_shards(directory, tokenizer=None, verify_checksums=False):
|
|
153
|
+
report = ValidationReport()
|
|
154
|
+
manifest_path = os.path.join(directory, MANIFEST_NAME)
|
|
155
|
+
|
|
156
|
+
if not os.path.isfile(manifest_path):
|
|
157
|
+
report.errors.append("manifest.json not found")
|
|
158
|
+
return report
|
|
159
|
+
|
|
160
|
+
with open(manifest_path) as f:
|
|
161
|
+
manifest = json.load(f)
|
|
162
|
+
|
|
163
|
+
itemsize = {"uint16": 2, "uint32": 4}.get(manifest["dtype"])
|
|
164
|
+
|
|
165
|
+
if itemsize is None:
|
|
166
|
+
report.errors.append(f"unsupported dtype {manifest['dtype']}")
|
|
167
|
+
return report
|
|
168
|
+
|
|
169
|
+
if tokenizer is not None:
|
|
170
|
+
if not identity_matches(manifest["tokenizer_identity"], tokenizer.identity):
|
|
171
|
+
report.errors.append("tokenizer identity does not match shard manifest")
|
|
172
|
+
|
|
173
|
+
if tokenizer.vocab_size != manifest["vocab_size"]:
|
|
174
|
+
report.errors.append("tokenizer vocab size does not match shard manifest")
|
|
175
|
+
|
|
176
|
+
total = 0
|
|
177
|
+
|
|
178
|
+
for shard in manifest["shards"]:
|
|
179
|
+
path = os.path.join(directory, shard["file"])
|
|
180
|
+
|
|
181
|
+
if not os.path.isfile(path):
|
|
182
|
+
report.errors.append(f"missing shard {shard['file']}")
|
|
183
|
+
continue
|
|
184
|
+
|
|
185
|
+
size = os.path.getsize(path)
|
|
186
|
+
|
|
187
|
+
if size != shard["token_count"] * itemsize:
|
|
188
|
+
report.errors.append(f"shard {shard['file']} size does not match token_count")
|
|
189
|
+
continue
|
|
190
|
+
|
|
191
|
+
total += shard["token_count"]
|
|
192
|
+
|
|
193
|
+
if verify_checksums and file_sha256(path) != shard["checksum"]:
|
|
194
|
+
report.errors.append(f"checksum mismatch in {shard['file']}")
|
|
195
|
+
|
|
196
|
+
if total != manifest["total_tokens"]:
|
|
197
|
+
report.errors.append("manifest total_tokens does not match shard token counts")
|
|
198
|
+
|
|
199
|
+
stale = [p for p in manifest.get("source_files", []) if not os.path.isfile(p)]
|
|
200
|
+
|
|
201
|
+
if stale:
|
|
202
|
+
report.warnings.append(f"{len(stale)} source files no longer exist; corpus_version cannot be re-verified")
|
|
203
|
+
elif manifest.get("source_files") and corpus_version(manifest["source_files"]) != manifest["corpus_version"]:
|
|
204
|
+
report.warnings.append("source corpus changed since shards were built")
|
|
205
|
+
|
|
206
|
+
report.stats.update({
|
|
207
|
+
"shards": len(manifest["shards"]),
|
|
208
|
+
"total_tokens": total,
|
|
209
|
+
"checksums_verified": verify_checksums,
|
|
210
|
+
})
|
|
211
|
+
|
|
212
|
+
return report
|
|
File without changes
|