From 1b55d9c473c98cc2e6bd008da531a160c65bc2b0 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Fri, 14 Aug 2026 07:23:48 -0700 Subject: [PATCH] Move ffi type definition before ffi target definitions Signed-off-by: Jeremy Berchtold --- transformer_engine/jax/cpp_extensions/base.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) 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.