google / google/flax

flax.nnx.cond causes tracing cache misses

Open
#5,512 3 comments 0 reactions 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
7.3k
Forks
833
Avg merge
5h 11m
Merged PRs (30d)
5

Description

Using `nnx.cond` causes tracing cache misses (and thus prevents persistent compilation cache loads) due to creation of transient SimpleCondFn passed to `jax.lax.cond`. This is very inconvenient when using nnx.cond with larger modules that take a while to compile.

```
TRACING CACHE MISS at .../flax/nnx/transforms/transforms.py:623:19 (cond):
never seen function:
true_fn id=17864732047504 defined at .../repro_cond_miss.py:15
but seen another function defined on the same line; maybe the function is
being re-defined repeatedly, preventing caching?
W0622 10:13:43.770388 1182679 partial_eval.py:2164] TRACING CACHE MISS at .../flax/nnx/transforms/transforms.py:623:19 (cond):
never seen function:
false_fn id=17864737580928 defined at .../repro_cond_miss.py:19
but seen another function defined on the same line; maybe the function is
being re-defined repeatedly, preventing caching?
W0622 10:13:43.771588 1182679 dispatch.py:204] Finished tracing cond for jit in 0.000418901 sec
```

Here's a repro code that prints the cache misses (including the jax-level prints)

```python
import weakref
from absl import app
from flax import nnx
from flax.nnx.transforms import transforms as nnx_transforms
import jax

jax.config.update("jax_log_compiles", True)
jax.config.update("jax_explain_cache_misses", True)

def true_fn(x):
return x + 1.0

def false_fn(x):
return x - 1.0

original_SimpleCondFn = nnx_transforms.SimpleCondFn

def run_test(name, patch_fn):
print(f"\n--- Running Test: {name} ---", flush=True)

patch_fn()

def test_fn(pred, x):
# Pass graph_updates=False to force the path that uses SimpleCondFn
return nnx.cond(pred, true_fn, false_fn, x, graph_updates=False)

for i in range(3):
print(f"\nIteration {i}:", flush=True)
pred = jax.numpy.array(i % 2 == 0)
res = test_fn(pred, 1.0)
jax.block_until_ready(res)

# Buggy implementation (original behavior, but with prints)
def apply_buggy_patch():
# We want to intercept calls to SimpleCondFn to print when they are created
def buggy_SimpleCondFn_factory(f, graph):
instance = original_SimpleCondFn(f, graph)
print(
f" [BUGGY] Created new wrapper for {f.__name__} (id={id(f)}):"
f" wrapper_id={id(instance)}",
flush=True,
)
return instance

nnx_transforms.SimpleCondFn = buggy_SimpleCondFn_factory

# Patched implementation (using cache)
_repro_cache = weakref.WeakKeyDictionary()

def apply_patched_patch():
_repro_cache.clear() # Clear cache for clean run

def patched_SimpleCondFn_factory(f, graph):
if f not in _repro_cache:
_repro_cache[f] = {}
if graph not in _repro_cache[f]:
instance = original_SimpleCondFn(f, graph)
_repro_cache[f][graph] = instance
print(
f" [PATCHED] Cache MISS for {f.__name__} (id={id(f)}). Created"
f" wrapper_id={id(instance)}",
flush=True,
)
else:
instance = _repro_cache[f][graph]
print(
f" [PATCHED] Cache HIT for {f.__name__} (id={id(f)}). Reused"
f" wrapper_id={id(instance)}",
flush=True,
)
return instance

nnx_transforms.SimpleCondFn = patched_SimpleCondFn_factory

def main(argv):
if len(argv) > 1:
raise app.UsageError("Too many command-line arguments.")

print(f"Global true_fn id: {id(true_fn)}")
print(f"Global false_fn id: {id(false_fn)}")

# Run buggy first
run_test("BUGGY (Original)", apply_buggy_patch)

# Run patched
run_test("PATCHED", apply_patched_patch)
print("\n--- Test Finished ---", flush=True)

if __name__ == "__main__":
app.run(main)
```

Contributor guide

Open the contributing guide

Research direction

Start in flax/nnx/transforms/transforms.py around line 623, where cond passes SimpleCondFn to jax.lax.cond, and compare that path with the provided reproduction. Verify repeated nnx.cond calls reuse the relevant wrappers rather than causing tracing cache misses, then rerun the reproduction to confirm the cache-miss messages no longer recur.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Bug
Difficulty
3/5
Estimated time
1-2 days
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
68/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.