{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "54cc2501",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-10-02T14:44:55.950746Z",
     "iopub.status.busy": "2026-10-02T14:44:55.950469Z",
     "iopub.status.idle": "2026-10-02T14:44:55.955719Z",
     "shell.execute_reply": "2026-10-02T14:44:55.955091Z"
    },
    "papermill": {
     "duration": 0.008437,
     "end_time": "2026-10-02T14:44:55.956484+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:55.948047+00:00",
     "status": "completed"
    },
    "tags": [
     "remove-input",
     "active-ipynb",
     "remove-output"
    ]
   },
   "outputs": [],
   "source": [
    "try:\n",
    "    from openmdao.utils.notebook_utils import notebook_mode  # noqa: F401\n",
    "except ImportError:\n",
    "    !python -m pip install openmdao[notebooks]"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "ab3413a7",
   "metadata": {
    "papermill": {
     "duration": 0.001245,
     "end_time": "2026-10-02T14:44:55.959307+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:55.958062+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "source": [
    "# Computing Partial Derivatives of Components Using JAX\n",
    "\n",
    "To truly take advantage of OpenMDAO, the user needs to compute partial derivatives for any `Component`\n",
    "that they write.  This can be done using finite difference, but that can have issues with accuracy\n",
    "and performance.  Using complex step is another option which has good accuracy but isn't always \n",
    "possible because it requires the component's computations to be compatible with complex numbers. \n",
    "In some cases, the user can provide analytic partial derivatives, which likely has good performance \n",
    "but can be difficult to determine depending on the complexity of the component.\n",
    "\n",
    "This notebook describes another method, which is to use the optional third-party \n",
    "[JAX](https://jax.readthedocs.io/en/latest/index.html) library, to \n",
    "automatically differentiate native Python and NumPy functions.  To simplify jax usage within OpenMDAO, \n",
    "we've created two component classes, [JaxExplicitComponent](jax_explicitcomp_api.ipynb) and \n",
    "[JaxImplicitComponent](jax_implicitcomp_api.ipynb).  These components require only the definition of \n",
    "a `compute_primal` method that replaces the `compute` method for `JaxExplicitComponent` and the\n",
    "`apply_nonlinear` method for `JaxImplicitComponent`.\n",
    "\n",
    "This notebook will describe in more detail how to create and use a JaxExplicitComponent or \n",
    "JaxImplicitComponent and will give examples.\n",
    "\n",
    "Before going further, it's a good idea to aquaint yourself with some of jax's 'sharp edges' \n",
    "[here](https://docs.jax.dev/en/latest/notebooks/Common_Gotchas_in_JAX.html). This will hopefully \n",
    "make the process of creating a `JaxExplicitComponent` or `JaxImplicitComponent` a less frustrating one."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "de11e0a3",
   "metadata": {
    "papermill": {
     "duration": 0.002199,
     "end_time": "2026-10-02T14:44:55.997510+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:55.995311+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "source": [
    "The use of JAX is optional for OpenMDAO so if not already installed, the user needs to install it by \n",
    "issuing *one* of the following commands at your operating system command prompt:\n",
    "```\n",
    "pip install jax jaxlib\n",
    "pip install openmdao[jax]\n",
    "pip install openmdao[all]\n",
    "```\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "bbb67dde",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-10-02T14:44:56.000812Z",
     "iopub.status.busy": "2026-10-02T14:44:56.000604Z",
     "iopub.status.idle": "2026-10-02T14:44:56.519757Z",
     "shell.execute_reply": "2026-10-02T14:44:56.519016Z"
    },
    "papermill": {
     "duration": 0.521934,
     "end_time": "2026-10-02T14:44:56.520674+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:55.998740+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Requirement already satisfied: jax in /home/runner/work/OpenMDAO/OpenMDAO/.pixi/envs/dev/lib/python3.13/site-packages (0.11.0)\r\n",
      "Requirement already satisfied: jaxlib<=0.11.0,>=0.11.0 in /home/runner/work/OpenMDAO/OpenMDAO/.pixi/envs/dev/lib/python3.13/site-packages (from jax) (0.11.0)\r\n",
      "Requirement already satisfied: ml_dtypes>=0.5.0 in /home/runner/work/OpenMDAO/OpenMDAO/.pixi/envs/dev/lib/python3.13/site-packages (from jax) (0.5.4)\r\n",
      "Requirement already satisfied: numpy>=2.1 in /home/runner/work/OpenMDAO/OpenMDAO/.pixi/envs/dev/lib/python3.13/site-packages (from jax) (2.4.6)\r\n",
      "Requirement already satisfied: opt_einsum in /home/runner/work/OpenMDAO/OpenMDAO/.pixi/envs/dev/lib/python3.13/site-packages (from jax) (3.4.0)\r\n",
      "Requirement already satisfied: scipy>=1.15 in /home/runner/work/OpenMDAO/OpenMDAO/.pixi/envs/dev/lib/python3.13/site-packages (from jax) (1.17.0)\r\n"
     ]
    }
   ],
   "source": [
    "!pip install jax"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "624bad71",
   "metadata": {
    "papermill": {
     "duration": 0.000947,
     "end_time": "2026-10-02T14:44:56.522805+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:56.521858+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "source": [
    "The JAX library includes a NumPy-like API, `jax.numpy`, which implements the NumPy API using the primitives in JAX. Almost anything that can be done with NumPy can be done with `jax.numpy`. JAX arrays are similar to NumPy arrays, but they are designed to work with accelerators such as GPUs and TPUs. `jax.numpy` is typically imported as `jnp`."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "b995e262",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-10-02T14:44:56.525840Z",
     "iopub.status.busy": "2026-10-02T14:44:56.525655Z",
     "iopub.status.idle": "2026-10-02T14:44:56.899801Z",
     "shell.execute_reply": "2026-10-02T14:44:56.898594Z"
    },
    "papermill": {
     "duration": 0.376763,
     "end_time": "2026-10-02T14:44:56.900538+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:56.523775+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "outputs": [],
   "source": [
    "import jax"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "77c4f56e",
   "metadata": {
    "papermill": {
     "duration": 0.001277,
     "end_time": "2026-10-02T14:44:56.903395+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:56.902118+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "source": [
    "The default for JAX is to do single precision computations. OpenMDAO uses double precision, so this line of code is needed."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "0631b35d",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-10-02T14:44:56.906976Z",
     "iopub.status.busy": "2026-10-02T14:44:56.906715Z",
     "iopub.status.idle": "2026-10-02T14:44:56.909107Z",
     "shell.execute_reply": "2026-10-02T14:44:56.908671Z"
    },
    "papermill": {
     "duration": 0.00484,
     "end_time": "2026-10-02T14:44:56.909484+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:56.904644+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "outputs": [],
   "source": [
    "jax.config.update(\"jax_enable_x64\", True)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "ee256559",
   "metadata": {
    "papermill": {
     "duration": 0.001539,
     "end_time": "2026-10-02T14:44:56.912290+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:56.910751+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "source": [
    "## Automatic Determination of Derivative Direction\n",
    "\n",
    "`JaxExplicitComponent` and `JaxImplicitComponent` automatically determine the direction they will\n",
    "use to compute their partial jacobians based on their jacobian's shape.  If there are more rows\n",
    "than columns in the jacobian, they'll use forward mode. Otherwise they'll use reverse mode.\n",
    "The number of columns in the `JaxExplicitComponent`'s jacobian is equal to the size of its inputs\n",
    "vector, and the number of columns in the `JaxImplicitComponent`'s jacobian is equal to combined size\n",
    "of its inputs **and** outputs vectors.\n",
    "Note that this automatic determination of derivative direction only occurs if the `matrix_free`\n",
    "attribute is False.\n",
    "\n",
    "\n",
    "## Matrix-Free Mode\n",
    "\n",
    "If `matrix_free` is True, `JaxExplicitComponent` computes derivatives via `jax.jvp` ('fwd' mode)\n",
    "or `jax.vjp` ('rev' mode) directly instead of forming an explicit Jacobian (see\n",
    "`compute_jacvec_product`). Both directions cache a compiled function so repeated calls don't\n",
    "retrace from scratch, but they invalidate on different triggers:\n",
    "\n",
    "- **'rev' mode** builds `jax.vjp` at the current linearization point, so it is rebuilt whenever\n",
    "  the component's continuous input values change (typically once per outer solver/optimizer\n",
    "  iteration), as well as whenever `get_self_statics()` or the discrete inputs change.\n",
    "- **'fwd' mode** passes the primal input values in as an ordinary call argument instead of\n",
    "  baking them into the compiled jvp, so it is rebuilt only when `get_self_statics()` or the\n",
    "  discrete inputs change -- the same compiled jvp is reused across the whole `Problem`, not\n",
    "  just within one linearization.\n",
    "\n",
    "\n",
    "## Self Statics\n",
    "\n",
    "When jax compiles a function, it assumes that the only variables that can change are those that are\n",
    "passed into the function as arguments and any internal variables that depend on those arguments.  \n",
    "All other variables are treated as static. But what if \n",
    "our jax component has an option or attribute that contributes to the output of our `compute_primal`\n",
    "function?  Since that option or attribute doesn't get passed into the function as an argument, jax\n",
    "doesn't know about it. In that case, we must be able to detect when those 'static' options or \n",
    "attributes change so that we can tell jax to recompile the function.  Otherwise the outputs of the \n",
    "function won't reflect the current values of the static options and attributes.\n",
    "\n",
    "In `JaxExplicitComponent` and `JaxImplicitComponent`, we add a method called `get_self_statics` to\n",
    "handle this situation.  `get_self_statics` is a simple method that returns a tuple containing\n",
    "any option or attribute in your component that will affect the output of your `compute_primal` method.\n",
    "If your component doesn't have any of these 'self static' variables then you don't have to define\n",
    "`get_self_statics`.\n",
    "\n",
    "Here's a simple example.  Suppose my component has an option called 'mult1' and an attribute called\n",
    "'mult2', and they're used in `compute_primal` as follows:\n",
    "\n",
    "\n",
    "```python\n",
    "def compute_primal(self, x):\n",
    "    return x * self.options['mult1'] * self.mult2\n",
    "```\n",
    "\n",
    "In this case, we would be required to define the `get_self_statics` method shown below:\n",
    "\n",
    "```python\n",
    "def get_self_statics(self):\n",
    "    return (self.options['mult1'], self.mult2)\n",
    "```\n",
    "\n",
    "Doing this will allow the `compute_primal` to be recompiled whenever `self.options['mult1']` or\n",
    "`self.mult2` change.\n",
    "\n",
    "\n",
    "Note that **not** all of a component's options and/or attributes need to be returned from `get_self_statics`.\n",
    "Only those that are referenced inside of `compute_primal` **and** affect its outputs should be returned.\n",
    "\n",
    "Note also that you **do not** return any variable that you've added to your component via `add_input`,\n",
    "`add_output`, `add_discrete_input`, or `add_discrete_output`, even if you think that they won't\n",
    "change during a run for some reason.  Jax already knows about all of them and can handle changes to\n",
    "them properly.\n",
    "\n",
    "\n",
    "## Configuration Options\n",
    "\n",
    "[JaxExplicitComponent](jax_explicitcomp_api.ipynb) and [JaxImplicitComponent](jax_implicitcomp_api.ipynb) \n",
    "both have the following options:\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "id": "9cdad094",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-10-02T14:44:56.929381Z",
     "iopub.status.busy": "2026-10-02T14:44:56.929208Z",
     "iopub.status.idle": "2026-10-02T14:44:57.863190Z",
     "shell.execute_reply": "2026-10-02T14:44:57.862361Z"
    },
    "papermill": {
     "duration": 0.950038,
     "end_time": "2026-10-02T14:44:57.863707+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:56.913669+00:00",
     "status": "completed"
    },
    "tags": [
     "remove-input",
     "remove-output"
    ]
   },
   "outputs": [
    {
     "data": {
      "text/html": [
       "\n",
       "<!DOCTYPE html>\n",
       "<html lang=\"en\">\n",
       "<head>\n",
       "    <style>\n",
       "        h2 {\n",
       "            text-align: center;\n",
       "        }\n",
       "    </style>\n",
       "</head>\n",
       "<body>\n",
       "    <h2></h2>\n",
       "        <table style=\"border: 1px solid #999; border-collapse: collapse;\">\n",
       "        <tr><th style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; background-color: #E9E9E9; text-align: left;\">Option</th><th style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; background-color: #E9E9E9; text-align: left;\">Default</th><th style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; background-color: #E9E9E9; text-align: left;\">Acceptable Values</th><th style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; background-color: #E9E9E9; text-align: left;\">Acceptable Types</th><th style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; background-color: #E9E9E9; text-align: left;\">Description</th></tr>\n",
       "       <tr style=\"background-color: ghostwhite;\"><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">always_opt</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">False</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[True, False]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[&#x27;bool&#x27;]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">If True, force nonlinear operations on this component to be included in the optimization loop even if this component is not relevant to the design variables and responses.</td></tr>\n",
       "       <tr style=\"background-color: #F3F3F3;\"><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">default_shape</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">(1,)</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">N/A</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[&#x27;tuple&#x27;]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">Default shape for variables that do not set val to a non-scalar value or set shape, shape_by_conn, copy_shape, or compute_shape. Default is (1,).</td></tr>\n",
       "       <tr style=\"background-color: ghostwhite;\"><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">default_to_dyn_shapes</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">False</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[True, False]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[&#x27;bool&#x27;]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">If True, use dynamic shaping for any variables whose value is scalar and whose shape is not explicitly set. Inputs will use shape_by_conn and outputs will use a compute_shape method based on jax.eval_shape. Default is False.</td></tr>\n",
       "       <tr style=\"background-color: #F3F3F3;\"><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">derivs_method</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">jax</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[&#x27;jax&#x27;, &#x27;cs&#x27;, &#x27;fd&#x27;, None]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">N/A</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">The method to use for computing derivatives</td></tr>\n",
       "       <tr style=\"background-color: ghostwhite;\"><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">run_root_only</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">False</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[True, False]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[&#x27;bool&#x27;]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">If True, call compute, compute_partials, linearize, apply_linear, apply_nonlinear, solve_linear, solve_nonlinear, and compute_jacvec_product only on rank 0 and broadcast the results to the other ranks.</td></tr>\n",
       "       <tr style=\"background-color: #F3F3F3;\"><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">use_jit</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">True</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[True, False]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">[&#x27;bool&#x27;]</td><td style=\"border: 1px solid #999; border-collapse: collapse; padding: 5px; text-align: left;\">If True, attempt to use jit on compute_primal, assuming jax or some other AD package capable of jitting is active.</td></tr>\n",
       "    </table>\n",
       "</body>\n",
       "</html>\n"
      ],
      "text/plain": [
       "<IPython.core.display.HTML object>"
      ]
     },
     "metadata": {},
     "output_type": "display_data"
    }
   ],
   "source": [
    "import openmdao.api as om\n",
    "class SimpleJaxComp(om.JaxExplicitComponent):\n",
    "    def compute_primal(self, x):\n",
    "        y = 2.*x\n",
    "        return y\n",
    "\n",
    "comp = SimpleJaxComp()\n",
    "om.show_options_table(comp)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "874a631d",
   "metadata": {
    "papermill": {
     "duration": 0.00107,
     "end_time": "2026-10-02T14:44:57.866042+00:00",
     "exception": false,
     "start_time": "2026-10-02T14:44:57.864972+00:00",
     "status": "completed"
    },
    "tags": []
   },
   "source": [
    "## Debugging\n",
    "\n",
    "While normally you want the `use_jit` option to be True for performance reasons, if you want to debug\n",
    "your `compute_primal` method it often helps to set `use_jit` to False. This will allow you to put\n",
    "print statements in your `compute_primal` or to set breakpoints inside it with a python debugger.\n",
    "\n",
    "\n",
    "## Examples\n",
    "\n",
    "- [JaxExplicitComponent Example](jax_explicitcomp_api)\n",
    "- [JaxImplicitComponent Example](jax_implicitcomp_api)\n"
   ]
  }
 ],
 "metadata": {
  "celltoolbar": "Tags",
  "kernelspec": {
   "display_name": "py311forge",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.13.14"
  },
  "orphan": true,
  "papermill": {
   "default_parameters": {},
   "duration": 3.122101,
   "end_time": "2026-10-02T14:44:58.384773+00:00",
   "environment_variables": {},
   "exception": null,
   "input_path": "/home/runner/work/OpenMDAO/OpenMDAO/openmdao/docs/openmdao_book/features/experimental/jax_partial_derivs.ipynb",
   "output_path": "/home/runner/work/OpenMDAO/OpenMDAO/openmdao/docs/_executed_book/features/experimental/jax_partial_derivs.ipynb",
   "parameters": {},
   "start_time": "2026-10-02T14:44:55.262672+00:00",
   "version": "2.7.0"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}