pytensorforge 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (146) hide show
  1. cli.py +604 -0
  2. pytensorforge-0.1.0.dist-info/METADATA +103 -0
  3. pytensorforge-0.1.0.dist-info/RECORD +146 -0
  4. pytensorforge-0.1.0.dist-info/WHEEL +5 -0
  5. pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
  6. pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
  7. src/__init__.py +0 -0
  8. src/activations/Activation.py +4 -0
  9. src/activations/ELU.py +11 -0
  10. src/activations/GELU.py +6 -0
  11. src/activations/ReLU.py +27 -0
  12. src/activations/SELU.py +14 -0
  13. src/activations/Sigmoid.py +27 -0
  14. src/activations/Softmax.py +84 -0
  15. src/activations/Tanh.py +29 -0
  16. src/activations/__init__.py +17 -0
  17. src/config.py +120 -0
  18. src/core/Matrix.py +3 -0
  19. src/core/Scalar.py +18 -0
  20. src/core/Tensor.py +866 -0
  21. src/core/Vector.py +31 -0
  22. src/core/__init__.py +0 -0
  23. src/data/__init__.py +0 -0
  24. src/data/chat_dataset.py +188 -0
  25. src/data/corpus.py +104 -0
  26. src/data/document_stream.py +178 -0
  27. src/data/parallel_encode.py +86 -0
  28. src/data/prefetch.py +62 -0
  29. src/data/shard_builder.py +119 -0
  30. src/data/shard_writer.py +81 -0
  31. src/data/sharded_dataset.py +112 -0
  32. src/data/streaming_dataset.py +132 -0
  33. src/data/validation.py +212 -0
  34. src/inference/__init__.py +0 -0
  35. src/inference/chat_template.py +384 -0
  36. src/inference/config.py +48 -0
  37. src/inference/engine.py +241 -0
  38. src/inference/export.py +133 -0
  39. src/inference/kv_cache.py +65 -0
  40. src/inference/runtime.py +161 -0
  41. src/inference/sampling.py +42 -0
  42. src/inference/scheduler.py +473 -0
  43. src/inference/text.py +67 -0
  44. src/initializers/Constant.py +9 -0
  45. src/initializers/GlorotNormal.py +15 -0
  46. src/initializers/GlorotUniform.py +26 -0
  47. src/initializers/HeNormal.py +15 -0
  48. src/initializers/HeUniform.py +14 -0
  49. src/initializers/Initializer.py +4 -0
  50. src/initializers/LecunNormal.py +16 -0
  51. src/initializers/LecunUniform.py +14 -0
  52. src/initializers/Ones.py +6 -0
  53. src/initializers/Orthogonal.py +14 -0
  54. src/initializers/RandomNormal.py +14 -0
  55. src/initializers/RandomUniform.py +14 -0
  56. src/initializers/Zeros.py +8 -0
  57. src/initializers/__init__.py +17 -0
  58. src/loss/CategoricalCrossEntropy.py +9 -0
  59. src/loss/CrossEntropyLoss.py +34 -0
  60. src/loss/CrossEntropyWithLogitsLoss.py +59 -0
  61. src/loss/Hinge.py +5 -0
  62. src/loss/Huber.py +22 -0
  63. src/loss/Loss.py +6 -0
  64. src/loss/MSE.py +7 -0
  65. src/loss/MSELoss.py +10 -0
  66. src/loss/SparseCategoricalCrossEntropy.py +15 -0
  67. src/loss/__init__.py +18 -0
  68. src/loss/bce.py +34 -0
  69. src/loss/mae.py +16 -0
  70. src/math/__init__.py +0 -0
  71. src/math/clip.py +37 -0
  72. src/math/exp.py +27 -0
  73. src/math/log.py +25 -0
  74. src/math/sigmoid.py +5 -0
  75. src/models/__init__.py +0 -0
  76. src/models/embedding/Embedding.py +65 -0
  77. src/models/embedding/__init__.py +0 -0
  78. src/models/gpt/__init__.py +0 -0
  79. src/models/gpt/attention.py +158 -0
  80. src/models/gpt/block.py +74 -0
  81. src/models/gpt/config.py +103 -0
  82. src/models/gpt/context.py +44 -0
  83. src/models/gpt/model.py +165 -0
  84. src/models/gpt/recompute.py +35 -0
  85. src/models/gpt/rope.py +84 -0
  86. src/models/regression/Linear.py +51 -0
  87. src/models/regression/Logistic.py +36 -0
  88. src/models/regression/__init__.py +0 -0
  89. src/models/seq/Sequential.py +297 -0
  90. src/models/seq/__init__.py +0 -0
  91. src/models/svm/__init__.py +0 -0
  92. src/models/tokenizer/BPETokenizer.py +228 -0
  93. src/models/tokenizer/__init__.py +0 -0
  94. src/models/transformers/Dropout.py +35 -0
  95. src/models/transformers/LastToken.py +10 -0
  96. src/models/transformers/LayerNorm.py +54 -0
  97. src/models/transformers/Linear.py +18 -0
  98. src/models/transformers/MultiHeadAttention.py +130 -0
  99. src/models/transformers/TransformerBlock.py +79 -0
  100. src/models/transformers/__init__.py +0 -0
  101. src/neural/Dense.py +58 -0
  102. src/neural/LSTM.py +167 -0
  103. src/neural/Layer.py +72 -0
  104. src/neural/Parameter.py +30 -0
  105. src/neural/RNN.py +83 -0
  106. src/neural/__init__.py +0 -0
  107. src/ops/__init__.py +0 -0
  108. src/ops/stack.py +40 -0
  109. src/optimizers/Adagrad.py +31 -0
  110. src/optimizers/Adam.py +98 -0
  111. src/optimizers/AdamW.py +84 -0
  112. src/optimizers/Batch.py +11 -0
  113. src/optimizers/Nesterov.py +35 -0
  114. src/optimizers/Optimizer.py +18 -0
  115. src/optimizers/RMSProp.py +35 -0
  116. src/optimizers/SGD.py +30 -0
  117. src/optimizers/SGDMomentum.py +28 -0
  118. src/optimizers/__init__.py +9 -0
  119. src/scaling/StandardScaler.py +15 -0
  120. src/scaling/__init__.py +0 -0
  121. src/serialization/__init__.py +0 -0
  122. src/serialization/checkpoint.py +58 -0
  123. src/serialization/modelio.py +132 -0
  124. src/serving/__init__.py +0 -0
  125. src/serving/app.py +792 -0
  126. src/serving/config.py +216 -0
  127. src/serving/errors.py +51 -0
  128. src/serving/http.py +599 -0
  129. src/serving/metrics.py +293 -0
  130. src/serving/model_server.py +287 -0
  131. src/serving/protocol.py +377 -0
  132. src/serving/security.py +200 -0
  133. src/serving/server.py +121 -0
  134. src/tokenization/__init__.py +0 -0
  135. src/tokenization/base.py +75 -0
  136. src/tokenization/bpe.py +190 -0
  137. src/tokenization/bytebpe.py +476 -0
  138. src/tokenization/registry.py +28 -0
  139. src/training/__init__.py +0 -0
  140. src/training/checkpoint_manager.py +101 -0
  141. src/training/experiment.py +71 -0
  142. src/training/losses.py +42 -0
  143. src/training/precision.py +141 -0
  144. src/training/profiler.py +38 -0
  145. src/training/scheduler.py +50 -0
  146. src/training/trainer.py +594 -0
@@ -0,0 +1,133 @@
1
+ import hashlib
2
+ import json
3
+ import os
4
+ import shutil
5
+
6
+ import numpy as np
7
+
8
+ from src.inference.config import GenerationConfig
9
+ from src.models.gpt.config import GPTConfig
10
+ from src.models.gpt.context import extend_context
11
+ from src.tokenization.base import identity_matches
12
+ from src.training.checkpoint_manager import CheckpointManager
13
+
14
+ EXPORT_FORMAT_VERSION = 1
15
+ WEIGHTS_FILE = os.path.join("weights", "model.npz")
16
+
17
+
18
+ def flatten_state(state, prefix=""):
19
+ out = {}
20
+
21
+ for key, value in state.items():
22
+ if isinstance(value, dict):
23
+ out.update(flatten_state(value, f"{prefix}{key}."))
24
+ elif isinstance(value, list):
25
+ for i, item in enumerate(value):
26
+ out.update(flatten_state(item, f"{prefix}{key}.{i}."))
27
+ else:
28
+ out[f"{prefix}{key}"] = np.asarray(value)
29
+
30
+ return out
31
+
32
+
33
+ def _sha256(path):
34
+ digest = hashlib.sha256()
35
+
36
+ with open(path, "rb") as f:
37
+ for block in iter(lambda: f.read(1 << 20), b""):
38
+ digest.update(block)
39
+
40
+ return digest.hexdigest()
41
+
42
+
43
+ def _write_json(path, payload):
44
+ tmp = path + ".tmp"
45
+
46
+ with open(tmp, "w") as f:
47
+ json.dump(payload, f, indent=2, sort_keys=True)
48
+
49
+ os.replace(tmp, path)
50
+
51
+
52
+ def export_model(checkpoint_path, output_dir, tokenizer_path, generation=None, dtype="float32", chat_template=None,
53
+ context_length=None, context_extension=None):
54
+ if dtype not in ("float32", "float16"):
55
+ raise ValueError("dtype must be float32 or float16")
56
+
57
+ payload = CheckpointManager(os.path.dirname(os.path.abspath(checkpoint_path))).load(checkpoint_path)
58
+
59
+ if payload is None:
60
+ raise FileNotFoundError(f"checkpoint not found: {checkpoint_path}")
61
+
62
+ from src.tokenization.registry import load_tokenizer
63
+
64
+ tokenizer = load_tokenizer(tokenizer_path)
65
+
66
+ if not identity_matches(payload["tokenizer_identity"], tokenizer.identity):
67
+ raise ValueError("tokenizer does not match the one this checkpoint was trained with")
68
+
69
+ config = GPTConfig.from_dict(payload["model_config"])
70
+
71
+ if context_length is not None or context_extension is not None:
72
+ if context_length is None or context_extension is None:
73
+ raise ValueError("context extension needs both context_length and context_extension")
74
+
75
+ config = extend_context(config, context_length, context_extension)
76
+
77
+ from src.inference.chat_template import ChatTemplate
78
+
79
+ data_config = payload.get("data_config") or {}
80
+
81
+ if chat_template is None and data_config.get("format") == "chat":
82
+ chat_template = data_config.get("chat_template")
83
+
84
+ template = ChatTemplate.resolve(chat_template)
85
+
86
+ if template is not None:
87
+ template.bind(tokenizer)
88
+ weights = {k: v.astype(dtype) for k, v in flatten_state(payload["model_state"]).items()}
89
+
90
+ os.makedirs(os.path.join(output_dir, "weights"), exist_ok=True)
91
+ os.makedirs(os.path.join(output_dir, "tokenizer"), exist_ok=True)
92
+
93
+ weights_path = os.path.join(output_dir, WEIGHTS_FILE)
94
+ tmp_weights = weights_path + ".tmp"
95
+
96
+ with open(tmp_weights, "wb") as f:
97
+ np.savez(f, **weights)
98
+
99
+ os.replace(tmp_weights, weights_path)
100
+
101
+ shutil.copyfile(tokenizer_path, os.path.join(output_dir, "tokenizer", "tokenizer.json"))
102
+
103
+ generation = generation or GenerationConfig(
104
+ max_new_tokens=128, temperature=0.8, top_k=40, top_p=0.95,
105
+ )
106
+
107
+ _write_json(os.path.join(output_dir, "generation.json"), generation.to_dict())
108
+
109
+ meta = {
110
+ "format_version": EXPORT_FORMAT_VERSION,
111
+ "architecture": config.architecture,
112
+ "model": config.to_dict(),
113
+ "tokenizer": {"file": "tokenizer/tokenizer.json", "identity": tokenizer.identity},
114
+ "weights": {
115
+ "file": WEIGHTS_FILE,
116
+ "dtype": dtype,
117
+ "sha256": _sha256(weights_path),
118
+ "tensors": {k: list(v.shape) for k, v in weights.items()},
119
+ },
120
+ "source": {
121
+ "framework_version": payload.get("framework_version"),
122
+ "run_id": payload.get("run_id"),
123
+ "optimizer_steps": payload.get("optimizer_steps"),
124
+ "tokens_processed": payload.get("tokens_processed"),
125
+ },
126
+ }
127
+
128
+ if template is not None:
129
+ meta["chat_template"] = template.to_dict()
130
+
131
+ _write_json(os.path.join(output_dir, "config.json"), meta)
132
+
133
+ return output_dir
@@ -0,0 +1,65 @@
1
+ import numpy as np
2
+
3
+
4
+ class CacheFull(Exception):
5
+ pass
6
+
7
+
8
+ class KVCache:
9
+
10
+ def __init__(self, n_layers, n_heads, head_dim, capacity, dtype=np.float32):
11
+ self.capacity = capacity
12
+ self.k = np.empty((n_layers, n_heads, capacity, head_dim), dtype=dtype)
13
+ self.v = np.empty((n_layers, n_heads, capacity, head_dim), dtype=dtype)
14
+ self.length = 0
15
+
16
+ @property
17
+ def nbytes(self):
18
+ return self.k.nbytes + self.v.nbytes
19
+
20
+ def reset(self):
21
+ self.length = 0
22
+
23
+
24
+ class KVCacheManager:
25
+
26
+ def __init__(self, n_layers, n_heads, head_dim, max_bytes=None, dtype=np.float32):
27
+ self.n_layers = n_layers
28
+ self.n_heads = n_heads
29
+ self.head_dim = head_dim
30
+ self.max_bytes = max_bytes
31
+ self.dtype = np.dtype(dtype)
32
+ self._caches = {}
33
+ self.used_bytes = 0
34
+
35
+ def bytes_for(self, capacity):
36
+ return 2 * self.n_layers * self.n_heads * capacity * self.head_dim * self.dtype.itemsize
37
+
38
+ def fits_at_all(self, capacity):
39
+ return self.max_bytes is None or self.bytes_for(capacity) <= self.max_bytes
40
+
41
+ def can_allocate(self, capacity):
42
+ return self.max_bytes is None or self.used_bytes + self.bytes_for(capacity) <= self.max_bytes
43
+
44
+ def allocate(self, key, capacity):
45
+ if key in self._caches:
46
+ raise KeyError(f"cache already allocated for {key}")
47
+
48
+ if not self.can_allocate(capacity):
49
+ raise CacheFull(f"kv cache budget of {self.max_bytes} bytes exhausted")
50
+
51
+ cache = KVCache(self.n_layers, self.n_heads, self.head_dim, capacity, self.dtype)
52
+ self._caches[key] = cache
53
+ self.used_bytes += cache.nbytes
54
+
55
+ return cache
56
+
57
+ def release(self, key):
58
+ cache = self._caches.pop(key, None)
59
+
60
+ if cache is not None:
61
+ self.used_bytes -= cache.nbytes
62
+
63
+ @property
64
+ def active(self):
65
+ return len(self._caches)
@@ -0,0 +1,161 @@
1
+ import json
2
+ import os
3
+
4
+ import numpy as np
5
+
6
+ from src.inference.chat_template import ChatTemplate
7
+ from src.inference.config import GenerationConfig
8
+ from src.inference.engine import InferenceModel
9
+ from src.inference.export import EXPORT_FORMAT_VERSION, _sha256
10
+ from src.inference.scheduler import BatchScheduler, GenerationRequest
11
+ from src.models.gpt.config import GPTConfig
12
+ from src.models.gpt.context import describe_context, extend_context
13
+ from src.tokenization.base import identity_matches
14
+ from src.tokenization.registry import load_tokenizer
15
+
16
+
17
+ class TextGenerator:
18
+
19
+ def __init__(
20
+ self,
21
+ engine,
22
+ tokenizer,
23
+ default_config=None,
24
+ max_batch_size=8,
25
+ cache_budget_bytes=None,
26
+ prefill_chunk=256,
27
+ chat_template=None,
28
+ metadata=None,
29
+ ):
30
+ self.engine = engine
31
+ self.tokenizer = tokenizer
32
+ self.default_config = default_config or GenerationConfig()
33
+ self.chat_template = chat_template
34
+ self.metadata = metadata or {}
35
+ self.scheduler = BatchScheduler(
36
+ engine,
37
+ tokenizer,
38
+ max_batch_size=max_batch_size,
39
+ cache_budget_bytes=cache_budget_bytes,
40
+ prefill_chunk=prefill_chunk,
41
+ )
42
+
43
+ @property
44
+ def config(self):
45
+ return self.engine.config
46
+
47
+ @property
48
+ def weights_nbytes(self):
49
+ return self.engine.nbytes
50
+
51
+ def bound_chat_template(self, fallback=None):
52
+ template = self.chat_template or ChatTemplate.resolve(fallback)
53
+
54
+ if template is None:
55
+ raise ValueError("model has no chat template and no fallback was given")
56
+
57
+ return template.bind(self.tokenizer)
58
+
59
+ def _request(self, prompt, overrides, timeout_s=None, truncate_prompt=False, request_id=None):
60
+ kwargs = {}
61
+
62
+ if request_id:
63
+ kwargs["request_id"] = request_id
64
+
65
+ return GenerationRequest(
66
+ prompt=prompt,
67
+ config=self.default_config.updated(**overrides),
68
+ timeout_s=timeout_s,
69
+ truncate_prompt=truncate_prompt,
70
+ **kwargs,
71
+ )
72
+
73
+ def submit(self, request):
74
+ return self.scheduler.submit(request)
75
+
76
+ def generate(self, prompt, timeout_s=None, truncate_prompt=False, **overrides):
77
+ handle = self.submit(self._request(prompt, overrides, timeout_s, truncate_prompt))
78
+ return handle.result(drive=not self.scheduler.background)
79
+
80
+ def generate_stream(self, prompt, timeout_s=None, truncate_prompt=False, events=False, **overrides):
81
+ handle = self.submit(self._request(prompt, overrides, timeout_s, truncate_prompt))
82
+
83
+ stream = handle.events(drive=not self.scheduler.background)
84
+
85
+ try:
86
+ for ev in stream:
87
+ if events:
88
+ yield ev
89
+ elif ev.text:
90
+ yield ev.text
91
+ finally:
92
+ stream.close()
93
+ handle.cancel()
94
+
95
+ def start(self):
96
+ self.scheduler.start()
97
+
98
+ def stop(self):
99
+ self.scheduler.stop()
100
+
101
+
102
+ def load_model(path, max_batch_size=8, cache_budget_bytes=None, verify=True, prefill_chunk=256, chat_template=None,
103
+ context_length=None, context_extension=None):
104
+ with open(os.path.join(path, "config.json")) as f:
105
+ meta = json.load(f)
106
+
107
+ if meta["format_version"] != EXPORT_FORMAT_VERSION:
108
+ raise ValueError(f"unsupported export format version {meta['format_version']}")
109
+
110
+ weights_path = os.path.join(path, meta["weights"]["file"])
111
+
112
+ if verify and _sha256(weights_path) != meta["weights"]["sha256"]:
113
+ raise ValueError("model weights failed checksum verification")
114
+
115
+ config = GPTConfig.from_dict(meta["model"])
116
+
117
+ if context_length is not None or context_extension is not None:
118
+ if context_length is None or context_extension is None:
119
+ raise ValueError("context extension needs both context_length and context_extension")
120
+
121
+ config = extend_context(config, context_length, context_extension)
122
+
123
+ with np.load(weights_path, allow_pickle=False) as archive:
124
+ weights = {k: archive[k] for k in archive.files}
125
+
126
+ tokenizer = load_tokenizer(os.path.join(path, meta["tokenizer"]["file"]))
127
+
128
+ if not identity_matches(meta["tokenizer"]["identity"], tokenizer.identity):
129
+ raise ValueError("bundled tokenizer does not match the exported model")
130
+
131
+ gen_path = os.path.join(path, "generation.json")
132
+ default = GenerationConfig()
133
+
134
+ if os.path.isfile(gen_path):
135
+ with open(gen_path) as f:
136
+ default = GenerationConfig.from_dict(json.load(f))
137
+
138
+ template = ChatTemplate.resolve(chat_template if chat_template is not None else meta.get("chat_template"))
139
+
140
+ if template is not None:
141
+ template.bind(tokenizer)
142
+
143
+ engine = InferenceModel(config, weights)
144
+ del weights
145
+
146
+ return TextGenerator(
147
+ engine,
148
+ tokenizer,
149
+ default,
150
+ max_batch_size=max_batch_size,
151
+ cache_budget_bytes=cache_budget_bytes,
152
+ prefill_chunk=prefill_chunk,
153
+ chat_template=template,
154
+ metadata={
155
+ "weights_sha256": meta["weights"]["sha256"],
156
+ "weights_dtype": meta["weights"]["dtype"],
157
+ "source": meta.get("source", {}),
158
+ "path": os.path.abspath(path),
159
+ "context": describe_context(config),
160
+ },
161
+ )
@@ -0,0 +1,42 @@
1
+ import numpy as np
2
+
3
+
4
+ def _softmax(x):
5
+ x = x - x.max()
6
+ e = np.exp(x)
7
+ return e / e.sum()
8
+
9
+
10
+ def sample_token(logits, config, rng, seen_ids=None):
11
+ logits = np.array(logits, dtype=np.float64, copy=True)
12
+
13
+ penalty = config.repetition_penalty
14
+
15
+ if penalty != 1.0 and seen_ids:
16
+ idx = np.fromiter(seen_ids, dtype=np.int64, count=len(seen_ids))
17
+ vals = logits[idx]
18
+ logits[idx] = np.where(vals > 0, vals / penalty, vals * penalty)
19
+
20
+ if config.greedy:
21
+ return int(np.argmax(logits))
22
+
23
+ logits /= config.temperature
24
+
25
+ candidates = None
26
+
27
+ if 0 < config.top_k < logits.size:
28
+ candidates = np.argpartition(logits, -config.top_k)[-config.top_k:]
29
+ logits = logits[candidates]
30
+
31
+ if config.top_p < 1.0:
32
+ order = np.argsort(-logits, kind="stable")
33
+ probs = _softmax(logits[order])
34
+ cutoff = int(np.searchsorted(np.cumsum(probs), config.top_p)) + 1
35
+ keep = order[:cutoff]
36
+ candidates = keep if candidates is None else candidates[keep]
37
+ logits = logits[keep]
38
+
39
+ probs = _softmax(logits)
40
+ choice = int(rng.choice(probs.size, p=probs))
41
+
42
+ return int(candidates[choice]) if candidates is not None else choice