google / google/jax-cfd

JaxNumPy functions for GridArrays/GridVariables

Open
#110 2 comments 0 reactions 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
964
Forks
142
Avg merge
3h 9m
Merged PRs (30d)
1

Description

Hi! This is a great project, and I'm a big fan of both the machine learning applications here and also some of the smaller, helpful structures, in particular base.grids.

Currently, it is possible to add two GridArrays, but it is not possible to add two GridVariables. So this works fine:

```
import jax_cfd.base.grids as gd
import jax.numpy as jnp

grid = gd.Grid([4,], domain = [(0, 1),])

array_of_values = jnp.array([2.0, 2.0, 3.0, 4.0])

centered_array = grid.center(array_of_values)

print(centered_array + centered_array)
```

But this throws an exception:

```
bc = gd.BoundaryConditions((gd.PERIODIC,))

centered_variable = gd.GridVariable(centered, bc)

print(centered_variable + centered_variable)
```

I'm happy to have a go at implementing this myself, if someone isn't already working on it.

Also, am I correct in thinking that the way to use a JaxNumPy function on a GridArray is to call it via NumPy? For example, this throws an exception:

```
print(jnp.abs(centered_array))
```

But this works:
```
import numpy as np
print(np.abs(centered_array))
```

I assume it's implemented this way because NumPy has an automatic mix-in that we can employ to funnel things to the appropriate JaxNumPy function, but JaxNumPy does not.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.