Skip to content

Commit 8747006

Browse files
committed
Implement dim-aware vectorize_graph
1 parent 9d99267 commit 8747006

11 files changed

Lines changed: 599 additions & 58 deletions

File tree

pytensor/graph/replace.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -283,6 +283,13 @@ def vectorize_graph(
283283
# [array([-10., -11.]), array([10., 11.])]
284284
285285
"""
286+
# TODO: Move this to tensor.vectorize, and make this helper type agnostic.
287+
#
288+
# This helper may dispatch to tensor.vectorize_graph or xtensor.vectorize_graph depending on the replacement types
289+
# The behavior is distinct, because tensor vectorization depends on axis-position while xtensor depends on dimension labels
290+
#
291+
# xtensor.vectorize_graph will be able to handle batched inner tensor operations, while tensor.vectorize_graph won't,
292+
# as it is by design unaware of xtensors and their semantics.
286293
if isinstance(outputs, Sequence):
287294
seq_outputs = outputs
288295
else:

pytensor/tensor/extra_ops.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
import warnings
2-
from collections.abc import Collection, Iterable
2+
from collections.abc import Collection, Iterable, Sequence
33
from textwrap import dedent
44

55
import numpy as np
@@ -1926,7 +1926,7 @@ def logspace(
19261926

19271927

19281928
def broadcast_to(
1929-
x: TensorVariable, shape: TensorVariable | tuple[Variable, ...]
1929+
x: TensorLike, shape: TensorLike | Sequence[TensorLike]
19301930
) -> TensorVariable:
19311931
"""Broadcast an array to a new shape.
19321932

pytensor/xtensor/basic.py

Lines changed: 19 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,9 @@ def perform(self, node, inputs, outputs):
1818
def do_constant_folding(self, fgraph, node):
1919
return False
2020

21-
def vectorize_node(self, node, *new_inputs) -> Sequence[Variable]:
21+
def vectorize_node(
22+
self, node, *new_inputs, new_dim: str | None
23+
) -> Sequence[Variable]:
2224
raise NotImplementedError(f"Vectorized node not implemented for {self}")
2325

2426

@@ -31,7 +33,9 @@ class XTypeCastOp(TypeCastingOp):
3133
def infer_shape(self, fgraph, node, input_shapes):
3234
return input_shapes
3335

34-
def vectorize_node(self, node, *new_inputs) -> Sequence[Variable]:
36+
def vectorize_node(
37+
self, node, *new_inputs, new_dim: str | None
38+
) -> Sequence[Variable]:
3539
raise NotImplementedError(f"Vectorized node not implemented for {self}")
3640

3741

@@ -49,12 +53,13 @@ def L_op(self, inputs, outs, g_outs):
4953
[g_out] = g_outs
5054
return [xtensor_from_tensor(g_out, dims=x.type.dims)]
5155

52-
def vectorize_node(self, node, new_x):
56+
def vectorize_node(self, node, new_x, new_dim):
5357
[old_x] = node.inputs
5458
if (new_x.ndim - old_x.ndim) > 1:
5559
raise NotImplementedError(
5660
f"Vectorization of {self} cannot guarantee correct placement of multiple batch dimensions. "
57-
"You can call vectorize_graph one batch dimension at a time."
61+
"You can call vectorize_graph one batch dimension at a time, "
62+
"or pytensor.xtensor.vectorization.vectorize_graph instead."
5863
)
5964
new_x = new_x.transpose(..., *old_x.dims)
6065
return [self(new_x)]
@@ -80,14 +85,17 @@ def L_op(self, inputs, outs, g_outs):
8085
[g_out] = g_outs
8186
return [tensor_from_xtensor(g_out)]
8287

83-
def vectorize_node(self, node, new_x):
88+
def vectorize_node(self, node, new_x, new_dim):
8489
[old_x] = node.inputs
8590
if new_x.ndim != old_x.ndim:
86-
raise NotImplementedError(
87-
f"Vectorization of {self} with batched inputs not implemented, "
88-
"as it can't infer new dimension labels"
89-
)
90-
return [self(new_x)]
91+
if new_dim is None:
92+
raise NotImplementedError(
93+
f"Vectorization of {self} cannot infer the new dimension labels. "
94+
"Use pytensor.xtensor.vectorization.vectorize_graph instead."
95+
)
96+
return [type(self)(dims=(new_dim, *self.dims))(new_x)]
97+
else:
98+
return [self(new_x)]
9199

92100

93101
def xtensor_from_tensor(x, dims, name=None):
@@ -111,7 +119,7 @@ def L_op(self, inputs, outs, g_outs):
111119
[g_out] = g_outs
112120
return [rename(g_out, dims=x.type.dims)]
113121

114-
def vectorize_node(self, node, new_x):
122+
def vectorize_node(self, node, new_x, new_dim):
115123
[old_x] = node.inputs
116124
old_dim_mapping = dict(zip(old_x.dims, self.new_dims, strict=True))
117125

pytensor/xtensor/indexing.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -197,7 +197,7 @@ def combine_dim_info(idx_dim, idx_dim_shape):
197197
output = xtensor(dtype=x.type.dtype, shape=out_shape, dims=out_dims)
198198
return Apply(self, [x, *idxs], [output])
199199

200-
def vectorize_node(self, node, new_x, *new_idxs):
200+
def vectorize_node(self, node, new_x, *new_idxs, new_dim):
201201
# new_x may have dims in different order
202202
# we pair each pre-existing dim to the respective index
203203
# with new dims having simply a slice(None)
@@ -237,7 +237,7 @@ def make_node(self, x, y, *idxs):
237237
out = x.type()
238238
return Apply(self, [x, y, *idxs], [out])
239239

240-
def vectorize_node(self, node, *new_inputs):
240+
def vectorize_node(self, node, *new_inputs, new_dim):
241241
# If y or the indices have new dimensions we need to broadcast_x
242242
exclude: set[str] = set(
243243
chain.from_iterable(

pytensor/xtensor/reduction.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ def make_node(self, x):
4646
output = xtensor(dtype=x.type.dtype, shape=out_shape, dims=out_dims)
4747
return Apply(self, [x], [output])
4848

49-
def vectorize_node(self, node, new_x):
49+
def vectorize_node(self, node, new_x, new_dim):
5050
return [self(new_x)]
5151

5252

@@ -120,7 +120,7 @@ def make_node(self, x):
120120
out = x.type()
121121
return Apply(self, [x], [out])
122122

123-
def vectorize_node(self, node, new_x):
123+
def vectorize_node(self, node, new_x, new_dim):
124124
return [self(new_x)]
125125

126126

pytensor/xtensor/shape.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -68,7 +68,7 @@ def make_node(self, x):
6868
)
6969
return Apply(self, [x], [output])
7070

71-
def vectorize_node(self, node, new_x):
71+
def vectorize_node(self, node, new_x, new_dim):
7272
return [self(new_x)]
7373

7474

@@ -149,7 +149,7 @@ def make_node(self, x, *unstacked_length):
149149
)
150150
return Apply(self, [x, *unstacked_lengths], [output])
151151

152-
def vectorize_node(self, node, new_x, *new_unstacked_length):
152+
def vectorize_node(self, node, new_x, *new_unstacked_length, new_dim):
153153
new_unstacked_length = [ul.squeeze() for ul in new_unstacked_length]
154154
if not all(ul.type.ndim == 0 for ul in new_unstacked_length):
155155
raise NotImplementedError(
@@ -200,7 +200,7 @@ def make_node(self, x):
200200
)
201201
return Apply(self, [x], [output])
202202

203-
def vectorize_node(self, node, new_x):
203+
def vectorize_node(self, node, new_x, new_dim):
204204
old_dims = self.dims
205205
new_dims = tuple(dim for dim in new_x.dims if dim not in old_dims)
206206
return [type(self)(dims=(*new_dims, *old_dims))(new_x)]
@@ -318,7 +318,7 @@ def make_node(self, *inputs):
318318
output = xtensor(dtype=dtype, dims=dims, shape=shape)
319319
return Apply(self, inputs, [output])
320320

321-
def vectorize_node(self, node, *new_inputs):
321+
def vectorize_node(self, node, *new_inputs, new_dim):
322322
return [self(*new_inputs)]
323323

324324

@@ -402,7 +402,7 @@ def make_node(self, x):
402402
)
403403
return Apply(self, [x], [out])
404404

405-
def vectorize_node(self, node, new_x):
405+
def vectorize_node(self, node, new_x, new_dim):
406406
return [self(new_x)]
407407

408408

@@ -464,7 +464,7 @@ def make_node(self, x, size):
464464
)
465465
return Apply(self, [x, size], [out])
466466

467-
def vectorize_node(self, node, new_x, new_size):
467+
def vectorize_node(self, node, new_x, new_size, new_dim):
468468
new_size = new_size.squeeze()
469469
if new_size.type.ndim != 0:
470470
raise NotImplementedError(
@@ -567,7 +567,7 @@ def make_node(self, *inputs):
567567

568568
return Apply(self, inputs, outputs)
569569

570-
def vectorize_node(self, node, *new_inputs):
570+
def vectorize_node(self, node, *new_inputs, new_dim):
571571
if exclude_set := set(self.exclude):
572572
for new_x, old_x in zip(node.inputs, new_inputs, strict=True):
573573
if invalid_excluded := (

0 commit comments

Comments
 (0)