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
src/core/Vector.py
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
class Vector:
|
|
2
|
+
def __init__(self, data):
|
|
3
|
+
self.data = data
|
|
4
|
+
|
|
5
|
+
def set(self, index, value):
|
|
6
|
+
self.data[index] = value
|
|
7
|
+
|
|
8
|
+
def get(self, index):
|
|
9
|
+
return self.data[index]
|
|
10
|
+
|
|
11
|
+
@staticmethod
|
|
12
|
+
def load(self, data):
|
|
13
|
+
self.data = data
|
|
14
|
+
|
|
15
|
+
def dump(self):
|
|
16
|
+
return self.data
|
|
17
|
+
|
|
18
|
+
@staticmethod
|
|
19
|
+
def Zero(self, data):
|
|
20
|
+
vec = Vector(data)
|
|
21
|
+
for i in vec.data:
|
|
22
|
+
vec.set(i, 0)
|
|
23
|
+
return vec
|
|
24
|
+
|
|
25
|
+
def zero(self):
|
|
26
|
+
for i in range(len(self.data)):
|
|
27
|
+
self.data[i] = 0
|
|
28
|
+
|
|
29
|
+
def log(self):
|
|
30
|
+
for i in range(len(self.data)):
|
|
31
|
+
print(self.data[i])
|
src/core/__init__.py
ADDED
|
File without changes
|
src/data/__init__.py
ADDED
|
File without changes
|
src/data/chat_dataset.py
ADDED
|
@@ -0,0 +1,188 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from collections import deque
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
PACKING_MODES = ("pack", "pad")
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ChatRecordError(ValueError):
|
|
10
|
+
pass
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class ChatSFTDataset:
|
|
14
|
+
|
|
15
|
+
def __init__(self, corpus, template, sequence_length, packing="pack", messages_field="messages", state=None):
|
|
16
|
+
if packing not in PACKING_MODES:
|
|
17
|
+
raise ValueError(f"unknown packing '{packing}'; choose one of {PACKING_MODES}")
|
|
18
|
+
|
|
19
|
+
bad = [f for f in corpus.files if not f.lower().endswith(".jsonl")]
|
|
20
|
+
|
|
21
|
+
if bad:
|
|
22
|
+
raise ValueError(f"chat datasets must be .jsonl files; got {bad[0]}")
|
|
23
|
+
|
|
24
|
+
self.corpus = corpus
|
|
25
|
+
self.template = template
|
|
26
|
+
self.tokenizer = template.tokenizer
|
|
27
|
+
self.sequence_length = int(sequence_length)
|
|
28
|
+
self.packing = packing
|
|
29
|
+
self.messages_field = messages_field
|
|
30
|
+
self.roles = set(template.template.roles)
|
|
31
|
+
self.pad_id = self.tokenizer.eos_id if self.tokenizer.eos_id is not None else 0
|
|
32
|
+
self.load_state_dict(state or {})
|
|
33
|
+
|
|
34
|
+
def _records(self):
|
|
35
|
+
for file_index in range(self.file_index, len(self.corpus.files)):
|
|
36
|
+
path = self.corpus.files[file_index]
|
|
37
|
+
offset = self.byte_offset if file_index == self.file_index else 0
|
|
38
|
+
|
|
39
|
+
with open(path, "rb") as f:
|
|
40
|
+
f.seek(offset)
|
|
41
|
+
|
|
42
|
+
while True:
|
|
43
|
+
raw = f.readline()
|
|
44
|
+
|
|
45
|
+
if not raw:
|
|
46
|
+
break
|
|
47
|
+
|
|
48
|
+
end = f.tell()
|
|
49
|
+
line = raw.strip()
|
|
50
|
+
|
|
51
|
+
if line:
|
|
52
|
+
yield file_index, end, self._parse(line, path, end)
|
|
53
|
+
else:
|
|
54
|
+
yield file_index, end, None
|
|
55
|
+
|
|
56
|
+
def _parse(self, line, path, end):
|
|
57
|
+
where = f"{path} (record ending at byte {end})"
|
|
58
|
+
|
|
59
|
+
try:
|
|
60
|
+
record = json.loads(line)
|
|
61
|
+
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
|
|
62
|
+
raise ChatRecordError(f"{where}: invalid JSON: {exc}") from None
|
|
63
|
+
|
|
64
|
+
messages = record.get(self.messages_field) if isinstance(record, dict) else None
|
|
65
|
+
|
|
66
|
+
if not isinstance(messages, list) or not messages:
|
|
67
|
+
raise ChatRecordError(f"{where}: expected a non-empty '{self.messages_field}' list")
|
|
68
|
+
|
|
69
|
+
for i, m in enumerate(messages):
|
|
70
|
+
if not isinstance(m, dict) or not isinstance(m.get("role"), str) or not isinstance(m.get("content"), str):
|
|
71
|
+
raise ChatRecordError(f"{where}: message {i} needs string 'role' and 'content'")
|
|
72
|
+
|
|
73
|
+
if m["role"] not in self.roles:
|
|
74
|
+
raise ChatRecordError(
|
|
75
|
+
f"{where}: message {i} has role '{m['role']}', which the chat template does not define"
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
return [{"role": m["role"], "content": m["content"]} for m in messages]
|
|
79
|
+
|
|
80
|
+
def _conversations(self):
|
|
81
|
+
for file_index, end, messages in self._records():
|
|
82
|
+
self.file_index = file_index
|
|
83
|
+
self.byte_offset = end
|
|
84
|
+
|
|
85
|
+
if messages is None:
|
|
86
|
+
continue
|
|
87
|
+
|
|
88
|
+
ids, learn = self.template.training_tokens(messages)
|
|
89
|
+
|
|
90
|
+
if not any(learn):
|
|
91
|
+
self.skipped += 1
|
|
92
|
+
continue
|
|
93
|
+
|
|
94
|
+
self.records += 1
|
|
95
|
+
yield ids, learn
|
|
96
|
+
|
|
97
|
+
def _pack(self, conversations, needed):
|
|
98
|
+
while True:
|
|
99
|
+
while len(self.tokens) < needed:
|
|
100
|
+
item = next(conversations, None)
|
|
101
|
+
|
|
102
|
+
if item is None:
|
|
103
|
+
return None
|
|
104
|
+
|
|
105
|
+
ids, learn = item
|
|
106
|
+
self.tokens.extend(ids)
|
|
107
|
+
self.learn.extend(learn)
|
|
108
|
+
|
|
109
|
+
block = [self.tokens.popleft() for _ in range(needed)]
|
|
110
|
+
mask = [self.learn.popleft() for _ in range(needed)]
|
|
111
|
+
return block, mask
|
|
112
|
+
|
|
113
|
+
def _pad(self, conversations, needed):
|
|
114
|
+
while True:
|
|
115
|
+
item = next(conversations, None)
|
|
116
|
+
|
|
117
|
+
if item is None:
|
|
118
|
+
return None
|
|
119
|
+
|
|
120
|
+
ids, learn = item
|
|
121
|
+
ids, learn = ids[:needed], learn[:needed]
|
|
122
|
+
|
|
123
|
+
if len(ids) < 2 or not any(learn[1:]):
|
|
124
|
+
self.truncated_away += 1
|
|
125
|
+
continue
|
|
126
|
+
|
|
127
|
+
fill = needed - len(ids)
|
|
128
|
+
return ids + [self.pad_id] * fill, learn + [0] * fill
|
|
129
|
+
|
|
130
|
+
def batches(self, batch_size):
|
|
131
|
+
needed = self.sequence_length + 1
|
|
132
|
+
conversations = self._conversations()
|
|
133
|
+
take = self._pack if self.packing == "pack" else self._pad
|
|
134
|
+
|
|
135
|
+
try:
|
|
136
|
+
while True:
|
|
137
|
+
rows = []
|
|
138
|
+
|
|
139
|
+
for _ in range(batch_size):
|
|
140
|
+
row = take(conversations, needed)
|
|
141
|
+
|
|
142
|
+
if row is None:
|
|
143
|
+
break
|
|
144
|
+
|
|
145
|
+
rows.append(row)
|
|
146
|
+
|
|
147
|
+
if not rows:
|
|
148
|
+
return
|
|
149
|
+
|
|
150
|
+
ids = np.array([r[0] for r in rows], dtype=np.int64)
|
|
151
|
+
learn = np.array([r[1] for r in rows], dtype=np.float32)
|
|
152
|
+
yield ids[:, :-1], ids[:, 1:], learn[:, 1:]
|
|
153
|
+
finally:
|
|
154
|
+
conversations.close()
|
|
155
|
+
|
|
156
|
+
def start_new_epoch(self):
|
|
157
|
+
self.load_state_dict({})
|
|
158
|
+
|
|
159
|
+
def state_dict(self):
|
|
160
|
+
return {
|
|
161
|
+
"kind": "chat",
|
|
162
|
+
"packing": self.packing,
|
|
163
|
+
"file_index": self.file_index,
|
|
164
|
+
"byte_offset": self.byte_offset,
|
|
165
|
+
"token_buffer": list(self.tokens),
|
|
166
|
+
"learn_buffer": list(self.learn),
|
|
167
|
+
"records": self.records,
|
|
168
|
+
"skipped": self.skipped,
|
|
169
|
+
"truncated_away": self.truncated_away,
|
|
170
|
+
}
|
|
171
|
+
|
|
172
|
+
def load_state_dict(self, data):
|
|
173
|
+
if data and data.get("kind") != "chat":
|
|
174
|
+
raise ValueError("dataset state is not from a chat dataset")
|
|
175
|
+
|
|
176
|
+
if data and data.get("packing", self.packing) != self.packing:
|
|
177
|
+
raise ValueError("dataset state was written with a different packing mode")
|
|
178
|
+
|
|
179
|
+
self.file_index = int(data.get("file_index", 0))
|
|
180
|
+
self.byte_offset = int(data.get("byte_offset", 0))
|
|
181
|
+
self.tokens = deque(data.get("token_buffer", []))
|
|
182
|
+
self.learn = deque(data.get("learn_buffer", []))
|
|
183
|
+
self.records = int(data.get("records", 0))
|
|
184
|
+
self.skipped = int(data.get("skipped", 0))
|
|
185
|
+
self.truncated_away = int(data.get("truncated_away", 0))
|
|
186
|
+
|
|
187
|
+
if len(self.tokens) != len(self.learn):
|
|
188
|
+
raise ValueError("corrupt chat dataset state: token and mask buffers differ in length")
|
src/data/corpus.py
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import random
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class CorpusIndex:
|
|
6
|
+
|
|
7
|
+
def __init__(
|
|
8
|
+
self,
|
|
9
|
+
paths,
|
|
10
|
+
extensions=(".txt", ".jsonl", ".json"),
|
|
11
|
+
recursive=True,
|
|
12
|
+
shuffle=False,
|
|
13
|
+
seed=0,
|
|
14
|
+
):
|
|
15
|
+
self.extensions = tuple(e.lower() for e in extensions)
|
|
16
|
+
self.recursive = recursive
|
|
17
|
+
self.shuffle = shuffle
|
|
18
|
+
self.seed = seed
|
|
19
|
+
|
|
20
|
+
self.files = self._discover(paths)
|
|
21
|
+
|
|
22
|
+
if shuffle:
|
|
23
|
+
rng = random.Random(seed)
|
|
24
|
+
rng.shuffle(self.files)
|
|
25
|
+
|
|
26
|
+
def _discover(self, paths):
|
|
27
|
+
if isinstance(paths, str):
|
|
28
|
+
paths = [paths]
|
|
29
|
+
|
|
30
|
+
found = []
|
|
31
|
+
|
|
32
|
+
for path in paths:
|
|
33
|
+
path = os.path.abspath(path)
|
|
34
|
+
|
|
35
|
+
if os.path.isfile(path):
|
|
36
|
+
if path.lower().endswith(self.extensions):
|
|
37
|
+
found.append(path)
|
|
38
|
+
continue
|
|
39
|
+
|
|
40
|
+
if not os.path.isdir(path):
|
|
41
|
+
raise FileNotFoundError(f"corpus path not found: {path}")
|
|
42
|
+
|
|
43
|
+
if self.recursive:
|
|
44
|
+
for root, _, names in os.walk(path):
|
|
45
|
+
for name in sorted(names):
|
|
46
|
+
if name.lower().endswith(self.extensions):
|
|
47
|
+
found.append(os.path.join(root, name))
|
|
48
|
+
else:
|
|
49
|
+
for name in sorted(os.listdir(path)):
|
|
50
|
+
full = os.path.join(path, name)
|
|
51
|
+
if os.path.isfile(full) and name.lower().endswith(self.extensions):
|
|
52
|
+
found.append(full)
|
|
53
|
+
|
|
54
|
+
found.sort()
|
|
55
|
+
|
|
56
|
+
return found
|
|
57
|
+
|
|
58
|
+
def __len__(self):
|
|
59
|
+
return len(self.files)
|
|
60
|
+
|
|
61
|
+
def state_dict(self):
|
|
62
|
+
return {
|
|
63
|
+
"files": list(self.files),
|
|
64
|
+
"shuffle": self.shuffle,
|
|
65
|
+
"seed": self.seed,
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
def load_state_dict(self, state):
|
|
69
|
+
if state["files"] != self.files:
|
|
70
|
+
raise ValueError(
|
|
71
|
+
"checkpoint corpus file list does not match the current "
|
|
72
|
+
"corpus; resuming would silently skip or duplicate data"
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
def split(self, validation_fraction, seed=0):
|
|
76
|
+
if not 0 < validation_fraction < 1:
|
|
77
|
+
raise ValueError("validation_fraction must be between 0 and 1")
|
|
78
|
+
|
|
79
|
+
import hashlib
|
|
80
|
+
|
|
81
|
+
train, val = [], []
|
|
82
|
+
|
|
83
|
+
for path in self.files:
|
|
84
|
+
h = hashlib.sha256(f"{seed}:{os.path.basename(path)}".encode()).digest()
|
|
85
|
+
bucket = int.from_bytes(h[:8], "big") / 2**64
|
|
86
|
+
(val if bucket < validation_fraction else train).append(path)
|
|
87
|
+
|
|
88
|
+
if not val and len(train) > 1:
|
|
89
|
+
val.append(train.pop())
|
|
90
|
+
|
|
91
|
+
if not train:
|
|
92
|
+
raise ValueError("split left no training files; use a smaller validation_fraction or more files")
|
|
93
|
+
|
|
94
|
+
a = CorpusIndex.__new__(CorpusIndex)
|
|
95
|
+
b = CorpusIndex.__new__(CorpusIndex)
|
|
96
|
+
|
|
97
|
+
for obj, files in ((a, train), (b, val)):
|
|
98
|
+
obj.extensions = self.extensions
|
|
99
|
+
obj.recursive = self.recursive
|
|
100
|
+
obj.shuffle = self.shuffle
|
|
101
|
+
obj.seed = self.seed
|
|
102
|
+
obj.files = files
|
|
103
|
+
|
|
104
|
+
return a, b
|
|
@@ -0,0 +1,178 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
|
|
4
|
+
JSON_EXTENSIONS = (".jsonl", ".json")
|
|
5
|
+
ASCII_WS = frozenset(b" \t\n\r\x0b\x0c")
|
|
6
|
+
WS_BYTES = (b" ", b"\t", b"\n", b"\r", b"\x0b", b"\x0c")
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def _complete_utf8_len(buf):
|
|
10
|
+
n = len(buf)
|
|
11
|
+
i = n - 1
|
|
12
|
+
back = 0
|
|
13
|
+
|
|
14
|
+
while i >= 0 and back < 4 and (buf[i] & 0xC0) == 0x80:
|
|
15
|
+
i -= 1
|
|
16
|
+
back += 1
|
|
17
|
+
|
|
18
|
+
if i < 0:
|
|
19
|
+
return n
|
|
20
|
+
|
|
21
|
+
lead = buf[i]
|
|
22
|
+
|
|
23
|
+
if lead < 0x80:
|
|
24
|
+
need = 1
|
|
25
|
+
elif lead >> 5 == 0b110:
|
|
26
|
+
need = 2
|
|
27
|
+
elif lead >> 4 == 0b1110:
|
|
28
|
+
need = 3
|
|
29
|
+
elif lead >> 3 == 0b11110:
|
|
30
|
+
need = 4
|
|
31
|
+
else:
|
|
32
|
+
need = 1
|
|
33
|
+
|
|
34
|
+
return n if n - i >= need else i
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def _prev_char_is_space(buf, end):
|
|
38
|
+
k = end - 1
|
|
39
|
+
|
|
40
|
+
while k > 0 and (buf[k] & 0xC0) == 0x80:
|
|
41
|
+
k -= 1
|
|
42
|
+
|
|
43
|
+
try:
|
|
44
|
+
return buf[k:end].decode("utf-8").isspace()
|
|
45
|
+
except UnicodeDecodeError:
|
|
46
|
+
return False
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _find_cut(buf):
|
|
50
|
+
limit = len(buf)
|
|
51
|
+
|
|
52
|
+
while True:
|
|
53
|
+
last = max(buf.rfind(c, 0, limit) for c in WS_BYTES)
|
|
54
|
+
|
|
55
|
+
if last <= 0:
|
|
56
|
+
return -1
|
|
57
|
+
|
|
58
|
+
start = last
|
|
59
|
+
|
|
60
|
+
while start > 0 and buf[start - 1] in ASCII_WS:
|
|
61
|
+
start -= 1
|
|
62
|
+
|
|
63
|
+
if start == 0:
|
|
64
|
+
return -1
|
|
65
|
+
|
|
66
|
+
if not _prev_char_is_space(buf, start):
|
|
67
|
+
return start
|
|
68
|
+
|
|
69
|
+
limit = start
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
class DocumentReader:
|
|
73
|
+
|
|
74
|
+
def __init__(self, files, read_buffer_size=1 << 20, text_field="text"):
|
|
75
|
+
self.files = files
|
|
76
|
+
self.read_buffer_size = read_buffer_size
|
|
77
|
+
self.text_field = text_field
|
|
78
|
+
|
|
79
|
+
def _read_txt(self, f, start_offset):
|
|
80
|
+
buf = b""
|
|
81
|
+
offset = start_offset
|
|
82
|
+
|
|
83
|
+
while True:
|
|
84
|
+
chunk = f.read(self.read_buffer_size)
|
|
85
|
+
|
|
86
|
+
if not chunk:
|
|
87
|
+
if buf:
|
|
88
|
+
offset += len(buf)
|
|
89
|
+
yield buf.decode("utf-8", errors="ignore"), offset, True
|
|
90
|
+
return
|
|
91
|
+
|
|
92
|
+
buf += chunk
|
|
93
|
+
|
|
94
|
+
cut = _find_cut(buf)
|
|
95
|
+
|
|
96
|
+
if cut <= 0:
|
|
97
|
+
if len(buf) < 4 * self.read_buffer_size:
|
|
98
|
+
continue
|
|
99
|
+
|
|
100
|
+
cut = _complete_utf8_len(buf)
|
|
101
|
+
|
|
102
|
+
if cut >= len(buf):
|
|
103
|
+
cut = _complete_utf8_len(buf[:-1])
|
|
104
|
+
|
|
105
|
+
if cut <= 0:
|
|
106
|
+
continue
|
|
107
|
+
|
|
108
|
+
segment = buf[:cut]
|
|
109
|
+
buf = buf[cut:]
|
|
110
|
+
offset += len(segment)
|
|
111
|
+
|
|
112
|
+
text = segment.decode("utf-8", errors="ignore")
|
|
113
|
+
|
|
114
|
+
if text:
|
|
115
|
+
yield text, offset, False
|
|
116
|
+
|
|
117
|
+
def _read_jsonl(self, f, start_offset):
|
|
118
|
+
buf = b""
|
|
119
|
+
offset = start_offset
|
|
120
|
+
|
|
121
|
+
while True:
|
|
122
|
+
chunk = f.read(self.read_buffer_size)
|
|
123
|
+
|
|
124
|
+
if not chunk:
|
|
125
|
+
if buf.strip():
|
|
126
|
+
text = self._extract_json(buf)
|
|
127
|
+
offset += len(buf)
|
|
128
|
+
if text:
|
|
129
|
+
yield text, offset, True
|
|
130
|
+
break
|
|
131
|
+
|
|
132
|
+
buf += chunk
|
|
133
|
+
|
|
134
|
+
while b"\n" in buf:
|
|
135
|
+
line, buf = buf.split(b"\n", 1)
|
|
136
|
+
offset += len(line) + 1
|
|
137
|
+
|
|
138
|
+
if line.strip():
|
|
139
|
+
text = self._extract_json(line)
|
|
140
|
+
if text:
|
|
141
|
+
yield text, offset, True
|
|
142
|
+
|
|
143
|
+
def _extract_json(self, raw_line):
|
|
144
|
+
try:
|
|
145
|
+
obj = json.loads(raw_line.decode("utf-8", errors="ignore"))
|
|
146
|
+
except (json.JSONDecodeError, UnicodeDecodeError):
|
|
147
|
+
return None
|
|
148
|
+
|
|
149
|
+
if isinstance(obj, dict):
|
|
150
|
+
value = obj.get(self.text_field, "")
|
|
151
|
+
return value if isinstance(value, str) else None
|
|
152
|
+
|
|
153
|
+
return None
|
|
154
|
+
|
|
155
|
+
def read_file(self, path, start_offset=0):
|
|
156
|
+
ext = os.path.splitext(path)[1].lower()
|
|
157
|
+
|
|
158
|
+
with open(path, "rb") as f:
|
|
159
|
+
if start_offset:
|
|
160
|
+
f.seek(start_offset)
|
|
161
|
+
|
|
162
|
+
if ext in JSON_EXTENSIONS:
|
|
163
|
+
yield from self._read_jsonl(f, start_offset)
|
|
164
|
+
else:
|
|
165
|
+
yield from self._read_txt(f, start_offset)
|
|
166
|
+
|
|
167
|
+
def iter_documents(self, start_file_index=0, start_byte_offset=0):
|
|
168
|
+
file_index = start_file_index
|
|
169
|
+
resume_offset = start_byte_offset
|
|
170
|
+
|
|
171
|
+
while file_index < len(self.files):
|
|
172
|
+
path = self.files[file_index]
|
|
173
|
+
offset = resume_offset if file_index == start_file_index else 0
|
|
174
|
+
|
|
175
|
+
for text, next_offset, doc_end in self.read_file(path, offset):
|
|
176
|
+
yield file_index, next_offset, text, doc_end
|
|
177
|
+
|
|
178
|
+
file_index += 1
|
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
import multiprocessing as mp
|
|
2
|
+
from collections import deque
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
_EMPTY = np.zeros(0, dtype=np.int32)
|
|
7
|
+
_worker_tokenizer = None
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _init_worker(tokenizer):
|
|
11
|
+
global _worker_tokenizer
|
|
12
|
+
_worker_tokenizer = tokenizer
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _encode_texts(texts):
|
|
16
|
+
tok = _worker_tokenizer
|
|
17
|
+
return [np.asarray(tok.encode(t), dtype=np.int32) if t else _EMPTY for t in texts]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def encode_serial(tokenizer, doc_iter):
|
|
21
|
+
for file_index, offset, text, doc_end in doc_iter:
|
|
22
|
+
ids = tokenizer.encode(text) if text else ()
|
|
23
|
+
yield file_index, offset, ids, doc_end
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class ParallelEncoder:
|
|
27
|
+
|
|
28
|
+
def __init__(self, tokenizer, workers, batch_chars=1 << 18, max_batch_docs=4096, inflight_per_worker=2,
|
|
29
|
+
start_method="spawn"):
|
|
30
|
+
if workers < 2:
|
|
31
|
+
raise ValueError("ParallelEncoder needs at least 2 workers; use encode_serial otherwise")
|
|
32
|
+
|
|
33
|
+
if tokenizer.vocab_size > np.iinfo(np.int32).max:
|
|
34
|
+
raise ValueError("vocabulary too large for int32 token transport")
|
|
35
|
+
|
|
36
|
+
self.tokenizer = tokenizer
|
|
37
|
+
self.workers = int(workers)
|
|
38
|
+
self.batch_chars = max(1, int(batch_chars))
|
|
39
|
+
self.max_batch_docs = max(1, int(max_batch_docs))
|
|
40
|
+
self.max_inflight = max(1, self.workers * int(inflight_per_worker))
|
|
41
|
+
self.start_method = start_method
|
|
42
|
+
|
|
43
|
+
def _next_batch(self, it):
|
|
44
|
+
meta = []
|
|
45
|
+
texts = []
|
|
46
|
+
chars = 0
|
|
47
|
+
|
|
48
|
+
for file_index, offset, text, doc_end in it:
|
|
49
|
+
meta.append((file_index, offset, doc_end))
|
|
50
|
+
texts.append(text)
|
|
51
|
+
chars += len(text)
|
|
52
|
+
|
|
53
|
+
if chars >= self.batch_chars or len(texts) >= self.max_batch_docs:
|
|
54
|
+
break
|
|
55
|
+
|
|
56
|
+
return meta, texts
|
|
57
|
+
|
|
58
|
+
def encode(self, doc_iter):
|
|
59
|
+
ctx = mp.get_context(self.start_method)
|
|
60
|
+
pool = ctx.Pool(self.workers, initializer=_init_worker, initargs=(self.tokenizer,))
|
|
61
|
+
pending = deque()
|
|
62
|
+
it = iter(doc_iter)
|
|
63
|
+
exhausted = False
|
|
64
|
+
|
|
65
|
+
try:
|
|
66
|
+
while True:
|
|
67
|
+
while not exhausted and len(pending) < self.max_inflight:
|
|
68
|
+
meta, texts = self._next_batch(it)
|
|
69
|
+
|
|
70
|
+
if not meta:
|
|
71
|
+
exhausted = True
|
|
72
|
+
break
|
|
73
|
+
|
|
74
|
+
pending.append((meta, pool.apply_async(_encode_texts, (texts,))))
|
|
75
|
+
|
|
76
|
+
if not pending:
|
|
77
|
+
return
|
|
78
|
+
|
|
79
|
+
meta, result = pending.popleft()
|
|
80
|
+
|
|
81
|
+
for (file_index, offset, doc_end), ids in zip(meta, result.get()):
|
|
82
|
+
yield file_index, offset, ids, doc_end
|
|
83
|
+
finally:
|
|
84
|
+
pending.clear()
|
|
85
|
+
pool.terminate()
|
|
86
|
+
pool.join()
|
src/data/prefetch.py
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
import queue
|
|
2
|
+
import threading
|
|
3
|
+
|
|
4
|
+
_STOP = object()
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class PrefetchLoader:
|
|
8
|
+
|
|
9
|
+
def __init__(self, iterable_factory, prefetch_size=4):
|
|
10
|
+
self.iterable_factory = iterable_factory
|
|
11
|
+
self.prefetch_size = max(1, prefetch_size)
|
|
12
|
+
self._queue = None
|
|
13
|
+
self._thread = None
|
|
14
|
+
self._error = None
|
|
15
|
+
self._halt = threading.Event()
|
|
16
|
+
|
|
17
|
+
def _put(self, item):
|
|
18
|
+
while not self._halt.is_set():
|
|
19
|
+
try:
|
|
20
|
+
self._queue.put(item, timeout=0.1)
|
|
21
|
+
return True
|
|
22
|
+
except queue.Full:
|
|
23
|
+
continue
|
|
24
|
+
|
|
25
|
+
return False
|
|
26
|
+
|
|
27
|
+
def _worker(self):
|
|
28
|
+
try:
|
|
29
|
+
for item in self.iterable_factory():
|
|
30
|
+
if not self._put(item):
|
|
31
|
+
return
|
|
32
|
+
except Exception as exc:
|
|
33
|
+
self._error = exc
|
|
34
|
+
finally:
|
|
35
|
+
self._put(_STOP)
|
|
36
|
+
|
|
37
|
+
def __iter__(self):
|
|
38
|
+
self._queue = queue.Queue(maxsize=self.prefetch_size)
|
|
39
|
+
self._error = None
|
|
40
|
+
self._halt.clear()
|
|
41
|
+
|
|
42
|
+
self._thread = threading.Thread(target=self._worker, daemon=True)
|
|
43
|
+
self._thread.start()
|
|
44
|
+
|
|
45
|
+
try:
|
|
46
|
+
while True:
|
|
47
|
+
item = self._queue.get()
|
|
48
|
+
|
|
49
|
+
if item is _STOP:
|
|
50
|
+
if self._error is not None:
|
|
51
|
+
raise self._error
|
|
52
|
+
return
|
|
53
|
+
|
|
54
|
+
yield item
|
|
55
|
+
finally:
|
|
56
|
+
self.stop()
|
|
57
|
+
|
|
58
|
+
def stop(self):
|
|
59
|
+
self._halt.set()
|
|
60
|
+
|
|
61
|
+
if self._thread is not None and self._thread.is_alive() and self._thread is not threading.current_thread():
|
|
62
|
+
self._thread.join(timeout=1.0)
|