vertebrae.extractors.graph

Optional graph-model extractor for PyG and DGL workflows.

Classes

GraphModelExtractor

Wrap a graph model that returns one embedding row per graph sample.

Module Contents

class vertebrae.extractors.graph.GraphModelExtractor(name, model, collate_fn, output_fn=None, device=None, framework=None, recipe_data=None, allow_sparse=False, streaming_safe=True, output_level='graph', move_batch_to_device=True, move_model_to_device=True, checkpoint_paths=None, cache_identity=None)[source]

Wrap a graph model that returns one embedding row per graph sample.

Parameters:
  • name (str)

  • model (Any)

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

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

  • device (Optional[str])

  • framework (Optional[str])

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

  • allow_sparse (bool)

  • streaming_safe (bool)

  • output_level (str)

  • move_batch_to_device (bool)

  • move_model_to_device (bool)

  • checkpoint_paths (Optional[Sequence[str]])

  • cache_identity (Optional[str])