From 2d791bb1458b01697298fcc88409eea6575a0b73 Mon Sep 17 00:00:00 2001 From: Yash Katariya Date: Thu, 20 Aug 2026 21:28:22 -0700 Subject: [PATCH] [No functional change] Delete VJPHiPrimitive and replace with HiPrim. This is a simple replace. PiperOrigin-RevId: 968242680 --- flax/nnx/variablelib.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/flax/nnx/variablelib.py b/flax/nnx/variablelib.py index 7b2b06cd3..2c29c2bca 100644 --- a/flax/nnx/variablelib.py +++ b/flax/nnx/variablelib.py @@ -294,8 +294,10 @@ def _new_hijax_from_variable(variable: Variable) -> HijaxVariable: ) return hijax_var +HiPrim = (hjx.VJPHiPrimitive if jax.__version_info__ <= (0, 11, 1) else + hjx.HiPrim) -class NewVariable(hjx.VJPHiPrimitive): +class NewVariable(HiPrim): def __init__(self, *leaf_avals, treedef, var_type, ref=False): self.in_avals = tuple(leaf_avals) self.out_aval = AbstractVariable( @@ -354,7 +356,7 @@ def bind(self, *leaves, treedef, var_type, ref=False): new_variable_p = _NewVariableShim() -class SetVariable(hjx.VJPHiPrimitive): +class SetVariable(HiPrim): def __init__(self, hijax_var_aval, *leaf_avals, treedef, var_type): self.in_avals = (hijax_var_aval, *leaf_avals) self.out_aval = () @@ -433,7 +435,7 @@ def _set_hijax_state(hijax_var, variable: Variable): ) -class GetVariable(hjx.VJPHiPrimitive): +class GetVariable(HiPrim): def __init__(self, abstract_var, *, treedef, avals, var_type): self.in_avals = (abstract_var,) self.out_aval = tuple(avals)