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/inference/export.py
ADDED
|
@@ -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)
|
src/inference/runtime.py
ADDED
|
@@ -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
|