NVIDIA / NVIDIA/warp

[QUESTION] Optimization example with JAX

Open
#607 2 comments 0 reactions 2 assignees View on GitHub

@daedalus5 is already working on this.

Since Mar 31, 2025.

interop question
Dominant language
Python
Stars
7.1k
Forks
624
Avg merge
3d 17h
Merged PRs (30d)
5

Description

Hi,
It's very exciting to see the recent improvements in JAX integration! I noticed there is an example showing how to use Warp with PyTorch for optimization (defining a loss in Warp and using PyTorch optimizers).
Could you provide a similar example but using optax in JAX instead? Specifically, I would like to know how to:

  • Define a loss function in Warp
  • Compute gradients using Warp's tape
  • Optimize parameters using JAX's optimizers

Thank you!

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.