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
22 changes: 20 additions & 2 deletions .github/workflows/run_tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ jobs:
matrix:
os: [ubuntu-latest]
python-version: ["3.12"]
pytensor-mode: [FAST_RUN]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Delete this axis. Line 118 already defaults the mode with || 'FAST_RUN', and floatX gets no axis of its own.

# Off by default, flipped to 1 by the per-backend jobs under `include`. A core job installs no
# backend at all, which is what proves the library runs without one.
install-jax: [0]
Expand All @@ -35,6 +36,7 @@ jobs:
paths: |
tests/test_transformer.py
tests/test_layers.py
tests/test_positional.py
tests/test_recurrent.py
tests/test_conv.py
tests/test_conv_transpose.py
Expand Down Expand Up @@ -64,8 +66,8 @@ jobs:
tests/test_workflow_groups.py
# Windows runners are roughly twice as slow, and almost everything here is platform independent.
# It runs one job over the parts that are not: file IO, paths, and compiling a graph to train.
# One job per backend, each installing only that backend. They live here rather than in
# `test-subset` because a matrix entry cannot change `os`, and mlx needs macOS.
# Backend jobs each install only their backend. Dispatch jobs keep the default mode; positional
# jobs run the same general tests in the backend's mode. These live here because mlx needs macOS.
include:
- os: ubuntu-latest
python-version: "3.12"
Expand All @@ -79,6 +81,21 @@ jobs:
test-subset:
name: dispatch-mlx
paths: tests/dispatch/mlx/
- os: ubuntu-latest
python-version: "3.12"
install-jax: 1
pytensor-mode: JAX
test-subset:
name: positional-jax
paths: tests/test_positional.py
- os: macos-15
python-version: "3.12"
install-mlx: 1
pytensor-mode: MLX
floatX: float32
test-subset:
name: positional-mlx
paths: tests/test_positional.py
- os: ubuntu-latest
python-version: "3.12"
install-torch: 1
Expand All @@ -98,6 +115,7 @@ jobs:

env:
TEST_SUBSET: ${{ matrix.test-subset.paths }}
PYTENSOR_FLAGS: mode=${{ matrix.pytensor-mode || 'FAST_RUN' }},floatX=${{ matrix.floatX || 'float64' }}

defaults:
run:
Expand Down
4 changes: 4 additions & 0 deletions docs/source/api/layers.rst
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ Base
:toctree: generated/

Layer
pytensor_ml.base.VariadicLayer

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Import VariadicLayer in pytensor_ml/layers/__init__.py and list it here as a bare name. Every other entry in this file resolves against the currentmodule on line 4.


Combinators
-----------
Expand Down Expand Up @@ -72,6 +73,7 @@ Normalization and regularization

BatchNorm
LayerNorm
RMSNorm
GroupNorm
Dropout

Expand Down Expand Up @@ -102,3 +104,5 @@ Attention and transformers
FeedForward
TransformerBlock
scaled_dot_product_attention
RotaryEmbedding
rotary_embedding
28 changes: 27 additions & 1 deletion pytensor_ml/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,32 @@ def __call__(self, X):
def __call__(self, x: pt.TensorLike) -> pt.TensorVariable: ...


class VariadicLayer(ABC):
"""
Base class for graph-building layers that consume multiple tensors.

Examples
--------
Declare a layer whose inputs are supplied together:

.. code-block:: python

import pytensor.tensor as pt

from pytensor_ml.base import VariadicLayer

class Add(VariadicLayer):
def __call__(self, *inputs: pt.TensorLike) -> pt.TensorVariable:
left, right = inputs
return pt.as_tensor(left) + pt.as_tensor(right)

total = Add()(pt.vector("left"), pt.vector("right"))
"""

@abstractmethod
def __call__(self, *inputs: pt.TensorLike) -> pt.TensorVariable: ...


class LayerOp(SymbolicOp):
"""Base class for the library's neural-network ops.

Expand Down Expand Up @@ -157,4 +183,4 @@ def update_chain_root(variable: Variable) -> tuple[Variable, int] | None:
depth += 1


__all__ = ["Layer", "LayerOp", "StatefulOp", "UnaryLayerOp", "update_chain_root"]
__all__ = ["Layer", "LayerOp", "StatefulOp", "UnaryLayerOp", "VariadicLayer", "update_chain_root"]
3 changes: 3 additions & 0 deletions pytensor_ml/layers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@
ZeroPad1D,
ZeroPad2D,
)
from pytensor_ml.layers.positional import RotaryEmbedding, RotaryEmbeddingLayer, rotary_embedding
from pytensor_ml.layers.recurrent import (
GRU,
LSTM,
Expand Down Expand Up @@ -108,12 +109,14 @@
"ReflectionPad2D",
"ReplicationPad1D",
"ReplicationPad2D",
"RotaryEmbedding",
"Sequential",
"Squeeze",
"TransformerBlock",
"Upsample1D",
"Upsample2D",
"ZeroPad1D",
"ZeroPad2D",
"rotary_embedding",
"scaled_dot_product_attention",
]
Loading
Loading