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