minigrad-framework 1.0.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.
- minigrad/__init__.py +155 -0
- minigrad/autograd.py +379 -0
- minigrad/avyaya.py +376 -0
- minigrad/cli.py +51 -0
- minigrad/compiler.py +732 -0
- minigrad/data/__init__.py +18 -0
- minigrad/data/dataloader.py +69 -0
- minigrad/data/dataset.py +111 -0
- minigrad/data/transforms.py +67 -0
- minigrad/dp.py +256 -0
- minigrad/glassbox.py +601 -0
- minigrad/graph.py +174 -0
- minigrad/graph_opt.py +609 -0
- minigrad/nn/__init__.py +86 -0
- minigrad/nn/activations.py +150 -0
- minigrad/nn/attention.py +189 -0
- minigrad/nn/batchnorm.py +240 -0
- minigrad/nn/conv.py +217 -0
- minigrad/nn/dropout.py +103 -0
- minigrad/nn/embedding.py +65 -0
- minigrad/nn/flatten.py +15 -0
- minigrad/nn/layernorm.py +91 -0
- minigrad/nn/linear.py +80 -0
- minigrad/nn/lora.py +194 -0
- minigrad/nn/loss.py +262 -0
- minigrad/nn/module.py +359 -0
- minigrad/nn/sequential.py +50 -0
- minigrad/nn/utils.py +59 -0
- minigrad/ops.py +350 -0
- minigrad/optim/__init__.py +13 -0
- minigrad/optim/adam.py +161 -0
- minigrad/optim/base.py +35 -0
- minigrad/optim/rmsprop.py +94 -0
- minigrad/optim/schedulers.py +103 -0
- minigrad/optim/sgd.py +82 -0
- minigrad/pramana.py +492 -0
- minigrad/safetensors.py +264 -0
- minigrad/spanda.py +480 -0
- minigrad/sutra.py +899 -0
- minigrad/tarka.py +627 -0
- minigrad/tensor.py +771 -0
- minigrad/utils.py +270 -0
- minigrad/vmap.py +411 -0
- minigrad_framework-1.0.0.dist-info/METADATA +409 -0
- minigrad_framework-1.0.0.dist-info/RECORD +49 -0
- minigrad_framework-1.0.0.dist-info/WHEEL +5 -0
- minigrad_framework-1.0.0.dist-info/entry_points.txt +6 -0
- minigrad_framework-1.0.0.dist-info/licenses/LICENSE +21 -0
- minigrad_framework-1.0.0.dist-info/top_level.txt +1 -0
minigrad/__init__.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""
|
|
2
|
+
miniGrad — A deep learning framework built from scratch.
|
|
3
|
+
|
|
4
|
+
Zero dependencies except NumPy. Provides:
|
|
5
|
+
- Autograd engine with dynamic computation graphs
|
|
6
|
+
- Neural network layers (Linear, Conv2D, BatchNorm, etc.)
|
|
7
|
+
- Optimizers (SGD, RMSprop, Adam)
|
|
8
|
+
- Loss functions (MSE, CrossEntropy, BCE)
|
|
9
|
+
- Data loading utilities
|
|
10
|
+
|
|
11
|
+
Usage:
|
|
12
|
+
from minigrad import Tensor
|
|
13
|
+
from minigrad.nn import Sequential, Linear, ReLU
|
|
14
|
+
from minigrad.optim import Adam
|
|
15
|
+
|
|
16
|
+
model = Sequential([Linear(784, 128), ReLU(), Linear(128, 10)])
|
|
17
|
+
optimizer = Adam(model.parameters(), lr=1e-3)
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
__version__ = "1.0.0"
|
|
21
|
+
|
|
22
|
+
from minigrad.tensor import Tensor
|
|
23
|
+
from minigrad.graph import topological_sort, trace, print_graph
|
|
24
|
+
from minigrad.safetensors import save_file, load_file, safe_open
|
|
25
|
+
from minigrad.autograd import grad, hessian
|
|
26
|
+
from minigrad.glassbox import (
|
|
27
|
+
detect_anomaly,
|
|
28
|
+
explain_gradients,
|
|
29
|
+
visualize,
|
|
30
|
+
collect_telemetry,
|
|
31
|
+
GradientAnomalyError,
|
|
32
|
+
NodeTelemetry,
|
|
33
|
+
)
|
|
34
|
+
from minigrad.compiler import export_c, to_c
|
|
35
|
+
from minigrad.graph_opt import optimize, optimize_graph, OptimizationReport
|
|
36
|
+
from minigrad.vmap import (
|
|
37
|
+
vmap,
|
|
38
|
+
make_functional,
|
|
39
|
+
per_sample_gradients,
|
|
40
|
+
jacrev,
|
|
41
|
+
batched_jacobian,
|
|
42
|
+
)
|
|
43
|
+
from minigrad.dp import (
|
|
44
|
+
clip_per_sample_gradients,
|
|
45
|
+
add_dp_noise,
|
|
46
|
+
apply_dp_gradients,
|
|
47
|
+
compute_dp_sgd_step,
|
|
48
|
+
PrivacyTelemetry,
|
|
49
|
+
)
|
|
50
|
+
from minigrad.sutra import (
|
|
51
|
+
SUTRA,
|
|
52
|
+
NeuralODE,
|
|
53
|
+
odeint,
|
|
54
|
+
AdaptiveStepTelemetry,
|
|
55
|
+
)
|
|
56
|
+
from minigrad.avyaya import (
|
|
57
|
+
AVYAYA,
|
|
58
|
+
ReversibleBlock,
|
|
59
|
+
ReversibleSequential,
|
|
60
|
+
ReconstructionTelemetry,
|
|
61
|
+
)
|
|
62
|
+
from minigrad.pramana import (
|
|
63
|
+
PRAMANA,
|
|
64
|
+
DistributionalTensor,
|
|
65
|
+
DistributionalLinear,
|
|
66
|
+
DistributionalSequential,
|
|
67
|
+
GaussianNLLLoss,
|
|
68
|
+
PramanaTelemetry,
|
|
69
|
+
)
|
|
70
|
+
from minigrad.tarka import (
|
|
71
|
+
TARKA,
|
|
72
|
+
LogicTensor,
|
|
73
|
+
NeuralPredicate,
|
|
74
|
+
NeuralRelation,
|
|
75
|
+
SemanticLoss,
|
|
76
|
+
TarkaTelemetry,
|
|
77
|
+
)
|
|
78
|
+
from minigrad.spanda import (
|
|
79
|
+
SPANDA,
|
|
80
|
+
surrogate_spike,
|
|
81
|
+
LIFCell,
|
|
82
|
+
LIFLayer,
|
|
83
|
+
SpikingLinear,
|
|
84
|
+
SpikingSequential,
|
|
85
|
+
RateEncoder,
|
|
86
|
+
DirectEncoder,
|
|
87
|
+
RateDecoder,
|
|
88
|
+
MembraneDecoder,
|
|
89
|
+
SpandaTelemetry,
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
__all__ = [
|
|
93
|
+
"Tensor",
|
|
94
|
+
"topological_sort",
|
|
95
|
+
"trace",
|
|
96
|
+
"print_graph",
|
|
97
|
+
"save_file",
|
|
98
|
+
"load_file",
|
|
99
|
+
"safe_open",
|
|
100
|
+
"grad",
|
|
101
|
+
"hessian",
|
|
102
|
+
"detect_anomaly",
|
|
103
|
+
"explain_gradients",
|
|
104
|
+
"visualize",
|
|
105
|
+
"collect_telemetry",
|
|
106
|
+
"GradientAnomalyError",
|
|
107
|
+
"NodeTelemetry",
|
|
108
|
+
"export_c",
|
|
109
|
+
"to_c",
|
|
110
|
+
"optimize",
|
|
111
|
+
"optimize_graph",
|
|
112
|
+
"OptimizationReport",
|
|
113
|
+
"vmap",
|
|
114
|
+
"make_functional",
|
|
115
|
+
"per_sample_gradients",
|
|
116
|
+
"jacrev",
|
|
117
|
+
"batched_jacobian",
|
|
118
|
+
"clip_per_sample_gradients",
|
|
119
|
+
"add_dp_noise",
|
|
120
|
+
"apply_dp_gradients",
|
|
121
|
+
"compute_dp_sgd_step",
|
|
122
|
+
"PrivacyTelemetry",
|
|
123
|
+
"SUTRA",
|
|
124
|
+
"NeuralODE",
|
|
125
|
+
"odeint",
|
|
126
|
+
"AdaptiveStepTelemetry",
|
|
127
|
+
"AVYAYA",
|
|
128
|
+
"ReversibleBlock",
|
|
129
|
+
"ReversibleSequential",
|
|
130
|
+
"ReconstructionTelemetry",
|
|
131
|
+
"PRAMANA",
|
|
132
|
+
"DistributionalTensor",
|
|
133
|
+
"DistributionalLinear",
|
|
134
|
+
"DistributionalSequential",
|
|
135
|
+
"GaussianNLLLoss",
|
|
136
|
+
"PramanaTelemetry",
|
|
137
|
+
"TARKA",
|
|
138
|
+
"LogicTensor",
|
|
139
|
+
"NeuralPredicate",
|
|
140
|
+
"NeuralRelation",
|
|
141
|
+
"SemanticLoss",
|
|
142
|
+
"TarkaTelemetry",
|
|
143
|
+
"SPANDA",
|
|
144
|
+
"surrogate_spike",
|
|
145
|
+
"LIFCell",
|
|
146
|
+
"LIFLayer",
|
|
147
|
+
"SpikingLinear",
|
|
148
|
+
"SpikingSequential",
|
|
149
|
+
"RateEncoder",
|
|
150
|
+
"DirectEncoder",
|
|
151
|
+
"RateDecoder",
|
|
152
|
+
"MembraneDecoder",
|
|
153
|
+
"SpandaTelemetry",
|
|
154
|
+
"__version__",
|
|
155
|
+
]
|
minigrad/autograd.py
ADDED
|
@@ -0,0 +1,379 @@
|
|
|
1
|
+
"""
|
|
2
|
+
autograd.py — Higher-order automatic differentiation engine and functional autograd API.
|
|
3
|
+
|
|
4
|
+
Provides:
|
|
5
|
+
- grad(): Computes and returns the sum of gradients of outputs with respect to the inputs.
|
|
6
|
+
When create_graph=True, constructs a differentiable computation graph for higher-order derivatives.
|
|
7
|
+
- hessian(): Computes the full Hessian matrix H_ij = ∂²y / (∂x_i ∂x_j).
|
|
8
|
+
|
|
9
|
+
Unlocks Physics-Informed Neural Networks (PINNs), gradient penalties (WGAN-GP),
|
|
10
|
+
MAML meta-learning, and curvature/Hessian-vector products.
|
|
11
|
+
"""
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
from typing import Any, Dict, List, Optional, Sequence, Set, Tuple, Union
|
|
15
|
+
|
|
16
|
+
import numpy as np
|
|
17
|
+
|
|
18
|
+
from minigrad.tensor import Tensor
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
# ── VJP Unbroadcasting Helper ────────────────────────────────────────
|
|
22
|
+
|
|
23
|
+
def unbroadcast_tensor(t: Tensor, target_shape: Tuple[int, ...]) -> Tensor:
|
|
24
|
+
"""
|
|
25
|
+
Unbroadcast a Tensor back to target_shape using differentiable Tensor operations.
|
|
26
|
+
"""
|
|
27
|
+
if t.data.shape == target_shape:
|
|
28
|
+
return t
|
|
29
|
+
|
|
30
|
+
# 1. Reduce extra leading dimensions
|
|
31
|
+
while t.data.ndim > len(target_shape):
|
|
32
|
+
t = t.sum(axis=0, keepdims=False)
|
|
33
|
+
|
|
34
|
+
# 2. Reduce dimensions that were broadcast from size 1
|
|
35
|
+
for i, (dim, target_dim) in enumerate(zip(t.data.shape, target_shape)):
|
|
36
|
+
if dim != target_dim and target_dim == 1:
|
|
37
|
+
t = t.sum(axis=i, keepdims=True)
|
|
38
|
+
|
|
39
|
+
if t.data.shape != target_shape:
|
|
40
|
+
t = t.reshape(*target_shape)
|
|
41
|
+
|
|
42
|
+
return t
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
# ── Symbolic Vector-Jacobian Product (VJP) Dispatch ─────────────────
|
|
46
|
+
|
|
47
|
+
def _compute_vjp(node: Tensor, g: Tensor) -> Tuple[Optional[Tensor], ...]:
|
|
48
|
+
"""
|
|
49
|
+
Compute Vector-Jacobian Products for each parent of node w.r.t incoming gradient g.
|
|
50
|
+
Returns a tuple of Tensors (or None) corresponding to node._prev.
|
|
51
|
+
"""
|
|
52
|
+
op = node._op
|
|
53
|
+
children = node._prev
|
|
54
|
+
|
|
55
|
+
if op == "add":
|
|
56
|
+
a, b = children
|
|
57
|
+
vjp_a = unbroadcast_tensor(g, a.shape) if a.requires_grad else None
|
|
58
|
+
vjp_b = unbroadcast_tensor(g, b.shape) if b.requires_grad else None
|
|
59
|
+
return (vjp_a, vjp_b)
|
|
60
|
+
|
|
61
|
+
elif op == "mul":
|
|
62
|
+
a, b = children
|
|
63
|
+
vjp_a = unbroadcast_tensor(g * b, a.shape) if a.requires_grad else None
|
|
64
|
+
vjp_b = unbroadcast_tensor(g * a, b.shape) if b.requires_grad else None
|
|
65
|
+
return (vjp_a, vjp_b)
|
|
66
|
+
|
|
67
|
+
elif op == "neg":
|
|
68
|
+
(a,) = children
|
|
69
|
+
return (-g if a.requires_grad else None,)
|
|
70
|
+
|
|
71
|
+
elif op.startswith("pow"):
|
|
72
|
+
if len(children) == 1:
|
|
73
|
+
(a,) = children
|
|
74
|
+
p = getattr(node, "_ctx", None)
|
|
75
|
+
if p is None and "^" in op:
|
|
76
|
+
try:
|
|
77
|
+
p = float(op.split("^", 1)[1])
|
|
78
|
+
except (ValueError, IndexError):
|
|
79
|
+
p = 1.0
|
|
80
|
+
if p == 0:
|
|
81
|
+
return (None,)
|
|
82
|
+
deriv = (a ** (p - 1)) * p
|
|
83
|
+
vjp_a = unbroadcast_tensor(g * deriv, a.shape) if a.requires_grad else None
|
|
84
|
+
return (vjp_a,)
|
|
85
|
+
else:
|
|
86
|
+
a, b = children
|
|
87
|
+
# d(a^b)/da = b * a^(b-1)
|
|
88
|
+
# d(a^b)/db = a^b * ln(a)
|
|
89
|
+
vjp_a = unbroadcast_tensor(g * b * (a ** (b - 1)), a.shape) if a.requires_grad else None
|
|
90
|
+
vjp_b = unbroadcast_tensor(g * node * a.log(), b.shape) if b.requires_grad else None
|
|
91
|
+
return (vjp_a, vjp_b)
|
|
92
|
+
|
|
93
|
+
elif op == "matmul":
|
|
94
|
+
a, b = children
|
|
95
|
+
if a.data.ndim == 1 and b.data.ndim == 2:
|
|
96
|
+
g_2d = g.reshape(1, -1) if g.data.ndim == 1 else g
|
|
97
|
+
a_col = a.reshape(-1, 1)
|
|
98
|
+
vjp_a = (g_2d @ b.transpose()).reshape(a.shape) if a.requires_grad else None
|
|
99
|
+
vjp_b = (a_col @ g_2d) if b.requires_grad else None
|
|
100
|
+
return (vjp_a, vjp_b)
|
|
101
|
+
elif a.data.ndim == 2 and b.data.ndim == 1:
|
|
102
|
+
g_2d = g.reshape(-1, 1) if g.data.ndim == 1 else g
|
|
103
|
+
b_row = b.reshape(1, -1)
|
|
104
|
+
vjp_a = (g_2d @ b_row) if a.requires_grad else None
|
|
105
|
+
vjp_b = (a.transpose() @ g_2d).reshape(b.shape) if b.requires_grad else None
|
|
106
|
+
return (vjp_a, vjp_b)
|
|
107
|
+
else:
|
|
108
|
+
vjp_a = (g @ b.transpose()) if a.requires_grad else None
|
|
109
|
+
vjp_b = (a.transpose() @ g) if b.requires_grad else None
|
|
110
|
+
return (vjp_a, vjp_b)
|
|
111
|
+
|
|
112
|
+
elif op == "tanh":
|
|
113
|
+
(a,) = children
|
|
114
|
+
one = Tensor(np.ones_like(node.data))
|
|
115
|
+
deriv = one - (node ** 2)
|
|
116
|
+
return (g * deriv if a.requires_grad else None,)
|
|
117
|
+
|
|
118
|
+
elif op == "sin":
|
|
119
|
+
(a,) = children
|
|
120
|
+
return (g * a.cos() if a.requires_grad else None,)
|
|
121
|
+
|
|
122
|
+
elif op == "cos":
|
|
123
|
+
(a,) = children
|
|
124
|
+
return (-g * a.sin() if a.requires_grad else None,)
|
|
125
|
+
|
|
126
|
+
elif op == "exp":
|
|
127
|
+
(a,) = children
|
|
128
|
+
return (g * node if a.requires_grad else None,)
|
|
129
|
+
|
|
130
|
+
elif op == "log":
|
|
131
|
+
(a,) = children
|
|
132
|
+
return (g / a if a.requires_grad else None,)
|
|
133
|
+
|
|
134
|
+
elif op == "relu":
|
|
135
|
+
(a,) = children
|
|
136
|
+
mask = Tensor((a.data > 0).astype(np.float64))
|
|
137
|
+
return (g * mask if a.requires_grad else None,)
|
|
138
|
+
|
|
139
|
+
elif op == "sigmoid":
|
|
140
|
+
(a,) = children
|
|
141
|
+
one = Tensor(np.ones_like(node.data))
|
|
142
|
+
deriv = node * (one - node)
|
|
143
|
+
return (g * deriv if a.requires_grad else None,)
|
|
144
|
+
|
|
145
|
+
elif op == "sum":
|
|
146
|
+
(a,) = children
|
|
147
|
+
ctx = getattr(node, "_ctx", None)
|
|
148
|
+
if ctx is not None and isinstance(ctx, tuple) and len(ctx) == 3:
|
|
149
|
+
axes, keepdims, orig_shape = ctx
|
|
150
|
+
else:
|
|
151
|
+
axes, keepdims, orig_shape = None, False, a.shape
|
|
152
|
+
|
|
153
|
+
if axes is not None and not keepdims:
|
|
154
|
+
expanded_shape = list(orig_shape)
|
|
155
|
+
for ax in axes:
|
|
156
|
+
expanded_shape[ax] = 1
|
|
157
|
+
g_expanded = g.reshape(*expanded_shape)
|
|
158
|
+
else:
|
|
159
|
+
g_expanded = g
|
|
160
|
+
|
|
161
|
+
vjp_a = g_expanded * Tensor(np.ones(orig_shape)) if a.requires_grad else None
|
|
162
|
+
return (vjp_a,)
|
|
163
|
+
|
|
164
|
+
elif op == "mean":
|
|
165
|
+
(a,) = children
|
|
166
|
+
ctx = getattr(node, "_ctx", None)
|
|
167
|
+
if ctx is not None and isinstance(ctx, tuple) and len(ctx) == 4:
|
|
168
|
+
axes, keepdims, orig_shape, n = ctx
|
|
169
|
+
else:
|
|
170
|
+
axes, keepdims, orig_shape, n = None, False, a.shape, a.data.size
|
|
171
|
+
|
|
172
|
+
if axes is not None and not keepdims:
|
|
173
|
+
expanded_shape = list(orig_shape)
|
|
174
|
+
for ax in axes:
|
|
175
|
+
expanded_shape[ax] = 1
|
|
176
|
+
g_expanded = g.reshape(*expanded_shape)
|
|
177
|
+
else:
|
|
178
|
+
g_expanded = g
|
|
179
|
+
|
|
180
|
+
vjp_a = (g_expanded * (1.0 / n)) * Tensor(np.ones(orig_shape)) if a.requires_grad else None
|
|
181
|
+
return (vjp_a,)
|
|
182
|
+
|
|
183
|
+
elif op == "reshape":
|
|
184
|
+
(a,) = children
|
|
185
|
+
orig_shape = getattr(node, "_ctx", a.shape)
|
|
186
|
+
return (g.reshape(*orig_shape) if a.requires_grad else None,)
|
|
187
|
+
|
|
188
|
+
elif op == "transpose":
|
|
189
|
+
(a,) = children
|
|
190
|
+
axes_info = getattr(node, "_ctx", None)
|
|
191
|
+
if axes_info is not None:
|
|
192
|
+
_, inv_axes = axes_info
|
|
193
|
+
return (g.transpose(*inv_axes) if a.requires_grad else None,)
|
|
194
|
+
return (g.transpose() if a.requires_grad else None,)
|
|
195
|
+
|
|
196
|
+
elif op == "getitem":
|
|
197
|
+
(a,) = children
|
|
198
|
+
idx = getattr(node, "_ctx", None)
|
|
199
|
+
grad_a = np.zeros(a.shape, dtype=np.float64)
|
|
200
|
+
np.add.at(grad_a, idx, g.data)
|
|
201
|
+
vjp_a = Tensor(grad_a, requires_grad=True) if a.requires_grad else None
|
|
202
|
+
return (vjp_a,)
|
|
203
|
+
|
|
204
|
+
elif op == "stack":
|
|
205
|
+
axis = getattr(node, "_ctx", 0)
|
|
206
|
+
vjps = []
|
|
207
|
+
for i, child in enumerate(children):
|
|
208
|
+
if child.requires_grad:
|
|
209
|
+
idx = [slice(None)] * g.data.ndim
|
|
210
|
+
idx[axis] = i
|
|
211
|
+
vjps.append(g[tuple(idx)])
|
|
212
|
+
else:
|
|
213
|
+
vjps.append(None)
|
|
214
|
+
return tuple(vjps)
|
|
215
|
+
|
|
216
|
+
elif op == "concat":
|
|
217
|
+
axis = getattr(node, "_ctx", 0)
|
|
218
|
+
vjps = []
|
|
219
|
+
offset = 0
|
|
220
|
+
for child in children:
|
|
221
|
+
length = child.shape[axis]
|
|
222
|
+
if child.requires_grad:
|
|
223
|
+
idx = [slice(None)] * g.data.ndim
|
|
224
|
+
idx[axis] = slice(offset, offset + length)
|
|
225
|
+
vjps.append(g[tuple(idx)])
|
|
226
|
+
else:
|
|
227
|
+
vjps.append(None)
|
|
228
|
+
offset += length
|
|
229
|
+
return tuple(vjps)
|
|
230
|
+
|
|
231
|
+
# Fallback: default to None for unhandled operations
|
|
232
|
+
return tuple(None for _ in children)
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
# ── Functional Autograd API ──────────────────────────────────────────
|
|
236
|
+
|
|
237
|
+
def grad(
|
|
238
|
+
outputs: Union[Tensor, Sequence[Tensor]],
|
|
239
|
+
inputs: Union[Tensor, Sequence[Tensor]],
|
|
240
|
+
grad_outputs: Optional[Union[Tensor, Sequence[Tensor]]] = None,
|
|
241
|
+
retain_graph: bool = False,
|
|
242
|
+
create_graph: bool = False,
|
|
243
|
+
allow_unused: bool = False,
|
|
244
|
+
) -> Tuple[Tensor, ...]:
|
|
245
|
+
"""
|
|
246
|
+
Computes and returns the sum of gradients of outputs with respect to the inputs.
|
|
247
|
+
|
|
248
|
+
Args:
|
|
249
|
+
outputs: Output tensor(s) to differentiate.
|
|
250
|
+
inputs: Input tensor(s) w.r.t which gradients are computed.
|
|
251
|
+
grad_outputs: Initial vector in the Vector-Jacobian Product. If None,
|
|
252
|
+
defaults to ones of shape like output (outputs must be scalar).
|
|
253
|
+
retain_graph: If False, the graph used to compute the grads will be freed.
|
|
254
|
+
create_graph: If True, graph of the derivative will be constructed,
|
|
255
|
+
allowing higher-order derivative products to be computed.
|
|
256
|
+
allow_unused: If False, raising an error when an input is unused in the graph.
|
|
257
|
+
|
|
258
|
+
Returns:
|
|
259
|
+
Tuple of Tensors containing gradients for each input.
|
|
260
|
+
"""
|
|
261
|
+
outputs_list: List[Tensor] = [outputs] if isinstance(outputs, Tensor) else list(outputs)
|
|
262
|
+
inputs_list: List[Tensor] = [inputs] if isinstance(inputs, Tensor) else list(inputs)
|
|
263
|
+
|
|
264
|
+
# Initialize grad_outputs
|
|
265
|
+
if grad_outputs is None:
|
|
266
|
+
grad_outputs_list: List[Tensor] = []
|
|
267
|
+
for out in outputs_list:
|
|
268
|
+
if out.data.size != 1:
|
|
269
|
+
raise RuntimeError("grad can only be implicitly created for scalar outputs")
|
|
270
|
+
grad_outputs_list.append(
|
|
271
|
+
Tensor(np.ones_like(out.data), requires_grad=create_graph)
|
|
272
|
+
)
|
|
273
|
+
else:
|
|
274
|
+
if isinstance(grad_outputs, Tensor):
|
|
275
|
+
grad_outputs_list = [grad_outputs]
|
|
276
|
+
else:
|
|
277
|
+
grad_outputs_list = [
|
|
278
|
+
g if isinstance(g, Tensor) else Tensor(g, requires_grad=create_graph)
|
|
279
|
+
for g in grad_outputs
|
|
280
|
+
]
|
|
281
|
+
|
|
282
|
+
# Build topological sort from outputs
|
|
283
|
+
topo: List[Tensor] = []
|
|
284
|
+
visited: Set[int] = set()
|
|
285
|
+
|
|
286
|
+
def build_topo(node: Tensor) -> None:
|
|
287
|
+
if id(node) not in visited:
|
|
288
|
+
visited.add(id(node))
|
|
289
|
+
for child in node._prev:
|
|
290
|
+
build_topo(child)
|
|
291
|
+
topo.append(node)
|
|
292
|
+
|
|
293
|
+
for out in outputs_list:
|
|
294
|
+
build_topo(out)
|
|
295
|
+
|
|
296
|
+
# Seed gradient map
|
|
297
|
+
grad_map: Dict[int, Tensor] = {}
|
|
298
|
+
for out, gout in zip(outputs_list, grad_outputs_list):
|
|
299
|
+
if id(out) in grad_map:
|
|
300
|
+
grad_map[id(out)] = grad_map[id(out)] + gout
|
|
301
|
+
else:
|
|
302
|
+
grad_map[id(out)] = gout
|
|
303
|
+
|
|
304
|
+
# Reverse topological order traversal
|
|
305
|
+
for node in reversed(topo):
|
|
306
|
+
if id(node) not in grad_map:
|
|
307
|
+
continue
|
|
308
|
+
g = grad_map[id(node)]
|
|
309
|
+
if not node._prev or not node._op:
|
|
310
|
+
continue
|
|
311
|
+
|
|
312
|
+
vjps = _compute_vjp(node, g)
|
|
313
|
+
for child, vjp in zip(node._prev, vjps):
|
|
314
|
+
if vjp is None or not child.requires_grad:
|
|
315
|
+
continue
|
|
316
|
+
if id(child) in grad_map:
|
|
317
|
+
grad_map[id(child)] = grad_map[id(child)] + vjp
|
|
318
|
+
else:
|
|
319
|
+
grad_map[id(child)] = vjp
|
|
320
|
+
|
|
321
|
+
# Collect gradients for requested inputs
|
|
322
|
+
result: List[Tensor] = []
|
|
323
|
+
for inp in inputs_list:
|
|
324
|
+
res = grad_map.get(id(inp), None)
|
|
325
|
+
if res is None:
|
|
326
|
+
if not allow_unused:
|
|
327
|
+
raise RuntimeError(
|
|
328
|
+
f"One of the differentiated Tensors appears to not have been used in the graph. "
|
|
329
|
+
f"Set allow_unused=True if this is the desired behavior."
|
|
330
|
+
)
|
|
331
|
+
res = Tensor(np.zeros_like(inp.data), requires_grad=create_graph)
|
|
332
|
+
result.append(res)
|
|
333
|
+
|
|
334
|
+
return tuple(result)
|
|
335
|
+
|
|
336
|
+
|
|
337
|
+
def hessian(
|
|
338
|
+
output: Tensor,
|
|
339
|
+
inputs: Union[Tensor, Sequence[Tensor]],
|
|
340
|
+
create_graph: bool = False,
|
|
341
|
+
) -> Tensor:
|
|
342
|
+
"""
|
|
343
|
+
Compute the full Hessian matrix H_ij = ∂²y / (∂x_i ∂x_j) of a scalar output.
|
|
344
|
+
|
|
345
|
+
Args:
|
|
346
|
+
output: Scalar output tensor.
|
|
347
|
+
inputs: Input tensor(s).
|
|
348
|
+
create_graph: If True, the returned Hessian will be differentiable.
|
|
349
|
+
|
|
350
|
+
Returns:
|
|
351
|
+
2D Tensor of shape (total_input_elements, total_input_elements).
|
|
352
|
+
"""
|
|
353
|
+
if output.data.size != 1:
|
|
354
|
+
raise ValueError("hessian requires a scalar output tensor")
|
|
355
|
+
|
|
356
|
+
inputs_list: List[Tensor] = [inputs] if isinstance(inputs, Tensor) else list(inputs)
|
|
357
|
+
|
|
358
|
+
# First gradient vector (with create_graph=True)
|
|
359
|
+
first_grads = grad(output, inputs_list, create_graph=True)
|
|
360
|
+
|
|
361
|
+
rows: List[Tensor] = []
|
|
362
|
+
for g_k, inp_k in zip(first_grads, inputs_list):
|
|
363
|
+
k_size = inp_k.data.size
|
|
364
|
+
for local_i in range(k_size):
|
|
365
|
+
# Basis vector matching shape of inp_k
|
|
366
|
+
basis = np.zeros(inp_k.data.shape, dtype=np.float64)
|
|
367
|
+
basis.flat[local_i] = 1.0
|
|
368
|
+
basis_t = Tensor(basis, requires_grad=False)
|
|
369
|
+
|
|
370
|
+
# Project first gradient: scalar g_proj = (g_k * basis_t).sum()
|
|
371
|
+
g_proj = (g_k * basis_t).sum()
|
|
372
|
+
|
|
373
|
+
# Backprop to get the row across all inputs
|
|
374
|
+
row_grads = grad(g_proj, inputs_list, retain_graph=True, create_graph=create_graph)
|
|
375
|
+
flat_row = np.concatenate([rg.data.flatten() for rg in row_grads])
|
|
376
|
+
rows.append(Tensor(flat_row, requires_grad=create_graph))
|
|
377
|
+
|
|
378
|
+
hessian_matrix = np.stack([r.data for r in rows])
|
|
379
|
+
return Tensor(hessian_matrix, requires_grad=create_graph)
|