Source code for gigl.distributed.utils.serialized_graph_metadata_translator
from typing import Optional, Tuple, Union
from gigl.common import UriFactory
from gigl.common.data.dataloaders import SerializedTFRecordInfo
from gigl.common.data.load_torch_tensors import SerializedGraphMetadata
from gigl.common.utils.tensorflow_schema import feature_spec_to_feature_index_map
from gigl.src.common.types.graph_data import EdgeType, NodeType
from gigl.src.common.types.pb_wrappers.graph_metadata import GraphMetadataPbWrapper
from gigl.src.common.types.pb_wrappers.preprocessed_metadata import (
PreprocessedMetadataPbWrapper,
)
from gigl.src.data_preprocessor.lib.types import FeatureSpecDict
from gigl.types.graph import FeatureQuantizationMetadata, to_homogeneous
from snapchat.research.gbml.preprocessed_metadata_pb2 import PreprocessedMetadata
def _build_serialized_tfrecord_entity_info(
preprocessed_metadata: Union[
PreprocessedMetadata.NodeMetadataOutput, PreprocessedMetadata.EdgeMetadataInfo
],
feature_spec_dict: FeatureSpecDict,
entity_key: Union[str, Tuple[str, str]],
tfrecord_uri_pattern: str,
quantization_metadata: Optional[FeatureQuantizationMetadata] = None,
) -> SerializedTFRecordInfo:
"""
Populates a SerializedTFRecordInfo field from provided arguments for either a node or edge entity of a single node/edge type.
Args:
preprocessed_metadata(Union[
PreprocessedMetadata.NodeMetadataOutput, PreprocessedMetadata.EdgeMetadataInfo
]): Preprocessed metadata pb for either NodeMetadataOutput or EdgeMetadataInfo
feature_spec_dict (FeatureSpecDict): Feature spec to register to SerializedTFRecordInfo
entity_key (Union[str, Tuple[str, str]]): Entity key to register to SerializedTFRecordInfo, is a str if Node entity or Tuple[str, str] if Edge entity
tfrecord_uri_pattern (str): Regex pattern for loading serialized tf records
quantization_metadata (Optional[FeatureQuantizationMetadata]): Quantization
metadata for a node or main-edge entity when its features are quantized.
Returns:
SerializedTFRecordInfo: Stored metadata for current entity
"""
if quantization_metadata is not None:
packed_feature_key = (
preprocessed_metadata.quantized_feature_metadata.packed_feature_key
)
packed_feature_dim = quantization_metadata.packed_feature_dim
quantized_indices = set(quantization_metadata.quantized_feature_indices)
feature_index = feature_spec_to_feature_index_map(
{key: feature_spec_dict[key] for key in preprocessed_metadata.feature_keys}
)
feature_keys_to_load: list[str] = []
# Identify non-quantized feature fields to load from float storage.
for key in preprocessed_metadata.feature_keys:
key_indices = set(range(*feature_index[key]))
quantized_key_indices = key_indices.intersection(quantized_indices)
if not quantized_key_indices:
feature_keys_to_load.append(key)
elif len(quantized_key_indices) != len(key_indices):
# TFRecord decoding cannot split a field between float and packed storage.
raise ValueError(f"Partial quantization not supported for {key}")
feature_dim = quantization_metadata.raw_feature_dim
else:
packed_feature_key = None
packed_feature_dim = 0
feature_keys_to_load = list(preprocessed_metadata.feature_keys)
feature_dim = preprocessed_metadata.feature_dim
serialized_keys = set(feature_keys_to_load)
serialized_keys.update(preprocessed_metadata.label_keys)
if packed_feature_key is not None:
serialized_keys.add(packed_feature_key)
feature_spec_dict = {
key: spec for key, spec in feature_spec_dict.items() if key in serialized_keys
}
return SerializedTFRecordInfo(
tfrecord_uri_prefix=UriFactory.create_uri(
preprocessed_metadata.tfrecord_uri_prefix
),
feature_keys=feature_keys_to_load,
feature_spec=feature_spec_dict,
feature_dim=feature_dim,
entity_key=entity_key,
packed_feature_key=packed_feature_key,
packed_feature_dim=packed_feature_dim,
label_keys=list(preprocessed_metadata.label_keys),
tfrecord_uri_pattern=tfrecord_uri_pattern,
)
def _build_feature_quantization_metadata(
quantized_metadata: PreprocessedMetadata.FeatureQuantizationMetadata,
feature_dim: int,
) -> FeatureQuantizationMetadata:
state = quantized_metadata.WhichOneof("state")
neg_mean: Optional[float] = None
pos_mean: Optional[float] = None
clip_min: Optional[float] = None
clip_max: Optional[float] = None
if state == "single_bit_state":
bits = 1
neg_mean = quantized_metadata.single_bit_state.neg_mean
pos_mean = quantized_metadata.single_bit_state.pos_mean
elif state == "multi_bit_state":
bits = quantized_metadata.multi_bit_state.bits
clip_min = quantized_metadata.multi_bit_state.clip_min
clip_max = quantized_metadata.multi_bit_state.clip_max
else:
raise ValueError("Expected quantization state to be set.")
return FeatureQuantizationMetadata(
bits=bits,
feature_dim=feature_dim,
quantized_feature_indices=tuple(quantized_metadata.quantized_feature_indices),
clip_min=clip_min,
clip_max=clip_max,
neg_mean=neg_mean,
pos_mean=pos_mean,
)
[docs]
def convert_pb_to_serialized_graph_metadata(
preprocessed_metadata_pb_wrapper: PreprocessedMetadataPbWrapper,
graph_metadata_pb_wrapper: GraphMetadataPbWrapper,
tfrecord_uri_pattern: str = ".*-of-.*\.tfrecord(\.gz)?$",
) -> SerializedGraphMetadata:
"""
Populates a SerializedGraphMetadata field from PreprocessedMetadataPbWrapper and GraphMetadataPbWrapper, containing information for loading tensors for all entities and node/edge types.
Args:
preprocessed_metadata_pb_wrapper (PreprocessedMetadataPbWrapper): Preprocessed Metadata Pb Wrapper to translate into SerializedGraphMetadata
graph_metadata_pb_wrapper (GraphMetadataPbWrapper): Graph Metadata Pb Wrapper to translate into Dataset Metadata
tfrecord_uri_pattern (str): Regex pattern for loading serialized tf records
Returns:
SerializedGraphMetadata: Dataset Metadata for all entity and node/edge types.
"""
node_entity_info: dict[NodeType, SerializedTFRecordInfo] = {}
edge_entity_info: dict[EdgeType, SerializedTFRecordInfo] = {}
positive_label_entity_info: dict[EdgeType, SerializedTFRecordInfo] = {}
negative_label_entity_info: dict[EdgeType, SerializedTFRecordInfo] = {}
node_quantization_metadata: dict[NodeType, FeatureQuantizationMetadata] = {}
edge_quantization_metadata: dict[EdgeType, FeatureQuantizationMetadata] = {}
preprocessed_metadata_pb = preprocessed_metadata_pb_wrapper.preprocessed_metadata_pb
for node_type in graph_metadata_pb_wrapper.node_types:
condensed_node_type = (
graph_metadata_pb_wrapper.node_type_to_condensed_node_type_map[node_type]
)
node_metadata = (
preprocessed_metadata_pb.condensed_node_type_to_preprocessed_metadata[
condensed_node_type
]
)
node_feature_spec_dict = (
preprocessed_metadata_pb_wrapper.condensed_node_type_to_feature_schema_map[
condensed_node_type
].feature_spec
)
node_key = node_metadata.node_id_key
if node_metadata.HasField("quantized_feature_metadata"):
node_quantization_metadata[node_type] = (
_build_feature_quantization_metadata(
quantized_metadata=node_metadata.quantized_feature_metadata,
feature_dim=node_metadata.feature_dim,
)
)
node_entity_info[node_type] = _build_serialized_tfrecord_entity_info(
preprocessed_metadata=node_metadata,
feature_spec_dict=node_feature_spec_dict,
entity_key=node_key,
tfrecord_uri_pattern=tfrecord_uri_pattern,
quantization_metadata=node_quantization_metadata.get(node_type),
)
for edge_type in graph_metadata_pb_wrapper.edge_types:
condensed_edge_type = (
graph_metadata_pb_wrapper.edge_type_to_condensed_edge_type_map[edge_type]
)
edge_metadata = (
preprocessed_metadata_pb.condensed_edge_type_to_preprocessed_metadata[
condensed_edge_type
]
)
edge_key = (
edge_metadata.src_node_id_key,
edge_metadata.dst_node_id_key,
)
if edge_metadata.HasField("main_edge_info"):
edge_feature_spec_dict = preprocessed_metadata_pb_wrapper.condensed_edge_type_to_feature_schema_map[
condensed_edge_type
].feature_spec
if edge_metadata.main_edge_info.HasField("quantized_feature_metadata"):
edge_quantization_metadata[edge_type] = (
_build_feature_quantization_metadata(
quantized_metadata=edge_metadata.main_edge_info.quantized_feature_metadata,
feature_dim=edge_metadata.main_edge_info.feature_dim,
)
)
edge_entity_info[edge_type] = _build_serialized_tfrecord_entity_info(
preprocessed_metadata=edge_metadata.main_edge_info,
feature_spec_dict=edge_feature_spec_dict,
entity_key=edge_key,
tfrecord_uri_pattern=tfrecord_uri_pattern,
quantization_metadata=edge_quantization_metadata.get(edge_type),
)
if edge_metadata.HasField("positive_edge_info"):
pos_edge_feature_spec_dict = preprocessed_metadata_pb_wrapper.condensed_edge_type_to_pos_edge_feature_schema_map[
condensed_edge_type
].feature_spec
positive_label_entity_info[edge_type] = (
_build_serialized_tfrecord_entity_info(
preprocessed_metadata=edge_metadata.positive_edge_info,
feature_spec_dict=pos_edge_feature_spec_dict,
entity_key=edge_key,
tfrecord_uri_pattern=tfrecord_uri_pattern,
)
)
if edge_metadata.HasField("negative_edge_info"):
hard_neg_edge_feature_spec_dict = preprocessed_metadata_pb_wrapper.condensed_edge_type_to_hard_neg_edge_feature_schema_map[
condensed_edge_type
].feature_spec
negative_label_entity_info[edge_type] = (
_build_serialized_tfrecord_entity_info(
preprocessed_metadata=edge_metadata.negative_edge_info,
feature_spec_dict=hard_neg_edge_feature_spec_dict,
entity_key=edge_key,
tfrecord_uri_pattern=tfrecord_uri_pattern,
)
)
if not graph_metadata_pb_wrapper.is_heterogeneous:
# If our input is homogeneous, we remove the node/edge type component of the metadata fields.
return SerializedGraphMetadata(
node_entity_info=to_homogeneous(node_entity_info),
edge_entity_info=to_homogeneous(edge_entity_info),
positive_label_entity_info=to_homogeneous(positive_label_entity_info)
if len(positive_label_entity_info) > 0
else None,
negative_label_entity_info=to_homogeneous(negative_label_entity_info)
if len(negative_label_entity_info) > 0
else None,
node_quantization_metadata=to_homogeneous(node_quantization_metadata)
if len(node_quantization_metadata) > 0
else None,
edge_quantization_metadata=to_homogeneous(edge_quantization_metadata)
if len(edge_quantization_metadata) > 0
else None,
)
else:
return SerializedGraphMetadata(
node_entity_info=node_entity_info,
edge_entity_info=edge_entity_info,
positive_label_entity_info=positive_label_entity_info
if len(positive_label_entity_info) > 0
else None,
negative_label_entity_info=negative_label_entity_info
if len(negative_label_entity_info) > 0
else None,
node_quantization_metadata=node_quantization_metadata
if len(node_quantization_metadata) > 0
else None,
edge_quantization_metadata=edge_quantization_metadata
if len(edge_quantization_metadata) > 0
else None,
)