google-deepmind / google-deepmind/optax

Type errors for general pytrees

Open
#384 2 comments 1 reaction 0 assignees View on GitHub
type:support
Dominant language
Python
Stars
2.3k
Forks
369
Avg merge
10h 15m
Merged PRs (30d)
7

Description

Hello!

When it comes to annotations, `optax` currently relies heavily on `optax.Updates` and `optax.Params`, which are all aliases for `chex.ArrayTree`.

This makes sense, but for folks who run type checkers means that a lot of type errors happen when working with pytrees that aren't strictly nested `Iterable` or `Mapping` types as specified in `chex`. For example:
```python
from typing import Tuple
import optax
from jax import numpy as jnp
import flax.struct

@flax.struct.dataclass
class Params:
weights: jnp.ndarray
bias: jnp.ndarray

def make_optimizer(
params: Params,
) -> Tuple[optax.GradientTransformation, optax.OptState]:
"""Make an optimizer."""
optimizer = optax.sgd(learning_rate=1e-3)
state = optimizer.init(params) # Type error.
return optimizer, state
```

A few questions from this:
- Is this considered a bug, or something that the optax team would be open to supporting? Are there better solutions for suppressing this error than simply adding a `# type: ignore`?
- It seems like type safety with optax could benefit immensely from support for generics, which have been present since Python 3.5 (`typing.Generic`, `typing.TypeVar`). Any chance this would be something that optax would be open to supporting?
- Simple example: with `ArrayTreeT = TypeVar("ArrayTreeT", bound=chex.ArrayTree)`, `optax.apply_updates()` could be annotated as `optax.apply_updates(params: ArrayTreeT, updates: ArrayTreeT) -> ArrayTreeT` to indicate that the argument and return types should all be the same.

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.