vertebrae.extractors.jax_flax

Optional JAX/Flax extractor.

Classes

JAXFlaxExtractor

Wrap a JAX/Flax apply function or model object as a vertebrae extractor.

Module Contents

class vertebrae.extractors.jax_flax.JAXFlaxExtractor(name, input_fn, output_fn=None, apply_fn=None, model=None, params=None, outputs=None, structured_outputs=None, modality='unknown', jit=True, apply_kwargs=None, cache_identity=None)[source]

Wrap a JAX/Flax apply function or model object as a vertebrae extractor.

Parameters:
  • name (str)

  • input_fn (Callable[[Any], Any])

  • output_fn (Optional[Callable[[Any], Any]])

  • apply_fn (Optional[Callable[Ellipsis, Any]])

  • model (Any)

  • params (Any)

  • outputs (Optional[Sequence[Dict[str, Any]]])

  • structured_outputs (Optional[Sequence[Dict[str, Any]]])

  • modality (str)

  • jit (bool)

  • apply_kwargs (Optional[Dict[str, Any]])

  • cache_identity (Optional[str])

get_resource_profile_adapter()[source]

Return JAX synchronization, device, and parameter-footprint hooks.

Return type:

Any