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
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
from .GlorotUniform import GlorotUniform
|
|
2
|
+
from .HeUniform import HeUniform
|
|
3
|
+
from .LecunUniform import LecunUniform
|
|
4
|
+
from .Ones import Ones
|
|
5
|
+
from .Orthogonal import Orthogonal
|
|
6
|
+
from .RandomNormal import RandomNormal
|
|
7
|
+
from .Zeros import Zeros
|
|
8
|
+
|
|
9
|
+
initializer_fns = {
|
|
10
|
+
"glorot_uniform": GlorotUniform(),
|
|
11
|
+
"he_uniform": HeUniform(),
|
|
12
|
+
"lecun_uniform": LecunUniform(),
|
|
13
|
+
"zeros": Zeros(),
|
|
14
|
+
"orthogonal": Orthogonal(),
|
|
15
|
+
"random_normal": RandomNormal(),
|
|
16
|
+
"ones": Ones(),
|
|
17
|
+
}
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
from src.core.Tensor import Tensor
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class CrossEntropyLoss:
|
|
5
|
+
|
|
6
|
+
def __call__(self, logits, target):
|
|
7
|
+
|
|
8
|
+
probs = logits.softmax()
|
|
9
|
+
|
|
10
|
+
batch = logits.shape[0]
|
|
11
|
+
|
|
12
|
+
indices = target.data.astype(int)
|
|
13
|
+
|
|
14
|
+
p = probs[
|
|
15
|
+
Tensor.arange(batch),
|
|
16
|
+
indices,
|
|
17
|
+
]
|
|
18
|
+
|
|
19
|
+
return -(p.log()).mean()
|
|
20
|
+
|
|
21
|
+
# class CrossEntropyLoss:
|
|
22
|
+
#
|
|
23
|
+
# def __call__(self, logits, target):
|
|
24
|
+
#
|
|
25
|
+
# probs = logits.softmax()
|
|
26
|
+
#
|
|
27
|
+
# batch = logits.shape[0]
|
|
28
|
+
#
|
|
29
|
+
# loss = 0
|
|
30
|
+
#
|
|
31
|
+
# for i in range(batch):
|
|
32
|
+
# loss += -probs[i][int(target.data[i])].log()
|
|
33
|
+
#
|
|
34
|
+
# return loss / batch
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.core.Tensor import Tensor
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class CrossEntropyWithLogitsLoss:
|
|
7
|
+
|
|
8
|
+
def __call__(self, logits, target):
|
|
9
|
+
|
|
10
|
+
batch = logits.shape[0]
|
|
11
|
+
|
|
12
|
+
shifted = logits.data - np.max(
|
|
13
|
+
logits.data,
|
|
14
|
+
axis=-1,
|
|
15
|
+
keepdims=True,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
logsumexp = np.log(
|
|
19
|
+
np.sum(np.exp(shifted), axis=-1)
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
indices = target.data.astype(np.int64)
|
|
23
|
+
|
|
24
|
+
loss = (
|
|
25
|
+
-shifted[np.arange(batch), indices]
|
|
26
|
+
+ logsumexp
|
|
27
|
+
).mean()
|
|
28
|
+
|
|
29
|
+
out = Tensor(
|
|
30
|
+
loss,
|
|
31
|
+
requires_grad=logits.requires_grad,
|
|
32
|
+
parents=(logits,),
|
|
33
|
+
op="CrossEntropyWithLogits",
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
def _backward():
|
|
37
|
+
|
|
38
|
+
if not logits.requires_grad:
|
|
39
|
+
return
|
|
40
|
+
|
|
41
|
+
exp = np.exp(shifted)
|
|
42
|
+
|
|
43
|
+
probs = exp / np.sum(
|
|
44
|
+
exp,
|
|
45
|
+
axis=-1,
|
|
46
|
+
keepdims=True,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
grad = probs
|
|
50
|
+
|
|
51
|
+
grad[np.arange(batch), indices] -= 1
|
|
52
|
+
|
|
53
|
+
grad /= batch
|
|
54
|
+
|
|
55
|
+
logits.grad += grad * out.grad
|
|
56
|
+
|
|
57
|
+
out._backward = _backward
|
|
58
|
+
|
|
59
|
+
return out
|
src/loss/Hinge.py
ADDED
src/loss/Huber.py
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
class Huber:
|
|
2
|
+
|
|
3
|
+
def __init__(self, delta=1.0):
|
|
4
|
+
self.delta = delta
|
|
5
|
+
|
|
6
|
+
def __call__(self, prediction, target):
|
|
7
|
+
|
|
8
|
+
error = prediction - target
|
|
9
|
+
|
|
10
|
+
abs_error = error.abs()
|
|
11
|
+
|
|
12
|
+
quadratic = 0.5 * error * error
|
|
13
|
+
|
|
14
|
+
linear = self.delta * (
|
|
15
|
+
abs_error - 0.5 * self.delta
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
return abs_error.where(
|
|
19
|
+
abs_error < self.delta,
|
|
20
|
+
quadratic,
|
|
21
|
+
linear,
|
|
22
|
+
).mean()
|
src/loss/Loss.py
ADDED
src/loss/MSE.py
ADDED
src/loss/MSELoss.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class SparseCategoricalCrossEntropy:
|
|
5
|
+
|
|
6
|
+
def __call__(self, prediction, target):
|
|
7
|
+
|
|
8
|
+
prediction = prediction.softmax()
|
|
9
|
+
|
|
10
|
+
log_probs = prediction.log()
|
|
11
|
+
|
|
12
|
+
return -log_probs[
|
|
13
|
+
np.arange(len(target)),
|
|
14
|
+
target,
|
|
15
|
+
].mean()
|
src/loss/__init__.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
from .CrossEntropyWithLogitsLoss import CrossEntropyWithLogitsLoss
|
|
2
|
+
from .MSE import MSE
|
|
3
|
+
from .CategoricalCrossEntropy import CategoricalCrossEntropy
|
|
4
|
+
from .SparseCategoricalCrossEntropy import SparseCategoricalCrossEntropy
|
|
5
|
+
from .bce import BinaryCrossEntropy
|
|
6
|
+
from .mae import MAE
|
|
7
|
+
from .CrossEntropyLoss import CrossEntropyLoss
|
|
8
|
+
|
|
9
|
+
losses = {
|
|
10
|
+
"mse": MSE(),
|
|
11
|
+
"mean_squared_error": MSE(),
|
|
12
|
+
"mean_absolute_error": MAE(),
|
|
13
|
+
"categorical_crossentropy": CategoricalCrossEntropy(),
|
|
14
|
+
"binary_crossentropy": BinaryCrossEntropy(),
|
|
15
|
+
"cross_entropy": CrossEntropyLoss(),
|
|
16
|
+
"sparse_categorical_crossentropy": SparseCategoricalCrossEntropy(),
|
|
17
|
+
"cross_entropy_with_logits": CrossEntropyWithLogitsLoss()
|
|
18
|
+
}
|
src/loss/bce.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
from src.loss.Loss import Loss
|
|
2
|
+
import numpy as np
|
|
3
|
+
|
|
4
|
+
class BCELoss(Loss):
|
|
5
|
+
|
|
6
|
+
def forward(self, pred, target):
|
|
7
|
+
eps = 1e-8
|
|
8
|
+
pred = np.clip(pred, eps, 1 - eps)
|
|
9
|
+
|
|
10
|
+
return -np.mean(
|
|
11
|
+
target * np.log(pred) +
|
|
12
|
+
(1 - target) * np.log(1 - pred)
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
def backward(self, pred, target):
|
|
16
|
+
eps = 1e-8
|
|
17
|
+
pred = np.clip(pred, eps, 1 - eps)
|
|
18
|
+
|
|
19
|
+
return (pred - target) / (pred * (1 - pred))
|
|
20
|
+
|
|
21
|
+
class BinaryCrossEntropy:
|
|
22
|
+
|
|
23
|
+
def __call__(self, prediction, target):
|
|
24
|
+
|
|
25
|
+
eps = 1e-7
|
|
26
|
+
|
|
27
|
+
prediction = prediction.clip(eps, 1 - eps)
|
|
28
|
+
|
|
29
|
+
loss = -(
|
|
30
|
+
target * prediction.log()
|
|
31
|
+
+ (1 - target) * (1 - prediction).log()
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
return loss.mean()
|
src/loss/mae.py
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
from src.loss.Loss import Loss
|
|
2
|
+
import numpy as np
|
|
3
|
+
|
|
4
|
+
class MAELoss(Loss):
|
|
5
|
+
|
|
6
|
+
def forward(self, pred, target):
|
|
7
|
+
return np.mean(np.abs(pred - target))
|
|
8
|
+
|
|
9
|
+
def backward(self, pred, target):
|
|
10
|
+
return np.sign(pred - target)
|
|
11
|
+
|
|
12
|
+
class MAE:
|
|
13
|
+
|
|
14
|
+
def __call__(self, prediction, target):
|
|
15
|
+
|
|
16
|
+
return (prediction - target).abs().mean()
|
src/math/__init__.py
ADDED
|
File without changes
|
src/math/clip.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.core.Tensor import Tensor
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class Clip:
|
|
7
|
+
|
|
8
|
+
@staticmethod
|
|
9
|
+
def forward(x, min_value, max_value):
|
|
10
|
+
|
|
11
|
+
clipped = np.clip(x.data, min_value, max_value)
|
|
12
|
+
|
|
13
|
+
out = Tensor(
|
|
14
|
+
clipped,
|
|
15
|
+
requires_grad=x.requires_grad,
|
|
16
|
+
parents=(x,),
|
|
17
|
+
op="Clip",
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
def _backward():
|
|
21
|
+
|
|
22
|
+
if not x.requires_grad:
|
|
23
|
+
return
|
|
24
|
+
|
|
25
|
+
mask = (
|
|
26
|
+
(x.data >= min_value)
|
|
27
|
+
& (x.data <= max_value)
|
|
28
|
+
).astype(np.float32)
|
|
29
|
+
|
|
30
|
+
x.grad += out.grad * mask
|
|
31
|
+
|
|
32
|
+
out._backward = _backward
|
|
33
|
+
|
|
34
|
+
return out
|
|
35
|
+
|
|
36
|
+
def __call__(self, x, min_value, max_value):
|
|
37
|
+
return self.forward(x, min_value, max_value)
|
src/math/exp.py
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.core.Tensor import Tensor
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class Exp:
|
|
7
|
+
|
|
8
|
+
@staticmethod
|
|
9
|
+
def forward(x):
|
|
10
|
+
|
|
11
|
+
e = np.exp(x.data)
|
|
12
|
+
|
|
13
|
+
out = Tensor(
|
|
14
|
+
e,
|
|
15
|
+
requires_grad=x.requires_grad,
|
|
16
|
+
parents=(x,),
|
|
17
|
+
op="Exp",
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
def _backward():
|
|
21
|
+
|
|
22
|
+
if x.requires_grad:
|
|
23
|
+
x.grad += out.grad * e
|
|
24
|
+
|
|
25
|
+
out._backward = _backward
|
|
26
|
+
|
|
27
|
+
return out
|
src/math/log.py
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.core.Tensor import Tensor
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class Log:
|
|
7
|
+
|
|
8
|
+
@staticmethod
|
|
9
|
+
def forward(x):
|
|
10
|
+
|
|
11
|
+
out = Tensor(
|
|
12
|
+
np.log(x.data),
|
|
13
|
+
requires_grad=x.requires_grad,
|
|
14
|
+
parents=(x,),
|
|
15
|
+
op="Log",
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
def _backward():
|
|
19
|
+
|
|
20
|
+
if x.requires_grad:
|
|
21
|
+
x.grad += out.grad / x.data
|
|
22
|
+
|
|
23
|
+
out._backward = _backward
|
|
24
|
+
|
|
25
|
+
return out
|
src/math/sigmoid.py
ADDED
src/models/__init__.py
ADDED
|
File without changes
|
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.neural.Layer import Layer
|
|
4
|
+
from src.core.Tensor import Tensor
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class Embedding(Layer):
|
|
8
|
+
|
|
9
|
+
def __init__(
|
|
10
|
+
self,
|
|
11
|
+
vocab_size,
|
|
12
|
+
embedding_dim,
|
|
13
|
+
):
|
|
14
|
+
super().__init__()
|
|
15
|
+
|
|
16
|
+
self.vocab_size = vocab_size
|
|
17
|
+
self.embedding_dim = embedding_dim
|
|
18
|
+
|
|
19
|
+
self.weight = None
|
|
20
|
+
|
|
21
|
+
def build(self, input_shape):
|
|
22
|
+
|
|
23
|
+
self.weight = self.add_weight(
|
|
24
|
+
shape=(
|
|
25
|
+
self.vocab_size,
|
|
26
|
+
self.embedding_dim,
|
|
27
|
+
),
|
|
28
|
+
initializer="random_normal",
|
|
29
|
+
name="embedding",
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
def call(self, x):
|
|
33
|
+
|
|
34
|
+
indices = x.data.astype(np.int64)
|
|
35
|
+
|
|
36
|
+
out = Tensor(
|
|
37
|
+
self.weight.data[indices],
|
|
38
|
+
requires_grad=True,
|
|
39
|
+
parents=(self.weight,),
|
|
40
|
+
op="Embedding",
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
def _backward():
|
|
44
|
+
|
|
45
|
+
if not self.weight.requires_grad:
|
|
46
|
+
return
|
|
47
|
+
|
|
48
|
+
np.add.at(
|
|
49
|
+
self.weight.grad,
|
|
50
|
+
indices,
|
|
51
|
+
out.grad,
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
out._backward = _backward
|
|
55
|
+
|
|
56
|
+
return out
|
|
57
|
+
|
|
58
|
+
def compute_output_shape(self, input_shape):
|
|
59
|
+
batch, seq = input_shape
|
|
60
|
+
|
|
61
|
+
return (
|
|
62
|
+
batch,
|
|
63
|
+
seq,
|
|
64
|
+
self.embedding_dim,
|
|
65
|
+
)
|
|
File without changes
|
|
File without changes
|
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
import threading
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from src.core.Tensor import Tensor, matmul_policy
|
|
6
|
+
from src.neural.Layer import Layer
|
|
7
|
+
|
|
8
|
+
_mask_cache = {}
|
|
9
|
+
_mask_lock = threading.Lock()
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _future_mask(T):
|
|
13
|
+
with _mask_lock:
|
|
14
|
+
mask = _mask_cache.get(T)
|
|
15
|
+
|
|
16
|
+
if mask is None:
|
|
17
|
+
mask = np.triu(np.ones((T, T), dtype=bool), k=1)
|
|
18
|
+
_mask_cache[T] = mask
|
|
19
|
+
|
|
20
|
+
return mask
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _identity(x):
|
|
24
|
+
return x
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _to_heads(x, B, T, H, hd):
|
|
28
|
+
return x.reshape(B, T, H, hd).transpose(0, 2, 1, 3)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _from_heads(x, B, T, d):
|
|
32
|
+
return np.ascontiguousarray(x.transpose(0, 2, 1, 3)).reshape(B, T, d)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def causal_attention(q, k, v, n_heads, rope=None):
|
|
36
|
+
B, T, d = q.shape
|
|
37
|
+
H = n_heads
|
|
38
|
+
hd = d // H
|
|
39
|
+
|
|
40
|
+
policy = matmul_policy()
|
|
41
|
+
rnd = _identity if policy is None else policy
|
|
42
|
+
scale = 1.0 / float(np.sqrt(hd))
|
|
43
|
+
|
|
44
|
+
qh = _to_heads(q.data, B, T, H, hd)
|
|
45
|
+
kh = _to_heads(k.data, B, T, H, hd)
|
|
46
|
+
vh = _to_heads(v.data, B, T, H, hd)
|
|
47
|
+
|
|
48
|
+
cos = sin = None
|
|
49
|
+
|
|
50
|
+
if rope is not None:
|
|
51
|
+
cos, sin = rope.tables(np.arange(T))
|
|
52
|
+
qh = rope.rotate(qh, cos, sin)
|
|
53
|
+
kh = rope.rotate(kh, cos, sin)
|
|
54
|
+
|
|
55
|
+
qs = rnd(qh)
|
|
56
|
+
ks = rnd(kh)
|
|
57
|
+
vs = rnd(vh)
|
|
58
|
+
|
|
59
|
+
scores = rnd(np.matmul(qs, ks.transpose(0, 1, 3, 2)))
|
|
60
|
+
scores *= scale
|
|
61
|
+
|
|
62
|
+
if T > 1:
|
|
63
|
+
scores[..., _future_mask(T)] = -np.inf
|
|
64
|
+
|
|
65
|
+
scores -= scores.max(axis=-1, keepdims=True)
|
|
66
|
+
np.exp(scores, out=scores)
|
|
67
|
+
scores /= scores.sum(axis=-1, keepdims=True)
|
|
68
|
+
probs = scores
|
|
69
|
+
|
|
70
|
+
ps = rnd(probs)
|
|
71
|
+
context = rnd(np.matmul(ps, vs))
|
|
72
|
+
|
|
73
|
+
out = Tensor(
|
|
74
|
+
_from_heads(context, B, T, d),
|
|
75
|
+
requires_grad=q.requires_grad or k.requires_grad or v.requires_grad,
|
|
76
|
+
parents=(q, k, v),
|
|
77
|
+
op="CausalAttention",
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
def _backward():
|
|
81
|
+
d_out = rnd(_to_heads(out.grad, B, T, H, hd))
|
|
82
|
+
|
|
83
|
+
if v.requires_grad:
|
|
84
|
+
d_v = rnd(np.matmul(ps.transpose(0, 1, 3, 2), d_out))
|
|
85
|
+
v.grad += _from_heads(d_v, B, T, d)
|
|
86
|
+
|
|
87
|
+
if not (q.requires_grad or k.requires_grad):
|
|
88
|
+
return
|
|
89
|
+
|
|
90
|
+
d_p = rnd(np.matmul(d_out, vs.transpose(0, 1, 3, 2)))
|
|
91
|
+
d_s = probs * (d_p - np.sum(d_p * probs, axis=-1, keepdims=True))
|
|
92
|
+
d_s *= scale
|
|
93
|
+
d_s = rnd(d_s)
|
|
94
|
+
|
|
95
|
+
if q.requires_grad:
|
|
96
|
+
d_q = rnd(np.matmul(d_s, ks))
|
|
97
|
+
|
|
98
|
+
if rope is not None:
|
|
99
|
+
d_q = rope.rotate(d_q, cos, sin, inverse=True)
|
|
100
|
+
|
|
101
|
+
q.grad += _from_heads(d_q, B, T, d)
|
|
102
|
+
|
|
103
|
+
if k.requires_grad:
|
|
104
|
+
d_k = rnd(np.matmul(d_s.transpose(0, 1, 3, 2), qs))
|
|
105
|
+
|
|
106
|
+
if rope is not None:
|
|
107
|
+
d_k = rope.rotate(d_k, cos, sin, inverse=True)
|
|
108
|
+
|
|
109
|
+
k.grad += _from_heads(d_k, B, T, d)
|
|
110
|
+
|
|
111
|
+
out._backward = _backward
|
|
112
|
+
|
|
113
|
+
return out
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class CausalSelfAttention(Layer):
|
|
117
|
+
|
|
118
|
+
def __init__(self, d_model, num_heads, rope=None):
|
|
119
|
+
super().__init__()
|
|
120
|
+
|
|
121
|
+
if d_model % num_heads != 0:
|
|
122
|
+
raise ValueError("d_model must be divisible by num_heads")
|
|
123
|
+
|
|
124
|
+
self.d_model = d_model
|
|
125
|
+
self.num_heads = num_heads
|
|
126
|
+
self.head_dim = d_model // num_heads
|
|
127
|
+
self.rope = rope
|
|
128
|
+
|
|
129
|
+
self.Wq = None
|
|
130
|
+
self.Wk = None
|
|
131
|
+
self.Wv = None
|
|
132
|
+
self.Wo = None
|
|
133
|
+
|
|
134
|
+
def build(self, input_shape):
|
|
135
|
+
shape = (self.d_model, self.d_model)
|
|
136
|
+
self.Wq = self.add_weight(shape=shape, initializer="glorot_uniform", name="Wq")
|
|
137
|
+
self.Wk = self.add_weight(shape=shape, initializer="glorot_uniform", name="Wk")
|
|
138
|
+
self.Wv = self.add_weight(shape=shape, initializer="glorot_uniform", name="Wv")
|
|
139
|
+
self.Wo = self.add_weight(shape=shape, initializer="glorot_uniform", name="Wo")
|
|
140
|
+
self.built = True
|
|
141
|
+
|
|
142
|
+
def call(self, x, mask=None):
|
|
143
|
+
context = causal_attention(x @ self.Wq, x @ self.Wk, x @ self.Wv, self.num_heads, self.rope)
|
|
144
|
+
return context @ self.Wo
|
|
145
|
+
|
|
146
|
+
def state_dict(self):
|
|
147
|
+
return {
|
|
148
|
+
"Wq": self.Wq.data.copy(),
|
|
149
|
+
"Wk": self.Wk.data.copy(),
|
|
150
|
+
"Wv": self.Wv.data.copy(),
|
|
151
|
+
"Wo": self.Wo.data.copy(),
|
|
152
|
+
}
|
|
153
|
+
|
|
154
|
+
def load_state_dict(self, state):
|
|
155
|
+
self.Wq.data[:] = state["Wq"]
|
|
156
|
+
self.Wk.data[:] = state["Wk"]
|
|
157
|
+
self.Wv.data[:] = state["Wv"]
|
|
158
|
+
self.Wo.data[:] = state["Wo"]
|