vertebrae.extractors.jax_flax
Optional JAX/Flax extractor.
Classes
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])