Skip to content

Commit 77aad72

Browse files
giotremuhrin
authored andcommitted
Fix GCNN graph reductions for IrrepsArray inputs
1 parent ec0e8ba commit 77aad72

2 files changed

Lines changed: 242 additions & 3 deletions

File tree

src/tensorial/gcnn/graph_ops.py

Lines changed: 71 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from functools import singledispatch
12
from typing import Literal
23

34
import e3nn_jax as e3j
@@ -35,6 +36,7 @@ def _prepare_segments(
3536
return num_segments, segment_ids
3637

3738

39+
@singledispatch
3840
def segment_sum(
3941
data: Float[jax.Array, "N ..."],
4042
segment_sizes: Int[jax.Array, "num_segments"],
@@ -85,6 +87,23 @@ def segment_sum(
8587
return result
8688

8789

90+
@segment_sum.register
91+
def _(
92+
data: e3j.IrrepsArray,
93+
segment_sizes: Int[jax.Array, "num_segments"],
94+
mask: Bool[jax.Array, "N ..."] | None = None,
95+
segment_mask: Bool[jax.Array, "num_segments"] | None = None,
96+
) -> e3j.IrrepsArray:
97+
result = segment_sum(
98+
data.array,
99+
segment_sizes,
100+
mask=mask,
101+
segment_mask=segment_mask,
102+
)
103+
return e3j.IrrepsArray(data.irreps, result)
104+
105+
106+
@singledispatch
88107
def segment_mean(
89108
data: Float[jax.Array, "N ..."],
90109
segment_sizes: Int[jax.Array, "num_segments"],
@@ -158,6 +177,23 @@ def segment_mean(
158177
return jnp.where(segment_mask, mean, jnp.zeros_like(mean))
159178

160179

180+
@segment_mean.register
181+
def _(
182+
data: e3j.IrrepsArray,
183+
segment_sizes: Int[jax.Array, "num_segments"],
184+
mask: Bool[jax.Array, "N ..."] | None = None,
185+
segment_mask: Bool[jax.Array, "num_segments"] | None = None,
186+
) -> e3j.IrrepsArray:
187+
result = segment_mean(
188+
data.array,
189+
segment_sizes,
190+
mask=mask,
191+
segment_mask=segment_mask,
192+
)
193+
return e3j.IrrepsArray(data.irreps, result)
194+
195+
196+
@singledispatch
161197
def segment_min(
162198
data: Float[jax.Array, "N ..."],
163199
segment_sizes: Int[jax.Array, "num_segments"],
@@ -210,6 +246,23 @@ def segment_min(
210246
return data_min
211247

212248

249+
@segment_min.register
250+
def _(
251+
data: e3j.IrrepsArray,
252+
segment_sizes: Int[jax.Array, "num_segments"],
253+
mask: Bool[jax.Array, "N ..."] | None = None,
254+
segment_mask: Bool[jax.Array, "num_segments"] | None = None,
255+
) -> e3j.IrrepsArray:
256+
result = segment_min(
257+
data.array,
258+
segment_sizes,
259+
mask=mask,
260+
segment_mask=segment_mask,
261+
)
262+
return e3j.IrrepsArray(data.irreps, result)
263+
264+
265+
@singledispatch
213266
def segment_max(
214267
data: Float[jax.Array, "N ..."],
215268
segment_sizes: Int[jax.Array, "num_segments"],
@@ -262,6 +315,22 @@ def segment_max(
262315
return data_max
263316

264317

318+
@segment_max.register
319+
def _(
320+
data: e3j.IrrepsArray,
321+
segment_sizes: Int[jax.Array, "num_segments"],
322+
mask: Bool[jax.Array, "N ..."] | None = None,
323+
segment_mask: Bool[jax.Array, "num_segments"] | None = None,
324+
) -> e3j.IrrepsArray:
325+
result = segment_max(
326+
data.array,
327+
segment_sizes,
328+
mask=mask,
329+
segment_mask=segment_mask,
330+
)
331+
return e3j.IrrepsArray(data.irreps, result)
332+
333+
265334
_REDUCTIONS = {
266335
"mean": segment_mean,
267336
"sum": segment_sum,
@@ -276,7 +345,7 @@ def segment_reduce(
276345
reduction: Literal["sum", "mean", "min", "max"],
277346
mask: Bool[jax.Array, "N ..."] | None = None,
278347
segment_mask: Bool[jax.Array, "num_segments"] | None = None,
279-
) -> Float[jax.Array, "num_segments ..."]:
348+
) -> Float[jax.Array, "num_segments ..."] | e3j.IrrepsArray:
280349
"""Performs a masked segment reduction over batched graph data.
281350
282351
This function is JAX-jittable and handles the logic for applying a mask
@@ -348,6 +417,6 @@ def _jraph_segment(
348417
except AttributeError:
349418
raise ValueError(f"Unknown reduction operation: {reduction}") from None
350419

351-
return jax.tree_util.tree_map(
420+
return jax.tree.map(
352421
lambda n: op(n, segment_ids, num_segments, indices_are_sorted, unique_indices), inputs
353422
)

tests/unit/gcnn/test_graph_ops.py

Lines changed: 171 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,9 @@
1+
import e3nn_jax as e3j
12
import jax
23
import jax.numpy as jnp
34
import pytest
45

5-
from tensorial.gcnn import graph_ops
6+
from tensorial.gcnn import graph_ops, keys
67

78

89
@pytest.mark.parametrize("jit", [False, True])
@@ -220,3 +221,172 @@ def test_segment_reduce_with_explicit_inf(jit, reduction, expected_inf):
220221

221222
res = op(data, segment_sizes, mask=mask, segment_mask=segment_mask)
222223
assert jnp.allclose(res, expected)
224+
225+
226+
@pytest.mark.parametrize("jit", [False, True])
227+
def test_segment_sum_irreps_array(jit):
228+
229+
op = jax.jit(graph_ops.segment_sum) if jit else graph_ops.segment_sum
230+
231+
x = e3j.IrrepsArray(
232+
"2x0e",
233+
jnp.array(
234+
[
235+
[1.0, 2.0],
236+
[3.0, 4.0],
237+
[5.0, 6.0],
238+
]
239+
),
240+
)
241+
segment_sizes = jnp.array([2, 1])
242+
# Case 1: no segment_mask
243+
res = op(x, segment_sizes)
244+
assert isinstance(res, e3j.IrrepsArray)
245+
assert res.irreps == x.irreps
246+
expected = jnp.array(
247+
[
248+
[4.0, 6.0],
249+
[5.0, 6.0],
250+
]
251+
)
252+
assert jnp.allclose(res.array, expected)
253+
254+
# Case 2: with segment_mask
255+
segment_mask = jnp.array([True, False])
256+
res_masked = op(x, segment_sizes, segment_mask=segment_mask)
257+
assert isinstance(res_masked, e3j.IrrepsArray)
258+
expected_masked = jnp.array(
259+
[
260+
[4.0, 6.0],
261+
[0.0, 0.0],
262+
]
263+
)
264+
assert jnp.allclose(res_masked.array, expected_masked)
265+
266+
267+
@pytest.mark.parametrize("jit", [False, True])
268+
def test_segment_mean_irreps_array(jit):
269+
270+
op = jax.jit(graph_ops.segment_mean) if jit else graph_ops.segment_mean
271+
272+
x = e3j.IrrepsArray(
273+
"2x0e",
274+
jnp.array(
275+
[
276+
[1.0, 2.0],
277+
[3.0, 4.0],
278+
[5.0, 6.0],
279+
]
280+
),
281+
)
282+
segment_sizes = jnp.array([2, 1])
283+
284+
# Case 1: no segment_mask
285+
res = op(x, segment_sizes)
286+
assert isinstance(res, e3j.IrrepsArray)
287+
assert res.irreps == x.irreps
288+
expected = jnp.array(
289+
[
290+
[2.0, 3.0],
291+
[5.0, 6.0],
292+
]
293+
)
294+
assert jnp.allclose(res.array, expected)
295+
296+
# Case 2: with segment_mask and data mask
297+
mask = jnp.array([True, False, True])
298+
segment_mask = jnp.array([True, False])
299+
res_masked = op(x, segment_sizes, mask=mask, segment_mask=segment_mask)
300+
assert isinstance(res_masked, e3j.IrrepsArray)
301+
expected_masked = jnp.array(
302+
[
303+
[1.0, 2.0],
304+
[0.0, 0.0],
305+
]
306+
)
307+
assert jnp.allclose(res_masked.array, expected_masked)
308+
309+
310+
@pytest.mark.parametrize("jit", [False, True])
311+
def test_segment_min_max_irreps_array(jit):
312+
313+
op_min = jax.jit(graph_ops.segment_min) if jit else graph_ops.segment_min
314+
op_max = jax.jit(graph_ops.segment_max) if jit else graph_ops.segment_max
315+
316+
x = e3j.IrrepsArray(
317+
"2x0e",
318+
jnp.array(
319+
[
320+
[1.0, 2.0],
321+
[3.0, 4.0],
322+
[5.0, 6.0],
323+
]
324+
),
325+
)
326+
segment_sizes = jnp.array([2, 1])
327+
mask = jnp.array([True, False, True])
328+
segment_mask = jnp.array([True, False])
329+
330+
# Min test
331+
res_min = op_min(x, segment_sizes, mask=mask, segment_mask=segment_mask)
332+
assert isinstance(res_min, e3j.IrrepsArray)
333+
expected_min = jnp.array(
334+
[
335+
[1.0, 2.0],
336+
[jnp.inf, jnp.inf],
337+
]
338+
)
339+
assert jnp.allclose(res_min.array, expected_min)
340+
341+
# Max test
342+
res_max = op_max(x, segment_sizes, mask=mask, segment_mask=segment_mask)
343+
assert isinstance(res_max, e3j.IrrepsArray)
344+
expected_max = jnp.array(
345+
[
346+
[1.0, 2.0],
347+
[-jnp.inf, -jnp.inf],
348+
]
349+
)
350+
assert jnp.allclose(res_max.array, expected_max)
351+
352+
353+
@pytest.mark.parametrize("jit", [False, True])
354+
def test_graph_segment_reduce_irreps_array_with_node_mask(jit):
355+
356+
op = (
357+
jax.jit(graph_ops.graph_segment_reduce, static_argnums=(1, 2))
358+
if jit
359+
else graph_ops.graph_segment_reduce
360+
)
361+
362+
x = e3j.IrrepsArray(
363+
"2x0e",
364+
jnp.array(
365+
[
366+
[1.0, 2.0],
367+
[3.0, 4.0],
368+
[5.0, 6.0],
369+
]
370+
),
371+
)
372+
373+
graph = {
374+
"nodes": {
375+
"features": x,
376+
keys.MASK: jnp.array([True, False, True]),
377+
},
378+
"n_node": jnp.array([2, 1]),
379+
}
380+
381+
res = op(graph, "nodes.features", "mean")
382+
383+
assert isinstance(res, e3j.IrrepsArray)
384+
assert res.irreps == x.irreps
385+
386+
expected = jnp.array(
387+
[
388+
[1.0, 2.0],
389+
[5.0, 6.0],
390+
]
391+
)
392+
assert jnp.allclose(res.array, expected)

0 commit comments

Comments
 (0)