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
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
@@ -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)