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.
Files changed (146) hide show
  1. cli.py +604 -0
  2. pytensorforge-0.1.0.dist-info/METADATA +103 -0
  3. pytensorforge-0.1.0.dist-info/RECORD +146 -0
  4. pytensorforge-0.1.0.dist-info/WHEEL +5 -0
  5. pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
  6. pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
  7. src/__init__.py +0 -0
  8. src/activations/Activation.py +4 -0
  9. src/activations/ELU.py +11 -0
  10. src/activations/GELU.py +6 -0
  11. src/activations/ReLU.py +27 -0
  12. src/activations/SELU.py +14 -0
  13. src/activations/Sigmoid.py +27 -0
  14. src/activations/Softmax.py +84 -0
  15. src/activations/Tanh.py +29 -0
  16. src/activations/__init__.py +17 -0
  17. src/config.py +120 -0
  18. src/core/Matrix.py +3 -0
  19. src/core/Scalar.py +18 -0
  20. src/core/Tensor.py +866 -0
  21. src/core/Vector.py +31 -0
  22. src/core/__init__.py +0 -0
  23. src/data/__init__.py +0 -0
  24. src/data/chat_dataset.py +188 -0
  25. src/data/corpus.py +104 -0
  26. src/data/document_stream.py +178 -0
  27. src/data/parallel_encode.py +86 -0
  28. src/data/prefetch.py +62 -0
  29. src/data/shard_builder.py +119 -0
  30. src/data/shard_writer.py +81 -0
  31. src/data/sharded_dataset.py +112 -0
  32. src/data/streaming_dataset.py +132 -0
  33. src/data/validation.py +212 -0
  34. src/inference/__init__.py +0 -0
  35. src/inference/chat_template.py +384 -0
  36. src/inference/config.py +48 -0
  37. src/inference/engine.py +241 -0
  38. src/inference/export.py +133 -0
  39. src/inference/kv_cache.py +65 -0
  40. src/inference/runtime.py +161 -0
  41. src/inference/sampling.py +42 -0
  42. src/inference/scheduler.py +473 -0
  43. src/inference/text.py +67 -0
  44. src/initializers/Constant.py +9 -0
  45. src/initializers/GlorotNormal.py +15 -0
  46. src/initializers/GlorotUniform.py +26 -0
  47. src/initializers/HeNormal.py +15 -0
  48. src/initializers/HeUniform.py +14 -0
  49. src/initializers/Initializer.py +4 -0
  50. src/initializers/LecunNormal.py +16 -0
  51. src/initializers/LecunUniform.py +14 -0
  52. src/initializers/Ones.py +6 -0
  53. src/initializers/Orthogonal.py +14 -0
  54. src/initializers/RandomNormal.py +14 -0
  55. src/initializers/RandomUniform.py +14 -0
  56. src/initializers/Zeros.py +8 -0
  57. src/initializers/__init__.py +17 -0
  58. src/loss/CategoricalCrossEntropy.py +9 -0
  59. src/loss/CrossEntropyLoss.py +34 -0
  60. src/loss/CrossEntropyWithLogitsLoss.py +59 -0
  61. src/loss/Hinge.py +5 -0
  62. src/loss/Huber.py +22 -0
  63. src/loss/Loss.py +6 -0
  64. src/loss/MSE.py +7 -0
  65. src/loss/MSELoss.py +10 -0
  66. src/loss/SparseCategoricalCrossEntropy.py +15 -0
  67. src/loss/__init__.py +18 -0
  68. src/loss/bce.py +34 -0
  69. src/loss/mae.py +16 -0
  70. src/math/__init__.py +0 -0
  71. src/math/clip.py +37 -0
  72. src/math/exp.py +27 -0
  73. src/math/log.py +25 -0
  74. src/math/sigmoid.py +5 -0
  75. src/models/__init__.py +0 -0
  76. src/models/embedding/Embedding.py +65 -0
  77. src/models/embedding/__init__.py +0 -0
  78. src/models/gpt/__init__.py +0 -0
  79. src/models/gpt/attention.py +158 -0
  80. src/models/gpt/block.py +74 -0
  81. src/models/gpt/config.py +103 -0
  82. src/models/gpt/context.py +44 -0
  83. src/models/gpt/model.py +165 -0
  84. src/models/gpt/recompute.py +35 -0
  85. src/models/gpt/rope.py +84 -0
  86. src/models/regression/Linear.py +51 -0
  87. src/models/regression/Logistic.py +36 -0
  88. src/models/regression/__init__.py +0 -0
  89. src/models/seq/Sequential.py +297 -0
  90. src/models/seq/__init__.py +0 -0
  91. src/models/svm/__init__.py +0 -0
  92. src/models/tokenizer/BPETokenizer.py +228 -0
  93. src/models/tokenizer/__init__.py +0 -0
  94. src/models/transformers/Dropout.py +35 -0
  95. src/models/transformers/LastToken.py +10 -0
  96. src/models/transformers/LayerNorm.py +54 -0
  97. src/models/transformers/Linear.py +18 -0
  98. src/models/transformers/MultiHeadAttention.py +130 -0
  99. src/models/transformers/TransformerBlock.py +79 -0
  100. src/models/transformers/__init__.py +0 -0
  101. src/neural/Dense.py +58 -0
  102. src/neural/LSTM.py +167 -0
  103. src/neural/Layer.py +72 -0
  104. src/neural/Parameter.py +30 -0
  105. src/neural/RNN.py +83 -0
  106. src/neural/__init__.py +0 -0
  107. src/ops/__init__.py +0 -0
  108. src/ops/stack.py +40 -0
  109. src/optimizers/Adagrad.py +31 -0
  110. src/optimizers/Adam.py +98 -0
  111. src/optimizers/AdamW.py +84 -0
  112. src/optimizers/Batch.py +11 -0
  113. src/optimizers/Nesterov.py +35 -0
  114. src/optimizers/Optimizer.py +18 -0
  115. src/optimizers/RMSProp.py +35 -0
  116. src/optimizers/SGD.py +30 -0
  117. src/optimizers/SGDMomentum.py +28 -0
  118. src/optimizers/__init__.py +9 -0
  119. src/scaling/StandardScaler.py +15 -0
  120. src/scaling/__init__.py +0 -0
  121. src/serialization/__init__.py +0 -0
  122. src/serialization/checkpoint.py +58 -0
  123. src/serialization/modelio.py +132 -0
  124. src/serving/__init__.py +0 -0
  125. src/serving/app.py +792 -0
  126. src/serving/config.py +216 -0
  127. src/serving/errors.py +51 -0
  128. src/serving/http.py +599 -0
  129. src/serving/metrics.py +293 -0
  130. src/serving/model_server.py +287 -0
  131. src/serving/protocol.py +377 -0
  132. src/serving/security.py +200 -0
  133. src/serving/server.py +121 -0
  134. src/tokenization/__init__.py +0 -0
  135. src/tokenization/base.py +75 -0
  136. src/tokenization/bpe.py +190 -0
  137. src/tokenization/bytebpe.py +476 -0
  138. src/tokenization/registry.py +28 -0
  139. src/training/__init__.py +0 -0
  140. src/training/checkpoint_manager.py +101 -0
  141. src/training/experiment.py +71 -0
  142. src/training/losses.py +42 -0
  143. src/training/precision.py +141 -0
  144. src/training/profiler.py +38 -0
  145. src/training/scheduler.py +50 -0
  146. 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
@@ -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