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