Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion docs/api.rst
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,12 @@ The JAX API consists of ``TensorProduct`` and ``TensorProductConv``
classes that behave identically to their PyTorch counterparts. These classes
do not conform exactly to the e3nn-jax API, but perform the same computation.

JAX ``TensorProductConv`` uses the established loop-unroll implementation by
default. Select ``mode="streaming"`` to require receiver streaming, or
``mode="auto"`` to use it when the problem is supported. Streaming supports
trailing padded edges represented by out-of-bounds node indices. Receiver row
pointers may optionally be provided for receiver-sorted edges.

If you plan to use ``oeq.jax`` without PyTorch installed,
you need to set ``OEQ_NOTORCH=1`` in your local environment (within Python,
``os.environ["OEQ_NOTORCH"] = 1``). For the moment, we require this to avoid
Expand All @@ -54,6 +60,16 @@ breaking the PyTorch version of OpenEquivariance.
:exclude-members:

.. autoclass:: openequivariance.jax.TensorProductConv
:members: forward, reorder_weights_from_e3nn, reorder_weights_to_e3nn, implementation, uses_streaming_kernel
:undoc-members:
:exclude-members:

.. autoclass:: openequivariance.jax.LoopUnrollTensorProductConv
:members: forward, reorder_weights_from_e3nn, reorder_weights_to_e3nn
:undoc-members:
:exclude-members:

.. autoclass:: openequivariance.jax.StreamingTensorProductConv
:members: forward, reorder_weights_from_e3nn, reorder_weights_to_e3nn
:undoc-members:
:exclude-members:
Expand All @@ -76,4 +92,4 @@ both packages.

.. autoclass:: openequivariance.Irreps
:members:
:undoc-members:
:undoc-members:
Loading
Loading