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
cli.py
ADDED
|
@@ -0,0 +1,604 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import json
|
|
3
|
+
import os
|
|
4
|
+
import random
|
|
5
|
+
import sys
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
from src.config import TrainConfig
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
from src.data.corpus import CorpusIndex
|
|
13
|
+
from src.data.shard_builder import build_shards
|
|
14
|
+
from src.data.sharded_dataset import ShardedTokenDataset, is_shard_dir
|
|
15
|
+
from src.data.streaming_dataset import StreamingTextDataset
|
|
16
|
+
from src.data.validation import validate_corpus, validate_shards
|
|
17
|
+
from src.models.gpt.config import GPTConfig
|
|
18
|
+
from src.models.gpt.model import GPTModel
|
|
19
|
+
from src.tokenization.bpe import PTFBPETokenizer
|
|
20
|
+
from src.tokenization.bytebpe import train_byte_bpe
|
|
21
|
+
from src.tokenization.registry import load_tokenizer
|
|
22
|
+
from src.training.trainer import Trainer
|
|
23
|
+
|
|
24
|
+
NOT_YET_IMPLEMENTED = {}
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _build_model(config, tokenizer):
|
|
28
|
+
random.seed(config.training.seed)
|
|
29
|
+
np.random.seed(config.training.seed)
|
|
30
|
+
|
|
31
|
+
gpt_config = GPTConfig(
|
|
32
|
+
vocab_size=tokenizer.vocab_size,
|
|
33
|
+
context_length=config.model.context_length,
|
|
34
|
+
d_model=config.model.d_model,
|
|
35
|
+
n_layers=config.model.n_layers,
|
|
36
|
+
n_heads=config.model.n_heads,
|
|
37
|
+
ff_dim=config.model.ff_dim,
|
|
38
|
+
activation=config.model.activation,
|
|
39
|
+
dropout=config.model.dropout,
|
|
40
|
+
norm_eps=config.model.norm_eps,
|
|
41
|
+
tie_weights=config.model.tie_weights,
|
|
42
|
+
position_encoding=config.model.position_encoding,
|
|
43
|
+
rope_theta=config.model.rope_theta,
|
|
44
|
+
rope_scaling=config.model.rope_scaling,
|
|
45
|
+
rope_scaling_factor=config.model.rope_scaling_factor,
|
|
46
|
+
trained_context_length=config.model.trained_context_length,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
model = GPTModel(gpt_config)
|
|
50
|
+
model.build()
|
|
51
|
+
|
|
52
|
+
return model
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _text_dataset(corpus, tokenizer, config, workers=1):
|
|
56
|
+
return StreamingTextDataset(
|
|
57
|
+
corpus=corpus,
|
|
58
|
+
tokenizer=tokenizer,
|
|
59
|
+
sequence_length=config.data.sequence_length,
|
|
60
|
+
read_buffer_size=config.data.read_buffer_size,
|
|
61
|
+
insert_eos=config.data.insert_eos,
|
|
62
|
+
text_field=config.data.text_field,
|
|
63
|
+
workers=workers,
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _dataset_for(paths, tokenizer, config, shuffle=False, seed=0, workers=1):
|
|
68
|
+
if isinstance(paths, str):
|
|
69
|
+
paths = [paths]
|
|
70
|
+
|
|
71
|
+
if len(paths) == 1 and is_shard_dir(paths[0]):
|
|
72
|
+
return ShardedTokenDataset(paths[0], config.data.sequence_length, tokenizer)
|
|
73
|
+
|
|
74
|
+
return _text_dataset(CorpusIndex(paths, shuffle=shuffle, seed=seed), tokenizer, config, workers=workers)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def _chat_dataset(corpus, tokenizer, config):
|
|
78
|
+
from src.data.chat_dataset import ChatSFTDataset
|
|
79
|
+
from src.inference.chat_template import ChatTemplate
|
|
80
|
+
|
|
81
|
+
template = ChatTemplate.resolve(config.data.chat_template).bind(tokenizer)
|
|
82
|
+
return ChatSFTDataset(corpus, template, config.data.sequence_length, packing=config.data.chat_packing,
|
|
83
|
+
messages_field=config.data.messages_field)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _build_chat_datasets(config, tokenizer):
|
|
87
|
+
seed = config.training.seed
|
|
88
|
+
|
|
89
|
+
if config.data.workers > 1:
|
|
90
|
+
raise ValueError("data.workers applies to raw-text training; chat datasets tokenize in-process")
|
|
91
|
+
|
|
92
|
+
corpus = CorpusIndex(config.data.train, extensions=(".jsonl",), shuffle=config.data.shuffle_files, seed=seed)
|
|
93
|
+
|
|
94
|
+
if config.data.validation:
|
|
95
|
+
val = CorpusIndex(config.data.validation, extensions=(".jsonl",))
|
|
96
|
+
return _chat_dataset(corpus, tokenizer, config), _chat_dataset(val, tokenizer, config)
|
|
97
|
+
|
|
98
|
+
if config.data.val_split_fraction:
|
|
99
|
+
base = CorpusIndex(config.data.train, extensions=(".jsonl",), shuffle=False, seed=seed)
|
|
100
|
+
train_corpus, val_corpus = base.split(config.data.val_split_fraction, seed)
|
|
101
|
+
return _chat_dataset(train_corpus, tokenizer, config), _chat_dataset(val_corpus, tokenizer, config)
|
|
102
|
+
|
|
103
|
+
return _chat_dataset(corpus, tokenizer, config), None
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def _build_datasets(config, tokenizer):
|
|
107
|
+
if config.data.format not in ("text", "chat"):
|
|
108
|
+
raise ValueError(f"data.format must be 'text' or 'chat', got '{config.data.format}'")
|
|
109
|
+
|
|
110
|
+
if config.data.format == "chat":
|
|
111
|
+
return _build_chat_datasets(config, tokenizer)
|
|
112
|
+
|
|
113
|
+
seed = config.training.seed
|
|
114
|
+
train_paths = config.data.train
|
|
115
|
+
|
|
116
|
+
if config.data.validation:
|
|
117
|
+
train_dataset = _dataset_for(train_paths, tokenizer, config, config.data.shuffle_files, seed,
|
|
118
|
+
workers=config.data.workers)
|
|
119
|
+
val_dataset = _dataset_for(config.data.validation, tokenizer, config)
|
|
120
|
+
return train_dataset, val_dataset
|
|
121
|
+
|
|
122
|
+
if config.data.val_split_fraction:
|
|
123
|
+
if len(train_paths) == 1 and is_shard_dir(train_paths[0]):
|
|
124
|
+
raise ValueError("val_split_fraction applies to raw corpora; build separate shard sets for validation")
|
|
125
|
+
|
|
126
|
+
corpus = CorpusIndex(train_paths, shuffle=False, seed=seed)
|
|
127
|
+
train_corpus, val_corpus = corpus.split(config.data.val_split_fraction, seed)
|
|
128
|
+
|
|
129
|
+
if config.data.shuffle_files:
|
|
130
|
+
import random
|
|
131
|
+
random.Random(seed).shuffle(train_corpus.files)
|
|
132
|
+
train_corpus.shuffle = True
|
|
133
|
+
|
|
134
|
+
return (_text_dataset(train_corpus, tokenizer, config, workers=config.data.workers),
|
|
135
|
+
_text_dataset(val_corpus, tokenizer, config))
|
|
136
|
+
|
|
137
|
+
return _dataset_for(train_paths, tokenizer, config, config.data.shuffle_files, seed,
|
|
138
|
+
workers=config.data.workers), None
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def cmd_train(args):
|
|
142
|
+
config = TrainConfig.load(args.config)
|
|
143
|
+
tokenizer = load_tokenizer(config.data.tokenizer)
|
|
144
|
+
|
|
145
|
+
model = _build_model(config, tokenizer)
|
|
146
|
+
train_dataset, val_dataset = _build_datasets(config, tokenizer)
|
|
147
|
+
|
|
148
|
+
trainer = Trainer(model, train_dataset, config, tokenizer, val_dataset=val_dataset, save_on_exit=not getattr(args, 'no_save_on_exit', False))
|
|
149
|
+
|
|
150
|
+
print(f"model parameters: {model.num_parameters():,}")
|
|
151
|
+
|
|
152
|
+
resumed = False
|
|
153
|
+
|
|
154
|
+
if args.resume:
|
|
155
|
+
resumed = trainer.load_checkpoint()
|
|
156
|
+
|
|
157
|
+
if resumed:
|
|
158
|
+
print(f"resumed from step {trainer.global_step}")
|
|
159
|
+
else:
|
|
160
|
+
print("no checkpoint found, starting from scratch")
|
|
161
|
+
|
|
162
|
+
if not resumed and config.training.init_from:
|
|
163
|
+
trainer.initialize_from(config.training.init_from)
|
|
164
|
+
|
|
165
|
+
trainer.train()
|
|
166
|
+
|
|
167
|
+
return 0
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def cmd_resume(args):
|
|
171
|
+
config = TrainConfig.load(args.config)
|
|
172
|
+
tokenizer = load_tokenizer(config.data.tokenizer)
|
|
173
|
+
|
|
174
|
+
model = _build_model(config, tokenizer)
|
|
175
|
+
train_dataset, val_dataset = _build_datasets(config, tokenizer)
|
|
176
|
+
|
|
177
|
+
trainer = Trainer(model, train_dataset, config, tokenizer, val_dataset=val_dataset, save_on_exit=not getattr(args, 'no_save_on_exit', False))
|
|
178
|
+
|
|
179
|
+
if not trainer.load_checkpoint(args.checkpoint if args.checkpoint != "latest" else None):
|
|
180
|
+
print("no checkpoint found", file=sys.stderr)
|
|
181
|
+
return 1
|
|
182
|
+
|
|
183
|
+
print(f"resumed from step {trainer.global_step}, {trainer.tokens_processed:,} tokens")
|
|
184
|
+
trainer.train()
|
|
185
|
+
|
|
186
|
+
return 0
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def cmd_evaluate(args):
|
|
190
|
+
config = TrainConfig.load(args.config)
|
|
191
|
+
tokenizer = load_tokenizer(config.data.tokenizer)
|
|
192
|
+
|
|
193
|
+
model = _build_model(config, tokenizer)
|
|
194
|
+
train_dataset, val_dataset = _build_datasets(config, tokenizer)
|
|
195
|
+
|
|
196
|
+
if val_dataset is None:
|
|
197
|
+
print("no validation corpus configured", file=sys.stderr)
|
|
198
|
+
return 1
|
|
199
|
+
|
|
200
|
+
trainer = Trainer(model, train_dataset, config, tokenizer, val_dataset=val_dataset, save_on_exit=not getattr(args, 'no_save_on_exit', False))
|
|
201
|
+
|
|
202
|
+
if not trainer.load_checkpoint(args.checkpoint if args.checkpoint != "latest" else None):
|
|
203
|
+
print("no checkpoint found", file=sys.stderr)
|
|
204
|
+
return 1
|
|
205
|
+
|
|
206
|
+
result = trainer.evaluate()
|
|
207
|
+
|
|
208
|
+
if result is None:
|
|
209
|
+
print("evaluation produced no batches", file=sys.stderr)
|
|
210
|
+
return 1
|
|
211
|
+
|
|
212
|
+
return 0
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def cmd_inspect(args):
|
|
216
|
+
from src.training.checkpoint_manager import CheckpointManager
|
|
217
|
+
|
|
218
|
+
manager = CheckpointManager(args.checkpoint_dir if args.checkpoint_dir else ".")
|
|
219
|
+
payload = manager.load(args.path if args.path != "latest" else None)
|
|
220
|
+
|
|
221
|
+
if payload is None:
|
|
222
|
+
print("checkpoint not found", file=sys.stderr)
|
|
223
|
+
return 1
|
|
224
|
+
|
|
225
|
+
print(f"framework_version {payload['framework_version']}")
|
|
226
|
+
print(f"global_step {payload['global_step']}")
|
|
227
|
+
print(f"tokens_processed {payload['tokens_processed']:,}")
|
|
228
|
+
print(f"examples_processed {payload['examples_processed']:,}")
|
|
229
|
+
print(f"last_loss {payload['last_loss']}")
|
|
230
|
+
print(f"model_config {payload['model_config']}")
|
|
231
|
+
print(f"tokenizer_identity {payload['tokenizer_identity']}")
|
|
232
|
+
|
|
233
|
+
return 0
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def _sample_texts(corpus, text_field, read_buffer_size, max_chars):
|
|
237
|
+
from src.data.document_stream import DocumentReader
|
|
238
|
+
|
|
239
|
+
reader = DocumentReader(corpus.files, read_buffer_size=read_buffer_size, text_field=text_field)
|
|
240
|
+
per_file = max(max_chars // max(len(corpus.files), 1), 1 << 20)
|
|
241
|
+
total = 0
|
|
242
|
+
|
|
243
|
+
for path in corpus.files:
|
|
244
|
+
used = 0
|
|
245
|
+
|
|
246
|
+
for text, _, _ in reader.read_file(path):
|
|
247
|
+
yield text
|
|
248
|
+
used += len(text)
|
|
249
|
+
total += len(text)
|
|
250
|
+
|
|
251
|
+
if used >= per_file or total >= max_chars:
|
|
252
|
+
break
|
|
253
|
+
|
|
254
|
+
if total >= max_chars:
|
|
255
|
+
return
|
|
256
|
+
|
|
257
|
+
|
|
258
|
+
def cmd_tokenize(args):
|
|
259
|
+
corpus = CorpusIndex(args.corpus)
|
|
260
|
+
|
|
261
|
+
if not corpus.files:
|
|
262
|
+
print("no supported files found", file=sys.stderr)
|
|
263
|
+
return 1
|
|
264
|
+
|
|
265
|
+
texts = _sample_texts(corpus, args.text_field, args.read_buffer_size, args.max_chars)
|
|
266
|
+
|
|
267
|
+
if args.type == "word":
|
|
268
|
+
tokenizer = PTFBPETokenizer(vocab_size=args.vocab_size, lowercase=not args.cased)
|
|
269
|
+
tokenizer.fit(texts)
|
|
270
|
+
else:
|
|
271
|
+
def progress(done, total):
|
|
272
|
+
print(f" merges {done}/{total}", file=sys.stderr)
|
|
273
|
+
|
|
274
|
+
tokenizer = train_byte_bpe(
|
|
275
|
+
texts,
|
|
276
|
+
args.vocab_size,
|
|
277
|
+
special_tokens=args.special_token or [],
|
|
278
|
+
min_frequency=args.min_frequency,
|
|
279
|
+
max_unique_words=args.max_unique_words,
|
|
280
|
+
progress=progress,
|
|
281
|
+
)
|
|
282
|
+
|
|
283
|
+
if tokenizer.vocab_size < args.vocab_size:
|
|
284
|
+
print(
|
|
285
|
+
f"warning: corpus sample only supported {tokenizer.vocab_size} tokens "
|
|
286
|
+
f"(requested {args.vocab_size}); use more data or lower --min-frequency",
|
|
287
|
+
file=sys.stderr,
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
tokenizer.save(args.output)
|
|
291
|
+
|
|
292
|
+
print(f"tokenizer trained: type={tokenizer.identity['type']} vocab_size={tokenizer.vocab_size} -> {args.output}")
|
|
293
|
+
|
|
294
|
+
return 0
|
|
295
|
+
|
|
296
|
+
|
|
297
|
+
def cmd_prepare(args):
|
|
298
|
+
corpus = CorpusIndex(args.corpus)
|
|
299
|
+
manifest = build_shards(
|
|
300
|
+
corpus,
|
|
301
|
+
args.tokenizer,
|
|
302
|
+
args.output,
|
|
303
|
+
shard_tokens=args.shard_tokens,
|
|
304
|
+
workers=args.workers,
|
|
305
|
+
insert_eos=not args.no_eos,
|
|
306
|
+
text_field=args.text_field,
|
|
307
|
+
read_buffer_size=args.read_buffer_size,
|
|
308
|
+
)
|
|
309
|
+
|
|
310
|
+
print(
|
|
311
|
+
f"built {len(manifest['shards'])} shards, {manifest['total_tokens']:,} tokens "
|
|
312
|
+
f"in {manifest['build_seconds']:.1f}s -> {args.output}"
|
|
313
|
+
)
|
|
314
|
+
|
|
315
|
+
return 0
|
|
316
|
+
|
|
317
|
+
|
|
318
|
+
def cmd_validate(args):
|
|
319
|
+
tokenizer = load_tokenizer(args.tokenizer) if args.tokenizer else None
|
|
320
|
+
|
|
321
|
+
if is_shard_dir(args.path):
|
|
322
|
+
report = validate_shards(args.path, tokenizer, verify_checksums=args.checksums)
|
|
323
|
+
else:
|
|
324
|
+
report = validate_corpus(CorpusIndex(args.path), args.text_field, tokenizer=tokenizer)
|
|
325
|
+
|
|
326
|
+
print(json.dumps(report.to_dict(), indent=2))
|
|
327
|
+
|
|
328
|
+
return 0 if report.ok else 1
|
|
329
|
+
|
|
330
|
+
|
|
331
|
+
def cmd_export(args):
|
|
332
|
+
from src.inference.config import GenerationConfig
|
|
333
|
+
from src.inference.export import export_model
|
|
334
|
+
|
|
335
|
+
generation = None
|
|
336
|
+
|
|
337
|
+
if args.temperature is not None or args.top_k is not None or args.top_p is not None or args.max_new_tokens:
|
|
338
|
+
generation = GenerationConfig(
|
|
339
|
+
max_new_tokens=args.max_new_tokens or 128,
|
|
340
|
+
temperature=0.8 if args.temperature is None else args.temperature,
|
|
341
|
+
top_k=40 if args.top_k is None else args.top_k,
|
|
342
|
+
top_p=0.95 if args.top_p is None else args.top_p,
|
|
343
|
+
)
|
|
344
|
+
|
|
345
|
+
export_model(args.checkpoint, args.output, args.tokenizer, generation, dtype=args.dtype,
|
|
346
|
+
chat_template=args.chat_template, context_length=args.extend_context,
|
|
347
|
+
context_extension=args.context_extension)
|
|
348
|
+
print(f"exported model to {args.output}")
|
|
349
|
+
|
|
350
|
+
return 0
|
|
351
|
+
|
|
352
|
+
|
|
353
|
+
def cmd_generate(args):
|
|
354
|
+
import sys as _sys
|
|
355
|
+
|
|
356
|
+
from src.inference.runtime import load_model
|
|
357
|
+
|
|
358
|
+
model = load_model(args.model, cache_budget_bytes=args.cache_budget_mb * (1 << 20) if args.cache_budget_mb is not None else None)
|
|
359
|
+
|
|
360
|
+
overrides = {
|
|
361
|
+
"max_new_tokens": args.max_new_tokens,
|
|
362
|
+
"temperature": args.temperature,
|
|
363
|
+
"top_k": args.top_k,
|
|
364
|
+
"top_p": args.top_p,
|
|
365
|
+
"repetition_penalty": args.repetition_penalty,
|
|
366
|
+
"seed": args.seed,
|
|
367
|
+
}
|
|
368
|
+
|
|
369
|
+
if args.greedy:
|
|
370
|
+
overrides["do_sample"] = False
|
|
371
|
+
|
|
372
|
+
if args.stop:
|
|
373
|
+
overrides["stop_strings"] = args.stop
|
|
374
|
+
|
|
375
|
+
prompt = args.prompt if args.prompt is not None else ""
|
|
376
|
+
last = None
|
|
377
|
+
|
|
378
|
+
for ev in model.generate_stream(prompt, timeout_s=args.timeout, truncate_prompt=True, events=True, **overrides):
|
|
379
|
+
if ev.text:
|
|
380
|
+
_sys.stdout.write(ev.text)
|
|
381
|
+
_sys.stdout.flush()
|
|
382
|
+
last = ev
|
|
383
|
+
|
|
384
|
+
_sys.stdout.write("\n")
|
|
385
|
+
|
|
386
|
+
if last is not None and last.error:
|
|
387
|
+
print(f"generation failed: {last.error}", file=_sys.stderr)
|
|
388
|
+
return 1
|
|
389
|
+
|
|
390
|
+
if last is not None and last.usage and args.verbose:
|
|
391
|
+
print(json.dumps({"finish_reason": last.finish_reason, **last.usage}), file=_sys.stderr)
|
|
392
|
+
|
|
393
|
+
return 0
|
|
394
|
+
|
|
395
|
+
|
|
396
|
+
def _serve_config(args):
|
|
397
|
+
from src.serving.config import ModelEntry, load_server_config, server_config_from_dict
|
|
398
|
+
|
|
399
|
+
if args.config:
|
|
400
|
+
if args.model:
|
|
401
|
+
raise ValueError("pass either a model directory or --config, not both")
|
|
402
|
+
config = load_server_config(args.config)
|
|
403
|
+
elif args.model:
|
|
404
|
+
name = args.name or os.path.basename(os.path.normpath(args.model)) or "model"
|
|
405
|
+
config = server_config_from_dict({})
|
|
406
|
+
config.models = [ModelEntry(
|
|
407
|
+
name=name,
|
|
408
|
+
path=os.path.abspath(args.model),
|
|
409
|
+
chat_template=args.chat_template,
|
|
410
|
+
max_batch_size=args.max_batch_size or 8,
|
|
411
|
+
kv_cache_budget_mb=args.kv_cache_mb,
|
|
412
|
+
)]
|
|
413
|
+
else:
|
|
414
|
+
raise ValueError("pass a model directory or --config")
|
|
415
|
+
|
|
416
|
+
if args.host is not None:
|
|
417
|
+
config.host = args.host
|
|
418
|
+
if args.port is not None:
|
|
419
|
+
config.port = args.port
|
|
420
|
+
if args.api_key_file:
|
|
421
|
+
config.security.api_keys_file = args.api_key_file
|
|
422
|
+
if args.admin_key_file:
|
|
423
|
+
config.security.admin_keys_file = args.admin_key_file
|
|
424
|
+
if args.cors_origin:
|
|
425
|
+
config.security.cors_origins = list(args.cors_origin)
|
|
426
|
+
if args.insecure_no_auth:
|
|
427
|
+
config.security.allow_unauthenticated = True
|
|
428
|
+
if args.max_concurrent:
|
|
429
|
+
config.limits.max_concurrent_requests = args.max_concurrent
|
|
430
|
+
if args.max_tokens:
|
|
431
|
+
config.limits.max_generation_tokens = args.max_tokens
|
|
432
|
+
if args.timeout:
|
|
433
|
+
config.limits.request_timeout_s = args.timeout
|
|
434
|
+
if args.memory_limit_mb:
|
|
435
|
+
config.runtime.memory_limit_mb = args.memory_limit_mb
|
|
436
|
+
if args.device:
|
|
437
|
+
config.runtime.device = args.device
|
|
438
|
+
if args.access_log is not None:
|
|
439
|
+
config.logging.access_log = args.access_log or None
|
|
440
|
+
if args.no_ui:
|
|
441
|
+
config.ui.enabled = False
|
|
442
|
+
|
|
443
|
+
return config.validate()
|
|
444
|
+
|
|
445
|
+
|
|
446
|
+
def cmd_serve(args):
|
|
447
|
+
import logging
|
|
448
|
+
|
|
449
|
+
from src.serving.server import APIServer, InsecureConfiguration
|
|
450
|
+
|
|
451
|
+
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
|
|
452
|
+
|
|
453
|
+
try:
|
|
454
|
+
config = _serve_config(args)
|
|
455
|
+
server = APIServer(config)
|
|
456
|
+
except InsecureConfiguration as exc:
|
|
457
|
+
print(f"error: {exc}", file=sys.stderr)
|
|
458
|
+
return 2
|
|
459
|
+
except (ValueError, OSError) as exc:
|
|
460
|
+
print(f"error: invalid server configuration: {exc}", file=sys.stderr)
|
|
461
|
+
return 2
|
|
462
|
+
|
|
463
|
+
try:
|
|
464
|
+
server.serve_forever()
|
|
465
|
+
except OSError as exc:
|
|
466
|
+
print(f"error: could not start server: {exc}", file=sys.stderr)
|
|
467
|
+
return 1
|
|
468
|
+
except Exception as exc:
|
|
469
|
+
print(f"error: server failed: {exc}", file=sys.stderr)
|
|
470
|
+
return 1
|
|
471
|
+
|
|
472
|
+
return 0
|
|
473
|
+
|
|
474
|
+
|
|
475
|
+
def cmd_unimplemented(name):
|
|
476
|
+
def handler(args):
|
|
477
|
+
print(f"'{name}' is not implemented yet: {NOT_YET_IMPLEMENTED[name]}", file=sys.stderr)
|
|
478
|
+
return 2
|
|
479
|
+
|
|
480
|
+
return handler
|
|
481
|
+
|
|
482
|
+
|
|
483
|
+
def build_parser():
|
|
484
|
+
parser = argparse.ArgumentParser(prog="pytensorforge")
|
|
485
|
+
sub = parser.add_subparsers(dest="command", required=True)
|
|
486
|
+
|
|
487
|
+
p_train = sub.add_parser("train")
|
|
488
|
+
p_train.add_argument("config")
|
|
489
|
+
p_train.add_argument("--resume", action="store_true")
|
|
490
|
+
p_train.add_argument("--no-save-on-exit", action="store_true")
|
|
491
|
+
p_train.set_defaults(func=cmd_train)
|
|
492
|
+
|
|
493
|
+
p_resume = sub.add_parser("resume")
|
|
494
|
+
p_resume.add_argument("checkpoint")
|
|
495
|
+
p_resume.add_argument("--config", required=True)
|
|
496
|
+
p_resume.add_argument("--no-save-on-exit", action="store_true")
|
|
497
|
+
p_resume.set_defaults(func=cmd_resume)
|
|
498
|
+
|
|
499
|
+
p_eval = sub.add_parser("evaluate")
|
|
500
|
+
p_eval.add_argument("checkpoint")
|
|
501
|
+
p_eval.add_argument("--config", required=True)
|
|
502
|
+
p_eval.set_defaults(func=cmd_evaluate)
|
|
503
|
+
|
|
504
|
+
p_inspect = sub.add_parser("inspect")
|
|
505
|
+
p_inspect.add_argument("path")
|
|
506
|
+
p_inspect.add_argument("--checkpoint-dir", default=None)
|
|
507
|
+
p_inspect.set_defaults(func=cmd_inspect)
|
|
508
|
+
|
|
509
|
+
p_tokenize = sub.add_parser("tokenize")
|
|
510
|
+
p_tokenize.add_argument("corpus")
|
|
511
|
+
p_tokenize.add_argument("--output", required=True)
|
|
512
|
+
p_tokenize.add_argument("--vocab-size", type=int, default=32000)
|
|
513
|
+
p_tokenize.add_argument("--type", choices=["bytebpe", "word"], default="bytebpe")
|
|
514
|
+
p_tokenize.add_argument("--cased", action="store_true")
|
|
515
|
+
p_tokenize.add_argument("--special-token", action="append", default=None)
|
|
516
|
+
p_tokenize.add_argument("--min-frequency", type=int, default=2)
|
|
517
|
+
p_tokenize.add_argument("--max-unique-words", type=int, default=1_000_000)
|
|
518
|
+
p_tokenize.add_argument("--text-field", default="text")
|
|
519
|
+
p_tokenize.add_argument("--read-buffer-size", type=int, default=1 << 20)
|
|
520
|
+
p_tokenize.add_argument("--max-chars", type=int, default=50_000_000)
|
|
521
|
+
p_tokenize.set_defaults(func=cmd_tokenize)
|
|
522
|
+
|
|
523
|
+
p_export = sub.add_parser("export")
|
|
524
|
+
p_export.add_argument("checkpoint")
|
|
525
|
+
p_export.add_argument("--output", required=True)
|
|
526
|
+
p_export.add_argument("--tokenizer", required=True)
|
|
527
|
+
p_export.add_argument("--dtype", choices=["float32", "float16"], default="float32")
|
|
528
|
+
p_export.add_argument("--max-new-tokens", type=int, default=None)
|
|
529
|
+
p_export.add_argument("--temperature", type=float, default=None)
|
|
530
|
+
p_export.add_argument("--top-k", type=int, default=None)
|
|
531
|
+
p_export.add_argument("--top-p", type=float, default=None)
|
|
532
|
+
p_export.add_argument("--chat-template", default=None)
|
|
533
|
+
p_export.add_argument("--extend-context", type=int, default=None)
|
|
534
|
+
p_export.add_argument("--context-extension", choices=["extrapolate", "linear", "ntk"], default=None)
|
|
535
|
+
p_export.set_defaults(func=cmd_export)
|
|
536
|
+
|
|
537
|
+
p_generate = sub.add_parser("generate")
|
|
538
|
+
p_generate.add_argument("model")
|
|
539
|
+
p_generate.add_argument("--prompt", default=None)
|
|
540
|
+
p_generate.add_argument("--max-new-tokens", type=int, default=None)
|
|
541
|
+
p_generate.add_argument("--temperature", type=float, default=None)
|
|
542
|
+
p_generate.add_argument("--top-k", type=int, default=None)
|
|
543
|
+
p_generate.add_argument("--top-p", type=float, default=None)
|
|
544
|
+
p_generate.add_argument("--repetition-penalty", type=float, default=None)
|
|
545
|
+
p_generate.add_argument("--seed", type=int, default=None)
|
|
546
|
+
p_generate.add_argument("--greedy", action="store_true")
|
|
547
|
+
p_generate.add_argument("--stop", action="append", default=None)
|
|
548
|
+
p_generate.add_argument("--timeout", type=float, default=None)
|
|
549
|
+
p_generate.add_argument("--cache-budget-mb", type=int, default=None)
|
|
550
|
+
p_generate.add_argument("--verbose", action="store_true")
|
|
551
|
+
p_generate.set_defaults(func=cmd_generate)
|
|
552
|
+
|
|
553
|
+
p_prepare = sub.add_parser("prepare-dataset")
|
|
554
|
+
p_prepare.add_argument("corpus")
|
|
555
|
+
p_prepare.add_argument("--tokenizer", required=True)
|
|
556
|
+
p_prepare.add_argument("--output", required=True)
|
|
557
|
+
p_prepare.add_argument("--shard-tokens", type=int, default=50_000_000)
|
|
558
|
+
p_prepare.add_argument("--workers", type=int, default=1)
|
|
559
|
+
p_prepare.add_argument("--text-field", default="text")
|
|
560
|
+
p_prepare.add_argument("--read-buffer-size", type=int, default=1 << 20)
|
|
561
|
+
p_prepare.add_argument("--no-eos", action="store_true")
|
|
562
|
+
p_prepare.set_defaults(func=cmd_prepare)
|
|
563
|
+
|
|
564
|
+
p_validate = sub.add_parser("validate-dataset")
|
|
565
|
+
p_validate.add_argument("path")
|
|
566
|
+
p_validate.add_argument("--tokenizer", default=None)
|
|
567
|
+
p_validate.add_argument("--text-field", default="text")
|
|
568
|
+
p_validate.add_argument("--checksums", action="store_true")
|
|
569
|
+
p_validate.set_defaults(func=cmd_validate)
|
|
570
|
+
|
|
571
|
+
p_serve = sub.add_parser("serve")
|
|
572
|
+
p_serve.add_argument("model", nargs="?", default=None)
|
|
573
|
+
p_serve.add_argument("--config", default=None)
|
|
574
|
+
p_serve.add_argument("--name", default=None)
|
|
575
|
+
p_serve.add_argument("--host", default=None)
|
|
576
|
+
p_serve.add_argument("--port", type=int, default=None)
|
|
577
|
+
p_serve.add_argument("--chat-template", default=None)
|
|
578
|
+
p_serve.add_argument("--max-batch-size", type=int, default=None)
|
|
579
|
+
p_serve.add_argument("--kv-cache-mb", type=float, default=None)
|
|
580
|
+
p_serve.add_argument("--max-concurrent", type=int, default=None)
|
|
581
|
+
p_serve.add_argument("--max-tokens", type=int, default=None)
|
|
582
|
+
p_serve.add_argument("--timeout", type=float, default=None)
|
|
583
|
+
p_serve.add_argument("--memory-limit-mb", type=float, default=None)
|
|
584
|
+
p_serve.add_argument("--device", default=None)
|
|
585
|
+
p_serve.add_argument("--api-key-file", default=None)
|
|
586
|
+
p_serve.add_argument("--admin-key-file", default=None)
|
|
587
|
+
p_serve.add_argument("--cors-origin", action="append", default=None)
|
|
588
|
+
p_serve.add_argument("--access-log", default=None)
|
|
589
|
+
p_serve.add_argument("--insecure-no-auth", action="store_true")
|
|
590
|
+
p_serve.add_argument("--no-ui", action="store_true")
|
|
591
|
+
p_serve.set_defaults(func=cmd_serve)
|
|
592
|
+
|
|
593
|
+
return parser
|
|
594
|
+
|
|
595
|
+
|
|
596
|
+
def main():
|
|
597
|
+
parser = build_parser()
|
|
598
|
+
args = parser.parse_args()
|
|
599
|
+
|
|
600
|
+
sys.exit(args.func(args))
|
|
601
|
+
|
|
602
|
+
|
|
603
|
+
if __name__ == "__main__":
|
|
604
|
+
main()
|