Weekly Project News

Archives
Subscribe

Weekly GitHub Report for Jax: August 04, 2026 - August 11, 2026 (00:34:29)

Weekly GitHub Report for Jax

Thank you for subscribing to our weekly newsletter! Each week, we deliver a comprehensive summary of your GitHub project's latest activity right to your inbox, including an overview of your project's issues, pull requests, contributors, and commit activity.


Table of Contents

  • I. News
    • 1.1. Recent Version Releases
    • 1.2. Other Noteworthy Updates
  • II. Issues
    • 2.1. Top 5 Active Issues
    • 2.2. Top 5 Stale Issues
    • 2.3. Open Issues
    • 2.4. Closed Issues
    • 2.5. Issue Discussion Insights
  • III. Pull Requests
    • 3.1. Open Pull Requests
    • 3.2. Closed Pull Requests
    • 3.3. Pull Request Discussion Insights
  • IV. Contributors
    • 4.1. Contributors

I. News

1.1 Recent Version Releases:

The current version of this repository is jax-v0.9.0.1

1.2 Version Information:

On February 3, 2026, JAX v0.9.0.1 was released as a patch update to v0.9.0, incorporating four specific pull requests from the OpenXLA repository to address targeted improvements or fixes without introducing major changes.

Click here to view the full release notes!

II. Issues

2.1 Top 5 Active Issues:

We consider active issues to be issues that that have been commented on most frequently within the last week. Bot comments are omitted.

  1. [BUG] An edge-case false cache hit for jax.jit tracing.: This issue describes a long-standing bug in JAX where nested jax.jit calls cause a false cache hit during tracing, leading to leaked tracer errors and rendering jax.debug.breakpoint effectively unusable in complex scenarios. The user provides a minimal working example demonstrating the problem, explains the root cause related to cached jaxprs holding stale tracer references, and shares an AI-suggested patch that invalidates the cache when tracer consts from a different trace are detected, aiming to fix the issue.

    • The comments include an expression of interest to work on the issue, a suggestion to remove jax.debug.breakpoint entirely in favor of user-implemented alternatives, and a discussion about the challenges of debugging JAX code, especially the limitations of eager mode due to performance and precision differences; examples and personal experiences with debugging slowdowns are also shared.
    • Number of comments this week: 4
  2. [BUG] [XLA] float64 log1p on CPU is up to 129 ulp off near x ≈ -0.414 (also affects atanh): This issue reports a significant accuracy problem with the float64 implementation of log1p on the CPU backend in JAX, where the error can reach up to 129 units in the last place near x ≈ -0.414, affecting both log1p and atanh functions. The problem does not occur on GPU or with float32 precision, and the issue was verified on multiple machines, highlighting a discrepancy compared to NumPy's more accurate results.

    • The comments show a user expressing interest in fixing the issue and requesting assignment, followed by a quick proposal of a fix via a pull request, appreciation for the rapid response, and a request for approval of the PR from a project maintainer.
    • Number of comments this week: 4
  3. [BUG] Pallas Mosaic-TPU: DMA silently returns zeros when pallas_call has an unused operand computed inside the jit (libtpu 0.0.43 regression): This issue describes a regression in libtpu versions 0.0.43 to 0.0.45 where a DMA operation silently returns all zeros when a pallas_call kernel includes an unused operand computed inside a JIT-compiled function, a behavior that was correct in version 0.0.41. The root cause is identified as a mismatch in memory space assignment where the unused operand is physically placed in CMEM by XLA on TPU v4, but the Mosaic kernel expects it in HBM, leading to an inconsistent DMA path and incorrect results.

    • The comments discuss attempts to reproduce the issue, confirm it persists on TPU v4 with newer libtpu and JAX versions, and identify the problem as related to CMEM memory space assignment changes in recent libtpu releases; workarounds include disabling CMEM assignment or passing the unused operand externally, and it is noted that TPU v5e/v5p are unaffected due to lack of CMEM.
    • Number of comments this week: 4
  4. [BUG] jax.jit produces NaN/Inf from softmax when pre-activation values are very large; eager returns finite output: This issue reports a numerical instability in jax.jit where the softmax function produces NaN or Inf outputs when given very large pre-activation values, despite the eager execution mode returning finite results. The root cause is identified as XLA's fusion and optimization passes reordering or removing the numerical stabilization step in softmax, leading to overflow in the exponentiation before max subtraction, which does not occur in the eager path that explicitly applies stabilization.

    • The comments confirm the root cause as XLA fusion reordering arithmetic operations that break softmax stabilization under JIT, propose explicit data-dependency barriers to fix this, suggest adding regression tests, and offer to submit a pull request with the fix and tests.
    • Number of comments this week: 2
  5. [BUG] jax.nn.silu gradient returns zero for large negative float64 input: This issue reports that the gradient of the jax.nn.silu function incorrectly returns zero for large negative float64 inputs, specifically at x = -709.0, where the true derivative is a very small but finite number. The problem arises because JAX flushes subnormal floating-point numbers to zero during the backward pass, causing the loss of the tiny but nonzero gradient value.

    • The comments clarify that the zero gradient is due to subnormal flushing of the logistic function's output at large negative inputs, which is representable in float64 but treated as zero by JAX, confirming that the autodiff expression is correct but the intermediate value is flushed to zero.
    • Number of comments this week: 2

2.2 Top 5 Stale Issues:

We consider stale issues to be issues that has had no activity within the last 30 days. The team should work together to get these issues resolved and closed as soon as possible.

As of our latest update, there are no stale issues for the project this week.

2.3 Open Issues

This section lists, groups, and then summarizes issues that were created within the last week in the repository.

Issues Opened This Week: 19

Summarized Issues:

  • Hijax type handling and batching errors: Several issues report failures and errors when using hijax types with JAX transformations such as jacfwd, jacrev, and vmap over lax.scan. These problems include AttributeErrors, TypeErrors, and TracerArrayConversionError due to improper handling of tangent types and batch dimensions, which break vectorized gradient computations but not serial ones.
  • [issues/39701, issues/39727]
  • Numerical accuracy and overflow in special functions: Multiple issues highlight incorrect overflow to infinity or inaccurate results in special functions like jax.numpy.i0, jax.scipy.special.i0, and jax.scipy.special.i1 for large finite float64 inputs. These indicate internal overflows or implementation errors where finite results are expected, affecting reliability of these functions.
  • [issues/39770, issues/39771, issues/39772]
  • Gradient computation errors for small or large inputs: Several reports describe incorrect zero or infinite gradients returned by jax.grad for functions such as jax.numpy.expm1, jax.nn.elu, jax.nn.selu, jax.nn.silu, and the modified Bessel function i0 at very small or large negative float64 inputs. These issues reveal problems in backward pass computations, including subnormal number flushing and internal overflows, leading to loss of tiny but finite derivatives.
  • [issues/39794, issues/39795, issues/39796, issues/39797, issues/39798, issues/39799]
  • Higher-order derivative inaccuracies: There are issues reporting incorrect second derivatives for functions like jax.lax.lgamma at negative half-integers and the hyperbolic tangent function at very small positive inputs. These problems show that higher-order automatic differentiation produces wrong or zero values instead of expected finite results.
  • [issues/39800, issues/39801]
  • Backend nondeterminism and hardware-specific regressions: One issue describes nondeterministic floating-point results on the CPU backend when using multiple intra-op threads, causing numerical instability and NaNs during training. Another reports a regression in TPU libtpu versions where a DMA operation returns zeros due to memory space assignment mismatches on TPU v4 hardware.
  • [issues/39741, issues/39744]
  • Unsupported operations and kernel crashes on specialized backends: There are reports of jnp.cumsum failing to lower on Mosaic GPU kernels due to unsupported scan operations, and a segmentation fault crash caused by enabling SparseCore vector layout passes in a TPU kernel compilation. These indicate missing backend support and compiler pass issues affecting kernel execution.
  • [issues/39763, issues/39840]
  • Documentation example causing TypeError due to unhashable JAX arrays: The CustomClass example in the JAX documentation incorrectly implements a __hash__ method that hashes a JAX array, which is unhashable, leading to a TypeError when used with JAX's compilation caching.
  • [issues/39858]

2.4 Closed Issues

This section lists, groups, and then summarizes issues that were closed within the last week in the repository. This section also links the associated pull requests if applicable.

Issues Closed This Week: 3

Summarized Issues:

  • Tiled cp.async memory copy bug: The tiled cp.async copies from global to shared memory incorrectly pass the full array shape instead of the sub-block copy window shape to the transfer_tiled function, causing a shape mismatch error. This prevents the tiled cp.async path from lowering correctly in JAX's GPU compilation pipeline.
  • issues/39687
  • hijax type AttributeError under jax.jit: Using a hijax type as the carry in lax.scan or fori_loop works eagerly but raises an AttributeError during jax.jit because the hijax type lacks a required mat attribute. This causes failures during JIT compilation despite working in eager mode.
  • issues/39700
  • Return type change for jax.numpy.meshgrid: There is a proposal to change the return type of jax.numpy.meshgrid from a list to a tuple to align with the Array API standard and numpy/cupy behavior. This aims to improve compliance with minimal impact on existing code.
  • issues/39779

2.5 Issue Discussion Insights

This section will analyze the tone and sentiment of discussions within this project's open and closed issues that occurred within the past week. It aims to identify potentially heated exchanges and to maintain a constructive project environment.

Based on our analysis, there are no instances of toxic discussions in the project's open or closed issues from the past week.


III. Pull Requests

3.1 Open Pull Requests

This section provides a summary of pull requests that were opened in the repository over the past week. The top three pull requests with the highest number of commits are highlighted as 'key' pull requests. Other pull requests are grouped based on similar characteristics for easier analysis. Up to 25 pull requests are displayed in this section, while any remaining pull requests beyond this limit are omitted for brevity.

Pull Requests Opened This Week: 26

Key Open Pull Requests

1. Added matmul tests in mgpu_torch_test.py: This pull request adds matrix multiplication tests to the mgpu_torch_test.py file, along with fixes to kernel compilation and removal of deprecated calls to improve GPU-related functionality in the JAX project.

  • URL: pull/39782
  • Associated Commits: 7bede, e2c9e, b2f83

2. Add a benchmark smoke-test workflow: This pull request adds a benchmark smoke-test workflow that runs every benchmarks/*_benchmark.py file for a single iteration on CPU nightly and on pull requests affecting benchmarks, ensuring that each benchmark file registers at least one benchmark to prevent silent no-op passes, while also repairing the previously unrunnable shape_poly microbenchmark.

  • URL: pull/39842
  • Associated Commits: 0aa03, e408f

3. [shape_poly] Decide mod residues via equality gcd analysis: This pull request enhances the decision procedure in the shape polynomial analysis by implementing equality gcd analysis to determine modular residues, allowing it to infer residue class constraints from multi-term equalities and thereby improve reasoning about modular arithmetic expressions that were previously undecidable.

  • URL: pull/39843
  • Associated Commits: 4fc22, 2c532

Other Open Pull Requests

  • PyTreeDef deserialization robustness: This pull request improves the robustness of PyTreeDef deserialization by rejecting malformed PyTreeDefProto nodes with invalid or inconsistent arity values. This prevents crashes caused by stack underflows or out-of-bounds accesses during post-order traversal, fixing issue #37410.
    pull/39732
  • Batching rule fix for vmap over scan with Ref: This fix addresses crashes caused by improper handling of batch dimensions for Ref objects when applying vmap to a scan that closes over a Ref batched on a non-leading axis. The batching rule was modified to batch along the existing batch dimension of the Ref without moving its axis, ensuring correct behavior under eager and JIT compilation.
    pull/39745
  • Typing updates to use Python 3.12 syntax: Multiple pull requests update typing in various modules such as jax.util and custom_derivatives to use implicit generics, type variables, and the latest typing standards available in Python 3.12 and newer. These changes modernize the codebase's type annotations without altering functionality.
    pull/39759, pull/39787
  • ROCm test environment and CI improvements: Several pull requests improve ROCm testing by removing unnecessary environment variable modifications, simplifying pytest parallelism heuristics, enabling previously skipped ROCm tests, and allowing selective test subset runs via an environment variable. These changes optimize test reliability and resource allocation on ROCm platforms.
    pull/39765, pull/39845, pull/39849, pull/39822, pull/39848
  • Documentation updates: Documentation was updated for the squeeze function and argument docstrings in jax.numpy and jax.scipy to fix or add missing information without changing any functionality. These improvements enhance clarity and correctness of the user-facing docs.
    pull/39769, pull/39847
  • Constant lowering report frame handling fix: This fix corrects the handling of the jax_captured_constants_report_frames=-1 setting by treating negative frame limits as unbounded in both legacy and simplified constant-lowering paths. It also updates the report header to accurately reflect the existing five-allocation limit, preventing errors from passing -1 to itertools.islice.
    pull/39785
  • Gradient definition for jax.numpy.linalg.norm at zero: The gradient of jax.numpy.linalg.norm at zero is defined as zero to fix NaN gradients at that point. This subgradient convention maintains the gradient as an odd function and ensures compatibility with optimization problems involving unit norm normalizations.
    pull/39792
  • Bazel build and runtime library handling for ROCm: The Bazel build configuration was modified to include versioned ROCm runtime libraries directly in GPU test runfiles for relevant targets. This eliminates the need to export LD_LIBRARY_PATH in test scripts and ensures hermetic location of ROCm shared libraries during tests.
    pull/39807
  • Fix for reduce_window_jvp with symbolic zero tangents: This fix updates reduce_window_jvp to correctly handle symbolic zero tangents consistent with other JAX functions. It adds regression tests and verifies no regressions across CPU and GPU test suites.
    pull/39809
  • Gradient check improvements in check_grads: The default finite-difference step size in jax.test_util.check_grads was updated to be dtype-appropriate, using larger steps for sub-float32 types to avoid precision issues. Error messages were also improved by clearly labeling computed and numerical sides in gradient mismatch outputs.
    pull/39821
  • Pallas TPU backend memory-space constraint fix: This fix ensures that HBM BlockSpec memory-space constraints are properly propagated to custom call inputs in the Pallas TPU backend. It prevents incorrect operand placement and DMA path generation on TPU v4 hardware and includes regression tests verifying correct memory-space coloring and execution.
    pull/39838
  • Reenable Pallas fused attention tests on ROCm: Previously skipped Pallas fused attention tests on ROCm were reenabled by addressing a WAR hazard between VALU and GDS instructions fixed upstream in LLVM. This restores test coverage for these kernels on ROCm platforms.
    pull/39848
  • OneAPI platform support additions: Support for the OneAPI platform was added to multi-platform export tests, including updated guidance messages, test helpers, and MLIR lowering registration. Additionally, OneAPI devices were mapped to the existing GPU platform in jaxlib/py_device.cc with corresponding test updates.
    pull/39851, pull/39853
  • SparseCore indirect DMA documentation: Documentation was added describing hardware limitations and workaround patterns for SparseCore indirect DMA transfers supporting 16-bit data types in Pallas TPU kernels. It includes best practices for static pre-packing embedding tables to optimize performance and avoid runtime overhead.
    pull/39854
  • Shape_poly module improvement: The decision procedure in the shape_poly module was improved by enabling Euclidean decomposition recognition for mod factors with provably positive divisors, resolving previously inconclusive comparisons involving only mod expressions. The shape_poly benchmark script was also fixed to run correctly after recent internal changes.
    pull/39841
  • Intermittent ROCm blocking gate failure investigation: An empty commit was created to sample and investigate intermittent failures of the ROCm blocking gate job on the EngFlow RBE pool, which crashes with misleading 'out of memory' Hip errors during runtime compilation of blit kernels. This diagnostic effort will be closed once sufficient data is collected.
    pull/39822

3.2 Closed Pull Requests

This section provides a summary of pull requests that were closed in the repository over the past week. The top three pull requests with the highest number of commits are highlighted as 'key' pull requests. Other pull requests are grouped based on similar characteristics for easier analysis. Up to 25 pull requests are displayed in this section, while any remaining pull requests beyond this limit are omitted for brevity.

Pull Requests Closed This Week: 43

Key Closed Pull Requests

1. [ROCm] Add blocking Bazel PR gate and file-driven test selection: This pull request introduces a blocking Bazel presubmit gate for ROCm that runs a curated subset of tests on every PR using a file-driven test selection approach, replacing a hard-coded exclusion list with configurable target files, adds new Bazel configurations and GitHub workflow inputs to support multi-GPU testing and flexible target selection, and implements a pinned S3-based mechanism for sourcing ROCm plugin and PJRT wheels to ensure stable and reproducible CI gating.

  • URL: pull/38604
  • Associated Commits: 1f830, 0db0e, 9eb00, 127c0

2. Allocate triton's global scratch buffer instead of passing null: This pull request addresses the issue of Triton kernels requiring a global scratch buffer by implementing the allocation of a properly sized, stream-ordered buffer in KernelCall::Launch instead of passing a null pointer, thereby preventing illegal memory access errors during kernel launch when on-device tensor descriptors are used.

  • URL: pull/38912
  • Associated Commits: f292d, 38204

3. Fix float64 type degradation in reciprocal primitive lowering: This pull request addresses a precision degradation bug in the Pallas reciprocal primitive by preventing the downcasting of float64 inputs to float32 during MLIR lowering when approx=True, and includes added test coverage to verify the fix.

  • URL: pull/35791
  • Associated Commits: ab5ba

Other Closed Pull Requests

  • Random sampler validation improvements: Multiple pull requests enhance random sampler functions by introducing shape and dtype validation using _check_broadcast_shapes and _check_all_safe_to_cast. These changes ensure consistent promotion and broadcast semantics across uniform, truncated normal, gamma, Dirichlet, and binomial samplers as part of a random API refactor.
    • pull/37637, pull/37639, pull/37801, pull/37804, pull/37807
  • ROCm CI container teardown and workflow fixes: Several pull requests address issues in ROCm CI workflows by adding the --init flag to Docker containers to prevent hangs and improve shutdown times, and by adding a short cleanup step to avoid prolonged stalling during container stop phases. These changes improve the reliability and efficiency of ROCm CI job container management.
    • pull/39473, pull/39601
  • OneAPI GPU support and CI enhancements: Pull requests introduce foundational OneAPI GPU solver kernel support including Cholesky kernel implementation and asynchronous memory operations, alongside adding OneAPI continuous integration test scripts and infrastructure. These contributions enable Intel GPU automatic detection, parallel testing, and local installation patterns for OneAPI wheels.
    • pull/39285, [pull/39639](https://github.com/pull/39639]
  • Fixes for GPU kernel compilation and launch: A pull request fixes failing Torch kernel compilation and launch by registering the LLVM NVPTX target during MosaicGpuCompile and recoding the kernel launch as a C function to correctly pass data pointers. This resolves CUDA compute capability errors and kernel launch issues.
    • pull/39640
  • Bug fixes in special functions: Pull requests fix bugs in jax.scipy.special.digamma and jax.scipy.special.betaln functions by correcting return values for edge cases and adding guards for infinite arguments. These fixes align behavior with SciPy and handle indeterminate forms properly.
    • pull/37917, pull/39682
  • State leakage prevention in xla_bridge_test: A pull request introduces a restore_backends_and_env decorator to snapshot and restore environment variables and backend factories around plugin-registration tests. This prevents state leakage and interference across xdist workers and subprocess tests.
    • pull/37053
  • Removal of call_p and closed_call_p primitives: A pull request removes the call_p and closed_call_p primitives in favor of using only eval_jaxpr_p, eliminating final-style call processing functions while retaining stubs and aliases for compatibility with downstream users.
    • pull/39593
  • Fix for tiled cp.async copy operation: A pull request fixes a bug where sub-block copies in the tiled cp.async operation fail to lower correctly due to incorrect shape, causing errors in pipelined copies and making the tiled path unusable from plgpu.emit_pipeline.
    • pull/39689
  • Dependency management update: A pull request removes the use of ratchet versions in dependency management and instead pins specific versions, shifting reliance to dependabot for updates due to inconsistent ratchet-style comment maintenance.
    • pull/39685
  • Regression test for eval_jaxpr_p partial evaluation: A pull request adds a regression test for the eval_jaxpr_p partial evaluation rule to address a specific issue.
    • pull/39697
  • ROC nightly CI update to ROCm 10 wheels: A pull request updates TheRock nightly continuous integration configuration to use ROCm 10 wheels instead of ROCm 7, aligning with mainline release numbering for correct N+1 release flow execution.
    • pull/39698
  • XLA sha256 checksum correction: A pull request corrects the XLA sha256 checksum for a specific commit to ensure build integrity and consistency.
    • pull/39710
  • Notebook and markdown sync: A pull request proposes syncing the content of the matmul.ipynb notebook with the matmul.md markdown file to maintain consistency between the two documents.
    • pull/39717
  • ROCm wheel test dependency fix: A pull request adds missing libhipsparse staging to ROCm wheel test dependencies by including @local_config_rocm//rocm:hipsparse in rocm_libs_data. This prevents failures in tridiagonal solve tests caused by unregistered FFI handlers in hermetic environments.
    • pull/39724
  • Addition of jax.numpy.top_k function: A pull request adds the jax.numpy.top_k function implementing the new numpy.top_k feature from NumPy v2.6.0, including local testing for compatibility with NumPy nightly and fallback baseline.
    • pull/39729
  • Predicate argument added to copy_gmem_to_smem: A pull request adds a predicate argument to the copy_gmem_to_smem function in the Pallas Mosaic GPU module and includes a test to validate this enhancement.
    • pull/38699

3.3 Pull Request Discussion Insights

This section will analyze the tone and sentiment of discussions within this project's open and closed pull requests that occurred within the past week. It aims to identify potentially heated exchanges and to maintain a constructive project environment.

Based on our analysis, there are no instances of toxic discussions in the project's open or closed pull requests from the past week.


IV. Contributors

4.1 Contributors

Active Contributors:

We consider an active contributor in this project to be any contributor who has made at least 1 commit, opened at least 1 issue, created at least 1 pull request, or made more than 2 comments in the last month.

If there are more than 10 active contributors, the list is truncated to the top 10 based on contribution metrics for better clarity.

Contributor Commits Pull Requests Issues Comments
mattjj 44 8 0 3
jakevdp 19 8 0 28
kodlan 27 4 0 1
magaonka-amd 10 7 0 0
vfdev-5 9 4 0 0
ALinrunrun 0 0 12 0
teddytennant 8 3 0 0
NeilGirdhar 3 2 0 6
kranipa 5 4 0 0
ayaka14732 1 0 3 4

Don't miss what's next. Subscribe to Weekly Project News:
Powered by Buttondown, the easiest way to start and grow your newsletter.