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/serving/config.py ADDED
@@ -0,0 +1,216 @@
1
+ import json
2
+ import os
3
+ import re
4
+ from dataclasses import MISSING, asdict, dataclass, field, fields, is_dataclass
5
+ from typing import List, Optional
6
+
7
+ MODEL_NAME_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:-]{0,127}$")
8
+
9
+
10
+ @dataclass
11
+ class ModelEntry:
12
+ name: str
13
+ path: str
14
+ chat_template: Optional[str] = None
15
+ max_batch_size: int = 8
16
+ kv_cache_budget_mb: Optional[float] = None
17
+ prefill_chunk: int = 256
18
+ preload: bool = True
19
+ max_context_tokens: Optional[int] = None
20
+ max_generation_tokens: Optional[int] = None
21
+ extend_context_to: Optional[int] = None
22
+ context_extension: Optional[str] = None
23
+
24
+ def validate(self):
25
+ if (self.extend_context_to is None) != (self.context_extension is None):
26
+ raise ValueError(f"model '{self.name}': extend_context_to and context_extension go together")
27
+
28
+ if self.context_extension is not None and self.context_extension not in ("extrapolate", "linear", "ntk"):
29
+ raise ValueError(f"model '{self.name}': context_extension must be extrapolate, linear or ntk")
30
+
31
+ if not MODEL_NAME_RE.match(self.name):
32
+ raise ValueError(f"invalid model name '{self.name}' (letters, digits, . _ : - ; max 128)")
33
+
34
+ if self.max_batch_size < 1:
35
+ raise ValueError(f"model '{self.name}': max_batch_size must be >= 1")
36
+
37
+ if self.kv_cache_budget_mb is not None and self.kv_cache_budget_mb <= 0:
38
+ raise ValueError(f"model '{self.name}': kv_cache_budget_mb must be > 0")
39
+
40
+ if self.prefill_chunk < 1:
41
+ raise ValueError(f"model '{self.name}': prefill_chunk must be >= 1")
42
+
43
+
44
+ @dataclass
45
+ class LimitsConfig:
46
+ max_concurrent_requests: int = 64
47
+ max_body_bytes: int = 1 << 20
48
+ max_header_bytes: int = 16 << 10
49
+ max_messages: int = 256
50
+ max_prompt_chars: int = 200_000
51
+ max_prompt_tokens: Optional[int] = None
52
+ max_generation_tokens: int = 1024
53
+ max_stop_sequences: int = 4
54
+ max_stop_sequence_chars: int = 128
55
+ request_timeout_s: float = 300.0
56
+ max_connections: int = 512
57
+ header_timeout_s: float = 10.0
58
+ body_timeout_s: float = 30.0
59
+ keepalive_timeout_s: float = 30.0
60
+ max_requests_per_connection: int = 1000
61
+ sse_keepalive_s: float = 15.0
62
+ shutdown_drain_s: float = 30.0
63
+
64
+ def validate(self):
65
+ for name in ("max_concurrent_requests", "max_body_bytes", "max_header_bytes", "max_messages",
66
+ "max_prompt_chars", "max_generation_tokens", "max_connections"):
67
+ if getattr(self, name) < 1:
68
+ raise ValueError(f"limits.{name} must be >= 1")
69
+
70
+ for name in ("request_timeout_s", "header_timeout_s", "body_timeout_s", "keepalive_timeout_s", "sse_keepalive_s"):
71
+ if getattr(self, name) <= 0:
72
+ raise ValueError(f"limits.{name} must be > 0")
73
+
74
+
75
+ @dataclass
76
+ class RateLimitConfig:
77
+ requests_per_minute: float = 120.0
78
+ burst: int = 30
79
+ max_concurrent_per_principal: int = 8
80
+ max_tracked_principals: int = 10_000
81
+
82
+ def validate(self):
83
+ if self.requests_per_minute <= 0 or self.burst < 1 or self.max_concurrent_per_principal < 1:
84
+ raise ValueError("rate_limit values must be positive")
85
+
86
+
87
+ @dataclass
88
+ class SecurityConfig:
89
+ api_keys: List[str] = field(default_factory=list)
90
+ api_keys_file: Optional[str] = None
91
+ api_keys_env: Optional[str] = "PTF_API_KEYS"
92
+ admin_keys: List[str] = field(default_factory=list)
93
+ admin_keys_file: Optional[str] = None
94
+ allow_unauthenticated: bool = False
95
+ metrics_require_auth: bool = True
96
+ cors_origins: List[str] = field(default_factory=list)
97
+ trust_forwarded_for: bool = False
98
+ rate_limit: RateLimitConfig = field(default_factory=RateLimitConfig)
99
+
100
+
101
+ @dataclass
102
+ class RuntimeConfig:
103
+ device: str = "auto"
104
+ memory_limit_mb: Optional[float] = None
105
+ fallback_chat_template: str = "plain"
106
+ chat_truncation: str = "auto"
107
+
108
+
109
+ @dataclass
110
+ class LoggingConfig:
111
+ access_log: Optional[str] = "-"
112
+ log_prompts: bool = False
113
+
114
+
115
+ @dataclass
116
+ class UIConfig:
117
+ enabled: bool = True
118
+ directory: Optional[str] = None
119
+
120
+
121
+ @dataclass
122
+ class ServerConfig:
123
+ host: str = "127.0.0.1"
124
+ port: int = 8000
125
+ models: List[ModelEntry] = field(default_factory=list)
126
+ limits: LimitsConfig = field(default_factory=LimitsConfig)
127
+ security: SecurityConfig = field(default_factory=SecurityConfig)
128
+ runtime: RuntimeConfig = field(default_factory=RuntimeConfig)
129
+ logging: LoggingConfig = field(default_factory=LoggingConfig)
130
+ ui: UIConfig = field(default_factory=UIConfig)
131
+
132
+ def validate(self):
133
+ if not self.models:
134
+ raise ValueError("at least one model must be configured")
135
+
136
+ names = [m.name for m in self.models]
137
+
138
+ if len(set(names)) != len(names):
139
+ raise ValueError("model names must be unique")
140
+
141
+ for m in self.models:
142
+ m.validate()
143
+
144
+ if not 0 <= self.port < 65536:
145
+ raise ValueError("port must be in 0..65535 (0 picks a free port)")
146
+
147
+ if self.runtime.chat_truncation not in ("auto", "disabled"):
148
+ raise ValueError("runtime.chat_truncation must be 'auto' or 'disabled'")
149
+
150
+ self.limits.validate()
151
+ self.security.rate_limit.validate()
152
+
153
+ return self
154
+
155
+ def to_dict(self):
156
+ return asdict(self)
157
+
158
+
159
+ def _build(cls, data, where):
160
+ if data is None:
161
+ return cls()
162
+
163
+ if not isinstance(data, dict):
164
+ raise ValueError(f"'{where}' must be a mapping")
165
+
166
+ known = {f.name: f for f in fields(cls)}
167
+ unknown = sorted(set(data) - set(known))
168
+
169
+ if unknown:
170
+ raise ValueError(f"unknown keys in '{where}': {unknown}")
171
+
172
+ kwargs = {}
173
+
174
+ for name, value in data.items():
175
+ f = known[name]
176
+ default = f.default_factory() if f.default_factory is not MISSING else f.default
177
+
178
+ if is_dataclass(default) and not isinstance(default, type):
179
+ kwargs[name] = _build(type(default), value, f"{where}.{name}")
180
+ else:
181
+ kwargs[name] = value
182
+
183
+ return cls(**kwargs)
184
+
185
+
186
+ def server_config_from_dict(data):
187
+ data = dict(data or {})
188
+ models = data.pop("models", [])
189
+
190
+ if not isinstance(models, list):
191
+ raise ValueError("'models' must be a list")
192
+
193
+ cfg = _build(ServerConfig, data, "server")
194
+ cfg.models = [_build(ModelEntry, m, f"models[{i}]") for i, m in enumerate(models)]
195
+
196
+ return cfg
197
+
198
+
199
+ def load_server_config(path):
200
+ with open(path) as f:
201
+ text = f.read()
202
+
203
+ if path.endswith((".yaml", ".yml")):
204
+ import yaml
205
+ data = yaml.safe_load(text)
206
+ else:
207
+ data = json.loads(text)
208
+
209
+ cfg = server_config_from_dict(data)
210
+ base = os.path.dirname(os.path.abspath(path))
211
+
212
+ for m in cfg.models:
213
+ if not os.path.isabs(m.path):
214
+ m.path = os.path.join(base, m.path)
215
+
216
+ return cfg
src/serving/errors.py ADDED
@@ -0,0 +1,51 @@
1
+ class APIError(Exception):
2
+
3
+ def __init__(self, status, message, error_type="invalid_request_error", code=None, param=None, headers=None):
4
+ super().__init__(message)
5
+ self.status = status
6
+ self.message = message
7
+ self.error_type = error_type
8
+ self.code = code
9
+ self.param = param
10
+ self.headers = headers or {}
11
+
12
+ def to_dict(self):
13
+ return {
14
+ "error": {
15
+ "message": self.message,
16
+ "type": self.error_type,
17
+ "param": self.param,
18
+ "code": self.code,
19
+ }
20
+ }
21
+
22
+
23
+ def bad_request(message, param=None, code=None):
24
+ return APIError(400, message, "invalid_request_error", code=code, param=param)
25
+
26
+
27
+ def unauthorized(message="invalid or missing API key"):
28
+ return APIError(401, message, "authentication_error", code="invalid_api_key",
29
+ headers={"WWW-Authenticate": "Bearer"})
30
+
31
+
32
+ def forbidden(message="this API key is not allowed to perform this action"):
33
+ return APIError(403, message, "permission_error", code="forbidden")
34
+
35
+
36
+ def not_found(message, code="not_found"):
37
+ return APIError(404, message, "invalid_request_error", code=code)
38
+
39
+
40
+ def rate_limited(message, retry_after_s):
41
+ return APIError(429, message, "rate_limit_error", code="rate_limit_exceeded",
42
+ headers={"Retry-After": str(max(1, int(retry_after_s + 0.999)))})
43
+
44
+
45
+ def overloaded(message="the server is at capacity; retry shortly", retry_after_s=1):
46
+ return APIError(503, message, "server_error", code="server_overloaded",
47
+ headers={"Retry-After": str(max(1, int(retry_after_s)))})
48
+
49
+
50
+ def internal(request_id):
51
+ return APIError(500, f"internal server error (request id {request_id})", "server_error", code="internal_error")