pymc-devs / pymc-devs/pytensor
Add error message to `NonZero` Op telling users it cannot be jitted in JAX
Open
Nobody has claimed this yet.
backend compatibility
beginner friendly
help wanted
jax
- Dominant language
- Python
- Stars
- 644
- Forks
- 208
- Avg merge
- 2d 14h
- Merged PRs (30d)
- 16
Description
Description
jax mode is missing a dispatch for the NonZero Op.
There's a jnp.nonzero, so it should be easy to do.
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by locating the JAX-mode dispatch handling for the NonZero Op and compare it with the available jnp.nonzero operation mentioned in the issue. Confirm the expected behavior for a jitted expression, then add the requested user-facing error message and verify it with the relevant JAX-mode tests.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- backend
- Issue type
- Feature
- Difficulty
- 3/5
- Estimated time
- 1-2 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 42/100