@@ -426,31 +426,62 @@ class SimpleJitFn:
426426 donate_argnames : frozenset [str ]
427427 graph : bool
428428 update_shardings : tuple [tp .Any , ...]
429+ captured_info : tp .Optional [list [tp .Any ]]
429430
430431 def __post_init__ (self ):
431432 functools .update_wrapper (self , self .f , updated = ())
433+ # When captures are prepended as an extra argument, the original
434+ # function's signature no longer matches the physical call signature.
435+ # Remove __wrapped__ so inspect.signature (used by JAX for validation
436+ # of static_argnums / donate_argnums) sees *args instead.
437+ if self .captured_info :
438+ delattr (self , '__wrapped__' )
432439
433440 @extract .treemap_copy_args
434441 def __call__ (self , * args , ** kwargs ):
435- current , snapshot = extract .snapshot (
436- labeled (args = args , kwargs = kwargs )
437- )
442+ if self .captured_info :
443+ captured_args = args [0 ]
444+ user_args = args [1 :]
445+ captured_shardings = self .in_shardings [0 ] if isinstance (self .in_shardings , tuple ) else None
446+ user_shardings = self .in_shardings [1 :] if isinstance (self .in_shardings , tuple ) else self .in_shardings
447+ current , snapshot = extract .snapshot (
448+ labeled (captured_args = captured_args , args = user_args , kwargs = kwargs )
449+ )
450+ else :
451+ user_args = args
452+ current , snapshot = extract .snapshot (
453+ labeled (args = user_args , kwargs = kwargs )
454+ )
438455 if self .graph :
439456 args , kwargs = extract .from_tree2 ((args , kwargs ))
440- out = self .f (* args , ** kwargs )
457+ if self .captured_info :
458+ captured_args = args [0 ]
459+ user_args = args [1 :]
460+ f = extract .replace_closure_cells (
461+ self .f , self .captured_info , captured_args )
462+ else :
463+ user_args = args
464+ f = self .f
465+ out = f (* user_args , ** kwargs )
441466 if self .graph :
442467 out = extract .to_tree2 (out , prefix = self .out_shardings )
443468 extract .check_no_aliases ('jit' , ** current , out = out , check = ['out' ])
444469 def keep_fn (jax_path , prefix , c , s ):
445470 if extract .variable_changed (c , s ):
446471 return True
447472 arg_type , arg_key , * _ = graphlib .jax_to_nnx_path (jax_path )
473+ if arg_type == 'captured_args' :
474+ return False
448475 if arg_type == 'args' :
449476 return arg_key in self .donate_argnums
450477 else : # arg_type == 'kwargs':
451478 return arg_key in self .donate_argnames
479+ if self .captured_info :
480+ prefix = labeled (captured_args = captured_shardings , args = user_shardings , kwargs = None )
481+ else :
482+ prefix = labeled (args = self .in_shardings , kwargs = None )
452483 updates = extract .get_updates (
453- current , snapshot , prefix = labeled ( args = self . in_shardings , kwargs = None ) ,
484+ current , snapshot , prefix = prefix ,
454485 known_prefixes = self .update_shardings , keep_fn = keep_fn
455486 )
456487 return out , updates
@@ -533,6 +564,9 @@ def __init__(
533564 self .partial_args = partial_args
534565 self .graph = graph
535566
567+ # Capture closure nodes once at construction time
568+ self ._captured_info = extract .find_captured_nodes (fun )
569+
536570 resolved = _resolve_argnums (fun , static_argnums , static_argnames )
537571 if isinstance (in_shardings , (tuple , list )) and resolved :
538572 expanded = list (in_shardings )
@@ -542,6 +576,24 @@ def __init__(
542576 else :
543577 self .in_shardings = in_shardings
544578
579+ # Prepend None shardings for captured nodes (passed as a list)
580+ n_captured = len (self ._captured_info )
581+ if n_captured > 0 :
582+ if isinstance (in_shardings , (tuple , list )):
583+ jit_in_shardings = ([None ] * n_captured ,) + tuple (in_shardings )
584+ elif in_shardings is not None :
585+ # Expand scalar in_shardings to match the number of user parameters
586+ # so it doesn't broadcast over the captures.
587+ sig = _fun_signature (fun )
588+ n_params = len (sig .parameters ) if sig is not None else 1
589+ jit_in_shardings = ([None ] * n_captured ,) + tuple (
590+ in_shardings for _ in range (n_params )
591+ )
592+ else :
593+ jit_in_shardings = in_shardings
594+ else :
595+ jit_in_shardings = in_shardings
596+
545597 donate_argnums_set = frozenset (
546598 (donate_argnums ,) if isinstance (donate_argnums , int )
547599 else donate_argnums or ()
@@ -550,21 +602,33 @@ def __init__(
550602 (donate_argnames ,) if isinstance (donate_argnames , str )
551603 else donate_argnames or ()
552604 )
605+
606+ # When captures are prepended as an extra first argument, offset
607+ # index-based jax.jit parameters so they still refer to the correct
608+ # physical arguments.
609+ if n_captured > 0 :
610+ jit_static_argnums = _offset_argnums (static_argnums , 1 )
611+ jit_donate_argnums = _offset_argnums (donate_argnums , 1 )
612+ else :
613+ jit_static_argnums = static_argnums
614+ jit_donate_argnums = donate_argnums
615+
553616 self .jitted_fn = jax .jit (
554617 SimpleJitFn (
555618 fun ,
556- self .in_shardings ,
619+ jit_in_shardings if n_captured > 0 else self .in_shardings ,
557620 out_shardings ,
558621 donate_argnums_set ,
559622 donate_argnames_set ,
560623 graph ,
561624 tuple (update_shardings ),
625+ self ._captured_info ,
562626 ),
563- in_shardings = in_shardings ,
627+ in_shardings = jit_in_shardings ,
564628 out_shardings = (out_shardings , update_shardings ),
565- static_argnums = static_argnums ,
629+ static_argnums = jit_static_argnums ,
566630 static_argnames = static_argnames ,
567- donate_argnums = donate_argnums ,
631+ donate_argnums = jit_donate_argnums ,
568632 donate_argnames = donate_argnames ,
569633 keep_unused = keep_unused ,
570634 device = device ,
@@ -595,8 +659,13 @@ def _maybe_from_tree(self, out):
595659
596660 def __call__ (self , * args : P .args , ** kwargs : P .kwargs ) -> R :
597661 args = (* self .partial_args , * args ) # type: ignore[assignment]
662+ args = (self ._captured_info , * args ) if self ._captured_info else args
598663 args , kwargs = self ._maybe_to_tree (args , kwargs )
599- variables = extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
664+ if self ._captured_info :
665+ variables = extract .check_no_aliases (
666+ 'jit' , captured_args = args [0 ], args = args [1 :], kwargs = kwargs )
667+ else :
668+ variables = extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
600669 out , updates = self .jitted_fn (* args , ** kwargs )
601670 extract .apply_updates (variables , updates )
602671 return self ._maybe_from_tree (out )
@@ -608,22 +677,33 @@ def __get__(self, obj, objtype=None):
608677
609678 def eval_shape (self , * args , ** kwargs ):
610679 args = (* self .partial_args , * args )
680+ args = (list (self ._captured_info ), * args ) if self ._captured_info else args
611681 args , kwargs = self ._maybe_to_tree (args , kwargs )
612682 if not self .graph :
613- extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
683+ if self ._captured_info :
684+ extract .check_no_aliases (
685+ 'jit' , captured_args = args [0 ], args = args [1 :], kwargs = kwargs )
686+ else :
687+ extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
614688 out , _ = self .jitted_fn .eval_shape (* args , ** kwargs )
615689 return self ._maybe_from_tree (out )
616690
617691 def trace (self , * args , ** kwargs ):
618692 args = (* self .partial_args , * args )
693+ args = (list (self ._captured_info ), * args ) if self ._captured_info else args
619694 args , kwargs = self ._maybe_to_tree (args , kwargs )
620695 if not self .graph :
621- extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
696+ if self ._captured_info :
697+ extract .check_no_aliases (
698+ 'jit' , captured_args = args [0 ], args = args [1 :], kwargs = kwargs )
699+ else :
700+ extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
622701 traced = self .jitted_fn .trace (* args , ** kwargs )
623702 return SimpleTraced (traced , self )
624703
625704 def lower (self , * args , ** kwargs ):
626705 args = (* self .partial_args , * args )
706+ args = (list (self ._captured_info ), * args ) if self ._captured_info else args
627707 args , kwargs = self ._maybe_to_tree (args , kwargs )
628708 if not self .graph :
629709 extract .check_no_aliases ('jit' , args = args , kwargs = kwargs )
@@ -1916,6 +1996,18 @@ def shard_map_wrapper(*args, **kwargs):
19161996 return shard_map_wrapper # type: ignore
19171997
19181998
1999+ def _offset_argnums (
2000+ argnums : int | tp .Sequence [int ] | None ,
2001+ offset : int ,
2002+ ) -> int | tuple [int , ...] | None :
2003+ """Shift argnum indices by ``offset`` (e.g. when prepending captures)."""
2004+ if argnums is None :
2005+ return None
2006+ if isinstance (argnums , int ):
2007+ return argnums + offset
2008+ return tuple (i + offset for i in argnums )
2009+
2010+
19192011# We can't use private methods from jax._src.api_util
19202012# We copy the function: api_util.fun_signature
19212013def _fun_signature (fun : tp .Callable ) -> inspect .Signature | None :
0 commit comments