From 5eab8bc0df8a688656a83182d9a1f375c5dadff8 Mon Sep 17 00:00:00 2001 From: jchmura Date: Sat, 15 Aug 2026 19:22:28 +0000 Subject: [PATCH 1/8] Add edge quantization metadata proto --- .../research/gbml/preprocessed_metadata.proto | 2 + .../PreprocessedMetadata.scala | 45 ++++++++++++++++--- .../PreprocessedMetadataProto.scala | 34 +++++++------- .../PreprocessedMetadata.scala | 45 ++++++++++++++++--- .../PreprocessedMetadataProto.scala | 34 +++++++------- .../gbml/preprocessed_metadata_pb2.py | 18 ++++---- .../gbml/preprocessed_metadata_pb2.pyi | 9 +++- 7 files changed, 134 insertions(+), 53 deletions(-) diff --git a/proto/snapchat/research/gbml/preprocessed_metadata.proto b/proto/snapchat/research/gbml/preprocessed_metadata.proto index d7dfe3469..661c91d85 100644 --- a/proto/snapchat/research/gbml/preprocessed_metadata.proto +++ b/proto/snapchat/research/gbml/preprocessed_metadata.proto @@ -71,6 +71,8 @@ message PreprocessedMetadata{ optional uint32 feature_dim = 6; // Contains categorical feature vocabularies string transform_fn_assets_uri = 7; + // Optional quantized main-edge feature metadata. + FeatureQuantizationMetadata quantized_feature_metadata = 8; } // Houses metadata about edge TFTransform output from DataPreprocessor. diff --git a/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala b/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala index 7a9012ffa..9ddc933a7 100644 --- a/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala +++ b/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala @@ -1125,6 +1125,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r * Feature dimension after preprocessing * @param transformFnAssetsUri * Contains categorical feature vocabularies + * @param quantizedFeatureMetadata + * Optional quantized main-edge feature metadata. */ @SerialVersionUID(0L) final case class EdgeMetadataInfo( @@ -1135,6 +1137,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedEdgeDataBqTable: _root_.scala.Predef.String = "", featureDim: _root_.scala.Option[_root_.scala.Int] = _root_.scala.None, transformFnAssetsUri: _root_.scala.Predef.String = "", + quantizedFeatureMetadata: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = _root_.scala.None, unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[EdgeMetadataInfo] { @transient @@ -1181,6 +1184,10 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __size += _root_.com.google.protobuf.CodedOutputStream.computeStringSize(7, __value) } }; + if (quantizedFeatureMetadata.isDefined) { + val __value = quantizedFeatureMetadata.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; __size += unknownFields.serializedSize __size } @@ -1230,6 +1237,12 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r _output__.writeString(7, __v) } }; + quantizedFeatureMetadata.foreach { __v => + val __m = __v + _output__.writeTag(8, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; unknownFields.writeTo(_output__) } def clearFeatureKeys = copy(featureKeys = _root_.scala.Seq.empty) @@ -1247,6 +1260,9 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r def clearFeatureDim: EdgeMetadataInfo = copy(featureDim = _root_.scala.None) def withFeatureDim(__v: _root_.scala.Int): EdgeMetadataInfo = copy(featureDim = Option(__v)) def withTransformFnAssetsUri(__v: _root_.scala.Predef.String): EdgeMetadataInfo = copy(transformFnAssetsUri = __v) + def getQuantizedFeatureMetadata: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = quantizedFeatureMetadata.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.defaultInstance) + def clearQuantizedFeatureMetadata: EdgeMetadataInfo = copy(quantizedFeatureMetadata = _root_.scala.None) + def withQuantizedFeatureMetadata(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata): EdgeMetadataInfo = copy(quantizedFeatureMetadata = Option(__v)) def withUnknownFields(__v: _root_.scalapb.UnknownFieldSet) = copy(unknownFields = __v) def discardUnknownFields = copy(unknownFields = _root_.scalapb.UnknownFieldSet.empty) def getFieldByNumber(__fieldNumber: _root_.scala.Int): _root_.scala.Any = { @@ -1270,6 +1286,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r val __t = transformFnAssetsUri if (__t != "") __t else null } + case 8 => quantizedFeatureMetadata.orNull } } def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { @@ -1282,6 +1299,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r case 5 => _root_.scalapb.descriptors.PString(enumeratedEdgeDataBqTable) case 6 => featureDim.map(_root_.scalapb.descriptors.PInt(_)).getOrElse(_root_.scalapb.descriptors.PEmpty) case 7 => _root_.scalapb.descriptors.PString(transformFnAssetsUri) + case 8 => quantizedFeatureMetadata.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) } } def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) @@ -1299,6 +1317,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r var __enumeratedEdgeDataBqTable: _root_.scala.Predef.String = "" var __featureDim: _root_.scala.Option[_root_.scala.Int] = _root_.scala.None var __transformFnAssetsUri: _root_.scala.Predef.String = "" + var __quantizedFeatureMetadata: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = _root_.scala.None var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null var _done__ = false while (!_done__) { @@ -1319,6 +1338,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __featureDim = Option(_input__.readUInt32()) case 58 => __transformFnAssetsUri = _input__.readStringRequireUtf8() + case 66 => + __quantizedFeatureMetadata = Option(__quantizedFeatureMetadata.fold(_root_.scalapb.LiteParser.readMessage[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata](_input__))(_root_.scalapb.LiteParser.readMessage(_input__, _))) case tag => if (_unknownFields__ == null) { _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() @@ -1334,6 +1355,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedEdgeDataBqTable = __enumeratedEdgeDataBqTable, featureDim = __featureDim, transformFnAssetsUri = __transformFnAssetsUri, + quantizedFeatureMetadata = __quantizedFeatureMetadata, unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() ) } @@ -1347,13 +1369,20 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(4).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), enumeratedEdgeDataBqTable = __fieldsMap.get(scalaDescriptor.findFieldByNumber(5).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), featureDim = __fieldsMap.get(scalaDescriptor.findFieldByNumber(6).get).flatMap(_.as[_root_.scala.Option[_root_.scala.Int]]), - transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(7).get).map(_.as[_root_.scala.Predef.String]).getOrElse("") + transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(7).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), + quantizedFeatureMetadata = __fieldsMap.get(scalaDescriptor.findFieldByNumber(8).get).flatMap(_.as[_root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata]]) ) case _ => throw new RuntimeException("Expected PMessage") } def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(4) def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(4) - def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { + var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null + (__number: @_root_.scala.unchecked) match { + case 8 => __out = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata + } + __out + } lazy val nestedMessagesCompanions: Seq[_root_.scalapb.GeneratedMessageCompanion[_ <: _root_.scalapb.GeneratedMessage]] = Seq.empty def enumCompanionForFieldNumber(__fieldNumber: _root_.scala.Int): _root_.scalapb.GeneratedEnumCompanion[_] = throw new MatchError(__fieldNumber) lazy val defaultInstance = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo( @@ -1363,7 +1392,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri = "", enumeratedEdgeDataBqTable = "", featureDim = _root_.scala.None, - transformFnAssetsUri = "" + transformFnAssetsUri = "", + quantizedFeatureMetadata = _root_.scala.None ) implicit class EdgeMetadataInfoLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo](_l) { def featureKeys: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Seq[_root_.scala.Predef.String]] = field(_.featureKeys)((c_, f_) => c_.copy(featureKeys = f_)) @@ -1374,6 +1404,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r def featureDim: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Int] = field(_.getFeatureDim)((c_, f_) => c_.copy(featureDim = Option(f_))) def optionalFeatureDim: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Option[_root_.scala.Int]] = field(_.featureDim)((c_, f_) => c_.copy(featureDim = f_)) def transformFnAssetsUri: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Predef.String] = field(_.transformFnAssetsUri)((c_, f_) => c_.copy(transformFnAssetsUri = f_)) + def quantizedFeatureMetadata: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = field(_.getQuantizedFeatureMetadata)((c_, f_) => c_.copy(quantizedFeatureMetadata = Option(f_))) + def optionalQuantizedFeatureMetadata: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata]] = field(_.quantizedFeatureMetadata)((c_, f_) => c_.copy(quantizedFeatureMetadata = f_)) } final val FEATURE_KEYS_FIELD_NUMBER = 1 final val LABEL_KEYS_FIELD_NUMBER = 2 @@ -1382,6 +1414,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r final val ENUMERATED_EDGE_DATA_BQ_TABLE_FIELD_NUMBER = 5 final val FEATURE_DIM_FIELD_NUMBER = 6 final val TRANSFORM_FN_ASSETS_URI_FIELD_NUMBER = 7 + final val QUANTIZED_FEATURE_METADATA_FIELD_NUMBER = 8 def of( featureKeys: _root_.scala.Seq[_root_.scala.Predef.String], labelKeys: _root_.scala.Seq[_root_.scala.Predef.String], @@ -1389,7 +1422,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri: _root_.scala.Predef.String, enumeratedEdgeDataBqTable: _root_.scala.Predef.String, featureDim: _root_.scala.Option[_root_.scala.Int], - transformFnAssetsUri: _root_.scala.Predef.String + transformFnAssetsUri: _root_.scala.Predef.String, + quantizedFeatureMetadata: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo( featureKeys, labelKeys, @@ -1397,7 +1431,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri, enumeratedEdgeDataBqTable, featureDim, - transformFnAssetsUri + transformFnAssetsUri, + quantizedFeatureMetadata ) // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfo]) } diff --git a/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala b/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala index ad80de0ad..998cadd75 100644 --- a/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala +++ b/scala/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala @@ -14,7 +14,7 @@ object PreprocessedMetadataProto extends _root_.scalapb.GeneratedFileObject { private lazy val ProtoBytes: _root_.scala.Array[Byte] = scalapb.Encoding.fromBase64(scala.collection.immutable.Seq( """CjJzbmFwY2hhdC9yZXNlYXJjaC9nYm1sL3ByZXByb2Nlc3NlZF9tZXRhZGF0YS5wcm90bxIWc25hcGNoYXQucmVzZWFyY2guZ - 2JtbCL9GgoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY + 2JtbCKlHAoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY XRhGAEgAygLMlkuc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5Db25kZW5zZWROb2RlVHlwZVRvU HJlcHJvY2Vzc2VkTWV0YWRhdGFFbnRyeUIs4j8pEidjb25kZW5zZWROb2RlVHlwZVRvUHJlcHJvY2Vzc2VkTWV0YWRhdGFSJ2Nvb mRlbnNlZE5vZGVUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZGF0YRLmAQosY29uZGVuc2VkX2VkZ2VfdHlwZV90b19wcmVwcm9jZXNzZ @@ -41,26 +41,28 @@ object PreprocessedMetadataProto extends _root_.scalapb.GeneratedFileObject { hR0cmFuc2Zvcm1GbkFzc2V0c1VyaVIUdHJhbnNmb3JtRm5Bc3NldHNVcmkSpQEKGnF1YW50aXplZF9mZWF0dXJlX21ldGFkYXRhG AogASgLMkguc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5GZWF0dXJlUXVhbnRpemF0aW9uTWV0Y WRhdGFCHeI/GhIYcXVhbnRpemVkRmVhdHVyZU1ldGFkYXRhUhhxdWFudGl6ZWRGZWF0dXJlTWV0YWRhdGFCDgoMX2ZlYXR1cmVfZ - GltGugDChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZUtle + GltGpAFChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZUtle XMSLQoKbGFiZWxfa2V5cxgCIAMoCUIO4j8LEglsYWJlbEtleXNSCWxhYmVsS2V5cxJGChN0ZnJlY29yZF91cmlfcHJlZml4GAMgA SgJQhbiPxMSEXRmcmVjb3JkVXJpUHJlZml4UhF0ZnJlY29yZFVyaVByZWZpeBItCgpzY2hlbWFfdXJpGAQgASgJQg7iPwsSCXNja GVtYVVyaVIJc2NoZW1hVXJpEmAKHWVudW1lcmF0ZWRfZWRnZV9kYXRhX2JxX3RhYmxlGAUgASgJQh7iPxsSGWVudW1lcmF0ZWRFZ GdlRGF0YUJxVGFibGVSGWVudW1lcmF0ZWRFZGdlRGF0YUJxVGFibGUSNQoLZmVhdHVyZV9kaW0YBiABKA1CD+I/DBIKZmVhdHVyZ URpbUgAUgpmZWF0dXJlRGltiAEBElAKF3RyYW5zZm9ybV9mbl9hc3NldHNfdXJpGAcgASgJQhniPxYSFHRyYW5zZm9ybUZuQXNzZ - XRzVXJpUhR0cmFuc2Zvcm1GbkFzc2V0c1VyaUIOCgxfZmVhdHVyZV9kaW0awgQKEkVkZ2VNZXRhZGF0YU91dHB1dBI4Cg9zcmNfb - m9kZV9pZF9rZXkYASABKAlCEeI/DhIMc3JjTm9kZUlkS2V5UgxzcmNOb2RlSWRLZXkSOAoPZHN0X25vZGVfaWRfa2V5GAIgASgJQ - hHiPw4SDGRzdE5vZGVJZEtleVIMZHN0Tm9kZUlkS2V5EnYKDm1haW5fZWRnZV9pbmZvGAMgASgLMj0uc25hcGNoYXQucmVzZWFyY - 2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFJbmZvQhHiPw4SDG1haW5FZGdlSW5mb1IMbWFpbkVkZ2VJb - mZvEocBChJwb3NpdGl2ZV9lZGdlX2luZm8YBCABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkY - XRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQcG9zaXRpdmVFZGdlSW5mb0gAUhBwb3NpdGl2ZUVkZ2VJbmZviAEBEocBChJuZWdhd - Gl2ZV9lZGdlX2luZm8YBSABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkVkZ2VNZXRhZ - GF0YUluZm9CFeI/EhIQbmVnYXRpdmVFZGdlSW5mb0gBUhBuZWdhdGl2ZUVkZ2VJbmZviAEBQhUKE19wb3NpdGl2ZV9lZGdlX2luZ - m9CFQoTX25lZ2F0aXZlX2VkZ2VfaW5mbxqxAQosQ29uZGVuc2VkTm9kZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSG - goDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZ - XNzZWRNZXRhZGF0YS5Ob2RlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4ARqxAQosQ29uZGVuc2VkRWRnZVR5c - GVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc - 25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSB - XZhbHVlOgI4AWIGcHJvdG8z""" + XRzVXJpUhR0cmFuc2Zvcm1GbkFzc2V0c1VyaRKlAQoacXVhbnRpemVkX2ZlYXR1cmVfbWV0YWRhdGEYCCABKAsySC5zbmFwY2hhd + C5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkZlYXR1cmVRdWFudGl6YXRpb25NZXRhZGF0YUId4j8aEhhxdWFud + Gl6ZWRGZWF0dXJlTWV0YWRhdGFSGHF1YW50aXplZEZlYXR1cmVNZXRhZGF0YUIOCgxfZmVhdHVyZV9kaW0awgQKEkVkZ2VNZXRhZ + GF0YU91dHB1dBI4Cg9zcmNfbm9kZV9pZF9rZXkYASABKAlCEeI/DhIMc3JjTm9kZUlkS2V5UgxzcmNOb2RlSWRLZXkSOAoPZHN0X + 25vZGVfaWRfa2V5GAIgASgJQhHiPw4SDGRzdE5vZGVJZEtleVIMZHN0Tm9kZUlkS2V5EnYKDm1haW5fZWRnZV9pbmZvGAMgASgLM + j0uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFJbmZvQhHiPw4SDG1haW5FZ + GdlSW5mb1IMbWFpbkVkZ2VJbmZvEocBChJwb3NpdGl2ZV9lZGdlX2luZm8YBCABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sL + lByZXByb2Nlc3NlZE1ldGFkYXRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQcG9zaXRpdmVFZGdlSW5mb0gAUhBwb3NpdGl2ZUVkZ + 2VJbmZviAEBEocBChJuZWdhdGl2ZV9lZGdlX2luZm8YBSABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZ + E1ldGFkYXRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQbmVnYXRpdmVFZGdlSW5mb0gBUhBuZWdhdGl2ZUVkZ2VJbmZviAEBQhUKE + 19wb3NpdGl2ZV9lZGdlX2luZm9CFQoTX25lZ2F0aXZlX2VkZ2VfaW5mbxqxAQosQ29uZGVuc2VkTm9kZVR5cGVUb1ByZXByb2Nlc + 3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZ + WFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5Ob2RlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4ARqxA + QosQ29uZGVuc2VkRWRnZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5E + mEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFPd + XRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4AWIGcHJvdG8z""" ).mkString) lazy val scalaDescriptor: _root_.scalapb.descriptors.FileDescriptor = { val scalaProto = com.google.protobuf.descriptor.FileDescriptorProto.parseFrom(ProtoBytes) diff --git a/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala b/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala index 7a9012ffa..9ddc933a7 100644 --- a/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala +++ b/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadata.scala @@ -1125,6 +1125,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r * Feature dimension after preprocessing * @param transformFnAssetsUri * Contains categorical feature vocabularies + * @param quantizedFeatureMetadata + * Optional quantized main-edge feature metadata. */ @SerialVersionUID(0L) final case class EdgeMetadataInfo( @@ -1135,6 +1137,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedEdgeDataBqTable: _root_.scala.Predef.String = "", featureDim: _root_.scala.Option[_root_.scala.Int] = _root_.scala.None, transformFnAssetsUri: _root_.scala.Predef.String = "", + quantizedFeatureMetadata: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = _root_.scala.None, unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[EdgeMetadataInfo] { @transient @@ -1181,6 +1184,10 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __size += _root_.com.google.protobuf.CodedOutputStream.computeStringSize(7, __value) } }; + if (quantizedFeatureMetadata.isDefined) { + val __value = quantizedFeatureMetadata.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; __size += unknownFields.serializedSize __size } @@ -1230,6 +1237,12 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r _output__.writeString(7, __v) } }; + quantizedFeatureMetadata.foreach { __v => + val __m = __v + _output__.writeTag(8, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; unknownFields.writeTo(_output__) } def clearFeatureKeys = copy(featureKeys = _root_.scala.Seq.empty) @@ -1247,6 +1260,9 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r def clearFeatureDim: EdgeMetadataInfo = copy(featureDim = _root_.scala.None) def withFeatureDim(__v: _root_.scala.Int): EdgeMetadataInfo = copy(featureDim = Option(__v)) def withTransformFnAssetsUri(__v: _root_.scala.Predef.String): EdgeMetadataInfo = copy(transformFnAssetsUri = __v) + def getQuantizedFeatureMetadata: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = quantizedFeatureMetadata.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.defaultInstance) + def clearQuantizedFeatureMetadata: EdgeMetadataInfo = copy(quantizedFeatureMetadata = _root_.scala.None) + def withQuantizedFeatureMetadata(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata): EdgeMetadataInfo = copy(quantizedFeatureMetadata = Option(__v)) def withUnknownFields(__v: _root_.scalapb.UnknownFieldSet) = copy(unknownFields = __v) def discardUnknownFields = copy(unknownFields = _root_.scalapb.UnknownFieldSet.empty) def getFieldByNumber(__fieldNumber: _root_.scala.Int): _root_.scala.Any = { @@ -1270,6 +1286,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r val __t = transformFnAssetsUri if (__t != "") __t else null } + case 8 => quantizedFeatureMetadata.orNull } } def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { @@ -1282,6 +1299,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r case 5 => _root_.scalapb.descriptors.PString(enumeratedEdgeDataBqTable) case 6 => featureDim.map(_root_.scalapb.descriptors.PInt(_)).getOrElse(_root_.scalapb.descriptors.PEmpty) case 7 => _root_.scalapb.descriptors.PString(transformFnAssetsUri) + case 8 => quantizedFeatureMetadata.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) } } def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) @@ -1299,6 +1317,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r var __enumeratedEdgeDataBqTable: _root_.scala.Predef.String = "" var __featureDim: _root_.scala.Option[_root_.scala.Int] = _root_.scala.None var __transformFnAssetsUri: _root_.scala.Predef.String = "" + var __quantizedFeatureMetadata: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = _root_.scala.None var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null var _done__ = false while (!_done__) { @@ -1319,6 +1338,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __featureDim = Option(_input__.readUInt32()) case 58 => __transformFnAssetsUri = _input__.readStringRequireUtf8() + case 66 => + __quantizedFeatureMetadata = Option(__quantizedFeatureMetadata.fold(_root_.scalapb.LiteParser.readMessage[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata](_input__))(_root_.scalapb.LiteParser.readMessage(_input__, _))) case tag => if (_unknownFields__ == null) { _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() @@ -1334,6 +1355,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedEdgeDataBqTable = __enumeratedEdgeDataBqTable, featureDim = __featureDim, transformFnAssetsUri = __transformFnAssetsUri, + quantizedFeatureMetadata = __quantizedFeatureMetadata, unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() ) } @@ -1347,13 +1369,20 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(4).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), enumeratedEdgeDataBqTable = __fieldsMap.get(scalaDescriptor.findFieldByNumber(5).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), featureDim = __fieldsMap.get(scalaDescriptor.findFieldByNumber(6).get).flatMap(_.as[_root_.scala.Option[_root_.scala.Int]]), - transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(7).get).map(_.as[_root_.scala.Predef.String]).getOrElse("") + transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(7).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), + quantizedFeatureMetadata = __fieldsMap.get(scalaDescriptor.findFieldByNumber(8).get).flatMap(_.as[_root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata]]) ) case _ => throw new RuntimeException("Expected PMessage") } def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(4) def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(4) - def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { + var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null + (__number: @_root_.scala.unchecked) match { + case 8 => __out = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata + } + __out + } lazy val nestedMessagesCompanions: Seq[_root_.scalapb.GeneratedMessageCompanion[_ <: _root_.scalapb.GeneratedMessage]] = Seq.empty def enumCompanionForFieldNumber(__fieldNumber: _root_.scala.Int): _root_.scalapb.GeneratedEnumCompanion[_] = throw new MatchError(__fieldNumber) lazy val defaultInstance = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo( @@ -1363,7 +1392,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri = "", enumeratedEdgeDataBqTable = "", featureDim = _root_.scala.None, - transformFnAssetsUri = "" + transformFnAssetsUri = "", + quantizedFeatureMetadata = _root_.scala.None ) implicit class EdgeMetadataInfoLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo](_l) { def featureKeys: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Seq[_root_.scala.Predef.String]] = field(_.featureKeys)((c_, f_) => c_.copy(featureKeys = f_)) @@ -1374,6 +1404,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r def featureDim: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Int] = field(_.getFeatureDim)((c_, f_) => c_.copy(featureDim = Option(f_))) def optionalFeatureDim: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Option[_root_.scala.Int]] = field(_.featureDim)((c_, f_) => c_.copy(featureDim = f_)) def transformFnAssetsUri: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Predef.String] = field(_.transformFnAssetsUri)((c_, f_) => c_.copy(transformFnAssetsUri = f_)) + def quantizedFeatureMetadata: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = field(_.getQuantizedFeatureMetadata)((c_, f_) => c_.copy(quantizedFeatureMetadata = Option(f_))) + def optionalQuantizedFeatureMetadata: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata]] = field(_.quantizedFeatureMetadata)((c_, f_) => c_.copy(quantizedFeatureMetadata = f_)) } final val FEATURE_KEYS_FIELD_NUMBER = 1 final val LABEL_KEYS_FIELD_NUMBER = 2 @@ -1382,6 +1414,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r final val ENUMERATED_EDGE_DATA_BQ_TABLE_FIELD_NUMBER = 5 final val FEATURE_DIM_FIELD_NUMBER = 6 final val TRANSFORM_FN_ASSETS_URI_FIELD_NUMBER = 7 + final val QUANTIZED_FEATURE_METADATA_FIELD_NUMBER = 8 def of( featureKeys: _root_.scala.Seq[_root_.scala.Predef.String], labelKeys: _root_.scala.Seq[_root_.scala.Predef.String], @@ -1389,7 +1422,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri: _root_.scala.Predef.String, enumeratedEdgeDataBqTable: _root_.scala.Predef.String, featureDim: _root_.scala.Option[_root_.scala.Int], - transformFnAssetsUri: _root_.scala.Predef.String + transformFnAssetsUri: _root_.scala.Predef.String, + quantizedFeatureMetadata: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo( featureKeys, labelKeys, @@ -1397,7 +1431,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r schemaUri, enumeratedEdgeDataBqTable, featureDim, - transformFnAssetsUri + transformFnAssetsUri, + quantizedFeatureMetadata ) // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfo]) } diff --git a/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala b/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala index ad80de0ad..998cadd75 100644 --- a/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala +++ b/scala_spark35/common/src/main/scala/snapchat/research/gbml/preprocessed_metadata/PreprocessedMetadataProto.scala @@ -14,7 +14,7 @@ object PreprocessedMetadataProto extends _root_.scalapb.GeneratedFileObject { private lazy val ProtoBytes: _root_.scala.Array[Byte] = scalapb.Encoding.fromBase64(scala.collection.immutable.Seq( """CjJzbmFwY2hhdC9yZXNlYXJjaC9nYm1sL3ByZXByb2Nlc3NlZF9tZXRhZGF0YS5wcm90bxIWc25hcGNoYXQucmVzZWFyY2guZ - 2JtbCL9GgoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY + 2JtbCKlHAoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY XRhGAEgAygLMlkuc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5Db25kZW5zZWROb2RlVHlwZVRvU HJlcHJvY2Vzc2VkTWV0YWRhdGFFbnRyeUIs4j8pEidjb25kZW5zZWROb2RlVHlwZVRvUHJlcHJvY2Vzc2VkTWV0YWRhdGFSJ2Nvb mRlbnNlZE5vZGVUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZGF0YRLmAQosY29uZGVuc2VkX2VkZ2VfdHlwZV90b19wcmVwcm9jZXNzZ @@ -41,26 +41,28 @@ object PreprocessedMetadataProto extends _root_.scalapb.GeneratedFileObject { hR0cmFuc2Zvcm1GbkFzc2V0c1VyaVIUdHJhbnNmb3JtRm5Bc3NldHNVcmkSpQEKGnF1YW50aXplZF9mZWF0dXJlX21ldGFkYXRhG AogASgLMkguc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5GZWF0dXJlUXVhbnRpemF0aW9uTWV0Y WRhdGFCHeI/GhIYcXVhbnRpemVkRmVhdHVyZU1ldGFkYXRhUhhxdWFudGl6ZWRGZWF0dXJlTWV0YWRhdGFCDgoMX2ZlYXR1cmVfZ - GltGugDChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZUtle + GltGpAFChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZUtle XMSLQoKbGFiZWxfa2V5cxgCIAMoCUIO4j8LEglsYWJlbEtleXNSCWxhYmVsS2V5cxJGChN0ZnJlY29yZF91cmlfcHJlZml4GAMgA SgJQhbiPxMSEXRmcmVjb3JkVXJpUHJlZml4UhF0ZnJlY29yZFVyaVByZWZpeBItCgpzY2hlbWFfdXJpGAQgASgJQg7iPwsSCXNja GVtYVVyaVIJc2NoZW1hVXJpEmAKHWVudW1lcmF0ZWRfZWRnZV9kYXRhX2JxX3RhYmxlGAUgASgJQh7iPxsSGWVudW1lcmF0ZWRFZ GdlRGF0YUJxVGFibGVSGWVudW1lcmF0ZWRFZGdlRGF0YUJxVGFibGUSNQoLZmVhdHVyZV9kaW0YBiABKA1CD+I/DBIKZmVhdHVyZ URpbUgAUgpmZWF0dXJlRGltiAEBElAKF3RyYW5zZm9ybV9mbl9hc3NldHNfdXJpGAcgASgJQhniPxYSFHRyYW5zZm9ybUZuQXNzZ - XRzVXJpUhR0cmFuc2Zvcm1GbkFzc2V0c1VyaUIOCgxfZmVhdHVyZV9kaW0awgQKEkVkZ2VNZXRhZGF0YU91dHB1dBI4Cg9zcmNfb - m9kZV9pZF9rZXkYASABKAlCEeI/DhIMc3JjTm9kZUlkS2V5UgxzcmNOb2RlSWRLZXkSOAoPZHN0X25vZGVfaWRfa2V5GAIgASgJQ - hHiPw4SDGRzdE5vZGVJZEtleVIMZHN0Tm9kZUlkS2V5EnYKDm1haW5fZWRnZV9pbmZvGAMgASgLMj0uc25hcGNoYXQucmVzZWFyY - 2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFJbmZvQhHiPw4SDG1haW5FZGdlSW5mb1IMbWFpbkVkZ2VJb - mZvEocBChJwb3NpdGl2ZV9lZGdlX2luZm8YBCABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkY - XRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQcG9zaXRpdmVFZGdlSW5mb0gAUhBwb3NpdGl2ZUVkZ2VJbmZviAEBEocBChJuZWdhd - Gl2ZV9lZGdlX2luZm8YBSABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkVkZ2VNZXRhZ - GF0YUluZm9CFeI/EhIQbmVnYXRpdmVFZGdlSW5mb0gBUhBuZWdhdGl2ZUVkZ2VJbmZviAEBQhUKE19wb3NpdGl2ZV9lZGdlX2luZ - m9CFQoTX25lZ2F0aXZlX2VkZ2VfaW5mbxqxAQosQ29uZGVuc2VkTm9kZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSG - goDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZ - XNzZWRNZXRhZGF0YS5Ob2RlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4ARqxAQosQ29uZGVuc2VkRWRnZVR5c - GVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc - 25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSB - XZhbHVlOgI4AWIGcHJvdG8z""" + XRzVXJpUhR0cmFuc2Zvcm1GbkFzc2V0c1VyaRKlAQoacXVhbnRpemVkX2ZlYXR1cmVfbWV0YWRhdGEYCCABKAsySC5zbmFwY2hhd + C5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkZlYXR1cmVRdWFudGl6YXRpb25NZXRhZGF0YUId4j8aEhhxdWFud + Gl6ZWRGZWF0dXJlTWV0YWRhdGFSGHF1YW50aXplZEZlYXR1cmVNZXRhZGF0YUIOCgxfZmVhdHVyZV9kaW0awgQKEkVkZ2VNZXRhZ + GF0YU91dHB1dBI4Cg9zcmNfbm9kZV9pZF9rZXkYASABKAlCEeI/DhIMc3JjTm9kZUlkS2V5UgxzcmNOb2RlSWRLZXkSOAoPZHN0X + 25vZGVfaWRfa2V5GAIgASgJQhHiPw4SDGRzdE5vZGVJZEtleVIMZHN0Tm9kZUlkS2V5EnYKDm1haW5fZWRnZV9pbmZvGAMgASgLM + j0uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFJbmZvQhHiPw4SDG1haW5FZ + GdlSW5mb1IMbWFpbkVkZ2VJbmZvEocBChJwb3NpdGl2ZV9lZGdlX2luZm8YBCABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sL + lByZXByb2Nlc3NlZE1ldGFkYXRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQcG9zaXRpdmVFZGdlSW5mb0gAUhBwb3NpdGl2ZUVkZ + 2VJbmZviAEBEocBChJuZWdhdGl2ZV9lZGdlX2luZm8YBSABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZ + E1ldGFkYXRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQbmVnYXRpdmVFZGdlSW5mb0gBUhBuZWdhdGl2ZUVkZ2VJbmZviAEBQhUKE + 19wb3NpdGl2ZV9lZGdlX2luZm9CFQoTX25lZ2F0aXZlX2VkZ2VfaW5mbxqxAQosQ29uZGVuc2VkTm9kZVR5cGVUb1ByZXByb2Nlc + 3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZ + WFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5Ob2RlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4ARqxA + QosQ29uZGVuc2VkRWRnZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5E + mEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFPd + XRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4AWIGcHJvdG8z""" ).mkString) lazy val scalaDescriptor: _root_.scalapb.descriptors.FileDescriptor = { val scalaProto = com.google.protobuf.descriptor.FileDescriptorProto.parseFrom(ProtoBytes) diff --git a/snapchat/research/gbml/preprocessed_metadata_pb2.py b/snapchat/research/gbml/preprocessed_metadata_pb2.py index 2fac76d2f..5a1c46a5d 100644 --- a/snapchat/research/gbml/preprocessed_metadata_pb2.py +++ b/snapchat/research/gbml/preprocessed_metadata_pb2.py @@ -14,7 +14,7 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n2snapchat/research/gbml/preprocessed_metadata.proto\x12\x16snapchat.research.gbml\"\x9c\x10\n\x14PreprocessedMetadata\x12\x8f\x01\n,condensed_node_type_to_preprocessed_metadata\x18\x01 \x03(\x0b\x32Y.snapchat.research.gbml.PreprocessedMetadata.CondensedNodeTypeToPreprocessedMetadataEntry\x12\x8f\x01\n,condensed_edge_type_to_preprocessed_metadata\x18\x02 \x03(\x0b\x32Y.snapchat.research.gbml.PreprocessedMetadata.CondensedEdgeTypeToPreprocessedMetadataEntry\x1aM\n\x19MultiBitQuantizationState\x12\x10\n\x08\x63lip_min\x18\x01 \x01(\x02\x12\x10\n\x08\x63lip_max\x18\x02 \x01(\x02\x12\x0c\n\x04\x62its\x18\x03 \x01(\r\x1a@\n\x1aSingleBitQuantizationState\x12\x10\n\x08neg_mean\x18\x01 \x01(\x02\x12\x10\n\x08pos_mean\x18\x02 \x01(\x02\x1a\xad\x02\n\x1b\x46\x65\x61tureQuantizationMetadata\x12\x1a\n\x12packed_feature_key\x18\x01 \x01(\t\x12!\n\x19quantized_feature_indices\x18\x02 \x03(\r\x12\x61\n\x0fmulti_bit_state\x18\x04 \x01(\x0b\x32\x46.snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationStateH\x00\x12\x63\n\x10single_bit_state\x18\x05 \x01(\x0b\x32G.snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationStateH\x00\x42\x07\n\x05state\x1a\x8a\x03\n\x12NodeMetadataOutput\x12\x13\n\x0bnode_id_key\x18\x01 \x01(\t\x12\x14\n\x0c\x66\x65\x61ture_keys\x18\x02 \x03(\t\x12\x12\n\nlabel_keys\x18\x03 \x03(\t\x12\x1b\n\x13tfrecord_uri_prefix\x18\x04 \x01(\t\x12\x12\n\nschema_uri\x18\x05 \x01(\t\x12$\n\x1c\x65numerated_node_ids_bq_table\x18\x06 \x01(\t\x12%\n\x1d\x65numerated_node_data_bq_table\x18\x07 \x01(\t\x12\x18\n\x0b\x66\x65\x61ture_dim\x18\x08 \x01(\rH\x00\x88\x01\x01\x12\x1f\n\x17transform_fn_assets_uri\x18\t \x01(\t\x12l\n\x1aquantized_feature_metadata\x18\n \x01(\x0b\x32H.snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadataB\x0e\n\x0c_feature_dim\x1a\xdf\x01\n\x10\x45\x64geMetadataInfo\x12\x14\n\x0c\x66\x65\x61ture_keys\x18\x01 \x03(\t\x12\x12\n\nlabel_keys\x18\x02 \x03(\t\x12\x1b\n\x13tfrecord_uri_prefix\x18\x03 \x01(\t\x12\x12\n\nschema_uri\x18\x04 \x01(\t\x12%\n\x1d\x65numerated_edge_data_bq_table\x18\x05 \x01(\t\x12\x18\n\x0b\x66\x65\x61ture_dim\x18\x06 \x01(\rH\x00\x88\x01\x01\x12\x1f\n\x17transform_fn_assets_uri\x18\x07 \x01(\tB\x0e\n\x0c_feature_dim\x1a\x8b\x03\n\x12\x45\x64geMetadataOutput\x12\x17\n\x0fsrc_node_id_key\x18\x01 \x01(\t\x12\x17\n\x0f\x64st_node_id_key\x18\x02 \x01(\t\x12U\n\x0emain_edge_info\x18\x03 \x01(\x0b\x32=.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfo\x12^\n\x12positive_edge_info\x18\x04 \x01(\x0b\x32=.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfoH\x00\x88\x01\x01\x12^\n\x12negative_edge_info\x18\x05 \x01(\x0b\x32=.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfoH\x01\x88\x01\x01\x42\x15\n\x13_positive_edge_infoB\x15\n\x13_negative_edge_info\x1a\x8f\x01\n,CondensedNodeTypeToPreprocessedMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\r\x12N\n\x05value\x18\x02 \x01(\x0b\x32?.snapchat.research.gbml.PreprocessedMetadata.NodeMetadataOutput:\x02\x38\x01\x1a\x8f\x01\n,CondensedEdgeTypeToPreprocessedMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\r\x12N\n\x05value\x18\x02 \x01(\x0b\x32?.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataOutput:\x02\x38\x01\x62\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n2snapchat/research/gbml/preprocessed_metadata.proto\x12\x16snapchat.research.gbml\"\x8a\x11\n\x14PreprocessedMetadata\x12\x8f\x01\n,condensed_node_type_to_preprocessed_metadata\x18\x01 \x03(\x0b\x32Y.snapchat.research.gbml.PreprocessedMetadata.CondensedNodeTypeToPreprocessedMetadataEntry\x12\x8f\x01\n,condensed_edge_type_to_preprocessed_metadata\x18\x02 \x03(\x0b\x32Y.snapchat.research.gbml.PreprocessedMetadata.CondensedEdgeTypeToPreprocessedMetadataEntry\x1aM\n\x19MultiBitQuantizationState\x12\x10\n\x08\x63lip_min\x18\x01 \x01(\x02\x12\x10\n\x08\x63lip_max\x18\x02 \x01(\x02\x12\x0c\n\x04\x62its\x18\x03 \x01(\r\x1a@\n\x1aSingleBitQuantizationState\x12\x10\n\x08neg_mean\x18\x01 \x01(\x02\x12\x10\n\x08pos_mean\x18\x02 \x01(\x02\x1a\xad\x02\n\x1b\x46\x65\x61tureQuantizationMetadata\x12\x1a\n\x12packed_feature_key\x18\x01 \x01(\t\x12!\n\x19quantized_feature_indices\x18\x02 \x03(\r\x12\x61\n\x0fmulti_bit_state\x18\x04 \x01(\x0b\x32\x46.snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationStateH\x00\x12\x63\n\x10single_bit_state\x18\x05 \x01(\x0b\x32G.snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationStateH\x00\x42\x07\n\x05state\x1a\x8a\x03\n\x12NodeMetadataOutput\x12\x13\n\x0bnode_id_key\x18\x01 \x01(\t\x12\x14\n\x0c\x66\x65\x61ture_keys\x18\x02 \x03(\t\x12\x12\n\nlabel_keys\x18\x03 \x03(\t\x12\x1b\n\x13tfrecord_uri_prefix\x18\x04 \x01(\t\x12\x12\n\nschema_uri\x18\x05 \x01(\t\x12$\n\x1c\x65numerated_node_ids_bq_table\x18\x06 \x01(\t\x12%\n\x1d\x65numerated_node_data_bq_table\x18\x07 \x01(\t\x12\x18\n\x0b\x66\x65\x61ture_dim\x18\x08 \x01(\rH\x00\x88\x01\x01\x12\x1f\n\x17transform_fn_assets_uri\x18\t \x01(\t\x12l\n\x1aquantized_feature_metadata\x18\n \x01(\x0b\x32H.snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadataB\x0e\n\x0c_feature_dim\x1a\xcd\x02\n\x10\x45\x64geMetadataInfo\x12\x14\n\x0c\x66\x65\x61ture_keys\x18\x01 \x03(\t\x12\x12\n\nlabel_keys\x18\x02 \x03(\t\x12\x1b\n\x13tfrecord_uri_prefix\x18\x03 \x01(\t\x12\x12\n\nschema_uri\x18\x04 \x01(\t\x12%\n\x1d\x65numerated_edge_data_bq_table\x18\x05 \x01(\t\x12\x18\n\x0b\x66\x65\x61ture_dim\x18\x06 \x01(\rH\x00\x88\x01\x01\x12\x1f\n\x17transform_fn_assets_uri\x18\x07 \x01(\t\x12l\n\x1aquantized_feature_metadata\x18\x08 \x01(\x0b\x32H.snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadataB\x0e\n\x0c_feature_dim\x1a\x8b\x03\n\x12\x45\x64geMetadataOutput\x12\x17\n\x0fsrc_node_id_key\x18\x01 \x01(\t\x12\x17\n\x0f\x64st_node_id_key\x18\x02 \x01(\t\x12U\n\x0emain_edge_info\x18\x03 \x01(\x0b\x32=.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfo\x12^\n\x12positive_edge_info\x18\x04 \x01(\x0b\x32=.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfoH\x00\x88\x01\x01\x12^\n\x12negative_edge_info\x18\x05 \x01(\x0b\x32=.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataInfoH\x01\x88\x01\x01\x42\x15\n\x13_positive_edge_infoB\x15\n\x13_negative_edge_info\x1a\x8f\x01\n,CondensedNodeTypeToPreprocessedMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\r\x12N\n\x05value\x18\x02 \x01(\x0b\x32?.snapchat.research.gbml.PreprocessedMetadata.NodeMetadataOutput:\x02\x38\x01\x1a\x8f\x01\n,CondensedEdgeTypeToPreprocessedMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\r\x12N\n\x05value\x18\x02 \x01(\x0b\x32?.snapchat.research.gbml.PreprocessedMetadata.EdgeMetadataOutput:\x02\x38\x01\x62\x06proto3') @@ -106,7 +106,7 @@ _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._options = None _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_options = b'8\001' _PREPROCESSEDMETADATA._serialized_start=79 - _PREPROCESSEDMETADATA._serialized_end=2155 + _PREPROCESSEDMETADATA._serialized_end=2265 _PREPROCESSEDMETADATA_MULTIBITQUANTIZATIONSTATE._serialized_start=395 _PREPROCESSEDMETADATA_MULTIBITQUANTIZATIONSTATE._serialized_end=472 _PREPROCESSEDMETADATA_SINGLEBITQUANTIZATIONSTATE._serialized_start=474 @@ -116,11 +116,11 @@ _PREPROCESSEDMETADATA_NODEMETADATAOUTPUT._serialized_start=845 _PREPROCESSEDMETADATA_NODEMETADATAOUTPUT._serialized_end=1239 _PREPROCESSEDMETADATA_EDGEMETADATAINFO._serialized_start=1242 - _PREPROCESSEDMETADATA_EDGEMETADATAINFO._serialized_end=1465 - _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT._serialized_start=1468 - _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT._serialized_end=1863 - _PREPROCESSEDMETADATA_CONDENSEDNODETYPETOPREPROCESSEDMETADATAENTRY._serialized_start=1866 - _PREPROCESSEDMETADATA_CONDENSEDNODETYPETOPREPROCESSEDMETADATAENTRY._serialized_end=2009 - _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_start=2012 - _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_end=2155 + _PREPROCESSEDMETADATA_EDGEMETADATAINFO._serialized_end=1575 + _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT._serialized_start=1578 + _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT._serialized_end=1973 + _PREPROCESSEDMETADATA_CONDENSEDNODETYPETOPREPROCESSEDMETADATAENTRY._serialized_start=1976 + _PREPROCESSEDMETADATA_CONDENSEDNODETYPETOPREPROCESSEDMETADATAENTRY._serialized_end=2119 + _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_start=2122 + _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_end=2265 # @@protoc_insertion_point(module_scope) diff --git a/snapchat/research/gbml/preprocessed_metadata_pb2.pyi b/snapchat/research/gbml/preprocessed_metadata_pb2.pyi index 46b80c7bb..8ae9c79a5 100644 --- a/snapchat/research/gbml/preprocessed_metadata_pb2.pyi +++ b/snapchat/research/gbml/preprocessed_metadata_pb2.pyi @@ -154,6 +154,7 @@ class PreprocessedMetadata(google.protobuf.message.Message): ENUMERATED_EDGE_DATA_BQ_TABLE_FIELD_NUMBER: builtins.int FEATURE_DIM_FIELD_NUMBER: builtins.int TRANSFORM_FN_ASSETS_URI_FIELD_NUMBER: builtins.int + QUANTIZED_FEATURE_METADATA_FIELD_NUMBER: builtins.int @property def feature_keys(self) -> google.protobuf.internal.containers.RepeatedScalarFieldContainer[builtins.str]: """Fields in output TFRecords which reference features.""" @@ -170,6 +171,9 @@ class PreprocessedMetadata(google.protobuf.message.Message): """Feature dimension after preprocessing""" transform_fn_assets_uri: builtins.str """Contains categorical feature vocabularies""" + @property + def quantized_feature_metadata(self) -> global___PreprocessedMetadata.FeatureQuantizationMetadata: + """Optional quantized main-edge feature metadata.""" def __init__( self, *, @@ -180,9 +184,10 @@ class PreprocessedMetadata(google.protobuf.message.Message): enumerated_edge_data_bq_table: builtins.str = ..., feature_dim: builtins.int | None = ..., transform_fn_assets_uri: builtins.str = ..., + quantized_feature_metadata: global___PreprocessedMetadata.FeatureQuantizationMetadata | None = ..., ) -> None: ... - def HasField(self, field_name: typing_extensions.Literal["_feature_dim", b"_feature_dim", "feature_dim", b"feature_dim"]) -> builtins.bool: ... - def ClearField(self, field_name: typing_extensions.Literal["_feature_dim", b"_feature_dim", "enumerated_edge_data_bq_table", b"enumerated_edge_data_bq_table", "feature_dim", b"feature_dim", "feature_keys", b"feature_keys", "label_keys", b"label_keys", "schema_uri", b"schema_uri", "tfrecord_uri_prefix", b"tfrecord_uri_prefix", "transform_fn_assets_uri", b"transform_fn_assets_uri"]) -> None: ... + def HasField(self, field_name: typing_extensions.Literal["_feature_dim", b"_feature_dim", "feature_dim", b"feature_dim", "quantized_feature_metadata", b"quantized_feature_metadata"]) -> builtins.bool: ... + def ClearField(self, field_name: typing_extensions.Literal["_feature_dim", b"_feature_dim", "enumerated_edge_data_bq_table", b"enumerated_edge_data_bq_table", "feature_dim", b"feature_dim", "feature_keys", b"feature_keys", "label_keys", b"label_keys", "quantized_feature_metadata", b"quantized_feature_metadata", "schema_uri", b"schema_uri", "tfrecord_uri_prefix", b"tfrecord_uri_prefix", "transform_fn_assets_uri", b"transform_fn_assets_uri"]) -> None: ... def WhichOneof(self, oneof_group: typing_extensions.Literal["_feature_dim", b"_feature_dim"]) -> typing_extensions.Literal["feature_dim"] | None: ... class EdgeMetadataOutput(google.protobuf.message.Message): From 5a72f639c86279253160ebbe0190bfc50b1fd352 Mon Sep 17 00:00:00 2001 From: jchmura Date: Sat, 15 Aug 2026 19:24:35 +0000 Subject: [PATCH 2/8] Write quantized edge features during preprocessing --- .../data_preprocessor/data_preprocessor.py | 81 ++++++++++++------- .../lib/transform/feature_quantization.py | 51 +++++++++--- .../data_preprocessor/lib/transform/utils.py | 11 ++- gigl/src/data_preprocessor/lib/types.py | 1 + .../feature_quantization_transform_test.py | 36 +++++++++ 5 files changed, 139 insertions(+), 41 deletions(-) diff --git a/gigl/src/data_preprocessor/data_preprocessor.py b/gigl/src/data_preprocessor/data_preprocessor.py index 9d84a8b42..9c9151993 100644 --- a/gigl/src/data_preprocessor/data_preprocessor.py +++ b/gigl/src/data_preprocessor/data_preprocessor.py @@ -74,6 +74,38 @@ logger = Logger() +def _load_feature_quantization_metadata_pb( + metadata_path: str, entity_description: str +) -> preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata: + if not tf.io.gfile.exists(metadata_path): + raise RuntimeError( + f"Quantization metadata was expected for {entity_description}, " + f"but was not produced at {metadata_path}." + ) + logger.info( + f"Loading {entity_description} quantization metadata from {metadata_path}" + ) + with tf.io.gfile.GFile(metadata_path) as metadata_file: + metadata = json.loads(metadata_file.read()) + logger.info(f"Loaded {entity_description} quantization metadata {metadata}") + + quantization_metadata = ( + preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( + packed_feature_key=metadata["packed_feature_key"], + quantized_feature_indices=metadata["quantized_feature_indices"], + ) + ) + bits = metadata["bits"] + if bits == 1: + quantization_metadata.single_bit_state.neg_mean = metadata["neg_mean"] + quantization_metadata.single_bit_state.pos_mean = metadata["pos_mean"] + else: + quantization_metadata.multi_bit_state.bits = bits + quantization_metadata.multi_bit_state.clip_min = metadata["clip_min"] + quantization_metadata.multi_bit_state.clip_max = metadata["clip_max"] + return quantization_metadata + + class PreprocessedMetadataReferences(NamedTuple): node_data: dict[NodeDataReference, TransformedFeaturesInfo] edge_data: dict[EdgeDataReference, TransformedFeaturesInfo] @@ -216,13 +248,18 @@ def __preprocess_single_data_reference( f"Got {type(data_reference)}." ) - if isinstance(preprocessing_spec, NodeDataPreprocessingSpec): + if isinstance( + preprocessing_spec, (NodeDataPreprocessingSpec, EdgeDataPreprocessingSpec) + ): feature_quantization_enabled = ( preprocessing_spec.feature_quantization_spec is not None ) - else: - # TODO(quantization): Support quantization for edge features. - feature_quantization_enabled = False + if ( + isinstance(data_reference, EdgeDataReference) + and feature_quantization_enabled + and data_reference.edge_usage_type != EdgeUsageType.MAIN + ): + raise ValueError("Feature quantization is supported only for main edges.") transformed_features_info = TransformedFeaturesInfo( applied_task_identifier=self.applied_task_identifier, @@ -428,7 +465,7 @@ def _generate_edge_metadata_info_pb( transformed_features_info: TransformedFeaturesInfo, enumerated_edge_metadata: EnumeratorEdgeTypeMetadata, ) -> preprocessed_metadata_pb2.PreprocessedMetadata.EdgeMetadataInfo: - return preprocessed_metadata_pb2.PreprocessedMetadata.EdgeMetadataInfo( + output = preprocessed_metadata_pb2.PreprocessedMetadata.EdgeMetadataInfo( tfrecord_uri_prefix=transformed_features_info.transformed_features_file_prefix.uri, schema_uri=transformed_features_info.transformed_features_schema_path.uri, feature_keys=transformed_features_info.features_outputs, @@ -437,6 +474,13 @@ def _generate_edge_metadata_info_pb( feature_dim=transformed_features_info.feature_dim_output, transform_fn_assets_uri=transformed_features_info.transformed_features_transform_fn_assets_path.uri, ) + if transformed_features_info.feature_quantization_enabled: + quantization_metadata = _load_feature_quantization_metadata_pb( + metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, + entity_description=f"edge type {transformed_features_info.entity_type}", + ) + output.quantized_feature_metadata.CopyFrom(quantization_metadata) + return output def generate_preprocessed_metadata_pb( self, @@ -492,30 +536,10 @@ def generate_preprocessed_metadata_pb( transform_fn_assets_uri=node_transformed_features_info.transformed_features_transform_fn_assets_path.uri, ) if node_transformed_features_info.feature_quantization_enabled: - metadata_path = node_transformed_features_info.feature_quantization_metadata_path.uri - if not tf.io.gfile.exists(metadata_path): - raise RuntimeError( - f"Quantization metadata was expected for node type {node_type}, " - f"but was not produced at {metadata_path}." - ) - logger.info(f"Loading node quantization metadata from {metadata_path}") - with tf.io.gfile.GFile(metadata_path) as f: - metadata = json.loads(f.read()) - logger.info(f"Loaded node quantization metadata {metadata}") - bits = metadata["bits"] - quantized_feature_metadata_pb = preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( - packed_feature_key=metadata["packed_feature_key"], - quantized_feature_indices=metadata["quantized_feature_indices"], + quantized_feature_metadata_pb = _load_feature_quantization_metadata_pb( + metadata_path=node_transformed_features_info.feature_quantization_metadata_path.uri, + entity_description=f"node type {node_type}", ) - if bits == 1: - single_bit_state = quantized_feature_metadata_pb.single_bit_state - single_bit_state.neg_mean = metadata["neg_mean"] - single_bit_state.pos_mean = metadata["pos_mean"] - else: - multi_bit_state = quantized_feature_metadata_pb.multi_bit_state - multi_bit_state.bits = bits - multi_bit_state.clip_min = metadata["clip_min"] - multi_bit_state.clip_max = metadata["clip_max"] node_metadata_output_pb.quantized_feature_metadata.CopyFrom( quantized_feature_metadata_pb ) @@ -782,6 +806,7 @@ def inner() -> FeatureSpecDict: pretrained_tft_model_uri=input_edge_preprocessing_spec.pretrained_tft_model_uri, features_outputs=input_edge_preprocessing_spec.features_outputs, labels_outputs=input_edge_preprocessing_spec.labels_outputs, + feature_quantization_spec=input_edge_preprocessing_spec.feature_quantization_spec, ) enumerated_edge_refs_to_preprocessing_specs[ enumerated_edge_metadata.enumerated_edge_data_reference diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index db7ecf500..0e1b7b3b2 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -15,7 +15,8 @@ from gigl.src.data_preprocessor.lib.types import FeatureQuantizationSpec logger = Logger() -_NODE_PACKED_FEATURE_KEY: Final[str] = "node_packed_features" +NODE_PACKED_FEATURE_KEY: Final[str] = "node_packed_features" +EDGE_PACKED_FEATURE_KEY: Final[str] = "edge_packed_features" _SignStats: TypeAlias = tuple[float, int, float, int] @@ -25,11 +26,14 @@ def apply_feature_quantization_transform( logical_feature_keys: list[str], quantization_spec: FeatureQuantizationSpec, quantization_metadata_path: str, + packed_feature_key: str, ) -> tuple[beam.PCollection[pa.RecordBatch], DatasetMetadata | beam.pvalue.AsSingleton]: """Quantizes selected feature columns and bit-packs each record's values. - Stores the packed bytes in ``node_packed_features`` and computes global - quantization statistics with Beam. + Stores packed bytes under ``packed_feature_key`` and computes global + quantization statistics with Beam. Node preprocessing uses + ``node_packed_features``; main-edge preprocessing uses + ``edge_packed_features``. Side Effects: Writes the quantization statistics JSON that ``data_preprocessor.py`` @@ -46,12 +50,17 @@ def apply_feature_quantization_transform( logical_feature_keys: Logical feature columns in original feature-vector order. quantization_spec: Feature keys and bit width to quantize. quantization_metadata_path: Destination for the quantization statistics JSON. + packed_feature_key: Reserved physical field used for packed values. Returns: Quantized RecordBatches and eager or deferred physical I/O metadata. That metadata removes quantized feature columns and adds - ``node_packed_features``. It affects serialized-record I/O only; the + ``packed_feature_key``. It affects serialized-record I/O only; the logical model schema remains unchanged. + + Raises: + ValueError: If the reserved packed key already exists, a selected feature + is absent or non-scalar, or feature values cannot be quantized. """ missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) if missing: @@ -73,6 +82,7 @@ def apply_feature_quantization_transform( quantization_spec=quantization_spec, logical_feature_keys=logical_feature_keys, logical_metadata=metadata_for_json, + packed_feature_key=packed_feature_key, ) | "Write quantization stats" >> beam.io.WriteToText( @@ -86,19 +96,24 @@ def apply_feature_quantization_transform( _quantize_record_batch, quantization_spec=quantization_spec, quantization_stats=beam.pvalue.AsSingleton(quantization_stats), + packed_feature_key=packed_feature_key, ) ) if logical_metadata_is_eager: physical_feature_metadata = DatasetMetadata( - _apply_quantization_schema(logical_metadata.schema, quantization_spec) + _apply_quantization_schema( + logical_metadata.schema, quantization_spec, packed_feature_key + ) ) else: physical_feature_metadata = logical_metadata | ( "Apply feature quantization schema" >> beam.Map( lambda metadata, quantization_spec: DatasetMetadata( - _apply_quantization_schema(metadata.schema, quantization_spec) + _apply_quantization_schema( + metadata.schema, quantization_spec, packed_feature_key + ) ), quantization_spec=quantization_spec, ) @@ -139,7 +154,12 @@ def _quantize_record_batch( batch: pa.RecordBatch, quantization_spec: FeatureQuantizationSpec, quantization_stats: dict[str, float], + packed_feature_key: str, ) -> pa.RecordBatch: + if packed_feature_key in batch.schema.names: + raise ValueError( + f"Reserved packed feature key {packed_feature_key} already exists in the logical schema." + ) feature_matrix = _build_feature_matrix(batch, quantization_spec.feature_keys) if quantization_spec.bits == 1: packed = quantize_ndarray(feature_matrix, bits=quantization_spec.bits) @@ -161,7 +181,7 @@ def _quantize_record_batch( arrays.append( pa.array([[row.tobytes()] for row in packed], type=pa.list_(pa.binary())) ) - names.append(_NODE_PACKED_FEATURE_KEY) + names.append(packed_feature_key) return pa.RecordBatch.from_arrays(arrays, names=names) @@ -170,9 +190,10 @@ def _quantization_stats_to_json( quantization_spec: FeatureQuantizationSpec, logical_feature_keys: list[str], logical_metadata: DatasetMetadata, + packed_feature_key: str, ) -> str: metadata = { - "packed_feature_key": _NODE_PACKED_FEATURE_KEY, + "packed_feature_key": packed_feature_key, "quantized_feature_indices": _quantized_feature_indices( logical_metadata, logical_feature_keys, quantization_spec.feature_keys ), @@ -204,9 +225,15 @@ def _quantized_feature_indices( def _apply_quantization_schema( - schema: schema_pb2.Schema, quantization_spec: FeatureQuantizationSpec + schema: schema_pb2.Schema, + quantization_spec: FeatureQuantizationSpec, + packed_feature_key: str, ) -> schema_pb2.Schema: - drop_keys = set(quantization_spec.feature_keys) | {_NODE_PACKED_FEATURE_KEY} + if any(feature.name == packed_feature_key for feature in schema.feature): + raise ValueError( + f"Reserved packed feature key {packed_feature_key} already exists in the logical schema." + ) + drop_keys = set(quantization_spec.feature_keys) quantized_schema = schema_pb2.Schema() quantized_schema.CopyFrom(schema) del quantized_schema.feature[:] @@ -214,14 +241,14 @@ def _apply_quantization_schema( feature for feature in schema.feature if feature.name not in drop_keys ) packed_feature = quantized_schema.feature.add() - packed_feature.name = _NODE_PACKED_FEATURE_KEY + packed_feature.name = packed_feature_key packed_feature.type = schema_pb2.BYTES packed_feature.value_count.min = 1 packed_feature.value_count.max = 1 logger.info( f"Updated transformed schema for feature quantization: dropped " f"{len(quantization_spec.feature_keys)} features and added bytes feature " - f"{_NODE_PACKED_FEATURE_KEY}." + f"{packed_feature_key}." ) return quantized_schema diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index 07bfeaf7c..39a3ccc7b 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -28,6 +28,8 @@ NodeDataReference, ) from gigl.src.data_preprocessor.lib.transform.feature_quantization import ( + EDGE_PACKED_FEATURE_KEY, + NODE_PACKED_FEATURE_KEY, apply_feature_quantization_transform, ) from gigl.src.data_preprocessor.lib.transform.tf_value_encoder import TFValueEncoder @@ -372,9 +374,15 @@ def get_load_data_and_transform_pipeline_component( else analyzed_transform_fn[1].deferred_metadata # type: ignore ) quantization_spec: FeatureQuantizationSpec | None = None - if isinstance(preprocessing_spec, NodeDataPreprocessingSpec): + if isinstance( + preprocessing_spec, (NodeDataPreprocessingSpec, EdgeDataPreprocessingSpec) + ): quantization_spec = preprocessing_spec.feature_quantization_spec if quantization_spec is not None: + if isinstance(preprocessing_spec, EdgeDataPreprocessingSpec): + packed_feature_key = EDGE_PACKED_FEATURE_KEY + else: + packed_feature_key = NODE_PACKED_FEATURE_KEY transformed_features, resolved_transformed_metadata = ( apply_feature_quantization_transform( logical_features=transformed_features, @@ -384,6 +392,7 @@ def get_load_data_and_transform_pipeline_component( ), quantization_spec=quantization_spec, quantization_metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, + packed_feature_key=packed_feature_key, ) ) diff --git a/gigl/src/data_preprocessor/lib/types.py b/gigl/src/data_preprocessor/lib/types.py index 014f7cbc0..8af220e4d 100644 --- a/gigl/src/data_preprocessor/lib/types.py +++ b/gigl/src/data_preprocessor/lib/types.py @@ -120,6 +120,7 @@ class EdgeDataPreprocessingSpec(NamedTuple): pretrained_tft_model_uri: Optional[Uri] = None features_outputs: Optional[list[str]] = None labels_outputs: Optional[list[str]] = None + feature_quantization_spec: Optional[FeatureQuantizationSpec] = None def __repr__(self) -> str: return f"""EdgeDataPreprocessingSpec( diff --git a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py index 00cfc390c..540140be9 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -12,6 +12,7 @@ from tensorflow_transform.tf_metadata.dataset_metadata import DatasetMetadata from gigl.src.data_preprocessor.lib.transform.feature_quantization import ( + NODE_PACKED_FEATURE_KEY, apply_feature_quantization_transform, ) from gigl.src.data_preprocessor.lib.types import FeatureQuantizationSpec @@ -19,6 +20,40 @@ class FeatureQuantizationTransformTest(TestCase): + def test_apply_feature_quantization_transform_rejects_reserved_schema_key( + self, + ) -> None: + logical_metadata = DatasetMetadata.from_feature_spec( + { + "f0": tf.io.FixedLenFeature(shape=[], dtype=tf.float32), + "edge_packed_features": tf.io.FixedLenFeature( + shape=[], dtype=tf.string + ), + } + ) + + with ( + self.assertRaisesRegex(ValueError, "Reserved packed feature key"), + TestPipeline() as pipeline, + ): + apply_feature_quantization_transform( + logical_features=pipeline + | "Create collision input" + >> beam.Create( + [ + pa.RecordBatch.from_arrays( + [pa.array([1.0]), pa.array([b"existing"])], + names=["f0", "edge_packed_features"], + ) + ] + ), + logical_metadata=logical_metadata, + logical_feature_keys=["f0"], + quantization_spec=FeatureQuantizationSpec(feature_keys=["f0"], bits=2), + quantization_metadata_path="unused", + packed_feature_key="edge_packed_features", + ) + @parameterized.expand( [ ( @@ -87,6 +122,7 @@ def test_apply_feature_quantization_transform_writes_metadata( feature_keys=logical_feature_keys, bits=bits ), quantization_metadata_path=metadata_path, + packed_feature_key=NODE_PACKED_FEATURE_KEY, ) ) if use_deferred_metadata: From 381538d876ce98c7627809d730493f5c6fc7ddb8 Mon Sep 17 00:00:00 2001 From: jchmura Date: Sat, 15 Aug 2026 19:25:39 +0000 Subject: [PATCH 3/8] Load and partition quantized edge features --- gigl/common/data/dataloaders.py | 21 +-- gigl/common/data/load_torch_tensors.py | 149 ++++++++++++++++-- gigl/distributed/dataset_factory.py | 5 + gigl/distributed/dist_partitioner.py | 146 ++++++++++++++--- gigl/distributed/dist_range_partitioner.py | 105 ++++++++++-- .../serialized_graph_metadata_translator.py | 24 ++- gigl/types/graph.py | 9 ++ .../run_distributed_partitioner.py | 29 +++- tests/unit/common/data/dataloaders_test.py | 87 ++++++++++ .../distributed_partitioner_test.py | 48 +++++- .../distributed_weighted_sampling_test.py | 12 +- 11 files changed, 570 insertions(+), 65 deletions(-) diff --git a/gigl/common/data/dataloaders.py b/gigl/common/data/dataloaders.py index 824e5225d..67152c898 100644 --- a/gigl/common/data/dataloaders.py +++ b/gigl/common/data/dataloaders.py @@ -398,16 +398,6 @@ def load_as_torch_tensors( feature_spec_dict[entity_key] = tf.io.FixedLenFeature( shape=[], dtype=tf.int64 ) - if ( - packed_feature_key is not None - and packed_feature_key not in feature_spec_dict - ): - logger.info( - f"Injecting packed feature key {packed_feature_key} into feature spec dictionary with value `tf.io.FixedLenFeature(shape=[], dtype=tf.string)`" - ) - feature_spec_dict[packed_feature_key] = tf.io.FixedLenFeature( - shape=[], dtype=tf.string - ) else: id_concat_axis = 1 proccess_id_tensor = lambda t: tf.stack( @@ -433,6 +423,17 @@ def load_as_torch_tensors( shape=[], dtype=tf.int64 ) + if ( + packed_feature_key is not None + and packed_feature_key not in feature_spec_dict + ): + logger.info( + f"Injecting packed feature key {packed_feature_key} into feature spec dictionary with value `tf.io.FixedLenFeature(shape=[], dtype=tf.string)`" + ) + feature_spec_dict[packed_feature_key] = tf.io.FixedLenFeature( + shape=[], dtype=tf.string + ) + uris = self._partition_children_uris( serialized_tf_record_info.tfrecord_uri_prefix, serialized_tf_record_info.tfrecord_uri_pattern, diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 3a1174888..667bc348d 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -1,7 +1,7 @@ import time import traceback -from dataclasses import dataclass -from typing import MutableMapping, Optional, Union +from dataclasses import dataclass, replace +from typing import MutableMapping, Optional, Union, cast import torch import torch.multiprocessing as mp @@ -119,6 +119,134 @@ class SerializedGraphMetadata: node_quantization_metadata: Optional[ Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] ] = None + edge_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] + ] = None + + +def _validate_weight_edge_feature_name( + edge_entity_info: Union[ + SerializedTFRecordInfo, dict[EdgeType, SerializedTFRecordInfo] + ], + weight_edge_feat_name: Optional[Union[str, dict[EdgeType, str]]], +) -> None: + if weight_edge_feat_name is None: + return + + configured_weights: list[tuple[EdgeType, str, SerializedTFRecordInfo]] + if isinstance(edge_entity_info, SerializedTFRecordInfo): + if not isinstance(weight_edge_feat_name, str): + raise ValueError("weight_edge_feat_name must be str for homogeneous graph") + edge_type = DEFAULT_HOMOGENEOUS_EDGE_TYPE + configured_weights = [(edge_type, weight_edge_feat_name, edge_entity_info)] + else: + if isinstance(weight_edge_feat_name, str): + if len(edge_entity_info) != 1: + raise ValueError( + "weight_edge_feat_name must be dict[EdgeType, str] for heterogeneous graph with multiple edge types" + ) + edge_type, serialized_info = next(iter(edge_entity_info.items())) + configured_weights = [(edge_type, weight_edge_feat_name, serialized_info)] + else: + unknown_edge_types = set(weight_edge_feat_name) - set(edge_entity_info) + if unknown_edge_types: + raise ValueError( + f"weight_edge_feat_name contains unknown edge types: {unknown_edge_types}" + ) + configured_weights = [ + (edge_type, feature_name, edge_entity_info[edge_type]) + for edge_type, feature_name in weight_edge_feat_name.items() + ] + + for edge_type, feature_name, serialized_info in configured_weights: + if feature_name not in serialized_info.feature_keys: + raise ValueError( + f"Sampling-weight field '{feature_name}' for edge type {edge_type} must be an unquantized raw edge feature." + ) + + +def remove_sampling_weight_from_edge_quantization_metadata( + serialized_graph_metadata: SerializedGraphMetadata, + weight_edge_feat_name: Optional[Union[str, dict[EdgeType, str]]], +) -> Optional[ + Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] +]: + """Remove separately stored sampling weights from edge reconstruction metadata. + + TFRecord loading removes the sampling-weight column from raw edge features + before registering it with the weighted sampler. The resulting metadata + must describe the remaining model features so batch reconstruction scatters + raw and dequantized columns into the correct positions. + + Args: + serialized_graph_metadata: Serialized edge schema and quantization metadata. + weight_edge_feat_name: Raw scalar feature configured as sampling weights. + + Returns: + Quantization metadata for the model-facing edge features. + """ + quantization_metadata = serialized_graph_metadata.edge_quantization_metadata + if quantization_metadata is None or weight_edge_feat_name is None: + return quantization_metadata + + if isinstance(serialized_graph_metadata.edge_entity_info, SerializedTFRecordInfo): + assert isinstance(quantization_metadata, FeatureQuantizationMetadata) + assert isinstance(weight_edge_feat_name, str) + edge_info_by_type: dict[EdgeType, SerializedTFRecordInfo] = { + DEFAULT_HOMOGENEOUS_EDGE_TYPE: serialized_graph_metadata.edge_entity_info + } + metadata_by_type: dict[EdgeType, FeatureQuantizationMetadata] = { + DEFAULT_HOMOGENEOUS_EDGE_TYPE: quantization_metadata + } + weight_by_type: dict[EdgeType, str] = { + DEFAULT_HOMOGENEOUS_EDGE_TYPE: weight_edge_feat_name + } + is_homogeneous = True + else: + assert isinstance(quantization_metadata, dict) + edge_info_by_type: dict[EdgeType, SerializedTFRecordInfo] = ( + serialized_graph_metadata.edge_entity_info + ) + metadata_by_type: dict[EdgeType, FeatureQuantizationMetadata] = cast( + dict[EdgeType, FeatureQuantizationMetadata], quantization_metadata + ) + if isinstance(weight_edge_feat_name, str): + edge_type = next(iter(edge_info_by_type)) + weight_by_type: dict[EdgeType, str] = {edge_type: weight_edge_feat_name} + else: + weight_by_type: dict[EdgeType, str] = weight_edge_feat_name + is_homogeneous = False + + adjusted_metadata: dict[EdgeType, FeatureQuantizationMetadata] = {} + for edge_type, metadata in metadata_by_type.items(): + weight_feature_name = weight_by_type.get(edge_type) + if weight_feature_name is None: + adjusted_metadata[edge_type] = metadata + continue + + edge_info = edge_info_by_type[edge_type] + raw_column_offset = 0 + for feature_name in edge_info.feature_keys: + if feature_name == weight_feature_name: + break + feature_spec = edge_info.feature_spec[feature_name] + raw_column_offset += feature_spec.shape[-1] if feature_spec.shape else 1 + weight_logical_index = metadata.raw_feature_indices[raw_column_offset] + adjusted_quantized_feature_indices = tuple( + quantized_feature_index - 1 + if quantized_feature_index > weight_logical_index + else quantized_feature_index + for quantized_feature_index in metadata.quantized_feature_indices + ) + adjusted_metadata[edge_type] = replace( + metadata, + feature_dim=metadata.feature_dim - 1, + quantized_feature_indices=adjusted_quantized_feature_indices, + ) + + if is_homogeneous: + return adjusted_metadata[DEFAULT_HOMOGENEOUS_EDGE_TYPE] + return adjusted_metadata def _data_loading_process( @@ -199,14 +327,6 @@ def _data_loading_process( raise NotImplementedError( "Label keys are not supported for edge entities" ) - if ( - serialized_entity_tf_record_info.packed_feature_key is not None - and not serialized_entity_tf_record_info.is_node_entity - ): - # TODO(quantization): Support feature quantization for edge features. - raise NotImplementedError( - "Packed feature keys are not supported for edge entities" - ) loaded_entity = tf_record_dataloader.load_as_torch_tensors( serialized_tf_record_info=serialized_entity_tf_record_info, tf_dataset_options=tf_dataset_options, @@ -396,6 +516,11 @@ def load_torch_tensors_from_tf_record( loaded_graph_tensors (LoadedGraphTensors): Unpartitioned Graph Tensors """ + _validate_weight_edge_feature_name( + edge_entity_info=serialized_graph_metadata.edge_entity_info, + weight_edge_feat_name=weight_edge_feat_name, + ) + logger.info(f"Rank {rank} starting loading torch tensors from serialized info ...") start_time = time.time() @@ -525,6 +650,9 @@ def load_torch_tensors_from_tf_record( edge_index = edge_output_dict[_ID_FMT.format(entity=_EDGE_KEY)] edge_features = edge_output_dict.get(_FEATURE_FMT.format(entity=_EDGE_KEY), None) + edge_quantized_features = edge_output_dict.get( + _PACKED_FEATURE_FMT.format(entity=_EDGE_KEY), None + ) edge_weights = edge_output_dict.get(_EDGE_WEIGHTS_KEY, None) positive_labels = edge_output_dict.get( @@ -552,6 +680,7 @@ def load_torch_tensors_from_tf_record( node_labels=node_labels, edge_index=edge_index, edge_features=edge_features, + edge_quantized_features=edge_quantized_features, positive_label=positive_labels, negative_label=negative_labels, edge_weights=edge_weights, diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 1a5f859b5..14e874be3 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -194,6 +194,10 @@ def _load_and_build_partitioned_dataset( partitioner.register_edge_features( edge_features=loaded_graph_tensors.edge_features ) + if loaded_graph_tensors.edge_quantized_features is not None: + partitioner.register_edge_quantized_features( + edge_quantized_features=loaded_graph_tensors.edge_quantized_features + ) if loaded_graph_tensors.positive_label is not None: partitioner.register_labels( label_edge_index=loaded_graph_tensors.positive_label, is_positive=True @@ -212,6 +216,7 @@ def _load_and_build_partitioned_dataset( loaded_graph_tensors.node_quantized_features, loaded_graph_tensors.edge_index, loaded_graph_tensors.edge_features, + loaded_graph_tensors.edge_quantized_features, loaded_graph_tensors.edge_weights, loaded_graph_tensors.positive_label, loaded_graph_tensors.negative_label, diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 04de8ce72..9549c9e3f 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -208,6 +208,8 @@ def __init__( self._edge_ids: Optional[dict[EdgeType, tuple[int, int]]] = None self._edge_feat: Optional[dict[EdgeType, torch.Tensor]] = None self._edge_feat_dim: Optional[dict[EdgeType, int]] = None + self._edge_quantized_feat: Optional[dict[EdgeType, torch.Tensor]] = None + self._edge_quantized_feat_dim: Optional[dict[EdgeType, int]] = None self._edge_weights: Optional[dict[EdgeType, torch.Tensor]] = None # TODO (mkolodner-sc): Deprecate the need for explicitly storing labels are part of this class, leveraging @@ -669,6 +671,36 @@ def register_edge_features( for edge_type in input_edge_features: self._edge_feat_dim[edge_type] = input_edge_features[edge_type].shape[1] + def register_edge_quantized_features( + self, edge_quantized_features: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] + ) -> None: + """Register packed uint8 main-edge features for co-partitioning.""" + + self._assert_and_get_rpc_setup() + if self._edge_quantized_feat is not None: + raise ValueError( + "Edge quantized features have already been registered. Cannot re-register edge quantized feature data." + ) + logger.info("Registering Edge Quantized Features ...") + + input_edge_quantized_features = ( + self._convert_edge_entity_to_heterogeneous_format( + input_edge_entity=edge_quantized_features + ) + ) + + assert input_edge_quantized_features, ( + "Edge quantized features is an empty dictionary. Please provide edge quantized features to register." + ) + + self._edge_quantized_feat = convert_to_tensor( + input_edge_quantized_features, dtype=torch.uint8 + ) + self._edge_quantized_feat_dim = { + edge_type: features.shape[1] + for edge_type, features in input_edge_quantized_features.items() + } + def register_edge_weights( self, edge_weights: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] ) -> None: @@ -1201,7 +1233,10 @@ def _partition_edge_index_and_edge_features( node_partition_book: dict[NodeType, PartitionBook], edge_type: EdgeType, ) -> Tuple[ - GraphPartitionData, Optional[FeaturePartitionData], Optional[PartitionBook] + GraphPartitionData, + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + Optional[PartitionBook], ]: r"""Partition graph topology and edge features of a specific edge type. If there are no edge features for the current edge type, both the returned edge feature and edge partition book will be None. @@ -1213,6 +1248,7 @@ def _partition_edge_index_and_edge_features( Returns: GraphPartitionData: The graph data of the current partition. Optional[FeaturePartitionData]: The edge features on the current partition, will be None if there are no edge features for the current edge type + Optional[FeaturePartitionData]: The quantized edge features on the current partition, will be None if there are no quantized edge features for the current edge type Optional[PartitionBook]: The partition book of graph edges, will be None if there are no edge features for the current edge type """ @@ -1225,11 +1261,17 @@ def _partition_edge_index_and_edge_features( ), "Must have registered edges prior to partitioning them" has_edge_feats = self._edge_feat is not None and edge_type in self._edge_feat + has_edge_quantized_feats = ( + self._edge_quantized_feat is not None + and edge_type in self._edge_quantized_feat + ) has_weights_for_edge_type = ( self._edge_weights is not None and edge_type in self._edge_weights ) # Need a partition book if we have features or weights to reindex. - should_generate_partition_book = has_edge_feats or has_weights_for_edge_type + should_generate_partition_book = ( + has_edge_feats or has_edge_quantized_feats or has_weights_for_edge_type + ) # Partitioning Edge Indices @@ -1283,12 +1325,12 @@ def _edge_pfn(_, chunk_range): gc.collect() - # Partition edge features and weights together in a single pass, + # Partition edge features, packed features, and weights together in a single pass, # mirroring how node features and labels are co-partitioned. - # Input tuple layout: (edge_feat?, edge_weights?, edge_ids) - # IDs are always at r[-1]; features at r[0]; weights at r[1] when - # features are also present, else r[0]. + # Input tuple layout: (edge_feat?, edge_quantized_feat?, edge_weights?, edge_ids) + # IDs are always last; optional tensor indices are recorded when appended. current_feat_part: Optional[FeaturePartitionData] = None + current_quantized_feat_part: Optional[FeaturePartitionData] = None partitioned_weights: Optional[torch.Tensor] = None partitioned_edge_ids: Optional[torch.Tensor] = None @@ -1309,6 +1351,8 @@ def _edge_pfn(_, chunk_range): edge_feat: Optional[torch.Tensor] = None edge_feat_dim: Optional[int] = None edge_weights_tensor: Optional[torch.Tensor] = None + edge_quantized_features: Optional[torch.Tensor] = None + edge_quantized_feature_dim: Optional[int] = None if has_edge_feats: assert self._edge_feat is not None and edge_type in self._edge_feat assert ( @@ -1316,30 +1360,38 @@ def _edge_pfn(_, chunk_range): ) edge_feat = self._edge_feat[edge_type] edge_feat_dim = self._edge_feat_dim[edge_type] + if has_edge_quantized_feats: + assert self._edge_quantized_feat is not None + assert self._edge_quantized_feat_dim is not None + edge_quantized_features = self._edge_quantized_feat[edge_type] + edge_quantized_feature_dim = self._edge_quantized_feat_dim[edge_type] if has_weights_for_edge_type: assert self._edge_weights is not None edge_weights_tensor = self._edge_weights[edge_type] input_parts: list[torch.Tensor] = [] + feat_idx: Optional[int] = None if edge_feat is not None: + feat_idx = len(input_parts) input_parts.append(edge_feat) + quantized_feat_idx: Optional[int] = None + if edge_quantized_features is not None: + quantized_feat_idx = len(input_parts) + input_parts.append(edge_quantized_features) + weight_idx: Optional[int] = None if edge_weights_tensor is not None: + weight_idx = len(input_parts) input_parts.append(edge_weights_tensor) input_parts.append(edge_ids) - # Positional indices: features first, weights next, ids always last. - feat_idx: Optional[int] = 0 if has_edge_feats else None - weight_idx: Optional[int] = None - if has_weights_for_edge_type: - weight_idx = 1 if has_edge_feats else 0 - + # Recorded indices keep result unpacking aligned with optional inputs. def _edge_feat_weight_pfn( ids_chunk: torch.Tensor, _: object ) -> torch.Tensor: assert edge_partition_book is not None return edge_partition_book[ids_chunk] - # Each result tuple contains (edge_feat?, edge_weights?, edge_ids). + # Each result tuple preserves the input tuple layout. feat_weight_res_list, _ = self._partition_by_chunk( input_data=tuple(input_parts), rank_indices=edge_ids, @@ -1360,6 +1412,21 @@ def _edge_feat_weight_pfn( if len(self._edge_feat) == 0 and len(self._edge_feat_dim) == 0: self._edge_feat = None self._edge_feat_dim = None + if has_edge_quantized_feats: + assert edge_quantized_features is not None + assert self._edge_quantized_feat is not None + assert self._edge_quantized_feat_dim is not None + del edge_quantized_features + del ( + self._edge_quantized_feat[edge_type], + self._edge_quantized_feat_dim[edge_type], + ) + if ( + len(self._edge_quantized_feat) == 0 + and len(self._edge_quantized_feat_dim) == 0 + ): + self._edge_quantized_feat = None + self._edge_quantized_feat_dim = None if has_weights_for_edge_type: assert edge_weights_tensor is not None assert self._edge_weights is not None @@ -1377,6 +1444,14 @@ def _edge_feat_weight_pfn( feats=torch.empty(0, edge_feat_dim), ids=partitioned_edge_ids, ) + if has_edge_quantized_feats: + assert edge_quantized_feature_dim is not None + current_quantized_feat_part = FeaturePartitionData( + feats=torch.empty( + 0, edge_quantized_feature_dim, dtype=torch.uint8 + ), + ids=partitioned_edge_ids, + ) if has_weights_for_edge_type: partitioned_weights = torch.empty(0) else: @@ -1387,6 +1462,14 @@ def _edge_feat_weight_pfn( feats=torch.cat([r[feat_idx] for r in feat_weight_res_list]), ids=partitioned_edge_ids, ) + if has_edge_quantized_feats: + assert quantized_feat_idx is not None + current_quantized_feat_part = FeaturePartitionData( + feats=torch.cat( + [r[quantized_feat_idx] for r in feat_weight_res_list] + ), + ids=partitioned_edge_ids, + ) if has_weights_for_edge_type: assert weight_idx is not None partitioned_weights = torch.cat( @@ -1410,7 +1493,12 @@ def _edge_feat_weight_pfn( weights=partitioned_weights, ) - return current_graph_part, current_feat_part, edge_partition_book + return ( + current_graph_part, + current_feat_part, + current_quantized_feat_part, + edge_partition_book, + ) def _partition_label_edge_index( self, @@ -1683,11 +1771,15 @@ def partition_edge_index_and_edge_features( self, node_partition_book: Union[PartitionBook, dict[NodeType, PartitionBook]] ) -> Union[ Tuple[ - GraphPartitionData, Optional[FeaturePartitionData], Optional[PartitionBook] + GraphPartitionData, + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + Optional[PartitionBook], ], Tuple[ dict[EdgeType, GraphPartitionData], Optional[dict[EdgeType, FeaturePartitionData]], + Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, PartitionBook]], ], ]: @@ -1698,8 +1790,8 @@ def partition_edge_index_and_edge_features( node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book Returns: Union[ - Tuple[GraphPartitionData, FeaturePartitionData, PartitionBook], - Tuple[dict[EdgeType, GraphPartitionData], Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, PartitionBook]]], + Tuple[GraphPartitionData, Optional[FeaturePartitionData], Optional[FeaturePartitionData], Optional[PartitionBook]], + Tuple[dict[EdgeType, GraphPartitionData], Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, PartitionBook]]], ]: Partitioned Graph Data, Feature Data, and corresponding edge partition book, is a dictionary if heterogeneous. The second and third elements of this tuple are only present if there are edge features to partition, and are None otherwise. @@ -1748,21 +1840,27 @@ def partition_edge_index_and_edge_features( edge_partition_book: dict[EdgeType, PartitionBook] = {} partitioned_edge_index: dict[EdgeType, GraphPartitionData] = {} partitioned_edge_features: dict[EdgeType, FeaturePartitionData] = {} + partitioned_edge_quantized_features: dict[EdgeType, FeaturePartitionData] = {} for edge_type in self._edge_types: ( partitioned_edge_index_per_edge_type, partitioned_edge_features_per_edge_type, + partitioned_edge_quantized_features_per_edge_type, edge_partition_book_per_edge_type, ) = self._partition_edge_index_and_edge_features( node_partition_book=transformed_node_partition_book, edge_type=edge_type ) partitioned_edge_index[edge_type] = partitioned_edge_index_per_edge_type - if partitioned_edge_features_per_edge_type is not None: - assert edge_partition_book_per_edge_type is not None + if edge_partition_book_per_edge_type is not None: edge_partition_book[edge_type] = edge_partition_book_per_edge_type + if partitioned_edge_features_per_edge_type is not None: partitioned_edge_features[edge_type] = ( partitioned_edge_features_per_edge_type ) + if partitioned_edge_quantized_features_per_edge_type is not None: + partitioned_edge_quantized_features[edge_type] = ( + partitioned_edge_quantized_features_per_edge_type + ) elapsed_time = time.time() - start_time logger.info(f"Edge Partitioning finished, took {elapsed_time:.3f}s") @@ -1784,6 +1882,9 @@ def partition_edge_index_and_edge_features( to_homogeneous(partitioned_edge_features) if partitioned_edge_features else None, + to_homogeneous(partitioned_edge_quantized_features) + if partitioned_edge_quantized_features + else None, to_homogeneous(edge_partition_book) if edge_partition_book else None, ) else: @@ -1791,6 +1892,11 @@ def partition_edge_index_and_edge_features( return ( partitioned_edge_index, partitioned_edge_features if partitioned_edge_features else None, + ( + partitioned_edge_quantized_features + if partitioned_edge_quantized_features + else None + ), edge_partition_book if edge_partition_book else None, ) @@ -1889,6 +1995,7 @@ def partition( ( partitioned_edge_index, partitioned_edge_features, + partitioned_edge_quantized_features, edge_partition_book, ) = self.partition_edge_index_and_edge_features( node_partition_book=node_partition_book @@ -1936,6 +2043,7 @@ def partition( partitioned_node_features=partitioned_node_features, partitioned_node_quantized_features=partitioned_node_quantized_features, partitioned_edge_features=partitioned_edge_features, + partitioned_edge_quantized_features=partitioned_edge_quantized_features, partitioned_positive_labels=partitioned_positive_edge_index, partitioned_negative_labels=partitioned_negative_edge_index, partitioned_node_labels=partitioned_node_labels, diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index b7b0754f7..170d5de09 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -215,7 +215,10 @@ def _partition_edge_index_and_edge_features( node_partition_book: dict[NodeType, PartitionBook], edge_type: EdgeType, ) -> tuple[ - GraphPartitionData, Optional[FeaturePartitionData], Optional[PartitionBook] + GraphPartitionData, + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + Optional[PartitionBook], ]: """ Partition graph topology of a specific edge type. For range-based partitioning, we partition @@ -232,6 +235,7 @@ def _partition_edge_index_and_edge_features( Returns: GraphPartitionData: The graph data of the current partition. Optional[FeaturePartitionData]: The edge features on the current partition, will be None if there are no edge features for the current edge type + Optional[FeaturePartitionData]: The quantized edge features on the current partition, will be None if there are no quantized edge features for the current edge type Optional[PartitionBook]: The partition book of graph edges, will be None if there are no edge features for the current edge type """ @@ -243,6 +247,10 @@ def _partition_edge_index_and_edge_features( edge_index = self._edge_index[edge_type] has_edge_feats = self._edge_feat is not None and edge_type in self._edge_feat + has_edge_quantized_feats = ( + self._edge_quantized_feat is not None + and edge_type in self._edge_quantized_feat + ) has_edge_weights = ( self._edge_weights is not None and edge_type in self._edge_weights ) @@ -255,24 +263,35 @@ def _partition_edge_index_and_edge_features( edge_feat: Optional[torch.Tensor] = None edge_feat_dim: Optional[int] = None edge_weights_tensor: Optional[torch.Tensor] = None + edge_quantized_features: Optional[torch.Tensor] = None + edge_quantized_feature_dim: Optional[int] = None if has_edge_feats: assert self._edge_feat is not None and self._edge_feat_dim is not None assert edge_type in self._edge_feat_dim edge_feat = self._edge_feat[edge_type] edge_feat_dim = self._edge_feat_dim[edge_type] + if has_edge_quantized_feats: + assert self._edge_quantized_feat is not None + assert self._edge_quantized_feat_dim is not None + edge_quantized_features = self._edge_quantized_feat[edge_type] + edge_quantized_feature_dim = self._edge_quantized_feat_dim[edge_type] if has_edge_weights: assert self._edge_weights is not None edge_weights_tensor = self._edge_weights[edge_type] - # Build input_data tuple: (src, dst[, feat][, weights]) - # Track the index of each optional tensor so we can unpack res_list correctly. + # Build input_data as (src, dst[, feat][, packed feat][, weights]). + # Recorded indices keep result unpacking aligned with optional inputs. input_parts: list[torch.Tensor] = [edge_index[0], edge_index[1]] feat_idx: Optional[int] = None weight_idx: Optional[int] = None if edge_feat is not None: feat_idx = len(input_parts) input_parts.append(edge_feat) + quantized_feat_idx: Optional[int] = None + if edge_quantized_features is not None: + quantized_feat_idx = len(input_parts) + input_parts.append(edge_quantized_features) if edge_weights_tensor is not None: weight_idx = len(input_parts) input_parts.append(edge_weights_tensor) @@ -301,6 +320,15 @@ def edge_partition_fn(rank_indices, _): del self._edge_feat[edge_type], self._edge_feat_dim[edge_type] if self._edge_weights is not None and edge_type in self._edge_weights: del self._edge_weights[edge_type] + if ( + self._edge_quantized_feat is not None + and edge_type in self._edge_quantized_feat + ): + assert self._edge_quantized_feat_dim is not None + del ( + self._edge_quantized_feat[edge_type], + self._edge_quantized_feat_dim[edge_type], + ) # We check if edge_index or edge_feat dict is empty after deleting the tensor. If so, we set these fields to None. if not self._edge_index: @@ -310,6 +338,9 @@ def edge_partition_fn(rank_indices, _): self._edge_feat_dim = None if self._edge_weights is not None and not self._edge_weights: self._edge_weights = None + if self._edge_quantized_feat is not None and not self._edge_quantized_feat: + self._edge_quantized_feat = None + self._edge_quantized_feat_dim = None gc.collect() @@ -319,6 +350,11 @@ def edge_partition_fn(rank_indices, _): torch.empty(0, edge_feat_dim) if edge_feat_dim is not None else None ) partitioned_weights = torch.empty(0) if has_edge_weights else None + partitioned_edge_quantized_features = ( + torch.empty(0, edge_quantized_feature_dim, dtype=torch.uint8) + if edge_quantized_feature_dim is not None + else None + ) else: partitioned_edge_index = torch.stack( ( @@ -337,12 +373,17 @@ def edge_partition_fn(rank_indices, _): if weight_idx is not None else None ) + partitioned_edge_quantized_features = ( + torch.cat([r[quantized_feat_idx] for r in res_list]) + if quantized_feat_idx is not None + else None + ) res_list.clear() gc.collect() - # Generate range-based edge partition book and infer edge IDs. - # Only needed when edge features are present — weights use positional IDs. + # Generate range-based edge partition book and infer edge IDs for every + # sidecar that requires sampled edge lookup. num_edges_on_each_rank: list[tuple[int, int]] = sorted( all_gather((self._rank, partitioned_edge_index.size(1))).values(), key=lambda x: x[0], @@ -354,21 +395,26 @@ def edge_partition_fn(rank_indices, _): partition_ranges.append((start, end)) start = end - if edge_feat_dim is not None: + if ( + edge_feat_dim is not None + or edge_quantized_feature_dim is not None + or has_edge_weights + ): edge_partition_book = RangePartitionBook( partition_ranges=partition_ranges, partition_idx=self._rank ) partitioned_edge_ids = get_ids_on_rank( partition_book=edge_partition_book, rank=self._rank ) - assert partitioned_edge_features is not None current_graph_part = GraphPartitionData( edge_index=partitioned_edge_index, edge_ids=partitioned_edge_ids, weights=partitioned_weights, ) - current_feat_part = FeaturePartitionData( - feats=partitioned_edge_features, ids=None + current_feat_part = ( + FeaturePartitionData(feats=partitioned_edge_features, ids=None) + if partitioned_edge_features is not None + else None ) logger.info( f"Got edge range-based partition book for edge type {edge_type} on rank {self._rank} with partition bounds: {edge_partition_book.partition_bounds}" @@ -386,17 +432,31 @@ def edge_partition_fn(rank_indices, _): f"Edge Index and Feature Partitioning for edge type {edge_type} finished, took {time.time() - start_time:.3f}s" ) - return current_graph_part, current_feat_part, edge_partition_book + current_quantized_feat_part = ( + FeaturePartitionData(feats=partitioned_edge_quantized_features, ids=None) + if partitioned_edge_quantized_features is not None + else None + ) + return ( + current_graph_part, + current_feat_part, + current_quantized_feat_part, + edge_partition_book, + ) def partition_edge_index_and_edge_features( self, node_partition_book: Union[PartitionBook, dict[NodeType, PartitionBook]] ) -> Union[ tuple[ - GraphPartitionData, Optional[FeaturePartitionData], Optional[PartitionBook] + GraphPartitionData, + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + Optional[PartitionBook], ], tuple[ dict[EdgeType, GraphPartitionData], Optional[dict[EdgeType, FeaturePartitionData]], + Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, PartitionBook]], ], ]: @@ -408,10 +468,11 @@ def partition_edge_index_and_edge_features( Args: node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book + Returns: Union[ - Tuple[GraphPartitionData, Optional[FeaturePartitionData], Optional[PartitionBook]], - Tuple[dict[EdgeType, GraphPartitionData], Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, PartitionBook]]], + Tuple[GraphPartitionData, Optional[FeaturePartitionData], Optional[FeaturePartitionData], Optional[PartitionBook]], + Tuple[dict[EdgeType, GraphPartitionData], Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, FeaturePartitionData]], Optional[dict[EdgeType, PartitionBook]]], ]: Partitioned Graph Data, Feature Data, and corresponding edge partition book, is a dictionary if heterogeneous. """ @@ -448,21 +509,27 @@ def partition_edge_index_and_edge_features( edge_partition_book: dict[EdgeType, PartitionBook] = {} partitioned_edge_index: dict[EdgeType, GraphPartitionData] = {} partitioned_edge_features: dict[EdgeType, FeaturePartitionData] = {} + partitioned_edge_quantized_features: dict[EdgeType, FeaturePartitionData] = {} for edge_type in self._edge_types: ( partitioned_edge_index_per_edge_type, partitioned_edge_features_per_edge_type, + partitioned_edge_quantized_features_per_edge_type, edge_partition_book_per_edge_type, ) = self._partition_edge_index_and_edge_features( node_partition_book=transformed_node_partition_book, edge_type=edge_type ) partitioned_edge_index[edge_type] = partitioned_edge_index_per_edge_type - if partitioned_edge_features_per_edge_type is not None: - assert edge_partition_book_per_edge_type is not None + if edge_partition_book_per_edge_type is not None: edge_partition_book[edge_type] = edge_partition_book_per_edge_type + if partitioned_edge_features_per_edge_type is not None: partitioned_edge_features[edge_type] = ( partitioned_edge_features_per_edge_type ) + if partitioned_edge_quantized_features_per_edge_type is not None: + partitioned_edge_quantized_features[edge_type] = ( + partitioned_edge_quantized_features_per_edge_type + ) elapsed_time = time.time() - start_time logger.info(f"Edge Partitioning finished, took {elapsed_time:.3f}s") @@ -481,6 +548,9 @@ def partition_edge_index_and_edge_features( to_homogeneous(partitioned_edge_features) if partitioned_edge_features else None, + to_homogeneous(partitioned_edge_quantized_features) + if partitioned_edge_quantized_features + else None, to_homogeneous(edge_partition_book) if edge_partition_book else None, ) else: @@ -488,5 +558,10 @@ def partition_edge_index_and_edge_features( return ( partitioned_edge_index, partitioned_edge_features if partitioned_edge_features else None, + ( + partitioned_edge_quantized_features + if partitioned_edge_quantized_features + else None + ), edge_partition_book if edge_partition_book else None, ) diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index 36ad31c52..25fb26882 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -33,18 +33,11 @@ def _build_serialized_tfrecord_entity_info( 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 entity, when its features are quantized. + 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: - if not isinstance( - preprocessed_metadata, PreprocessedMetadata.NodeMetadataOutput - ): - # TODO(quantization): Support edge feature quantization. - raise NotImplementedError( - "Feature quantization is not supported for edge entities." - ) packed_feature_key = ( preprocessed_metadata.quantized_feature_metadata.packed_feature_key ) @@ -146,6 +139,7 @@ def convert_pb_to_serialized_graph_metadata( 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 @@ -202,11 +196,19 @@ def convert_pb_to_serialized_graph_metadata( 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"): @@ -251,6 +253,9 @@ def convert_pb_to_serialized_graph_metadata( 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( @@ -265,4 +270,7 @@ def convert_pb_to_serialized_graph_metadata( 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, ) diff --git a/gigl/types/graph.py b/gigl/types/graph.py index 849f7708a..eb501f0d7 100644 --- a/gigl/types/graph.py +++ b/gigl/types/graph.py @@ -105,6 +105,9 @@ class PartitionOutput: partitioned_node_quantized_features: Optional[ Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]] ] = None + partitioned_edge_quantized_features: Optional[ + Union[FeaturePartitionData, dict[EdgeType, FeaturePartitionData]] + ] = None @dataclass(frozen=True) @@ -236,6 +239,9 @@ class LoadedGraphTensors: node_quantized_features: Optional[ Union[torch.Tensor, dict[NodeType, torch.Tensor]] ] = None + edge_quantized_features: Optional[ + Union[torch.Tensor, dict[EdgeType, torch.Tensor]] + ] = None def treat_labels_as_edges(self, edge_dir: Literal["in", "out"]) -> None: """ @@ -337,6 +343,9 @@ def treat_labels_as_edges(self, edge_dir: Literal["in", "out"]) -> None: self.node_quantized_features = to_heterogeneous_node( self.node_quantized_features ) + self.edge_quantized_features = to_heterogeneous_edge( + self.edge_quantized_features + ) self.edge_index = edge_index_with_labels self.edge_features = to_heterogeneous_edge(self.edge_features) self.edge_weights = to_heterogeneous_edge(self.edge_weights) diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 046b8bf49..89c863d07 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -22,6 +22,9 @@ class InputDataStrategy(Enum): REGISTER_EDGE_WEIGHTS_WITHOUT_EDGE_FEATURES = ( "REGISTER_EDGE_WEIGHTS_WITHOUT_EDGE_FEATURES" ) + REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES = ( + "REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES" + ) def run_distributed_partitioner( @@ -95,7 +98,29 @@ def run_distributed_partitioner( init_rpc(master_addr=master_addr, master_port=master_port, num_rpc_threads=4) dist_partitioner: DistPartitioner - if input_data_strategy in ( + if ( + input_data_strategy + == InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES + ): + dist_partitioner = partitioner_class( + should_assign_edges_by_src_node=should_assign_edges_by_src_node, + ) + dist_partitioner.register_node_ids(node_ids=node_ids) + dist_partitioner.register_edge_index(edge_index=edge_index) + edge_quantized_features: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] + if isinstance(edge_index, dict): + edge_index_by_type = cast(dict[EdgeType, torch.Tensor], edge_index) + edge_quantized_features = { + edge_type: indices[0].to(torch.uint8).unsqueeze(1) + for edge_type, indices in edge_index_by_type.items() + } + else: + edge_quantized_features = edge_index[0].to(torch.uint8).unsqueeze(1) + dist_partitioner.register_edge_quantized_features( + edge_quantized_features=edge_quantized_features + ) + partition_output = dist_partitioner.partition() + elif input_data_strategy in ( InputDataStrategy.REGISTER_ALL_ENTITIES_SEPARATELY, InputDataStrategy.REGISTER_EDGE_WEIGHTS_WITHOUT_EDGE_FEATURES, ): @@ -119,6 +144,7 @@ def run_distributed_partitioner( ( output_edge_index, output_edge_features, + _, output_edge_partition_book, ) = dist_partitioner.partition_edge_index_and_edge_features( node_partition_book=output_node_partition_book @@ -181,6 +207,7 @@ def run_distributed_partitioner( ( output_graph, output_edge_features, + _, output_edge_partition_book, ) = dist_partitioner.partition_edge_index_and_edge_features( node_partition_book=output_node_partition_book diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 3bfaff851..12b5de2e5 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -19,6 +19,7 @@ from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, load_torch_tensors_from_tf_record, + remove_sampling_weight_from_edge_quantization_metadata, ) from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper from gigl.src.data_preprocessor.lib.types import FeatureSpecDict @@ -29,6 +30,7 @@ from gigl.src.mocking.mocking_assets.mocked_datasets_for_pipeline_tests import ( CORA_NODE_CLASSIFICATION_MOCKED_DATASET_INFO, ) +from gigl.types.graph import FeatureQuantizationMetadata from tests.test_assets.test_case import TestCase _FEATURE_SPEC_WITH_ENTITY_KEY: FeatureSpecDict = { @@ -644,6 +646,91 @@ def test_load_edge_weights_from_tf_record(self): torch.tensor(sorted(edge_feature_vals), dtype=torch.float32), ) + def test_load_edge_weights_rejects_non_raw_field_before_loading(self) -> None: + missing_path = UriFactory.create_uri("/does/not/exist") + serialized_graph_metadata = SerializedGraphMetadata( + node_entity_info=SerializedTFRecordInfo( + tfrecord_uri_prefix=missing_path, + feature_spec={"node_id": tf.io.FixedLenFeature([], tf.int64)}, + feature_keys=[], + feature_dim=0, + entity_key="node_id", + ), + edge_entity_info=SerializedTFRecordInfo( + tfrecord_uri_prefix=missing_path, + feature_spec={ + "src_id": tf.io.FixedLenFeature([], tf.int64), + "dst_id": tf.io.FixedLenFeature([], tf.int64), + "edge_packed_features": tf.io.FixedLenFeature([], tf.string), + }, + feature_keys=[], + feature_dim=0, + entity_key=("src_id", "dst_id"), + packed_feature_key="edge_packed_features", + packed_feature_dim=1, + ), + ) + + with self.assertRaises(ValueError): + load_torch_tensors_from_tf_record( + tf_record_dataloader=TFRecordDataLoader(rank=0, world_size=1), + serialized_graph_metadata=serialized_graph_metadata, + should_load_tensors_in_parallel=False, + weight_edge_feat_name="quantized_weight", + ) + + def test_sampling_weight_removal_updates_edge_quantization_metadata( + self, + ) -> None: + missing_path = UriFactory.create_uri("/does/not/exist") + serialized_graph_metadata = SerializedGraphMetadata( + node_entity_info=SerializedTFRecordInfo( + tfrecord_uri_prefix=missing_path, + feature_spec={"node_id": tf.io.FixedLenFeature([], tf.int64)}, + feature_keys=[], + feature_dim=0, + entity_key="node_id", + ), + edge_entity_info=SerializedTFRecordInfo( + tfrecord_uri_prefix=missing_path, + feature_spec={ + "src_id": tf.io.FixedLenFeature([], tf.int64), + "dst_id": tf.io.FixedLenFeature([], tf.int64), + "raw_embedding": tf.io.FixedLenFeature([2], tf.float32), + "weight": tf.io.FixedLenFeature([], tf.float32), + "edge_packed_features": tf.io.FixedLenFeature([], tf.string), + }, + feature_keys=["raw_embedding", "weight"], + feature_dim=3, + entity_key=("src_id", "dst_id"), + packed_feature_key="edge_packed_features", + packed_feature_dim=1, + ), + edge_quantization_metadata=FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(3,), + clip_min=0.0, + clip_max=3.0, + ), + ) + + adjusted_metadata = remove_sampling_weight_from_edge_quantization_metadata( + serialized_graph_metadata=serialized_graph_metadata, + weight_edge_feat_name="weight", + ) + + self.assertEqual( + adjusted_metadata, + FeatureQuantizationMetadata( + bits=2, + feature_dim=3, + quantized_feature_indices=(2,), + clip_min=0.0, + clip_max=3.0, + ), + ) + def test_load_edge_weights_multidim_feature(self): """Weight column offset is correct when a preceding feature key is multi-dimensional. diff --git a/tests/unit/distributed/distributed_partitioner_test.py b/tests/unit/distributed/distributed_partitioner_test.py index 0f817bafa..f3aaea091 100644 --- a/tests/unit/distributed/distributed_partitioner_test.py +++ b/tests/unit/distributed/distributed_partitioner_test.py @@ -686,6 +686,22 @@ def _assert_label_outputs( partitioner_class=DistRangePartitioner, expected_pb_dtype=torch.int64, ), + param( + "Homogeneous packed-edge-only tensor partitioning", + is_heterogeneous=False, + input_data_strategy=InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES, + should_assign_edges_by_src_node=True, + partitioner_class=DistPartitioner, + expected_pb_dtype=torch.uint8, + ), + param( + "Homogeneous packed-edge-only range partitioning", + is_heterogeneous=False, + input_data_strategy=InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES, + should_assign_edges_by_src_node=True, + partitioner_class=DistRangePartitioner, + expected_pb_dtype=torch.int64, + ), ] ) def test_partitioning_correctness( @@ -756,6 +772,11 @@ def test_partitioning_correctness( else: expected_edge_feat_types = [USER_TO_USER_EDGE_TYPE] + is_packed_edge_only = ( + input_data_strategy + == InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES + ) + for rank, partition_output in output_dict.items(): partitioned_edge_index = partition_output.partitioned_edge_index assert partitioned_edge_index is not None @@ -780,7 +801,32 @@ def test_partitioning_correctness( graph.edge_index ) - if ( + if is_packed_edge_only: + self.assertIsNotNone(partition_output.edge_partition_book) + self.assertIsNone(partition_output.partitioned_edge_features) + self.assertIsNotNone( + partition_output.partitioned_edge_quantized_features + ) + packed_features = partition_output.partitioned_edge_quantized_features + assert isinstance(packed_features, FeaturePartitionData) + assert isinstance(partitioned_edge_index, GraphPartitionData) + self.assertEqual(packed_features.feats.dtype, torch.uint8) + self.assertEqual( + packed_features.feats.size(0), + partitioned_edge_index.edge_index.size(1), + ) + assert partitioned_edge_index.edge_ids is not None + if packed_features.ids is not None: + self.assert_tensor_equality( + tensor_a=packed_features.ids, + tensor_b=partitioned_edge_index.edge_ids, + ) + for index, edge_id in enumerate(partitioned_edge_index.edge_ids): + self.assert_tensor_equality( + tensor_a=packed_features.feats[index], + tensor_b=edge_id.to(torch.uint8).unsqueeze(0), + ) + elif ( input_data_strategy == InputDataStrategy.REGISTER_MINIMAL_ENTITIES_SEPARATELY ): diff --git a/tests/unit/distributed/distributed_weighted_sampling_test.py b/tests/unit/distributed/distributed_weighted_sampling_test.py index 9b30acdd8..cb9386b6e 100644 --- a/tests/unit/distributed/distributed_weighted_sampling_test.py +++ b/tests/unit/distributed/distributed_weighted_sampling_test.py @@ -553,6 +553,11 @@ def test_weights_only_no_features_partitioned_correctly(self) -> None: ) assert edge_ids is not None + self.assertIsNotNone( + partition_output.edge_partition_book, + msg=f"Rank {rank}: edge partition book must be retained for weights", + ) + self.assertEqual(weights.shape, edge_ids.shape) expected_weights = edge_ids.float() * 0.1 torch.testing.assert_close( @@ -732,7 +737,7 @@ def test_range_partitioner_homogeneous_weights_partitioned_correctly(self) -> No True, # should_assign_edges_by_src_node self._master_ip_address, master_port, - InputDataStrategy.REGISTER_ALL_ENTITIES_SEPARATELY, + InputDataStrategy.REGISTER_EDGE_WEIGHTS_WITHOUT_EDGE_FEATURES, DistRangePartitioner, rank_to_edge_weights, ), @@ -759,6 +764,11 @@ def test_range_partitioner_homogeneous_weights_partitioned_correctly(self) -> No ) assert edge_ids is not None + self.assertIsNotNone( + partition_output.edge_partition_book, + msg=f"Rank {rank}: edge partition book must be retained for weights", + ) + self.assertEqual( weights.shape, edge_ids.shape, From 69cd87d003d98a438819bc1edf8fc1e6ba120e8f Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 17 Aug 2026 16:54:16 +0000 Subject: [PATCH 4/8] Address read path review feedback --- gigl/common/data/load_torch_tensors.py | 81 +++++++++---------- .../run_distributed_partitioner.py | 3 +- 2 files changed, 38 insertions(+), 46 deletions(-) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 667bc348d..74d98ff4b 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -124,47 +124,6 @@ class SerializedGraphMetadata: ] = None -def _validate_weight_edge_feature_name( - edge_entity_info: Union[ - SerializedTFRecordInfo, dict[EdgeType, SerializedTFRecordInfo] - ], - weight_edge_feat_name: Optional[Union[str, dict[EdgeType, str]]], -) -> None: - if weight_edge_feat_name is None: - return - - configured_weights: list[tuple[EdgeType, str, SerializedTFRecordInfo]] - if isinstance(edge_entity_info, SerializedTFRecordInfo): - if not isinstance(weight_edge_feat_name, str): - raise ValueError("weight_edge_feat_name must be str for homogeneous graph") - edge_type = DEFAULT_HOMOGENEOUS_EDGE_TYPE - configured_weights = [(edge_type, weight_edge_feat_name, edge_entity_info)] - else: - if isinstance(weight_edge_feat_name, str): - if len(edge_entity_info) != 1: - raise ValueError( - "weight_edge_feat_name must be dict[EdgeType, str] for heterogeneous graph with multiple edge types" - ) - edge_type, serialized_info = next(iter(edge_entity_info.items())) - configured_weights = [(edge_type, weight_edge_feat_name, serialized_info)] - else: - unknown_edge_types = set(weight_edge_feat_name) - set(edge_entity_info) - if unknown_edge_types: - raise ValueError( - f"weight_edge_feat_name contains unknown edge types: {unknown_edge_types}" - ) - configured_weights = [ - (edge_type, feature_name, edge_entity_info[edge_type]) - for edge_type, feature_name in weight_edge_feat_name.items() - ] - - for edge_type, feature_name, serialized_info in configured_weights: - if feature_name not in serialized_info.feature_keys: - raise ValueError( - f"Sampling-weight field '{feature_name}' for edge type {edge_type} must be an unquantized raw edge feature." - ) - - def remove_sampling_weight_from_edge_quantization_metadata( serialized_graph_metadata: SerializedGraphMetadata, weight_edge_feat_name: Optional[Union[str, dict[EdgeType, str]]], @@ -516,10 +475,42 @@ def load_torch_tensors_from_tf_record( loaded_graph_tensors (LoadedGraphTensors): Unpartitioned Graph Tensors """ - _validate_weight_edge_feature_name( - edge_entity_info=serialized_graph_metadata.edge_entity_info, - weight_edge_feat_name=weight_edge_feat_name, - ) + edge_entity_info = serialized_graph_metadata.edge_entity_info + if weight_edge_feat_name is not None: + if isinstance(edge_entity_info, SerializedTFRecordInfo): + if not isinstance(weight_edge_feat_name, str): + raise ValueError( + "weight_edge_feat_name must be str for homogeneous graph" + ) + if weight_edge_feat_name not in edge_entity_info.feature_keys: + raise ValueError( + f"Sampling-weight field '{weight_edge_feat_name}' for edge type " + f"{DEFAULT_HOMOGENEOUS_EDGE_TYPE} must be an unquantized raw edge feature." + ) + elif isinstance(weight_edge_feat_name, str): + if len(edge_entity_info) != 1: + raise ValueError( + "weight_edge_feat_name must be dict[EdgeType, str] for " + "heterogeneous graph with multiple edge types" + ) + edge_type, serialized_info = next(iter(edge_entity_info.items())) + if weight_edge_feat_name not in serialized_info.feature_keys: + raise ValueError( + f"Sampling-weight field '{weight_edge_feat_name}' for edge type " + f"{edge_type} must be an unquantized raw edge feature." + ) + else: + unknown_edge_types = set(weight_edge_feat_name) - set(edge_entity_info) + if unknown_edge_types: + raise ValueError( + f"weight_edge_feat_name contains unknown edge types: {unknown_edge_types}" + ) + for edge_type, feature_name in weight_edge_feat_name.items(): + if feature_name not in edge_entity_info[edge_type].feature_keys: + raise ValueError( + f"Sampling-weight field '{feature_name}' for edge type " + f"{edge_type} must be an unquantized raw edge feature." + ) logger.info(f"Rank {rank} starting loading torch tensors from serialized info ...") start_time = time.time() diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 89c863d07..3eaf8ad73 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -144,11 +144,12 @@ def run_distributed_partitioner( ( output_edge_index, output_edge_features, - _, + output_edge_quantized_features, output_edge_partition_book, ) = dist_partitioner.partition_edge_index_and_edge_features( node_partition_book=output_node_partition_book ) + assert output_edge_quantized_features is None dist_partitioner.register_node_features(node_features=node_features) dist_partitioner.register_node_quantized_features( From 2f5128d7168848e976770c2d1825658eecc2d6c4 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 17 Aug 2026 17:08:49 +0000 Subject: [PATCH 5/8] Strengthen edge quantization partitioner tests --- .../run_distributed_partitioner.py | 31 ++++---- .../distributed_partitioner_test.py | 73 ++++++++++++------- 2 files changed, 65 insertions(+), 39 deletions(-) diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 3eaf8ad73..c16a3469b 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -22,9 +22,7 @@ class InputDataStrategy(Enum): REGISTER_EDGE_WEIGHTS_WITHOUT_EDGE_FEATURES = ( "REGISTER_EDGE_WEIGHTS_WITHOUT_EDGE_FEATURES" ) - REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES = ( - "REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES" - ) + REGISTER_EDGE_QUANTIZED_FEATURES = "REGISTER_EDGE_QUANTIZED_FEATURES" def run_distributed_partitioner( @@ -98,10 +96,7 @@ def run_distributed_partitioner( init_rpc(master_addr=master_addr, master_port=master_port, num_rpc_threads=4) dist_partitioner: DistPartitioner - if ( - input_data_strategy - == InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES - ): + if input_data_strategy == InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES: dist_partitioner = partitioner_class( should_assign_edges_by_src_node=should_assign_edges_by_src_node, ) @@ -110,12 +105,19 @@ def run_distributed_partitioner( edge_quantized_features: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] if isinstance(edge_index, dict): edge_index_by_type = cast(dict[EdgeType, torch.Tensor], edge_index) + assert isinstance(edge_features, dict) edge_quantized_features = { - edge_type: indices[0].to(torch.uint8).unsqueeze(1) + edge_type: torch.stack( + (indices[0] * 3 + 17, indices[0] * 5 + 29), dim=1 + ).to(torch.uint8) for edge_type, indices in edge_index_by_type.items() + if edge_type in edge_features } else: - edge_quantized_features = edge_index[0].to(torch.uint8).unsqueeze(1) + edge_quantized_features = torch.stack( + (edge_index[0] * 3 + 17, edge_index[0] * 5 + 29), dim=1 + ).to(torch.uint8) + dist_partitioner.register_edge_features(edge_features=edge_features) dist_partitioner.register_edge_quantized_features( edge_quantized_features=edge_quantized_features ) @@ -149,7 +151,6 @@ def run_distributed_partitioner( ) = dist_partitioner.partition_edge_index_and_edge_features( node_partition_book=output_node_partition_book ) - assert output_edge_quantized_features is None dist_partitioner.register_node_features(node_features=node_features) dist_partitioner.register_node_quantized_features( @@ -191,6 +192,7 @@ def run_distributed_partitioner( partitioned_node_quantized_features=output_node_quantized_features, partitioned_node_labels=output_node_labels, partitioned_edge_features=output_edge_features, + partitioned_edge_quantized_features=output_edge_quantized_features, partitioned_positive_labels=output_positive_labels, partitioned_negative_labels=output_negative_labels, ) @@ -206,9 +208,9 @@ def run_distributed_partitioner( dist_partitioner.register_edge_index(edge_index=edge_index) del edge_index ( - output_graph, + output_edge_index, output_edge_features, - _, + output_edge_quantized_features, output_edge_partition_book, ) = dist_partitioner.partition_edge_index_and_edge_features( node_partition_book=output_node_partition_book @@ -217,10 +219,11 @@ def run_distributed_partitioner( partition_output = PartitionOutput( node_partition_book=output_node_partition_book, edge_partition_book=output_edge_partition_book, - partitioned_edge_index=output_graph, + partitioned_edge_index=output_edge_index, partitioned_node_features=None, partitioned_node_labels=None, - partitioned_edge_features=None, + partitioned_edge_features=output_edge_features, + partitioned_edge_quantized_features=output_edge_quantized_features, partitioned_positive_labels=None, partitioned_negative_labels=None, ) diff --git a/tests/unit/distributed/distributed_partitioner_test.py b/tests/unit/distributed/distributed_partitioner_test.py index f3aaea091..578b8814b 100644 --- a/tests/unit/distributed/distributed_partitioner_test.py +++ b/tests/unit/distributed/distributed_partitioner_test.py @@ -687,17 +687,17 @@ def _assert_label_outputs( expected_pb_dtype=torch.int64, ), param( - "Homogeneous packed-edge-only tensor partitioning", + "Homogeneous raw and quantized edge feature tensor partitioning", is_heterogeneous=False, - input_data_strategy=InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES, + input_data_strategy=InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES, should_assign_edges_by_src_node=True, partitioner_class=DistPartitioner, expected_pb_dtype=torch.uint8, ), param( - "Homogeneous packed-edge-only range partitioning", - is_heterogeneous=False, - input_data_strategy=InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES, + "Heterogeneous raw and quantized edge feature range partitioning", + is_heterogeneous=True, + input_data_strategy=InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES, should_assign_edges_by_src_node=True, partitioner_class=DistRangePartitioner, expected_pb_dtype=torch.int64, @@ -772,9 +772,8 @@ def test_partitioning_correctness( else: expected_edge_feat_types = [USER_TO_USER_EDGE_TYPE] - is_packed_edge_only = ( - input_data_strategy - == InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES_WITHOUT_EDGE_FEATURES + has_edge_quantized_features = ( + input_data_strategy == InputDataStrategy.REGISTER_EDGE_QUANTIZED_FEATURES ) for rank, partition_output in output_dict.items(): @@ -801,42 +800,66 @@ def test_partitioning_correctness( graph.edge_index ) - if is_packed_edge_only: + if has_edge_quantized_features: self.assertIsNotNone(partition_output.edge_partition_book) - self.assertIsNone(partition_output.partitioned_edge_features) + assert partition_output.partitioned_edge_features is not None + self._assert_edge_feature_outputs( + rank=rank, + is_heterogeneous=is_heterogeneous, + is_range_based_partition=is_range_based_partition, + should_assign_edges_by_src_node=should_assign_edges_by_src_node, + output_graph=partitioned_edge_index, + output_edge_feat=partition_output.partitioned_edge_features, + expected_edge_types=MOCKED_HETEROGENEOUS_EDGE_TYPES, + ) self.assertIsNotNone( partition_output.partitioned_edge_quantized_features ) packed_features = partition_output.partitioned_edge_quantized_features - assert isinstance(packed_features, FeaturePartitionData) - assert isinstance(partitioned_edge_index, GraphPartitionData) - self.assertEqual(packed_features.feats.dtype, torch.uint8) - self.assertEqual( - packed_features.feats.size(0), - partitioned_edge_index.edge_index.size(1), - ) - assert partitioned_edge_index.edge_ids is not None - if packed_features.ids is not None: - self.assert_tensor_equality( - tensor_a=packed_features.ids, - tensor_b=partitioned_edge_index.edge_ids, + if isinstance(packed_features, abc.Mapping): + assert isinstance(partitioned_edge_index, abc.Mapping) + self.assertEqual(set(packed_features), {USER_TO_USER_EDGE_TYPE}) + packed_feature_items = packed_features.items() + else: + assert isinstance(partitioned_edge_index, GraphPartitionData) + packed_feature_items = [(USER_TO_USER_EDGE_TYPE, packed_features)] + for edge_type, edge_type_features in packed_feature_items: + edge_type_graph = ( + partitioned_edge_index[edge_type] + if isinstance(partitioned_edge_index, abc.Mapping) + else partitioned_edge_index ) - for index, edge_id in enumerate(partitioned_edge_index.edge_ids): + self.assertEqual(edge_type_features.feats.dtype, torch.uint8) + expected_features = torch.stack( + ( + edge_type_graph.edge_index[0] * 3 + 17, + edge_type_graph.edge_index[0] * 5 + 29, + ), + dim=1, + ).to(torch.uint8) self.assert_tensor_equality( - tensor_a=packed_features.feats[index], - tensor_b=edge_id.to(torch.uint8).unsqueeze(0), + tensor_a=edge_type_features.feats, + tensor_b=expected_features, ) + if edge_type_features.ids is not None: + assert edge_type_graph.edge_ids is not None + self.assert_tensor_equality( + tensor_a=edge_type_features.ids, + tensor_b=edge_type_graph.edge_ids, + ) elif ( input_data_strategy == InputDataStrategy.REGISTER_MINIMAL_ENTITIES_SEPARATELY ): self.assertIsNone(partition_output.edge_partition_book) self.assertIsNone(partition_output.partitioned_edge_features) + self.assertIsNone(partition_output.partitioned_edge_quantized_features) self.assertIsNone(partition_output.partitioned_node_features) self.assertIsNone(partition_output.partitioned_node_labels) self.assertIsNone(partition_output.partitioned_positive_labels) self.assertIsNone(partition_output.partitioned_negative_labels) else: + self.assertIsNone(partition_output.partitioned_edge_quantized_features) assert partition_output.edge_partition_book is not None, ( f"Must create edge partition book for strategy {input_data_strategy.value}" ) From 788a04c442b0e186b878389822529b8fde5f2e16 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 18 Aug 2026 15:24:52 +0000 Subject: [PATCH 6/8] Release partitioner source tensors promptly --- gigl/distributed/dist_partitioner.py | 2 ++ gigl/distributed/dist_range_partitioner.py | 11 ++++++++++- 2 files changed, 12 insertions(+), 1 deletion(-) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 9549c9e3f..f706de514 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -1400,6 +1400,8 @@ def _edge_feat_weight_pfn( generate_pb=False, ) + input_parts.clear() + del input_parts del edge_ids del self._edge_ids[edge_type] if len(self._edge_ids) == 0: diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index 170d5de09..819e75ef1 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -313,7 +313,16 @@ def edge_partition_fn(rank_indices, _): partition_function=edge_partition_fn, ) - del input_data, edge_index, target_indices, edge_feat, edge_weights_tensor + input_parts.clear() + del ( + input_parts, + input_data, + edge_index, + target_indices, + edge_feat, + edge_quantized_features, + edge_weights_tensor, + ) del self._edge_index[edge_type] if self._edge_feat is not None and edge_type in self._edge_feat: assert self._edge_feat_dim is not None From f826e5c6485931cd8ba482f6d1e73fe414650890 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 18 Aug 2026 15:29:03 +0000 Subject: [PATCH 7/8] Document partition result dataclass follow-up --- gigl/distributed/dist_partitioner.py | 2 ++ gigl/distributed/dist_range_partitioner.py | 2 ++ 2 files changed, 4 insertions(+) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index f706de514..663645cb8 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -1001,6 +1001,7 @@ def _node_pfn(n_ids, _): return node_partition_book + # TODO: Create a dataclass for this positional three-value partition result. def _partition_node_features_and_labels( self, node_partition_book: dict[NodeType, PartitionBook], @@ -1228,6 +1229,7 @@ def _node_feature_partition_fn(node_feature_ids, _): node_label_partition_data, ) + # TODO: Create a dataclass for this positional four-value partition result. def _partition_edge_index_and_edge_features( self, node_partition_book: dict[NodeType, PartitionBook], diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index 819e75ef1..479e9902f 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -124,6 +124,7 @@ def _partition_node(self, node_type: NodeType) -> PartitionBook: return node_partition_book + # TODO: Create a dataclass for this positional three-value partition result. def _partition_node_features_and_labels( self, node_partition_book: dict[NodeType, PartitionBook], @@ -210,6 +211,7 @@ def _partition_node_features_and_labels( partitioned_node_label_data, ) + # TODO: Create a dataclass for this positional four-value partition result. def _partition_edge_index_and_edge_features( self, node_partition_book: dict[NodeType, PartitionBook], From 116dd1341ba33bd1f9845ef30fc3c93f2e29efcd Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 18 Aug 2026 15:35:06 +0000 Subject: [PATCH 8/8] Reject ambiguous heterogeneous sampling weights --- gigl/common/data/load_torch_tensors.py | 5 +++ tests/unit/common/data/dataloaders_test.py | 38 ++++++++++++++++++++++ 2 files changed, 43 insertions(+) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 74d98ff4b..309f21ea0 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -170,6 +170,11 @@ def remove_sampling_weight_from_edge_quantization_metadata( dict[EdgeType, FeatureQuantizationMetadata], quantization_metadata ) if isinstance(weight_edge_feat_name, str): + if len(edge_info_by_type) != 1: + raise ValueError( + "weight_edge_feat_name must be dict[EdgeType, str] for " + "heterogeneous graph with multiple edge types" + ) edge_type = next(iter(edge_info_by_type)) weight_by_type: dict[EdgeType, str] = {edge_type: weight_edge_feat_name} else: diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 12b5de2e5..b9d3b6319 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -21,6 +21,7 @@ load_torch_tensors_from_tf_record, remove_sampling_weight_from_edge_quantization_metadata, ) +from gigl.src.common.types.graph_data import EdgeType from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper from gigl.src.data_preprocessor.lib.types import FeatureSpecDict from gigl.src.mocking.lib.versioning import ( @@ -731,6 +732,43 @@ def test_sampling_weight_removal_updates_edge_quantization_metadata( ), ) + def test_sampling_weight_removal_rejects_scalar_name_for_multiple_edge_types( + self, + ) -> None: + missing_path = UriFactory.create_uri("/does/not/exist") + edge_type_a = EdgeType("node", "relation_a", "node") + edge_type_b = EdgeType("node", "relation_b", "node") + edge_info = SerializedTFRecordInfo( + tfrecord_uri_prefix=missing_path, + feature_spec={ + "src_id": tf.io.FixedLenFeature([], tf.int64), + "dst_id": tf.io.FixedLenFeature([], tf.int64), + "weight": tf.io.FixedLenFeature([], tf.float32), + }, + feature_keys=["weight"], + feature_dim=1, + entity_key=("src_id", "dst_id"), + ) + serialized_graph_metadata = SerializedGraphMetadata( + node_entity_info=edge_info, + edge_entity_info={edge_type_a: edge_info, edge_type_b: edge_info}, + edge_quantization_metadata={ + edge_type_b: FeatureQuantizationMetadata( + bits=2, + feature_dim=1, + quantized_feature_indices=(0,), + clip_min=0.0, + clip_max=3.0, + ) + }, + ) + + with self.assertRaises(ValueError): + remove_sampling_weight_from_edge_quantization_metadata( + serialized_graph_metadata=serialized_graph_metadata, + weight_edge_feat_name="weight", + ) + def test_load_edge_weights_multidim_feature(self): """Weight column offset is correct when a preceding feature key is multi-dimensional.