diff --git a/transformer_engine/jax/cpp_extensions/base.py b/transformer_engine/jax/cpp_extensions/base.py index 2cdef4bfe7..f45a77012c 100644 --- a/transformer_engine/jax/cpp_extensions/base.py +++ b/transformer_engine/jax/cpp_extensions/base.py @@ -263,10 +263,8 @@ def _gspmd_wrapper(*args, **kwargs): cls.outer_primitive = outer_p -for _name, _value in transformer_engine_jax.registrations().items(): - ffi.register_ffi_target(_name, _value, platform="CUDA") - # Register EpInstanceState (no-op when TE is built without NCCL EP). +# Custom types must be registered before any FFI handler that references them. if hasattr(transformer_engine_jax, "get_ep_instance_state_type_id"): ffi.register_ffi_type( "EpInstanceState", @@ -278,6 +276,10 @@ def _gspmd_wrapper(*args, **kwargs): ) +for _name, _value in transformer_engine_jax.registrations().items(): + ffi.register_ffi_target(_name, _value, platform="CUDA") + + def manage_primitives(enable_names=None, disable_names=None, disable_all_first=False): """ Helper function to manage primitive states by name without modifying environment variables.