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,14 @@
1
+ import numpy as np
2
+
3
+ class RandomUniform:
4
+
5
+ def __init__(self, low=-0.05, high=0.05):
6
+ self.low = low
7
+ self.high = high
8
+
9
+ def __call__(self, shape):
10
+ return np.random.uniform(
11
+ self.low,
12
+ self.high,
13
+ shape,
14
+ ).astype(np.float32)
@@ -0,0 +1,8 @@
1
+ import numpy as np
2
+
3
+ from src.initializers.Initializer import Initializer
4
+
5
+
6
+ class Zeros(Initializer):
7
+ def __call__(self, shape):
8
+ return np.zeros(shape, dtype=np.float32)
@@ -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,9 @@
1
+ class CategoricalCrossEntropy:
2
+
3
+ def __call__(self, prediction, target):
4
+
5
+ eps = 1e-7
6
+
7
+ prediction = prediction.clip(eps, 1)
8
+
9
+ return -(target * prediction.log()).sum(axis=1).mean()
@@ -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
@@ -0,0 +1,5 @@
1
+ class Hinge:
2
+
3
+ def __call__(self, prediction, target):
4
+
5
+ return (1 - target * prediction).maximum(0).mean()
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
@@ -0,0 +1,6 @@
1
+ class Loss:
2
+ def forward(self, pred, target):
3
+ raise NotImplementedError
4
+
5
+ def backward(self, pred, target):
6
+ raise NotImplementedError
src/loss/MSE.py ADDED
@@ -0,0 +1,7 @@
1
+ class MSE:
2
+
3
+ def __call__(self, prediction, target):
4
+
5
+ diff = prediction - target
6
+
7
+ return (diff * diff).mean()
src/loss/MSELoss.py ADDED
@@ -0,0 +1,10 @@
1
+ import numpy as np
2
+ from src.loss.Loss import Loss
3
+
4
+ class MSELoss(Loss):
5
+
6
+ def forward(self, pred, target):
7
+ return 0.5 * np.mean((pred - target) ** 2)
8
+
9
+ def backward(self, pred, target):
10
+ return pred - target
@@ -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
@@ -0,0 +1,5 @@
1
+ import numpy as np
2
+
3
+
4
+ def sigmoid(x):
5
+ return 1 / (1 + np.exp(-x))
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"]