From 97f938a637437a41853515ec3d38f8dd409c94df Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 20:42:29 +0000 Subject: [PATCH 01/78] Compile new protos --- .../research/gbml/preprocessed_metadata.proto | 30 + .../PreprocessedMetadata.scala | 639 +++++++++++++++++- .../PreprocessedMetadataProto.scala | 72 +- .../PreprocessedMetadata.scala | 639 +++++++++++++++++- .../PreprocessedMetadataProto.scala | 72 +- .../gbml/preprocessed_metadata_pb2.py | 57 +- .../gbml/preprocessed_metadata_pb2.pyi | 75 +- 7 files changed, 1480 insertions(+), 104 deletions(-) diff --git a/proto/snapchat/research/gbml/preprocessed_metadata.proto b/proto/snapchat/research/gbml/preprocessed_metadata.proto index b1b7c9e15..d7dfe3469 100644 --- a/proto/snapchat/research/gbml/preprocessed_metadata.proto +++ b/proto/snapchat/research/gbml/preprocessed_metadata.proto @@ -3,6 +3,34 @@ syntax = "proto3"; package snapchat.research.gbml; message PreprocessedMetadata{ + message MultiBitQuantizationState{ + // Lower clipping bound; dequantized value for linear code 0. + float clip_min = 1; + // Upper clipping bound; dequantized value for max linear code. + float clip_max = 2; + // Quantization level bit-width + uint32 bits = 3; + } + + message SingleBitQuantizationState{ + // Mean value for negative features, produced by packed bit/code 0. + float neg_mean = 1; + // Mean value for positive features, produced by packed bit/code 1. + float pos_mean = 2; + } + + message FeatureQuantizationMetadata{ + // Field in output TFRecords that stores packed uint8 features. + string packed_feature_key = 1; + // Original feature indices stored in packed_feature_key. + repeated uint32 quantized_feature_indices = 2; + // Stats required to dequantize features for the selected bit width. + oneof state { + MultiBitQuantizationState multi_bit_state = 4; + SingleBitQuantizationState single_bit_state = 5; + } + } + // Houses metadata about node TFTransform output from DataPreprocessor. message NodeMetadataOutput{ // The field in output TFRecords which references the node identifier. @@ -23,6 +51,8 @@ message PreprocessedMetadata{ optional uint32 feature_dim = 8; // Contains categorical feature vocabularies string transform_fn_assets_uri = 9; + // Optional quantized node feature metadata. + FeatureQuantizationMetadata quantized_feature_metadata = 10; } // Houses metadata of edge features 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 80160636b..7a9012ffa 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 @@ -132,6 +132,9 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r } lazy val nestedMessagesCompanions: Seq[_root_.scalapb.GeneratedMessageCompanion[_ <: _root_.scalapb.GeneratedMessage]] = Seq[_root_.scalapb.GeneratedMessageCompanion[_ <: _root_.scalapb.GeneratedMessage]]( + _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState, + _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState, + _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata, _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput, _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo, _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataOutput, @@ -143,6 +146,577 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r condensedNodeTypeToPreprocessedMetadata = _root_.scala.collection.immutable.Map.empty, condensedEdgeTypeToPreprocessedMetadata = _root_.scala.collection.immutable.Map.empty ) + /** @param clipMin + * Lower clipping bound; dequantized value for linear code 0. + * @param clipMax + * Upper clipping bound; dequantized value for max linear code. + * @param bits + * Quantization level bit-width + */ + @SerialVersionUID(0L) + final case class MultiBitQuantizationState( + clipMin: _root_.scala.Float = 0.0f, + clipMax: _root_.scala.Float = 0.0f, + bits: _root_.scala.Int = 0, + unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty + ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[MultiBitQuantizationState] { + @transient + private[this] var __serializedSizeMemoized: _root_.scala.Int = 0 + private[this] def __computeSerializedSize(): _root_.scala.Int = { + var __size = 0 + + { + val __value = clipMin + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(1, __value) + } + }; + + { + val __value = clipMax + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(2, __value) + } + }; + + { + val __value = bits + if (__value != 0) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeUInt32Size(3, __value) + } + }; + __size += unknownFields.serializedSize + __size + } + override def serializedSize: _root_.scala.Int = { + var __size = __serializedSizeMemoized + if (__size == 0) { + __size = __computeSerializedSize() + 1 + __serializedSizeMemoized = __size + } + __size - 1 + + } + def writeTo(`_output__`: _root_.com.google.protobuf.CodedOutputStream): _root_.scala.Unit = { + { + val __v = clipMin + if (__v != 0.0f) { + _output__.writeFloat(1, __v) + } + }; + { + val __v = clipMax + if (__v != 0.0f) { + _output__.writeFloat(2, __v) + } + }; + { + val __v = bits + if (__v != 0) { + _output__.writeUInt32(3, __v) + } + }; + unknownFields.writeTo(_output__) + } + def withClipMin(__v: _root_.scala.Float): MultiBitQuantizationState = copy(clipMin = __v) + def withClipMax(__v: _root_.scala.Float): MultiBitQuantizationState = copy(clipMax = __v) + def withBits(__v: _root_.scala.Int): MultiBitQuantizationState = copy(bits = __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 = { + (__fieldNumber: @_root_.scala.unchecked) match { + case 1 => { + val __t = clipMin + if (__t != 0.0f) __t else null + } + case 2 => { + val __t = clipMax + if (__t != 0.0f) __t else null + } + case 3 => { + val __t = bits + if (__t != 0) __t else null + } + } + } + def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { + _root_.scala.Predef.require(__field.containingMessage eq companion.scalaDescriptor) + (__field.number: @_root_.scala.unchecked) match { + case 1 => _root_.scalapb.descriptors.PFloat(clipMin) + case 2 => _root_.scalapb.descriptors.PFloat(clipMax) + case 3 => _root_.scalapb.descriptors.PInt(bits) + } + } + def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) + def companion: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState.type = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState + // @@protoc_insertion_point(GeneratedMessage[snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationState]) + } + + object MultiBitQuantizationState extends scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] { + implicit def messageCompanion: scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = this + def parseFrom(`_input__`: _root_.com.google.protobuf.CodedInputStream): snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState = { + var __clipMin: _root_.scala.Float = 0.0f + var __clipMax: _root_.scala.Float = 0.0f + var __bits: _root_.scala.Int = 0 + var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null + var _done__ = false + while (!_done__) { + val _tag__ = _input__.readTag() + _tag__ match { + case 0 => _done__ = true + case 13 => + __clipMin = _input__.readFloat() + case 21 => + __clipMax = _input__.readFloat() + case 24 => + __bits = _input__.readUInt32() + case tag => + if (_unknownFields__ == null) { + _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() + } + _unknownFields__.parseField(tag, _input__) + } + } + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState( + clipMin = __clipMin, + clipMax = __clipMax, + bits = __bits, + unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() + ) + } + implicit def messageReads: _root_.scalapb.descriptors.Reads[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = _root_.scalapb.descriptors.Reads{ + case _root_.scalapb.descriptors.PMessage(__fieldsMap) => + _root_.scala.Predef.require(__fieldsMap.keys.forall(_.containingMessage eq scalaDescriptor), "FieldDescriptor does not match message type.") + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState( + clipMin = __fieldsMap.get(scalaDescriptor.findFieldByNumber(1).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f), + clipMax = __fieldsMap.get(scalaDescriptor.findFieldByNumber(2).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f), + bits = __fieldsMap.get(scalaDescriptor.findFieldByNumber(3).get).map(_.as[_root_.scala.Int]).getOrElse(0) + ) + case _ => throw new RuntimeException("Expected PMessage") + } + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(0) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(0) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + 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.MultiBitQuantizationState( + clipMin = 0.0f, + clipMax = 0.0f, + bits = 0 + ) + implicit class MultiBitQuantizationStateLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState](_l) { + def clipMin: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.clipMin)((c_, f_) => c_.copy(clipMin = f_)) + def clipMax: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.clipMax)((c_, f_) => c_.copy(clipMax = f_)) + def bits: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Int] = field(_.bits)((c_, f_) => c_.copy(bits = f_)) + } + final val CLIP_MIN_FIELD_NUMBER = 1 + final val CLIP_MAX_FIELD_NUMBER = 2 + final val BITS_FIELD_NUMBER = 3 + def of( + clipMin: _root_.scala.Float, + clipMax: _root_.scala.Float, + bits: _root_.scala.Int + ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState( + clipMin, + clipMax, + bits + ) + // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationState]) + } + + /** @param negMean + * Mean value for negative features, produced by packed bit/code 0. + * @param posMean + * Mean value for positive features, produced by packed bit/code 1. + */ + @SerialVersionUID(0L) + final case class SingleBitQuantizationState( + negMean: _root_.scala.Float = 0.0f, + posMean: _root_.scala.Float = 0.0f, + unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty + ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[SingleBitQuantizationState] { + @transient + private[this] var __serializedSizeMemoized: _root_.scala.Int = 0 + private[this] def __computeSerializedSize(): _root_.scala.Int = { + var __size = 0 + + { + val __value = negMean + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(1, __value) + } + }; + + { + val __value = posMean + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(2, __value) + } + }; + __size += unknownFields.serializedSize + __size + } + override def serializedSize: _root_.scala.Int = { + var __size = __serializedSizeMemoized + if (__size == 0) { + __size = __computeSerializedSize() + 1 + __serializedSizeMemoized = __size + } + __size - 1 + + } + def writeTo(`_output__`: _root_.com.google.protobuf.CodedOutputStream): _root_.scala.Unit = { + { + val __v = negMean + if (__v != 0.0f) { + _output__.writeFloat(1, __v) + } + }; + { + val __v = posMean + if (__v != 0.0f) { + _output__.writeFloat(2, __v) + } + }; + unknownFields.writeTo(_output__) + } + def withNegMean(__v: _root_.scala.Float): SingleBitQuantizationState = copy(negMean = __v) + def withPosMean(__v: _root_.scala.Float): SingleBitQuantizationState = copy(posMean = __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 = { + (__fieldNumber: @_root_.scala.unchecked) match { + case 1 => { + val __t = negMean + if (__t != 0.0f) __t else null + } + case 2 => { + val __t = posMean + if (__t != 0.0f) __t else null + } + } + } + def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { + _root_.scala.Predef.require(__field.containingMessage eq companion.scalaDescriptor) + (__field.number: @_root_.scala.unchecked) match { + case 1 => _root_.scalapb.descriptors.PFloat(negMean) + case 2 => _root_.scalapb.descriptors.PFloat(posMean) + } + } + def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) + def companion: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState.type = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState + // @@protoc_insertion_point(GeneratedMessage[snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationState]) + } + + object SingleBitQuantizationState extends scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] { + implicit def messageCompanion: scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = this + def parseFrom(`_input__`: _root_.com.google.protobuf.CodedInputStream): snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState = { + var __negMean: _root_.scala.Float = 0.0f + var __posMean: _root_.scala.Float = 0.0f + var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null + var _done__ = false + while (!_done__) { + val _tag__ = _input__.readTag() + _tag__ match { + case 0 => _done__ = true + case 13 => + __negMean = _input__.readFloat() + case 21 => + __posMean = _input__.readFloat() + case tag => + if (_unknownFields__ == null) { + _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() + } + _unknownFields__.parseField(tag, _input__) + } + } + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState( + negMean = __negMean, + posMean = __posMean, + unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() + ) + } + implicit def messageReads: _root_.scalapb.descriptors.Reads[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = _root_.scalapb.descriptors.Reads{ + case _root_.scalapb.descriptors.PMessage(__fieldsMap) => + _root_.scala.Predef.require(__fieldsMap.keys.forall(_.containingMessage eq scalaDescriptor), "FieldDescriptor does not match message type.") + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState( + negMean = __fieldsMap.get(scalaDescriptor.findFieldByNumber(1).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f), + posMean = __fieldsMap.get(scalaDescriptor.findFieldByNumber(2).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f) + ) + case _ => throw new RuntimeException("Expected PMessage") + } + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(1) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(1) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + 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.SingleBitQuantizationState( + negMean = 0.0f, + posMean = 0.0f + ) + implicit class SingleBitQuantizationStateLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState](_l) { + def negMean: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.negMean)((c_, f_) => c_.copy(negMean = f_)) + def posMean: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.posMean)((c_, f_) => c_.copy(posMean = f_)) + } + final val NEG_MEAN_FIELD_NUMBER = 1 + final val POS_MEAN_FIELD_NUMBER = 2 + def of( + negMean: _root_.scala.Float, + posMean: _root_.scala.Float + ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState( + negMean, + posMean + ) + // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationState]) + } + + /** @param packedFeatureKey + * Field in output TFRecords that stores packed uint8 features. + * @param quantizedFeatureIndices + * Original feature indices stored in packed_feature_key. + */ + @SerialVersionUID(0L) + final case class FeatureQuantizationMetadata( + packedFeatureKey: _root_.scala.Predef.String = "", + quantizedFeatureIndices: _root_.scala.Seq[_root_.scala.Int] = _root_.scala.Seq.empty, + state: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty, + unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty + ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[FeatureQuantizationMetadata] { + private[this] def quantizedFeatureIndicesSerializedSize = { + if (__quantizedFeatureIndicesSerializedSizeField == 0) __quantizedFeatureIndicesSerializedSizeField = { + var __s: _root_.scala.Int = 0 + quantizedFeatureIndices.foreach(__i => __s += _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__i)) + __s + } + __quantizedFeatureIndicesSerializedSizeField + } + @transient private[this] var __quantizedFeatureIndicesSerializedSizeField: _root_.scala.Int = 0 + @transient + private[this] var __serializedSizeMemoized: _root_.scala.Int = 0 + private[this] def __computeSerializedSize(): _root_.scala.Int = { + var __size = 0 + + { + val __value = packedFeatureKey + if (!__value.isEmpty) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeStringSize(1, __value) + } + }; + if (quantizedFeatureIndices.nonEmpty) { + val __localsize = quantizedFeatureIndicesSerializedSize + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__localsize) + __localsize + } + if (state.multiBitState.isDefined) { + val __value = state.multiBitState.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; + if (state.singleBitState.isDefined) { + val __value = state.singleBitState.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; + __size += unknownFields.serializedSize + __size + } + override def serializedSize: _root_.scala.Int = { + var __size = __serializedSizeMemoized + if (__size == 0) { + __size = __computeSerializedSize() + 1 + __serializedSizeMemoized = __size + } + __size - 1 + + } + def writeTo(`_output__`: _root_.com.google.protobuf.CodedOutputStream): _root_.scala.Unit = { + { + val __v = packedFeatureKey + if (!__v.isEmpty) { + _output__.writeString(1, __v) + } + }; + if (quantizedFeatureIndices.nonEmpty) { + _output__.writeTag(2, 2) + _output__.writeUInt32NoTag(quantizedFeatureIndicesSerializedSize) + quantizedFeatureIndices.foreach(_output__.writeUInt32NoTag) + }; + state.multiBitState.foreach { __v => + val __m = __v + _output__.writeTag(4, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; + state.singleBitState.foreach { __v => + val __m = __v + _output__.writeTag(5, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; + unknownFields.writeTo(_output__) + } + def withPackedFeatureKey(__v: _root_.scala.Predef.String): FeatureQuantizationMetadata = copy(packedFeatureKey = __v) + def clearQuantizedFeatureIndices = copy(quantizedFeatureIndices = _root_.scala.Seq.empty) + def addQuantizedFeatureIndices(__vs: _root_.scala.Int *): FeatureQuantizationMetadata = addAllQuantizedFeatureIndices(__vs) + def addAllQuantizedFeatureIndices(__vs: Iterable[_root_.scala.Int]): FeatureQuantizationMetadata = copy(quantizedFeatureIndices = quantizedFeatureIndices ++ __vs) + def withQuantizedFeatureIndices(__v: _root_.scala.Seq[_root_.scala.Int]): FeatureQuantizationMetadata = copy(quantizedFeatureIndices = __v) + def getMultiBitState: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState = state.multiBitState.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState.defaultInstance) + def withMultiBitState(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState): FeatureQuantizationMetadata = copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(__v)) + def getSingleBitState: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState = state.singleBitState.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState.defaultInstance) + def withSingleBitState(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState): FeatureQuantizationMetadata = copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(__v)) + def clearState: FeatureQuantizationMetadata = copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty) + def withState(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State): FeatureQuantizationMetadata = copy(state = __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 = { + (__fieldNumber: @_root_.scala.unchecked) match { + case 1 => { + val __t = packedFeatureKey + if (__t != "") __t else null + } + case 2 => quantizedFeatureIndices + case 4 => state.multiBitState.orNull + case 5 => state.singleBitState.orNull + } + } + def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { + _root_.scala.Predef.require(__field.containingMessage eq companion.scalaDescriptor) + (__field.number: @_root_.scala.unchecked) match { + case 1 => _root_.scalapb.descriptors.PString(packedFeatureKey) + case 2 => _root_.scalapb.descriptors.PRepeated(quantizedFeatureIndices.iterator.map(_root_.scalapb.descriptors.PInt(_)).toVector) + case 4 => state.multiBitState.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) + case 5 => state.singleBitState.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) + } + } + def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) + def companion: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.type = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata + // @@protoc_insertion_point(GeneratedMessage[snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadata]) + } + + object FeatureQuantizationMetadata extends scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] { + implicit def messageCompanion: scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = this + def parseFrom(`_input__`: _root_.com.google.protobuf.CodedInputStream): snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = { + var __packedFeatureKey: _root_.scala.Predef.String = "" + val __quantizedFeatureIndices: _root_.scala.collection.immutable.VectorBuilder[_root_.scala.Int] = new _root_.scala.collection.immutable.VectorBuilder[_root_.scala.Int] + var __state: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty + var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null + var _done__ = false + while (!_done__) { + val _tag__ = _input__.readTag() + _tag__ match { + case 0 => _done__ = true + case 10 => + __packedFeatureKey = _input__.readStringRequireUtf8() + case 16 => + __quantizedFeatureIndices += _input__.readUInt32() + case 18 => { + val length = _input__.readRawVarint32() + val oldLimit = _input__.pushLimit(length) + while (_input__.getBytesUntilLimit > 0) { + __quantizedFeatureIndices += _input__.readUInt32() + } + _input__.popLimit(oldLimit) + } + case 34 => + __state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(__state.multiBitState.fold(_root_.scalapb.LiteParser.readMessage[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState](_input__))(_root_.scalapb.LiteParser.readMessage(_input__, _))) + case 42 => + __state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(__state.singleBitState.fold(_root_.scalapb.LiteParser.readMessage[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState](_input__))(_root_.scalapb.LiteParser.readMessage(_input__, _))) + case tag => + if (_unknownFields__ == null) { + _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() + } + _unknownFields__.parseField(tag, _input__) + } + } + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata( + packedFeatureKey = __packedFeatureKey, + quantizedFeatureIndices = __quantizedFeatureIndices.result(), + state = __state, + unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() + ) + } + implicit def messageReads: _root_.scalapb.descriptors.Reads[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = _root_.scalapb.descriptors.Reads{ + case _root_.scalapb.descriptors.PMessage(__fieldsMap) => + _root_.scala.Predef.require(__fieldsMap.keys.forall(_.containingMessage eq scalaDescriptor), "FieldDescriptor does not match message type.") + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata( + packedFeatureKey = __fieldsMap.get(scalaDescriptor.findFieldByNumber(1).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), + quantizedFeatureIndices = __fieldsMap.get(scalaDescriptor.findFieldByNumber(2).get).map(_.as[_root_.scala.Seq[_root_.scala.Int]]).getOrElse(_root_.scala.Seq.empty), + state = __fieldsMap.get(scalaDescriptor.findFieldByNumber(4).get).flatMap(_.as[_root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState]]).map(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(_)) + .orElse[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State](__fieldsMap.get(scalaDescriptor.findFieldByNumber(5).get).flatMap(_.as[_root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState]]).map(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(_))) + .getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty) + ) + case _ => throw new RuntimeException("Expected PMessage") + } + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(2) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(2) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { + var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null + (__number: @_root_.scala.unchecked) match { + case 4 => __out = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState + case 5 => __out = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState + } + __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.FeatureQuantizationMetadata( + packedFeatureKey = "", + quantizedFeatureIndices = _root_.scala.Seq.empty, + state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty + ) + sealed trait State extends _root_.scalapb.GeneratedOneof { + def isEmpty: _root_.scala.Boolean = false + def isDefined: _root_.scala.Boolean = true + def isMultiBitState: _root_.scala.Boolean = false + def isSingleBitState: _root_.scala.Boolean = false + def multiBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = _root_.scala.None + def singleBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = _root_.scala.None + } + object State { + @SerialVersionUID(0L) + case object Empty extends snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State { + type ValueType = _root_.scala.Nothing + override def isEmpty: _root_.scala.Boolean = true + override def isDefined: _root_.scala.Boolean = false + override def number: _root_.scala.Int = 0 + override def value: _root_.scala.Nothing = throw new java.util.NoSuchElementException("Empty.value") + } + + @SerialVersionUID(0L) + final case class MultiBitState(value: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState) extends snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State { + type ValueType = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState + override def isMultiBitState: _root_.scala.Boolean = true + override def multiBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = Some(value) + override def number: _root_.scala.Int = 4 + } + @SerialVersionUID(0L) + final case class SingleBitState(value: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState) extends snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State { + type ValueType = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState + override def isSingleBitState: _root_.scala.Boolean = true + override def singleBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = Some(value) + override def number: _root_.scala.Int = 5 + } + } + implicit class FeatureQuantizationMetadataLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata](_l) { + def packedFeatureKey: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Predef.String] = field(_.packedFeatureKey)((c_, f_) => c_.copy(packedFeatureKey = f_)) + def quantizedFeatureIndices: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Seq[_root_.scala.Int]] = field(_.quantizedFeatureIndices)((c_, f_) => c_.copy(quantizedFeatureIndices = f_)) + def multiBitState: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = field(_.getMultiBitState)((c_, f_) => c_.copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(f_))) + def singleBitState: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = field(_.getSingleBitState)((c_, f_) => c_.copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(f_))) + def state: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State] = field(_.state)((c_, f_) => c_.copy(state = f_)) + } + final val PACKED_FEATURE_KEY_FIELD_NUMBER = 1 + final val QUANTIZED_FEATURE_INDICES_FIELD_NUMBER = 2 + final val MULTI_BIT_STATE_FIELD_NUMBER = 4 + final val SINGLE_BIT_STATE_FIELD_NUMBER = 5 + def of( + packedFeatureKey: _root_.scala.Predef.String, + quantizedFeatureIndices: _root_.scala.Seq[_root_.scala.Int], + state: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State + ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata( + packedFeatureKey, + quantizedFeatureIndices, + state + ) + // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadata]) + } + /** Houses metadata about node TFTransform output from DataPreprocessor. * * @param nodeIdKey @@ -163,6 +737,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r * Feature dimension after preprocessing * @param transformFnAssetsUri * Contains categorical feature vocabularies + * @param quantizedFeatureMetadata + * Optional quantized node feature metadata. */ @SerialVersionUID(0L) final case class NodeMetadataOutput( @@ -175,6 +751,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeDataBqTable: _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[NodeMetadataOutput] { @transient @@ -235,6 +812,10 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __size += _root_.com.google.protobuf.CodedOutputStream.computeStringSize(9, __value) } }; + if (quantizedFeatureMetadata.isDefined) { + val __value = quantizedFeatureMetadata.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; __size += unknownFields.serializedSize __size } @@ -296,6 +877,12 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r _output__.writeString(9, __v) } }; + quantizedFeatureMetadata.foreach { __v => + val __m = __v + _output__.writeTag(10, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; unknownFields.writeTo(_output__) } def withNodeIdKey(__v: _root_.scala.Predef.String): NodeMetadataOutput = copy(nodeIdKey = __v) @@ -315,6 +902,9 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r def clearFeatureDim: NodeMetadataOutput = copy(featureDim = _root_.scala.None) def withFeatureDim(__v: _root_.scala.Int): NodeMetadataOutput = copy(featureDim = Option(__v)) def withTransformFnAssetsUri(__v: _root_.scala.Predef.String): NodeMetadataOutput = copy(transformFnAssetsUri = __v) + def getQuantizedFeatureMetadata: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = quantizedFeatureMetadata.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.defaultInstance) + def clearQuantizedFeatureMetadata: NodeMetadataOutput = copy(quantizedFeatureMetadata = _root_.scala.None) + def withQuantizedFeatureMetadata(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata): NodeMetadataOutput = 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 = { @@ -346,6 +936,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r val __t = transformFnAssetsUri if (__t != "") __t else null } + case 10 => quantizedFeatureMetadata.orNull } } def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { @@ -360,6 +951,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r case 7 => _root_.scalapb.descriptors.PString(enumeratedNodeDataBqTable) case 8 => featureDim.map(_root_.scalapb.descriptors.PInt(_)).getOrElse(_root_.scalapb.descriptors.PEmpty) case 9 => _root_.scalapb.descriptors.PString(transformFnAssetsUri) + case 10 => quantizedFeatureMetadata.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) } } def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) @@ -379,6 +971,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r var __enumeratedNodeDataBqTable: _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__) { @@ -403,6 +996,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __featureDim = Option(_input__.readUInt32()) case 74 => __transformFnAssetsUri = _input__.readStringRequireUtf8() + case 82 => + __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() @@ -420,6 +1015,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeDataBqTable = __enumeratedNodeDataBqTable, featureDim = __featureDim, transformFnAssetsUri = __transformFnAssetsUri, + quantizedFeatureMetadata = __quantizedFeatureMetadata, unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() ) } @@ -435,13 +1031,20 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable = __fieldsMap.get(scalaDescriptor.findFieldByNumber(6).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), enumeratedNodeDataBqTable = __fieldsMap.get(scalaDescriptor.findFieldByNumber(7).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), featureDim = __fieldsMap.get(scalaDescriptor.findFieldByNumber(8).get).flatMap(_.as[_root_.scala.Option[_root_.scala.Int]]), - transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(9).get).map(_.as[_root_.scala.Predef.String]).getOrElse("") + transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(9).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), + quantizedFeatureMetadata = __fieldsMap.get(scalaDescriptor.findFieldByNumber(10).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(0) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(0) - def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(3) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(3) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { + var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null + (__number: @_root_.scala.unchecked) match { + case 10 => __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.NodeMetadataOutput( @@ -453,7 +1056,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable = "", enumeratedNodeDataBqTable = "", featureDim = _root_.scala.None, - transformFnAssetsUri = "" + transformFnAssetsUri = "", + quantizedFeatureMetadata = _root_.scala.None ) implicit class NodeMetadataOutputLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput](_l) { def nodeIdKey: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Predef.String] = field(_.nodeIdKey)((c_, f_) => c_.copy(nodeIdKey = f_)) @@ -466,6 +1070,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 NODE_ID_KEY_FIELD_NUMBER = 1 final val FEATURE_KEYS_FIELD_NUMBER = 2 @@ -476,6 +1082,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r final val ENUMERATED_NODE_DATA_BQ_TABLE_FIELD_NUMBER = 7 final val FEATURE_DIM_FIELD_NUMBER = 8 final val TRANSFORM_FN_ASSETS_URI_FIELD_NUMBER = 9 + final val QUANTIZED_FEATURE_METADATA_FIELD_NUMBER = 10 def of( nodeIdKey: _root_.scala.Predef.String, featureKeys: _root_.scala.Seq[_root_.scala.Predef.String], @@ -485,7 +1092,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable: _root_.scala.Predef.String, enumeratedNodeDataBqTable: _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.NodeMetadataOutput = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput( nodeIdKey, featureKeys, @@ -495,7 +1103,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable, enumeratedNodeDataBqTable, featureDim, - transformFnAssetsUri + transformFnAssetsUri, + quantizedFeatureMetadata ) // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.NodeMetadataOutput]) } @@ -742,8 +1351,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) case _ => throw new RuntimeException("Expected PMessage") } - def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(1) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(1) + 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) 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) @@ -985,8 +1594,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) case _ => throw new RuntimeException("Expected PMessage") } - def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(2) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(2) + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(5) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(5) def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null (__number: @_root_.scala.unchecked) match { @@ -1148,8 +1757,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) case _ => throw new RuntimeException("Expected PMessage") } - def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(3) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(3) + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(6) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(6) def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null (__number: @_root_.scala.unchecked) match { @@ -1295,8 +1904,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) 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 javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(7) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(7) def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null (__number: @_root_.scala.unchecked) match { 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 becc2d068..ad80de0ad 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,41 +14,53 @@ object PreprocessedMetadataProto extends _root_.scalapb.GeneratedFileObject { private lazy val ProtoBytes: _root_.scala.Array[Byte] = scalapb.Encoding.fromBase64(scala.collection.immutable.Seq( """CjJzbmFwY2hhdC9yZXNlYXJjaC9nYm1sL3ByZXByb2Nlc3NlZF9tZXRhZGF0YS5wcm90bxIWc25hcGNoYXQucmVzZWFyY2guZ - 2JtbCL8EwoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY + 2JtbCL9GgoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY XRhGAEgAygLMlkuc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5Db25kZW5zZWROb2RlVHlwZVRvU HJlcHJvY2Vzc2VkTWV0YWRhdGFFbnRyeUIs4j8pEidjb25kZW5zZWROb2RlVHlwZVRvUHJlcHJvY2Vzc2VkTWV0YWRhdGFSJ2Nvb mRlbnNlZE5vZGVUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZGF0YRLmAQosY29uZGVuc2VkX2VkZ2VfdHlwZV90b19wcmVwcm9jZXNzZ WRfbWV0YWRhdGEYAiADKAsyWS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkNvbmRlbnNlZEVkZ 2VUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZGF0YUVudHJ5QiziPykSJ2NvbmRlbnNlZEVkZ2VUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZ - GF0YVInY29uZGVuc2VkRWRnZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhGvkEChJOb2RlTWV0YWRhdGFPdXRwdXQSLgoLbm9kZ - V9pZF9rZXkYASABKAlCDuI/CxIJbm9kZUlkS2V5Uglub2RlSWRLZXkSMwoMZmVhdHVyZV9rZXlzGAIgAygJQhDiPw0SC2ZlYXR1c - mVLZXlzUgtmZWF0dXJlS2V5cxItCgpsYWJlbF9rZXlzGAMgAygJQg7iPwsSCWxhYmVsS2V5c1IJbGFiZWxLZXlzEkYKE3RmcmVjb - 3JkX3VyaV9wcmVmaXgYBCABKAlCFuI/ExIRdGZyZWNvcmRVcmlQcmVmaXhSEXRmcmVjb3JkVXJpUHJlZml4Ei0KCnNjaGVtYV91c - mkYBSABKAlCDuI/CxIJc2NoZW1hVXJpUglzY2hlbWFVcmkSXQocZW51bWVyYXRlZF9ub2RlX2lkc19icV90YWJsZRgGIAEoCUId4 - j8aEhhlbnVtZXJhdGVkTm9kZUlkc0JxVGFibGVSGGVudW1lcmF0ZWROb2RlSWRzQnFUYWJsZRJgCh1lbnVtZXJhdGVkX25vZGVfZ - GF0YV9icV90YWJsZRgHIAEoCUIe4j8bEhllbnVtZXJhdGVkTm9kZURhdGFCcVRhYmxlUhllbnVtZXJhdGVkTm9kZURhdGFCcVRhY - mxlEjUKC2ZlYXR1cmVfZGltGAggASgNQg/iPwwSCmZlYXR1cmVEaW1IAFIKZmVhdHVyZURpbYgBARJQChd0cmFuc2Zvcm1fZm5fY - XNzZXRzX3VyaRgJIAEoCUIZ4j8WEhR0cmFuc2Zvcm1GbkFzc2V0c1VyaVIUdHJhbnNmb3JtRm5Bc3NldHNVcmlCDgoMX2ZlYXR1c - mVfZGltGugDChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZ - UtleXMSLQoKbGFiZWxfa2V5cxgCIAMoCUIO4j8LEglsYWJlbEtleXNSCWxhYmVsS2V5cxJGChN0ZnJlY29yZF91cmlfcHJlZml4G - AMgASgJQhbiPxMSEXRmcmVjb3JkVXJpUHJlZml4UhF0ZnJlY29yZFVyaVByZWZpeBItCgpzY2hlbWFfdXJpGAQgASgJQg7iPwsSC - XNjaGVtYVVyaVIJc2NoZW1hVXJpEmAKHWVudW1lcmF0ZWRfZWRnZV9kYXRhX2JxX3RhYmxlGAUgASgJQh7iPxsSGWVudW1lcmF0Z - WRFZGdlRGF0YUJxVGFibGVSGWVudW1lcmF0ZWRFZGdlRGF0YUJxVGFibGUSNQoLZmVhdHVyZV9kaW0YBiABKA1CD+I/DBIKZmVhd - HVyZURpbUgAUgpmZWF0dXJlRGltiAEBElAKF3RyYW5zZm9ybV9mbl9hc3NldHNfdXJpGAcgASgJQhniPxYSFHRyYW5zZm9ybUZuQ - XNzZXRzVXJpUhR0cmFuc2Zvcm1GbkFzc2V0c1VyaUIOCgxfZmVhdHVyZV9kaW0awgQKEkVkZ2VNZXRhZGF0YU91dHB1dBI4Cg9zc - mNfbm9kZV9pZF9rZXkYASABKAlCEeI/DhIMc3JjTm9kZUlkS2V5UgxzcmNOb2RlSWRLZXkSOAoPZHN0X25vZGVfaWRfa2V5GAIgA - SgJQhHiPw4SDGRzdE5vZGVJZEtleVIMZHN0Tm9kZUlkS2V5EnYKDm1haW5fZWRnZV9pbmZvGAMgASgLMj0uc25hcGNoYXQucmVzZ - WFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFJbmZvQhHiPw4SDG1haW5FZGdlSW5mb1IMbWFpbkVkZ - 2VJbmZvEocBChJwb3NpdGl2ZV9lZGdlX2luZm8YBCABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ld - GFkYXRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQcG9zaXRpdmVFZGdlSW5mb0gAUhBwb3NpdGl2ZUVkZ2VJbmZviAEBEocBChJuZ - WdhdGl2ZV9lZGdlX2luZm8YBSABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkVkZ2VNZ - XRhZGF0YUluZm9CFeI/EhIQbmVnYXRpdmVFZGdlSW5mb0gBUhBuZWdhdGl2ZUVkZ2VJbmZviAEBQhUKE19wb3NpdGl2ZV9lZGdlX - 2luZm9CFQoTX25lZ2F0aXZlX2VkZ2VfaW5mbxqxAQosQ29uZGVuc2VkTm9kZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50c - nkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwc - m9jZXNzZWRNZXRhZGF0YS5Ob2RlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4ARqxAQosQ29uZGVuc2VkRWRnZ - VR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLM - j8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsd - WVSBXZhbHVlOgI4AWIGcHJvdG8z""" + GF0YVInY29uZGVuc2VkRWRnZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhGowBChlNdWx0aUJpdFF1YW50aXphdGlvblN0YXRlE + icKCGNsaXBfbWluGAEgASgCQgziPwkSB2NsaXBNaW5SB2NsaXBNaW4SJwoIY2xpcF9tYXgYAiABKAJCDOI/CRIHY2xpcE1heFIHY + 2xpcE1heBIdCgRiaXRzGAMgASgNQgniPwYSBGJpdHNSBGJpdHMabgoaU2luZ2xlQml0UXVhbnRpemF0aW9uU3RhdGUSJwoIbmVnX + 21lYW4YASABKAJCDOI/CRIHbmVnTWVhblIHbmVnTWVhbhInCghwb3NfbWVhbhgCIAEoAkIM4j8JEgdwb3NNZWFuUgdwb3NNZWFuG + tcDChtGZWF0dXJlUXVhbnRpemF0aW9uTWV0YWRhdGESQwoScGFja2VkX2ZlYXR1cmVfa2V5GAEgASgJQhXiPxISEHBhY2tlZEZlY + XR1cmVLZXlSEHBhY2tlZEZlYXR1cmVLZXkSWAoZcXVhbnRpemVkX2ZlYXR1cmVfaW5kaWNlcxgCIAMoDUIc4j8ZEhdxdWFudGl6Z + WRGZWF0dXJlSW5kaWNlc1IXcXVhbnRpemVkRmVhdHVyZUluZGljZXMShAEKD211bHRpX2JpdF9zdGF0ZRgEIAEoCzJGLnNuYXBja + GF0LnJlc2VhcmNoLmdibWwuUHJlcHJvY2Vzc2VkTWV0YWRhdGEuTXVsdGlCaXRRdWFudGl6YXRpb25TdGF0ZUIS4j8PEg1tdWx0a + UJpdFN0YXRlSABSDW11bHRpQml0U3RhdGUSiAEKEHNpbmdsZV9iaXRfc3RhdGUYBSABKAsyRy5zbmFwY2hhdC5yZXNlYXJjaC5nY + m1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLlNpbmdsZUJpdFF1YW50aXphdGlvblN0YXRlQhPiPxASDnNpbmdsZUJpdFN0YXRlSABSD + nNpbmdsZUJpdFN0YXRlQgcKBXN0YXRlGqEGChJOb2RlTWV0YWRhdGFPdXRwdXQSLgoLbm9kZV9pZF9rZXkYASABKAlCDuI/CxIJb + m9kZUlkS2V5Uglub2RlSWRLZXkSMwoMZmVhdHVyZV9rZXlzGAIgAygJQhDiPw0SC2ZlYXR1cmVLZXlzUgtmZWF0dXJlS2V5cxItC + gpsYWJlbF9rZXlzGAMgAygJQg7iPwsSCWxhYmVsS2V5c1IJbGFiZWxLZXlzEkYKE3RmcmVjb3JkX3VyaV9wcmVmaXgYBCABKAlCF + uI/ExIRdGZyZWNvcmRVcmlQcmVmaXhSEXRmcmVjb3JkVXJpUHJlZml4Ei0KCnNjaGVtYV91cmkYBSABKAlCDuI/CxIJc2NoZW1hV + XJpUglzY2hlbWFVcmkSXQocZW51bWVyYXRlZF9ub2RlX2lkc19icV90YWJsZRgGIAEoCUId4j8aEhhlbnVtZXJhdGVkTm9kZUlkc + 0JxVGFibGVSGGVudW1lcmF0ZWROb2RlSWRzQnFUYWJsZRJgCh1lbnVtZXJhdGVkX25vZGVfZGF0YV9icV90YWJsZRgHIAEoCUIe4 + j8bEhllbnVtZXJhdGVkTm9kZURhdGFCcVRhYmxlUhllbnVtZXJhdGVkTm9kZURhdGFCcVRhYmxlEjUKC2ZlYXR1cmVfZGltGAggA + SgNQg/iPwwSCmZlYXR1cmVEaW1IAFIKZmVhdHVyZURpbYgBARJQChd0cmFuc2Zvcm1fZm5fYXNzZXRzX3VyaRgJIAEoCUIZ4j8WE + hR0cmFuc2Zvcm1GbkFzc2V0c1VyaVIUdHJhbnNmb3JtRm5Bc3NldHNVcmkSpQEKGnF1YW50aXplZF9mZWF0dXJlX21ldGFkYXRhG + AogASgLMkguc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5GZWF0dXJlUXVhbnRpemF0aW9uTWV0Y + WRhdGFCHeI/GhIYcXVhbnRpemVkRmVhdHVyZU1ldGFkYXRhUhhxdWFudGl6ZWRGZWF0dXJlTWV0YWRhdGFCDgoMX2ZlYXR1cmVfZ + GltGugDChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZUtle + 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""" ).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 80160636b..7a9012ffa 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 @@ -132,6 +132,9 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r } lazy val nestedMessagesCompanions: Seq[_root_.scalapb.GeneratedMessageCompanion[_ <: _root_.scalapb.GeneratedMessage]] = Seq[_root_.scalapb.GeneratedMessageCompanion[_ <: _root_.scalapb.GeneratedMessage]]( + _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState, + _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState, + _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata, _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput, _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataInfo, _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.EdgeMetadataOutput, @@ -143,6 +146,577 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r condensedNodeTypeToPreprocessedMetadata = _root_.scala.collection.immutable.Map.empty, condensedEdgeTypeToPreprocessedMetadata = _root_.scala.collection.immutable.Map.empty ) + /** @param clipMin + * Lower clipping bound; dequantized value for linear code 0. + * @param clipMax + * Upper clipping bound; dequantized value for max linear code. + * @param bits + * Quantization level bit-width + */ + @SerialVersionUID(0L) + final case class MultiBitQuantizationState( + clipMin: _root_.scala.Float = 0.0f, + clipMax: _root_.scala.Float = 0.0f, + bits: _root_.scala.Int = 0, + unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty + ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[MultiBitQuantizationState] { + @transient + private[this] var __serializedSizeMemoized: _root_.scala.Int = 0 + private[this] def __computeSerializedSize(): _root_.scala.Int = { + var __size = 0 + + { + val __value = clipMin + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(1, __value) + } + }; + + { + val __value = clipMax + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(2, __value) + } + }; + + { + val __value = bits + if (__value != 0) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeUInt32Size(3, __value) + } + }; + __size += unknownFields.serializedSize + __size + } + override def serializedSize: _root_.scala.Int = { + var __size = __serializedSizeMemoized + if (__size == 0) { + __size = __computeSerializedSize() + 1 + __serializedSizeMemoized = __size + } + __size - 1 + + } + def writeTo(`_output__`: _root_.com.google.protobuf.CodedOutputStream): _root_.scala.Unit = { + { + val __v = clipMin + if (__v != 0.0f) { + _output__.writeFloat(1, __v) + } + }; + { + val __v = clipMax + if (__v != 0.0f) { + _output__.writeFloat(2, __v) + } + }; + { + val __v = bits + if (__v != 0) { + _output__.writeUInt32(3, __v) + } + }; + unknownFields.writeTo(_output__) + } + def withClipMin(__v: _root_.scala.Float): MultiBitQuantizationState = copy(clipMin = __v) + def withClipMax(__v: _root_.scala.Float): MultiBitQuantizationState = copy(clipMax = __v) + def withBits(__v: _root_.scala.Int): MultiBitQuantizationState = copy(bits = __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 = { + (__fieldNumber: @_root_.scala.unchecked) match { + case 1 => { + val __t = clipMin + if (__t != 0.0f) __t else null + } + case 2 => { + val __t = clipMax + if (__t != 0.0f) __t else null + } + case 3 => { + val __t = bits + if (__t != 0) __t else null + } + } + } + def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { + _root_.scala.Predef.require(__field.containingMessage eq companion.scalaDescriptor) + (__field.number: @_root_.scala.unchecked) match { + case 1 => _root_.scalapb.descriptors.PFloat(clipMin) + case 2 => _root_.scalapb.descriptors.PFloat(clipMax) + case 3 => _root_.scalapb.descriptors.PInt(bits) + } + } + def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) + def companion: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState.type = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState + // @@protoc_insertion_point(GeneratedMessage[snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationState]) + } + + object MultiBitQuantizationState extends scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] { + implicit def messageCompanion: scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = this + def parseFrom(`_input__`: _root_.com.google.protobuf.CodedInputStream): snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState = { + var __clipMin: _root_.scala.Float = 0.0f + var __clipMax: _root_.scala.Float = 0.0f + var __bits: _root_.scala.Int = 0 + var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null + var _done__ = false + while (!_done__) { + val _tag__ = _input__.readTag() + _tag__ match { + case 0 => _done__ = true + case 13 => + __clipMin = _input__.readFloat() + case 21 => + __clipMax = _input__.readFloat() + case 24 => + __bits = _input__.readUInt32() + case tag => + if (_unknownFields__ == null) { + _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() + } + _unknownFields__.parseField(tag, _input__) + } + } + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState( + clipMin = __clipMin, + clipMax = __clipMax, + bits = __bits, + unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() + ) + } + implicit def messageReads: _root_.scalapb.descriptors.Reads[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = _root_.scalapb.descriptors.Reads{ + case _root_.scalapb.descriptors.PMessage(__fieldsMap) => + _root_.scala.Predef.require(__fieldsMap.keys.forall(_.containingMessage eq scalaDescriptor), "FieldDescriptor does not match message type.") + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState( + clipMin = __fieldsMap.get(scalaDescriptor.findFieldByNumber(1).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f), + clipMax = __fieldsMap.get(scalaDescriptor.findFieldByNumber(2).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f), + bits = __fieldsMap.get(scalaDescriptor.findFieldByNumber(3).get).map(_.as[_root_.scala.Int]).getOrElse(0) + ) + case _ => throw new RuntimeException("Expected PMessage") + } + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(0) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(0) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + 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.MultiBitQuantizationState( + clipMin = 0.0f, + clipMax = 0.0f, + bits = 0 + ) + implicit class MultiBitQuantizationStateLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState](_l) { + def clipMin: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.clipMin)((c_, f_) => c_.copy(clipMin = f_)) + def clipMax: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.clipMax)((c_, f_) => c_.copy(clipMax = f_)) + def bits: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Int] = field(_.bits)((c_, f_) => c_.copy(bits = f_)) + } + final val CLIP_MIN_FIELD_NUMBER = 1 + final val CLIP_MAX_FIELD_NUMBER = 2 + final val BITS_FIELD_NUMBER = 3 + def of( + clipMin: _root_.scala.Float, + clipMax: _root_.scala.Float, + bits: _root_.scala.Int + ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState( + clipMin, + clipMax, + bits + ) + // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationState]) + } + + /** @param negMean + * Mean value for negative features, produced by packed bit/code 0. + * @param posMean + * Mean value for positive features, produced by packed bit/code 1. + */ + @SerialVersionUID(0L) + final case class SingleBitQuantizationState( + negMean: _root_.scala.Float = 0.0f, + posMean: _root_.scala.Float = 0.0f, + unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty + ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[SingleBitQuantizationState] { + @transient + private[this] var __serializedSizeMemoized: _root_.scala.Int = 0 + private[this] def __computeSerializedSize(): _root_.scala.Int = { + var __size = 0 + + { + val __value = negMean + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(1, __value) + } + }; + + { + val __value = posMean + if (__value != 0.0f) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeFloatSize(2, __value) + } + }; + __size += unknownFields.serializedSize + __size + } + override def serializedSize: _root_.scala.Int = { + var __size = __serializedSizeMemoized + if (__size == 0) { + __size = __computeSerializedSize() + 1 + __serializedSizeMemoized = __size + } + __size - 1 + + } + def writeTo(`_output__`: _root_.com.google.protobuf.CodedOutputStream): _root_.scala.Unit = { + { + val __v = negMean + if (__v != 0.0f) { + _output__.writeFloat(1, __v) + } + }; + { + val __v = posMean + if (__v != 0.0f) { + _output__.writeFloat(2, __v) + } + }; + unknownFields.writeTo(_output__) + } + def withNegMean(__v: _root_.scala.Float): SingleBitQuantizationState = copy(negMean = __v) + def withPosMean(__v: _root_.scala.Float): SingleBitQuantizationState = copy(posMean = __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 = { + (__fieldNumber: @_root_.scala.unchecked) match { + case 1 => { + val __t = negMean + if (__t != 0.0f) __t else null + } + case 2 => { + val __t = posMean + if (__t != 0.0f) __t else null + } + } + } + def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { + _root_.scala.Predef.require(__field.containingMessage eq companion.scalaDescriptor) + (__field.number: @_root_.scala.unchecked) match { + case 1 => _root_.scalapb.descriptors.PFloat(negMean) + case 2 => _root_.scalapb.descriptors.PFloat(posMean) + } + } + def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) + def companion: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState.type = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState + // @@protoc_insertion_point(GeneratedMessage[snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationState]) + } + + object SingleBitQuantizationState extends scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] { + implicit def messageCompanion: scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = this + def parseFrom(`_input__`: _root_.com.google.protobuf.CodedInputStream): snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState = { + var __negMean: _root_.scala.Float = 0.0f + var __posMean: _root_.scala.Float = 0.0f + var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null + var _done__ = false + while (!_done__) { + val _tag__ = _input__.readTag() + _tag__ match { + case 0 => _done__ = true + case 13 => + __negMean = _input__.readFloat() + case 21 => + __posMean = _input__.readFloat() + case tag => + if (_unknownFields__ == null) { + _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() + } + _unknownFields__.parseField(tag, _input__) + } + } + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState( + negMean = __negMean, + posMean = __posMean, + unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() + ) + } + implicit def messageReads: _root_.scalapb.descriptors.Reads[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = _root_.scalapb.descriptors.Reads{ + case _root_.scalapb.descriptors.PMessage(__fieldsMap) => + _root_.scala.Predef.require(__fieldsMap.keys.forall(_.containingMessage eq scalaDescriptor), "FieldDescriptor does not match message type.") + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState( + negMean = __fieldsMap.get(scalaDescriptor.findFieldByNumber(1).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f), + posMean = __fieldsMap.get(scalaDescriptor.findFieldByNumber(2).get).map(_.as[_root_.scala.Float]).getOrElse(0.0f) + ) + case _ => throw new RuntimeException("Expected PMessage") + } + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(1) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(1) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + 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.SingleBitQuantizationState( + negMean = 0.0f, + posMean = 0.0f + ) + implicit class SingleBitQuantizationStateLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState](_l) { + def negMean: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.negMean)((c_, f_) => c_.copy(negMean = f_)) + def posMean: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Float] = field(_.posMean)((c_, f_) => c_.copy(posMean = f_)) + } + final val NEG_MEAN_FIELD_NUMBER = 1 + final val POS_MEAN_FIELD_NUMBER = 2 + def of( + negMean: _root_.scala.Float, + posMean: _root_.scala.Float + ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState( + negMean, + posMean + ) + // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationState]) + } + + /** @param packedFeatureKey + * Field in output TFRecords that stores packed uint8 features. + * @param quantizedFeatureIndices + * Original feature indices stored in packed_feature_key. + */ + @SerialVersionUID(0L) + final case class FeatureQuantizationMetadata( + packedFeatureKey: _root_.scala.Predef.String = "", + quantizedFeatureIndices: _root_.scala.Seq[_root_.scala.Int] = _root_.scala.Seq.empty, + state: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty, + unknownFields: _root_.scalapb.UnknownFieldSet = _root_.scalapb.UnknownFieldSet.empty + ) extends scalapb.GeneratedMessage with scalapb.lenses.Updatable[FeatureQuantizationMetadata] { + private[this] def quantizedFeatureIndicesSerializedSize = { + if (__quantizedFeatureIndicesSerializedSizeField == 0) __quantizedFeatureIndicesSerializedSizeField = { + var __s: _root_.scala.Int = 0 + quantizedFeatureIndices.foreach(__i => __s += _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__i)) + __s + } + __quantizedFeatureIndicesSerializedSizeField + } + @transient private[this] var __quantizedFeatureIndicesSerializedSizeField: _root_.scala.Int = 0 + @transient + private[this] var __serializedSizeMemoized: _root_.scala.Int = 0 + private[this] def __computeSerializedSize(): _root_.scala.Int = { + var __size = 0 + + { + val __value = packedFeatureKey + if (!__value.isEmpty) { + __size += _root_.com.google.protobuf.CodedOutputStream.computeStringSize(1, __value) + } + }; + if (quantizedFeatureIndices.nonEmpty) { + val __localsize = quantizedFeatureIndicesSerializedSize + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__localsize) + __localsize + } + if (state.multiBitState.isDefined) { + val __value = state.multiBitState.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; + if (state.singleBitState.isDefined) { + val __value = state.singleBitState.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; + __size += unknownFields.serializedSize + __size + } + override def serializedSize: _root_.scala.Int = { + var __size = __serializedSizeMemoized + if (__size == 0) { + __size = __computeSerializedSize() + 1 + __serializedSizeMemoized = __size + } + __size - 1 + + } + def writeTo(`_output__`: _root_.com.google.protobuf.CodedOutputStream): _root_.scala.Unit = { + { + val __v = packedFeatureKey + if (!__v.isEmpty) { + _output__.writeString(1, __v) + } + }; + if (quantizedFeatureIndices.nonEmpty) { + _output__.writeTag(2, 2) + _output__.writeUInt32NoTag(quantizedFeatureIndicesSerializedSize) + quantizedFeatureIndices.foreach(_output__.writeUInt32NoTag) + }; + state.multiBitState.foreach { __v => + val __m = __v + _output__.writeTag(4, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; + state.singleBitState.foreach { __v => + val __m = __v + _output__.writeTag(5, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; + unknownFields.writeTo(_output__) + } + def withPackedFeatureKey(__v: _root_.scala.Predef.String): FeatureQuantizationMetadata = copy(packedFeatureKey = __v) + def clearQuantizedFeatureIndices = copy(quantizedFeatureIndices = _root_.scala.Seq.empty) + def addQuantizedFeatureIndices(__vs: _root_.scala.Int *): FeatureQuantizationMetadata = addAllQuantizedFeatureIndices(__vs) + def addAllQuantizedFeatureIndices(__vs: Iterable[_root_.scala.Int]): FeatureQuantizationMetadata = copy(quantizedFeatureIndices = quantizedFeatureIndices ++ __vs) + def withQuantizedFeatureIndices(__v: _root_.scala.Seq[_root_.scala.Int]): FeatureQuantizationMetadata = copy(quantizedFeatureIndices = __v) + def getMultiBitState: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState = state.multiBitState.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState.defaultInstance) + def withMultiBitState(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState): FeatureQuantizationMetadata = copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(__v)) + def getSingleBitState: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState = state.singleBitState.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState.defaultInstance) + def withSingleBitState(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState): FeatureQuantizationMetadata = copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(__v)) + def clearState: FeatureQuantizationMetadata = copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty) + def withState(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State): FeatureQuantizationMetadata = copy(state = __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 = { + (__fieldNumber: @_root_.scala.unchecked) match { + case 1 => { + val __t = packedFeatureKey + if (__t != "") __t else null + } + case 2 => quantizedFeatureIndices + case 4 => state.multiBitState.orNull + case 5 => state.singleBitState.orNull + } + } + def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { + _root_.scala.Predef.require(__field.containingMessage eq companion.scalaDescriptor) + (__field.number: @_root_.scala.unchecked) match { + case 1 => _root_.scalapb.descriptors.PString(packedFeatureKey) + case 2 => _root_.scalapb.descriptors.PRepeated(quantizedFeatureIndices.iterator.map(_root_.scalapb.descriptors.PInt(_)).toVector) + case 4 => state.multiBitState.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) + case 5 => state.singleBitState.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) + } + } + def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) + def companion: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.type = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata + // @@protoc_insertion_point(GeneratedMessage[snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadata]) + } + + object FeatureQuantizationMetadata extends scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] { + implicit def messageCompanion: scalapb.GeneratedMessageCompanion[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = this + def parseFrom(`_input__`: _root_.com.google.protobuf.CodedInputStream): snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = { + var __packedFeatureKey: _root_.scala.Predef.String = "" + val __quantizedFeatureIndices: _root_.scala.collection.immutable.VectorBuilder[_root_.scala.Int] = new _root_.scala.collection.immutable.VectorBuilder[_root_.scala.Int] + var __state: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty + var `_unknownFields__`: _root_.scalapb.UnknownFieldSet.Builder = null + var _done__ = false + while (!_done__) { + val _tag__ = _input__.readTag() + _tag__ match { + case 0 => _done__ = true + case 10 => + __packedFeatureKey = _input__.readStringRequireUtf8() + case 16 => + __quantizedFeatureIndices += _input__.readUInt32() + case 18 => { + val length = _input__.readRawVarint32() + val oldLimit = _input__.pushLimit(length) + while (_input__.getBytesUntilLimit > 0) { + __quantizedFeatureIndices += _input__.readUInt32() + } + _input__.popLimit(oldLimit) + } + case 34 => + __state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(__state.multiBitState.fold(_root_.scalapb.LiteParser.readMessage[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState](_input__))(_root_.scalapb.LiteParser.readMessage(_input__, _))) + case 42 => + __state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(__state.singleBitState.fold(_root_.scalapb.LiteParser.readMessage[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState](_input__))(_root_.scalapb.LiteParser.readMessage(_input__, _))) + case tag => + if (_unknownFields__ == null) { + _unknownFields__ = new _root_.scalapb.UnknownFieldSet.Builder() + } + _unknownFields__.parseField(tag, _input__) + } + } + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata( + packedFeatureKey = __packedFeatureKey, + quantizedFeatureIndices = __quantizedFeatureIndices.result(), + state = __state, + unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() + ) + } + implicit def messageReads: _root_.scalapb.descriptors.Reads[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata] = _root_.scalapb.descriptors.Reads{ + case _root_.scalapb.descriptors.PMessage(__fieldsMap) => + _root_.scala.Predef.require(__fieldsMap.keys.forall(_.containingMessage eq scalaDescriptor), "FieldDescriptor does not match message type.") + snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata( + packedFeatureKey = __fieldsMap.get(scalaDescriptor.findFieldByNumber(1).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), + quantizedFeatureIndices = __fieldsMap.get(scalaDescriptor.findFieldByNumber(2).get).map(_.as[_root_.scala.Seq[_root_.scala.Int]]).getOrElse(_root_.scala.Seq.empty), + state = __fieldsMap.get(scalaDescriptor.findFieldByNumber(4).get).flatMap(_.as[_root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState]]).map(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(_)) + .orElse[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State](__fieldsMap.get(scalaDescriptor.findFieldByNumber(5).get).flatMap(_.as[_root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState]]).map(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(_))) + .getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty) + ) + case _ => throw new RuntimeException("Expected PMessage") + } + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(2) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(2) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { + var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null + (__number: @_root_.scala.unchecked) match { + case 4 => __out = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState + case 5 => __out = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState + } + __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.FeatureQuantizationMetadata( + packedFeatureKey = "", + quantizedFeatureIndices = _root_.scala.Seq.empty, + state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.Empty + ) + sealed trait State extends _root_.scalapb.GeneratedOneof { + def isEmpty: _root_.scala.Boolean = false + def isDefined: _root_.scala.Boolean = true + def isMultiBitState: _root_.scala.Boolean = false + def isSingleBitState: _root_.scala.Boolean = false + def multiBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = _root_.scala.None + def singleBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = _root_.scala.None + } + object State { + @SerialVersionUID(0L) + case object Empty extends snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State { + type ValueType = _root_.scala.Nothing + override def isEmpty: _root_.scala.Boolean = true + override def isDefined: _root_.scala.Boolean = false + override def number: _root_.scala.Int = 0 + override def value: _root_.scala.Nothing = throw new java.util.NoSuchElementException("Empty.value") + } + + @SerialVersionUID(0L) + final case class MultiBitState(value: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState) extends snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State { + type ValueType = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState + override def isMultiBitState: _root_.scala.Boolean = true + override def multiBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = Some(value) + override def number: _root_.scala.Int = 4 + } + @SerialVersionUID(0L) + final case class SingleBitState(value: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState) extends snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State { + type ValueType = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState + override def isSingleBitState: _root_.scala.Boolean = true + override def singleBitState: _root_.scala.Option[snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = Some(value) + override def number: _root_.scala.Int = 5 + } + } + implicit class FeatureQuantizationMetadataLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata](_l) { + def packedFeatureKey: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Predef.String] = field(_.packedFeatureKey)((c_, f_) => c_.copy(packedFeatureKey = f_)) + def quantizedFeatureIndices: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Seq[_root_.scala.Int]] = field(_.quantizedFeatureIndices)((c_, f_) => c_.copy(quantizedFeatureIndices = f_)) + def multiBitState: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.MultiBitQuantizationState] = field(_.getMultiBitState)((c_, f_) => c_.copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.MultiBitState(f_))) + def singleBitState: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.SingleBitQuantizationState] = field(_.getSingleBitState)((c_, f_) => c_.copy(state = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State.SingleBitState(f_))) + def state: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State] = field(_.state)((c_, f_) => c_.copy(state = f_)) + } + final val PACKED_FEATURE_KEY_FIELD_NUMBER = 1 + final val QUANTIZED_FEATURE_INDICES_FIELD_NUMBER = 2 + final val MULTI_BIT_STATE_FIELD_NUMBER = 4 + final val SINGLE_BIT_STATE_FIELD_NUMBER = 5 + def of( + packedFeatureKey: _root_.scala.Predef.String, + quantizedFeatureIndices: _root_.scala.Seq[_root_.scala.Int], + state: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.State + ): _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata( + packedFeatureKey, + quantizedFeatureIndices, + state + ) + // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadata]) + } + /** Houses metadata about node TFTransform output from DataPreprocessor. * * @param nodeIdKey @@ -163,6 +737,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r * Feature dimension after preprocessing * @param transformFnAssetsUri * Contains categorical feature vocabularies + * @param quantizedFeatureMetadata + * Optional quantized node feature metadata. */ @SerialVersionUID(0L) final case class NodeMetadataOutput( @@ -175,6 +751,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeDataBqTable: _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[NodeMetadataOutput] { @transient @@ -235,6 +812,10 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __size += _root_.com.google.protobuf.CodedOutputStream.computeStringSize(9, __value) } }; + if (quantizedFeatureMetadata.isDefined) { + val __value = quantizedFeatureMetadata.get + __size += 1 + _root_.com.google.protobuf.CodedOutputStream.computeUInt32SizeNoTag(__value.serializedSize) + __value.serializedSize + }; __size += unknownFields.serializedSize __size } @@ -296,6 +877,12 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r _output__.writeString(9, __v) } }; + quantizedFeatureMetadata.foreach { __v => + val __m = __v + _output__.writeTag(10, 2) + _output__.writeUInt32NoTag(__m.serializedSize) + __m.writeTo(_output__) + }; unknownFields.writeTo(_output__) } def withNodeIdKey(__v: _root_.scala.Predef.String): NodeMetadataOutput = copy(nodeIdKey = __v) @@ -315,6 +902,9 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r def clearFeatureDim: NodeMetadataOutput = copy(featureDim = _root_.scala.None) def withFeatureDim(__v: _root_.scala.Int): NodeMetadataOutput = copy(featureDim = Option(__v)) def withTransformFnAssetsUri(__v: _root_.scala.Predef.String): NodeMetadataOutput = copy(transformFnAssetsUri = __v) + def getQuantizedFeatureMetadata: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata = quantizedFeatureMetadata.getOrElse(snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata.defaultInstance) + def clearQuantizedFeatureMetadata: NodeMetadataOutput = copy(quantizedFeatureMetadata = _root_.scala.None) + def withQuantizedFeatureMetadata(__v: snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.FeatureQuantizationMetadata): NodeMetadataOutput = 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 = { @@ -346,6 +936,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r val __t = transformFnAssetsUri if (__t != "") __t else null } + case 10 => quantizedFeatureMetadata.orNull } } def getField(__field: _root_.scalapb.descriptors.FieldDescriptor): _root_.scalapb.descriptors.PValue = { @@ -360,6 +951,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r case 7 => _root_.scalapb.descriptors.PString(enumeratedNodeDataBqTable) case 8 => featureDim.map(_root_.scalapb.descriptors.PInt(_)).getOrElse(_root_.scalapb.descriptors.PEmpty) case 9 => _root_.scalapb.descriptors.PString(transformFnAssetsUri) + case 10 => quantizedFeatureMetadata.map(_.toPMessage).getOrElse(_root_.scalapb.descriptors.PEmpty) } } def toProtoString: _root_.scala.Predef.String = _root_.scalapb.TextFormat.printToUnicodeString(this) @@ -379,6 +971,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r var __enumeratedNodeDataBqTable: _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__) { @@ -403,6 +996,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r __featureDim = Option(_input__.readUInt32()) case 74 => __transformFnAssetsUri = _input__.readStringRequireUtf8() + case 82 => + __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() @@ -420,6 +1015,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeDataBqTable = __enumeratedNodeDataBqTable, featureDim = __featureDim, transformFnAssetsUri = __transformFnAssetsUri, + quantizedFeatureMetadata = __quantizedFeatureMetadata, unknownFields = if (_unknownFields__ == null) _root_.scalapb.UnknownFieldSet.empty else _unknownFields__.result() ) } @@ -435,13 +1031,20 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable = __fieldsMap.get(scalaDescriptor.findFieldByNumber(6).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), enumeratedNodeDataBqTable = __fieldsMap.get(scalaDescriptor.findFieldByNumber(7).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), featureDim = __fieldsMap.get(scalaDescriptor.findFieldByNumber(8).get).flatMap(_.as[_root_.scala.Option[_root_.scala.Int]]), - transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(9).get).map(_.as[_root_.scala.Predef.String]).getOrElse("") + transformFnAssetsUri = __fieldsMap.get(scalaDescriptor.findFieldByNumber(9).get).map(_.as[_root_.scala.Predef.String]).getOrElse(""), + quantizedFeatureMetadata = __fieldsMap.get(scalaDescriptor.findFieldByNumber(10).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(0) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(0) - def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = throw new MatchError(__number) + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(3) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(3) + def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { + var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null + (__number: @_root_.scala.unchecked) match { + case 10 => __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.NodeMetadataOutput( @@ -453,7 +1056,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable = "", enumeratedNodeDataBqTable = "", featureDim = _root_.scala.None, - transformFnAssetsUri = "" + transformFnAssetsUri = "", + quantizedFeatureMetadata = _root_.scala.None ) implicit class NodeMetadataOutputLens[UpperPB](_l: _root_.scalapb.lenses.Lens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput]) extends _root_.scalapb.lenses.ObjectLens[UpperPB, snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput](_l) { def nodeIdKey: _root_.scalapb.lenses.Lens[UpperPB, _root_.scala.Predef.String] = field(_.nodeIdKey)((c_, f_) => c_.copy(nodeIdKey = f_)) @@ -466,6 +1070,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 NODE_ID_KEY_FIELD_NUMBER = 1 final val FEATURE_KEYS_FIELD_NUMBER = 2 @@ -476,6 +1082,7 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r final val ENUMERATED_NODE_DATA_BQ_TABLE_FIELD_NUMBER = 7 final val FEATURE_DIM_FIELD_NUMBER = 8 final val TRANSFORM_FN_ASSETS_URI_FIELD_NUMBER = 9 + final val QUANTIZED_FEATURE_METADATA_FIELD_NUMBER = 10 def of( nodeIdKey: _root_.scala.Predef.String, featureKeys: _root_.scala.Seq[_root_.scala.Predef.String], @@ -485,7 +1092,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable: _root_.scala.Predef.String, enumeratedNodeDataBqTable: _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.NodeMetadataOutput = _root_.snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.NodeMetadataOutput( nodeIdKey, featureKeys, @@ -495,7 +1103,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r enumeratedNodeIdsBqTable, enumeratedNodeDataBqTable, featureDim, - transformFnAssetsUri + transformFnAssetsUri, + quantizedFeatureMetadata ) // @@protoc_insertion_point(GeneratedMessageCompanion[snapchat.research.gbml.PreprocessedMetadata.NodeMetadataOutput]) } @@ -742,8 +1351,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) case _ => throw new RuntimeException("Expected PMessage") } - def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(1) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(1) + 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) 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) @@ -985,8 +1594,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) case _ => throw new RuntimeException("Expected PMessage") } - def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(2) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(2) + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(5) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(5) def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null (__number: @_root_.scala.unchecked) match { @@ -1148,8 +1757,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) case _ => throw new RuntimeException("Expected PMessage") } - def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(3) - def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(3) + def javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(6) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(6) def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null (__number: @_root_.scala.unchecked) match { @@ -1295,8 +1904,8 @@ object PreprocessedMetadata extends scalapb.GeneratedMessageCompanion[snapchat.r ) 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 javaDescriptor: _root_.com.google.protobuf.Descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.javaDescriptor.getNestedTypes().get(7) + def scalaDescriptor: _root_.scalapb.descriptors.Descriptor = snapchat.research.gbml.preprocessed_metadata.PreprocessedMetadata.scalaDescriptor.nestedMessages(7) def messageCompanionForFieldNumber(__number: _root_.scala.Int): _root_.scalapb.GeneratedMessageCompanion[_] = { var __out: _root_.scalapb.GeneratedMessageCompanion[_] = null (__number: @_root_.scala.unchecked) match { 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 becc2d068..ad80de0ad 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,41 +14,53 @@ object PreprocessedMetadataProto extends _root_.scalapb.GeneratedFileObject { private lazy val ProtoBytes: _root_.scala.Array[Byte] = scalapb.Encoding.fromBase64(scala.collection.immutable.Seq( """CjJzbmFwY2hhdC9yZXNlYXJjaC9nYm1sL3ByZXByb2Nlc3NlZF9tZXRhZGF0YS5wcm90bxIWc25hcGNoYXQucmVzZWFyY2guZ - 2JtbCL8EwoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY + 2JtbCL9GgoUUHJlcHJvY2Vzc2VkTWV0YWRhdGES5gEKLGNvbmRlbnNlZF9ub2RlX3R5cGVfdG9fcHJlcHJvY2Vzc2VkX21ldGFkY XRhGAEgAygLMlkuc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5Db25kZW5zZWROb2RlVHlwZVRvU HJlcHJvY2Vzc2VkTWV0YWRhdGFFbnRyeUIs4j8pEidjb25kZW5zZWROb2RlVHlwZVRvUHJlcHJvY2Vzc2VkTWV0YWRhdGFSJ2Nvb mRlbnNlZE5vZGVUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZGF0YRLmAQosY29uZGVuc2VkX2VkZ2VfdHlwZV90b19wcmVwcm9jZXNzZ WRfbWV0YWRhdGEYAiADKAsyWS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkNvbmRlbnNlZEVkZ 2VUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZGF0YUVudHJ5QiziPykSJ2NvbmRlbnNlZEVkZ2VUeXBlVG9QcmVwcm9jZXNzZWRNZXRhZ - GF0YVInY29uZGVuc2VkRWRnZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhGvkEChJOb2RlTWV0YWRhdGFPdXRwdXQSLgoLbm9kZ - V9pZF9rZXkYASABKAlCDuI/CxIJbm9kZUlkS2V5Uglub2RlSWRLZXkSMwoMZmVhdHVyZV9rZXlzGAIgAygJQhDiPw0SC2ZlYXR1c - mVLZXlzUgtmZWF0dXJlS2V5cxItCgpsYWJlbF9rZXlzGAMgAygJQg7iPwsSCWxhYmVsS2V5c1IJbGFiZWxLZXlzEkYKE3RmcmVjb - 3JkX3VyaV9wcmVmaXgYBCABKAlCFuI/ExIRdGZyZWNvcmRVcmlQcmVmaXhSEXRmcmVjb3JkVXJpUHJlZml4Ei0KCnNjaGVtYV91c - mkYBSABKAlCDuI/CxIJc2NoZW1hVXJpUglzY2hlbWFVcmkSXQocZW51bWVyYXRlZF9ub2RlX2lkc19icV90YWJsZRgGIAEoCUId4 - j8aEhhlbnVtZXJhdGVkTm9kZUlkc0JxVGFibGVSGGVudW1lcmF0ZWROb2RlSWRzQnFUYWJsZRJgCh1lbnVtZXJhdGVkX25vZGVfZ - GF0YV9icV90YWJsZRgHIAEoCUIe4j8bEhllbnVtZXJhdGVkTm9kZURhdGFCcVRhYmxlUhllbnVtZXJhdGVkTm9kZURhdGFCcVRhY - mxlEjUKC2ZlYXR1cmVfZGltGAggASgNQg/iPwwSCmZlYXR1cmVEaW1IAFIKZmVhdHVyZURpbYgBARJQChd0cmFuc2Zvcm1fZm5fY - XNzZXRzX3VyaRgJIAEoCUIZ4j8WEhR0cmFuc2Zvcm1GbkFzc2V0c1VyaVIUdHJhbnNmb3JtRm5Bc3NldHNVcmlCDgoMX2ZlYXR1c - mVfZGltGugDChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZ - UtleXMSLQoKbGFiZWxfa2V5cxgCIAMoCUIO4j8LEglsYWJlbEtleXNSCWxhYmVsS2V5cxJGChN0ZnJlY29yZF91cmlfcHJlZml4G - AMgASgJQhbiPxMSEXRmcmVjb3JkVXJpUHJlZml4UhF0ZnJlY29yZFVyaVByZWZpeBItCgpzY2hlbWFfdXJpGAQgASgJQg7iPwsSC - XNjaGVtYVVyaVIJc2NoZW1hVXJpEmAKHWVudW1lcmF0ZWRfZWRnZV9kYXRhX2JxX3RhYmxlGAUgASgJQh7iPxsSGWVudW1lcmF0Z - WRFZGdlRGF0YUJxVGFibGVSGWVudW1lcmF0ZWRFZGdlRGF0YUJxVGFibGUSNQoLZmVhdHVyZV9kaW0YBiABKA1CD+I/DBIKZmVhd - HVyZURpbUgAUgpmZWF0dXJlRGltiAEBElAKF3RyYW5zZm9ybV9mbl9hc3NldHNfdXJpGAcgASgJQhniPxYSFHRyYW5zZm9ybUZuQ - XNzZXRzVXJpUhR0cmFuc2Zvcm1GbkFzc2V0c1VyaUIOCgxfZmVhdHVyZV9kaW0awgQKEkVkZ2VNZXRhZGF0YU91dHB1dBI4Cg9zc - mNfbm9kZV9pZF9rZXkYASABKAlCEeI/DhIMc3JjTm9kZUlkS2V5UgxzcmNOb2RlSWRLZXkSOAoPZHN0X25vZGVfaWRfa2V5GAIgA - SgJQhHiPw4SDGRzdE5vZGVJZEtleVIMZHN0Tm9kZUlkS2V5EnYKDm1haW5fZWRnZV9pbmZvGAMgASgLMj0uc25hcGNoYXQucmVzZ - WFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFJbmZvQhHiPw4SDG1haW5FZGdlSW5mb1IMbWFpbkVkZ - 2VJbmZvEocBChJwb3NpdGl2ZV9lZGdlX2luZm8YBCABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ld - GFkYXRhLkVkZ2VNZXRhZGF0YUluZm9CFeI/EhIQcG9zaXRpdmVFZGdlSW5mb0gAUhBwb3NpdGl2ZUVkZ2VJbmZviAEBEocBChJuZ - WdhdGl2ZV9lZGdlX2luZm8YBSABKAsyPS5zbmFwY2hhdC5yZXNlYXJjaC5nYm1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLkVkZ2VNZ - XRhZGF0YUluZm9CFeI/EhIQbmVnYXRpdmVFZGdlSW5mb0gBUhBuZWdhdGl2ZUVkZ2VJbmZviAEBQhUKE19wb3NpdGl2ZV9lZGdlX - 2luZm9CFQoTX25lZ2F0aXZlX2VkZ2VfaW5mbxqxAQosQ29uZGVuc2VkTm9kZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50c - nkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLMj8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwc - m9jZXNzZWRNZXRhZGF0YS5Ob2RlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsdWVSBXZhbHVlOgI4ARqxAQosQ29uZGVuc2VkRWRnZ - VR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhRW50cnkSGgoDa2V5GAEgASgNQgjiPwUSA2tleVIDa2V5EmEKBXZhbHVlGAIgASgLM - j8uc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5FZGdlTWV0YWRhdGFPdXRwdXRCCuI/BxIFdmFsd - WVSBXZhbHVlOgI4AWIGcHJvdG8z""" + GF0YVInY29uZGVuc2VkRWRnZVR5cGVUb1ByZXByb2Nlc3NlZE1ldGFkYXRhGowBChlNdWx0aUJpdFF1YW50aXphdGlvblN0YXRlE + icKCGNsaXBfbWluGAEgASgCQgziPwkSB2NsaXBNaW5SB2NsaXBNaW4SJwoIY2xpcF9tYXgYAiABKAJCDOI/CRIHY2xpcE1heFIHY + 2xpcE1heBIdCgRiaXRzGAMgASgNQgniPwYSBGJpdHNSBGJpdHMabgoaU2luZ2xlQml0UXVhbnRpemF0aW9uU3RhdGUSJwoIbmVnX + 21lYW4YASABKAJCDOI/CRIHbmVnTWVhblIHbmVnTWVhbhInCghwb3NfbWVhbhgCIAEoAkIM4j8JEgdwb3NNZWFuUgdwb3NNZWFuG + tcDChtGZWF0dXJlUXVhbnRpemF0aW9uTWV0YWRhdGESQwoScGFja2VkX2ZlYXR1cmVfa2V5GAEgASgJQhXiPxISEHBhY2tlZEZlY + XR1cmVLZXlSEHBhY2tlZEZlYXR1cmVLZXkSWAoZcXVhbnRpemVkX2ZlYXR1cmVfaW5kaWNlcxgCIAMoDUIc4j8ZEhdxdWFudGl6Z + WRGZWF0dXJlSW5kaWNlc1IXcXVhbnRpemVkRmVhdHVyZUluZGljZXMShAEKD211bHRpX2JpdF9zdGF0ZRgEIAEoCzJGLnNuYXBja + GF0LnJlc2VhcmNoLmdibWwuUHJlcHJvY2Vzc2VkTWV0YWRhdGEuTXVsdGlCaXRRdWFudGl6YXRpb25TdGF0ZUIS4j8PEg1tdWx0a + UJpdFN0YXRlSABSDW11bHRpQml0U3RhdGUSiAEKEHNpbmdsZV9iaXRfc3RhdGUYBSABKAsyRy5zbmFwY2hhdC5yZXNlYXJjaC5nY + m1sLlByZXByb2Nlc3NlZE1ldGFkYXRhLlNpbmdsZUJpdFF1YW50aXphdGlvblN0YXRlQhPiPxASDnNpbmdsZUJpdFN0YXRlSABSD + nNpbmdsZUJpdFN0YXRlQgcKBXN0YXRlGqEGChJOb2RlTWV0YWRhdGFPdXRwdXQSLgoLbm9kZV9pZF9rZXkYASABKAlCDuI/CxIJb + m9kZUlkS2V5Uglub2RlSWRLZXkSMwoMZmVhdHVyZV9rZXlzGAIgAygJQhDiPw0SC2ZlYXR1cmVLZXlzUgtmZWF0dXJlS2V5cxItC + gpsYWJlbF9rZXlzGAMgAygJQg7iPwsSCWxhYmVsS2V5c1IJbGFiZWxLZXlzEkYKE3RmcmVjb3JkX3VyaV9wcmVmaXgYBCABKAlCF + uI/ExIRdGZyZWNvcmRVcmlQcmVmaXhSEXRmcmVjb3JkVXJpUHJlZml4Ei0KCnNjaGVtYV91cmkYBSABKAlCDuI/CxIJc2NoZW1hV + XJpUglzY2hlbWFVcmkSXQocZW51bWVyYXRlZF9ub2RlX2lkc19icV90YWJsZRgGIAEoCUId4j8aEhhlbnVtZXJhdGVkTm9kZUlkc + 0JxVGFibGVSGGVudW1lcmF0ZWROb2RlSWRzQnFUYWJsZRJgCh1lbnVtZXJhdGVkX25vZGVfZGF0YV9icV90YWJsZRgHIAEoCUIe4 + j8bEhllbnVtZXJhdGVkTm9kZURhdGFCcVRhYmxlUhllbnVtZXJhdGVkTm9kZURhdGFCcVRhYmxlEjUKC2ZlYXR1cmVfZGltGAggA + SgNQg/iPwwSCmZlYXR1cmVEaW1IAFIKZmVhdHVyZURpbYgBARJQChd0cmFuc2Zvcm1fZm5fYXNzZXRzX3VyaRgJIAEoCUIZ4j8WE + hR0cmFuc2Zvcm1GbkFzc2V0c1VyaVIUdHJhbnNmb3JtRm5Bc3NldHNVcmkSpQEKGnF1YW50aXplZF9mZWF0dXJlX21ldGFkYXRhG + AogASgLMkguc25hcGNoYXQucmVzZWFyY2guZ2JtbC5QcmVwcm9jZXNzZWRNZXRhZGF0YS5GZWF0dXJlUXVhbnRpemF0aW9uTWV0Y + WRhdGFCHeI/GhIYcXVhbnRpemVkRmVhdHVyZU1ldGFkYXRhUhhxdWFudGl6ZWRGZWF0dXJlTWV0YWRhdGFCDgoMX2ZlYXR1cmVfZ + GltGugDChBFZGdlTWV0YWRhdGFJbmZvEjMKDGZlYXR1cmVfa2V5cxgBIAMoCUIQ4j8NEgtmZWF0dXJlS2V5c1ILZmVhdHVyZUtle + 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""" ).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 8f2f818f3..2fac76d2f 100644 --- a/snapchat/research/gbml/preprocessed_metadata_pb2.py +++ b/snapchat/research/gbml/preprocessed_metadata_pb2.py @@ -14,11 +14,14 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n2snapchat/research/gbml/preprocessed_metadata.proto\x12\x16snapchat.research.gbml\"\xed\x0b\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\x1a\x9c\x02\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(\tB\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\"\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') _PREPROCESSEDMETADATA = DESCRIPTOR.message_types_by_name['PreprocessedMetadata'] +_PREPROCESSEDMETADATA_MULTIBITQUANTIZATIONSTATE = _PREPROCESSEDMETADATA.nested_types_by_name['MultiBitQuantizationState'] +_PREPROCESSEDMETADATA_SINGLEBITQUANTIZATIONSTATE = _PREPROCESSEDMETADATA.nested_types_by_name['SingleBitQuantizationState'] +_PREPROCESSEDMETADATA_FEATUREQUANTIZATIONMETADATA = _PREPROCESSEDMETADATA.nested_types_by_name['FeatureQuantizationMetadata'] _PREPROCESSEDMETADATA_NODEMETADATAOUTPUT = _PREPROCESSEDMETADATA.nested_types_by_name['NodeMetadataOutput'] _PREPROCESSEDMETADATA_EDGEMETADATAINFO = _PREPROCESSEDMETADATA.nested_types_by_name['EdgeMetadataInfo'] _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT = _PREPROCESSEDMETADATA.nested_types_by_name['EdgeMetadataOutput'] @@ -26,6 +29,27 @@ _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY = _PREPROCESSEDMETADATA.nested_types_by_name['CondensedEdgeTypeToPreprocessedMetadataEntry'] PreprocessedMetadata = _reflection.GeneratedProtocolMessageType('PreprocessedMetadata', (_message.Message,), { + 'MultiBitQuantizationState' : _reflection.GeneratedProtocolMessageType('MultiBitQuantizationState', (_message.Message,), { + 'DESCRIPTOR' : _PREPROCESSEDMETADATA_MULTIBITQUANTIZATIONSTATE, + '__module__' : 'snapchat.research.gbml.preprocessed_metadata_pb2' + # @@protoc_insertion_point(class_scope:snapchat.research.gbml.PreprocessedMetadata.MultiBitQuantizationState) + }) + , + + 'SingleBitQuantizationState' : _reflection.GeneratedProtocolMessageType('SingleBitQuantizationState', (_message.Message,), { + 'DESCRIPTOR' : _PREPROCESSEDMETADATA_SINGLEBITQUANTIZATIONSTATE, + '__module__' : 'snapchat.research.gbml.preprocessed_metadata_pb2' + # @@protoc_insertion_point(class_scope:snapchat.research.gbml.PreprocessedMetadata.SingleBitQuantizationState) + }) + , + + 'FeatureQuantizationMetadata' : _reflection.GeneratedProtocolMessageType('FeatureQuantizationMetadata', (_message.Message,), { + 'DESCRIPTOR' : _PREPROCESSEDMETADATA_FEATUREQUANTIZATIONMETADATA, + '__module__' : 'snapchat.research.gbml.preprocessed_metadata_pb2' + # @@protoc_insertion_point(class_scope:snapchat.research.gbml.PreprocessedMetadata.FeatureQuantizationMetadata) + }) + , + 'NodeMetadataOutput' : _reflection.GeneratedProtocolMessageType('NodeMetadataOutput', (_message.Message,), { 'DESCRIPTOR' : _PREPROCESSEDMETADATA_NODEMETADATAOUTPUT, '__module__' : 'snapchat.research.gbml.preprocessed_metadata_pb2' @@ -65,6 +89,9 @@ # @@protoc_insertion_point(class_scope:snapchat.research.gbml.PreprocessedMetadata) }) _sym_db.RegisterMessage(PreprocessedMetadata) +_sym_db.RegisterMessage(PreprocessedMetadata.MultiBitQuantizationState) +_sym_db.RegisterMessage(PreprocessedMetadata.SingleBitQuantizationState) +_sym_db.RegisterMessage(PreprocessedMetadata.FeatureQuantizationMetadata) _sym_db.RegisterMessage(PreprocessedMetadata.NodeMetadataOutput) _sym_db.RegisterMessage(PreprocessedMetadata.EdgeMetadataInfo) _sym_db.RegisterMessage(PreprocessedMetadata.EdgeMetadataOutput) @@ -79,15 +106,21 @@ _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._options = None _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_options = b'8\001' _PREPROCESSEDMETADATA._serialized_start=79 - _PREPROCESSEDMETADATA._serialized_end=1596 - _PREPROCESSEDMETADATA_NODEMETADATAOUTPUT._serialized_start=396 - _PREPROCESSEDMETADATA_NODEMETADATAOUTPUT._serialized_end=680 - _PREPROCESSEDMETADATA_EDGEMETADATAINFO._serialized_start=683 - _PREPROCESSEDMETADATA_EDGEMETADATAINFO._serialized_end=906 - _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT._serialized_start=909 - _PREPROCESSEDMETADATA_EDGEMETADATAOUTPUT._serialized_end=1304 - _PREPROCESSEDMETADATA_CONDENSEDNODETYPETOPREPROCESSEDMETADATAENTRY._serialized_start=1307 - _PREPROCESSEDMETADATA_CONDENSEDNODETYPETOPREPROCESSEDMETADATAENTRY._serialized_end=1450 - _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_start=1453 - _PREPROCESSEDMETADATA_CONDENSEDEDGETYPETOPREPROCESSEDMETADATAENTRY._serialized_end=1596 + _PREPROCESSEDMETADATA._serialized_end=2155 + _PREPROCESSEDMETADATA_MULTIBITQUANTIZATIONSTATE._serialized_start=395 + _PREPROCESSEDMETADATA_MULTIBITQUANTIZATIONSTATE._serialized_end=472 + _PREPROCESSEDMETADATA_SINGLEBITQUANTIZATIONSTATE._serialized_start=474 + _PREPROCESSEDMETADATA_SINGLEBITQUANTIZATIONSTATE._serialized_end=538 + _PREPROCESSEDMETADATA_FEATUREQUANTIZATIONMETADATA._serialized_start=541 + _PREPROCESSEDMETADATA_FEATUREQUANTIZATIONMETADATA._serialized_end=842 + _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 # @@protoc_insertion_point(module_scope) diff --git a/snapchat/research/gbml/preprocessed_metadata_pb2.pyi b/snapchat/research/gbml/preprocessed_metadata_pb2.pyi index 2271a7fec..46b80c7bb 100644 --- a/snapchat/research/gbml/preprocessed_metadata_pb2.pyi +++ b/snapchat/research/gbml/preprocessed_metadata_pb2.pyi @@ -20,6 +20,72 @@ DESCRIPTOR: google.protobuf.descriptor.FileDescriptor class PreprocessedMetadata(google.protobuf.message.Message): DESCRIPTOR: google.protobuf.descriptor.Descriptor + class MultiBitQuantizationState(google.protobuf.message.Message): + DESCRIPTOR: google.protobuf.descriptor.Descriptor + + CLIP_MIN_FIELD_NUMBER: builtins.int + CLIP_MAX_FIELD_NUMBER: builtins.int + BITS_FIELD_NUMBER: builtins.int + clip_min: builtins.float + """Lower clipping bound; dequantized value for linear code 0.""" + clip_max: builtins.float + """Upper clipping bound; dequantized value for max linear code.""" + bits: builtins.int + """Quantization level bit-width""" + def __init__( + self, + *, + clip_min: builtins.float = ..., + clip_max: builtins.float = ..., + bits: builtins.int = ..., + ) -> None: ... + def ClearField(self, field_name: typing_extensions.Literal["bits", b"bits", "clip_max", b"clip_max", "clip_min", b"clip_min"]) -> None: ... + + class SingleBitQuantizationState(google.protobuf.message.Message): + DESCRIPTOR: google.protobuf.descriptor.Descriptor + + NEG_MEAN_FIELD_NUMBER: builtins.int + POS_MEAN_FIELD_NUMBER: builtins.int + neg_mean: builtins.float + """Mean value for negative features, produced by packed bit/code 0.""" + pos_mean: builtins.float + """Mean value for positive features, produced by packed bit/code 1.""" + def __init__( + self, + *, + neg_mean: builtins.float = ..., + pos_mean: builtins.float = ..., + ) -> None: ... + def ClearField(self, field_name: typing_extensions.Literal["neg_mean", b"neg_mean", "pos_mean", b"pos_mean"]) -> None: ... + + class FeatureQuantizationMetadata(google.protobuf.message.Message): + DESCRIPTOR: google.protobuf.descriptor.Descriptor + + PACKED_FEATURE_KEY_FIELD_NUMBER: builtins.int + QUANTIZED_FEATURE_INDICES_FIELD_NUMBER: builtins.int + MULTI_BIT_STATE_FIELD_NUMBER: builtins.int + SINGLE_BIT_STATE_FIELD_NUMBER: builtins.int + packed_feature_key: builtins.str + """Field in output TFRecords that stores packed uint8 features.""" + @property + def quantized_feature_indices(self) -> google.protobuf.internal.containers.RepeatedScalarFieldContainer[builtins.int]: + """Original feature indices stored in packed_feature_key.""" + @property + def multi_bit_state(self) -> global___PreprocessedMetadata.MultiBitQuantizationState: ... + @property + def single_bit_state(self) -> global___PreprocessedMetadata.SingleBitQuantizationState: ... + def __init__( + self, + *, + packed_feature_key: builtins.str = ..., + quantized_feature_indices: collections.abc.Iterable[builtins.int] | None = ..., + multi_bit_state: global___PreprocessedMetadata.MultiBitQuantizationState | None = ..., + single_bit_state: global___PreprocessedMetadata.SingleBitQuantizationState | None = ..., + ) -> None: ... + def HasField(self, field_name: typing_extensions.Literal["multi_bit_state", b"multi_bit_state", "single_bit_state", b"single_bit_state", "state", b"state"]) -> builtins.bool: ... + def ClearField(self, field_name: typing_extensions.Literal["multi_bit_state", b"multi_bit_state", "packed_feature_key", b"packed_feature_key", "quantized_feature_indices", b"quantized_feature_indices", "single_bit_state", b"single_bit_state", "state", b"state"]) -> None: ... + def WhichOneof(self, oneof_group: typing_extensions.Literal["state", b"state"]) -> typing_extensions.Literal["multi_bit_state", "single_bit_state"] | None: ... + class NodeMetadataOutput(google.protobuf.message.Message): """Houses metadata about node TFTransform output from DataPreprocessor.""" @@ -34,6 +100,7 @@ class PreprocessedMetadata(google.protobuf.message.Message): ENUMERATED_NODE_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 node_id_key: builtins.str """The field in output TFRecords which references the node identifier.""" @property @@ -54,6 +121,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 node feature metadata.""" def __init__( self, *, @@ -66,9 +136,10 @@ class PreprocessedMetadata(google.protobuf.message.Message): enumerated_node_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_node_data_bq_table", b"enumerated_node_data_bq_table", "enumerated_node_ids_bq_table", b"enumerated_node_ids_bq_table", "feature_dim", b"feature_dim", "feature_keys", b"feature_keys", "label_keys", b"label_keys", "node_id_key", b"node_id_key", "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_node_data_bq_table", b"enumerated_node_data_bq_table", "enumerated_node_ids_bq_table", b"enumerated_node_ids_bq_table", "feature_dim", b"feature_dim", "feature_keys", b"feature_keys", "label_keys", b"label_keys", "node_id_key", b"node_id_key", "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 EdgeMetadataInfo(google.protobuf.message.Message): From 18629384fc2f3d0de60849ff8e6d38958e97ec7a Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 21:04:59 +0000 Subject: [PATCH 02/78] Add quantization ops --- .../utils/feature_quantization/__init__.py | 0 .../utils/feature_quantization/numpy_ops.py | 38 ++++++++++ .../utils/feature_quantization/torch_ops.py | 42 +++++++++++ gigl/types/graph.py | 70 +++++++++++++++++++ 4 files changed, 150 insertions(+) create mode 100644 gigl/common/utils/feature_quantization/__init__.py create mode 100644 gigl/common/utils/feature_quantization/numpy_ops.py create mode 100644 gigl/common/utils/feature_quantization/torch_ops.py diff --git a/gigl/common/utils/feature_quantization/__init__.py b/gigl/common/utils/feature_quantization/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py new file mode 100644 index 000000000..bb8e30f28 --- /dev/null +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -0,0 +1,38 @@ +from collections.abc import Mapping + +import numpy as np + + +def quantize_ndarray( + features: np.ndarray, *, bits: int, stats: Mapping[str, float] +) -> np.ndarray: + """Quantize a 2D float array into packed uint8 codes.""" + if bits not in (1, 2, 4, 8): + raise ValueError(f"bits must be one of 1, 2, 4, or 8, got {bits}.") + if features.ndim != 2: + raise ValueError(f"Expected a 2D feature array, got shape {features.shape}.") + if bits == 1: + # 1-bit quantization keeps only sign; values restore from neg/pos means. + codes = (features > 0).astype(np.uint8) + else: + # Linearly map clipped values into integer buckets. + levels = (1 << bits) - 1 + lo, hi = stats["clip_min"], stats["clip_max"] + clipped = np.clip(features, lo, hi) + scaled = (clipped - lo) / (hi - lo) + codes = np.rint(scaled * levels).astype(np.uint8) + return pack_codes(codes, bits) + + +def pack_codes(codes: np.ndarray, bits: int) -> np.ndarray: + """Pack low-bit feature codes high-bits-first along the final dimension.""" + per_byte = 8 // bits + pad = (-codes.shape[-1]) % per_byte + if pad: + # Pad only the feature dimension of this 2D [row, feature] array. + codes = np.pad(codes, ((0, 0), (0, pad)), constant_values=0) + # Group the padded feature dimension into chunks that each form one byte. + codes = codes.reshape(codes.shape[0], -1, per_byte).astype(np.uint16) + shifts = bits * np.arange(per_byte - 1, -1, -1, dtype=np.uint16) + weights = (1 << shifts).astype(np.uint16) + return np.sum(codes * weights, axis=-1).astype(np.uint8) diff --git a/gigl/common/utils/feature_quantization/torch_ops.py b/gigl/common/utils/feature_quantization/torch_ops.py new file mode 100644 index 000000000..ab802a3b3 --- /dev/null +++ b/gigl/common/utils/feature_quantization/torch_ops.py @@ -0,0 +1,42 @@ +import torch + +from gigl.types.graph import FeatureQuantizationMetadata + + +def dequantize_torch_tensor( + packed_features: torch.Tensor, + metadata: FeatureQuantizationMetadata, +) -> torch.Tensor: + """Reconstruct approximate float features from packed uint8 codes.""" + q = metadata + + if packed_features.size(-1) != q.packed_feature_dim: + raise ValueError( + f"Expected packed feature dim {q.packed_feature_dim} for " + f"{q.quantized_feature_dim} {q.bits}-bit features, got " + f"{packed_features.size(-1)}." + ) + + codes = _unpack_torch_tensor( + packed_features, dim=q.quantized_feature_dim, bits=q.bits + ).float() + if q.bits == 1: + if q.neg_mean is None or q.pos_mean is None: + raise ValueError("1-bit dequantization requires pos_mean/neg_mean") + return torch.where(codes.bool(), q.pos_mean, q.neg_mean) + else: + if q.clip_min is None or q.clip_max is None: + raise ValueError(f"{q.bits}-bit dequantization requires clip_min/clip_max") + levels = (1 << q.bits) - 1 + return q.clip_min + (codes / levels) * (q.clip_max - q.clip_min) + + +def _unpack_torch_tensor( + packed_features: torch.Tensor, *, dim: int, bits: int +) -> torch.Tensor: + per_byte = 8 // bits + mask = (1 << bits) - 1 + # Extract high-bits-first codes from each packed byte. + shifts = bits * torch.arange(per_byte - 1, -1, -1, device=packed_features.device) + codes = (packed_features.unsqueeze(-1).to(torch.int16) >> shifts).bitwise_and(mask) + return codes.reshape(*packed_features.shape[:-1], -1)[..., :dim].to(torch.uint8) diff --git a/gigl/types/graph.py b/gigl/types/graph.py index e7323693e..1d4d208a0 100644 --- a/gigl/types/graph.py +++ b/gigl/types/graph.py @@ -1,6 +1,7 @@ import gc from collections import abc from dataclasses import dataclass +from functools import cached_property, lru_cache from typing import Literal, Optional, TypeVar, Union, overload import torch @@ -107,6 +108,75 @@ class FeatureInfo: dtype: torch.dtype +@dataclass(frozen=True) +class FeatureQuantizationIndexTensors: + """Device-local indices used to scatter dequantized and raw features.""" + + quantized: torch.Tensor + raw: torch.Tensor + + +@dataclass(frozen=True) +class FeatureQuantizationMetadata: + """Metadata needed to unpack/dequantize/scatter packed features.""" + + bits: int + feature_dim: int + quantized_feature_indices: tuple[int, ...] = () + clip_min: Optional[float] = None + clip_max: Optional[float] = None + neg_mean: Optional[float] = None + pos_mean: Optional[float] = None + + def __post_init__(self) -> None: + valid_bits = (1, 2, 4, 8) + if self.bits not in valid_bits: + raise ValueError(f"bits must be one of {valid_bits}, got {self.bits}") + if any(i < 0 or i >= self.feature_dim for i in self.quantized_feature_indices): + raise ValueError( + f"quantized_feature_indices must be in [0, {self.feature_dim}), got {self.quantized_feature_indices}" + ) + if len(set(self.quantized_feature_indices)) != len( + self.quantized_feature_indices + ): + raise ValueError( + f"quantized_feature_indices contains duplicates: {self.quantized_feature_indices}" + ) + + @property + def quantized_feature_dim(self) -> int: + """Number of logical features stored in packed quantized form.""" + return len(self.quantized_feature_indices) + + @property + def packed_feature_dim(self) -> int: + """Number of uint8 columns needed for the packed features.""" + per_byte = 8 // self.bits + return (self.quantized_feature_dim + per_byte - 1) // per_byte + + @cached_property + def raw_feature_indices(self) -> tuple[int, ...]: + """Logical feature positions that remain in raw form.""" + quantized_indices = set(self.quantized_feature_indices) + return tuple(i for i in range(self.feature_dim) if i not in quantized_indices) + + @property + def raw_feature_dim(self) -> int: + """Number of logical features that remain in raw form.""" + return len(self.raw_feature_indices) + + @lru_cache(maxsize=2) + def scatter_index_tensors( + self, device: torch.device + ) -> FeatureQuantizationIndexTensors: + """Device-local indices for scattering quantized and raw features.""" + quantized = torch.tensor( + self.quantized_feature_indices, dtype=torch.long, device=device + ) + raw = torch.tensor(self.raw_feature_indices, dtype=torch.long, device=device) + return FeatureQuantizationIndexTensors(quantized=quantized, raw=raw) + + def _get_label_edges( labeled_edge_index: torch.Tensor, edge_dir: Literal["in", "out"], From 3b9134e7261594f3ce91379ea81913a5ed925cd2 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 21:12:24 +0000 Subject: [PATCH 03/78] Minor format updates --- gigl/common/utils/feature_quantization/numpy_ops.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index bb8e30f28..47b090a95 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -7,8 +7,9 @@ def quantize_ndarray( features: np.ndarray, *, bits: int, stats: Mapping[str, float] ) -> np.ndarray: """Quantize a 2D float array into packed uint8 codes.""" - if bits not in (1, 2, 4, 8): - raise ValueError(f"bits must be one of 1, 2, 4, or 8, got {bits}.") + valid_bits = (1, 2, 4, 8) + if bits not in valid_bits: + raise ValueError(f"bits must be one of {valid_bits}, got {bits}") if features.ndim != 2: raise ValueError(f"Expected a 2D feature array, got shape {features.shape}.") if bits == 1: @@ -21,10 +22,10 @@ def quantize_ndarray( clipped = np.clip(features, lo, hi) scaled = (clipped - lo) / (hi - lo) codes = np.rint(scaled * levels).astype(np.uint8) - return pack_codes(codes, bits) + return _pack_codes(codes, bits) -def pack_codes(codes: np.ndarray, bits: int) -> np.ndarray: +def _pack_codes(codes: np.ndarray, bits: int) -> np.ndarray: """Pack low-bit feature codes high-bits-first along the final dimension.""" per_byte = 8 // bits pad = (-codes.shape[-1]) % per_byte From 17a9d84c4e7a76199aead6a62c3d4e9169f04f44 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 22:04:15 +0000 Subject: [PATCH 04/78] Add tests --- .../utils/feature_quantization/numpy_ops.py | 7 + .../utils/feature_quantization/torch_ops.py | 7 + .../utils/feature_quantization/__init__.py | 0 .../feature_quantization/numpy_ops_test.py | 125 ++++++++++ .../feature_quantization/torch_ops_test.py | 230 ++++++++++++++++++ tests/unit/types_tests/graph_test.py | 91 +++++++ 6 files changed, 460 insertions(+) create mode 100644 tests/unit/common/utils/feature_quantization/__init__.py create mode 100644 tests/unit/common/utils/feature_quantization/numpy_ops_test.py create mode 100644 tests/unit/common/utils/feature_quantization/torch_ops_test.py diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index 47b090a95..ccc355a7b 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -1,3 +1,10 @@ +"""NumPy feature quantization helpers for preprocessing. + +Quantization runs in the data preprocessor, where feature data is stored as CPU +arrays and torch is not available. Dequantization lives in torch_ops.py because +the dataloader collate path operates on torch tensors that may already be on GPU. +""" + from collections.abc import Mapping import numpy as np diff --git a/gigl/common/utils/feature_quantization/torch_ops.py b/gigl/common/utils/feature_quantization/torch_ops.py index ab802a3b3..934c15ff9 100644 --- a/gigl/common/utils/feature_quantization/torch_ops.py +++ b/gigl/common/utils/feature_quantization/torch_ops.py @@ -1,3 +1,10 @@ +"""Torch feature dequantization helpers for dataloader collation. + +Quantization lives in numpy_ops.py because preprocessing works with CPU arrays +in an environment without torch. Dequantization runs in the dataloader collate +path, where packed feature data is already represented as torch tensors. +""" + import torch from gigl.types.graph import FeatureQuantizationMetadata diff --git a/tests/unit/common/utils/feature_quantization/__init__.py b/tests/unit/common/utils/feature_quantization/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py new file mode 100644 index 000000000..23fa6ff26 --- /dev/null +++ b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py @@ -0,0 +1,125 @@ +import numpy as np + +from gigl.common.utils.feature_quantization.numpy_ops import quantize_ndarray +from tests.test_assets.test_case import TestCase + + +class NumpyFeatureQuantizationOpsTest(TestCase): + def test_quantize_ndarray_single_bit_packs_full_byte(self) -> None: + # Values > 0 become 1 and the rest become 0: + # [-1, 0, 0.5, 2, -0.5, 3, 4, -4] -> [0, 0, 1, 1, 0, 1, 1, 0]. + # High-bits-first packing gives 0b00110110 = 54. + features = np.array([[-1.0, 0.0, 0.5, 2.0, -0.5, 3.0, 4.0, -4.0]]) + + actual = quantize_ndarray(features, bits=1, stats={}) + + np.testing.assert_array_equal(actual, np.array([[54]], dtype=np.uint8)) + + def test_quantize_ndarray_single_bit_pads_final_byte(self) -> None: + # Values > 0 become 1 and the rest become 0: + # [1, -1, 2, -2, 3] becomes [1, 0, 1, 0, 1]. + # Padding fills the remaining bit slots with zeros: 0b10101000 = 168. + features = np.array([[1.0, -1.0, 2.0, -2.0, 3.0]]) + + actual = quantize_ndarray(features, bits=1, stats={}) + + np.testing.assert_array_equal(actual, np.array([[168]], dtype=np.uint8)) + + def test_quantize_ndarray_two_bit_packs_full_byte_ascending_codes(self) -> None: + # With clip range [0, 3], these values equal their 2-bit codes. + # [0, 1, 2, 3] packs as 00 01 10 11 = 0b00011011 = 27. + features = np.array([[0.0, 1.0, 2.0, 3.0]]) + + actual = quantize_ndarray( + features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} + ) + + np.testing.assert_array_equal(actual, np.array([[27]], dtype=np.uint8)) + + def test_quantize_ndarray_two_bit_packs_full_byte_descending_codes(self) -> None: + # With clip range [0, 3], these values equal their 2-bit codes. + # [3, 2, 1, 0] packs as 11 10 01 00 = 0b11100100 = 228. + features = np.array([[3.0, 2.0, 1.0, 0.0]]) + + actual = quantize_ndarray( + features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} + ) + + np.testing.assert_array_equal(actual, np.array([[228]], dtype=np.uint8)) + + def test_quantize_ndarray_two_bit_pads_final_byte(self) -> None: + # The first four codes [0, 1, 2, 3] pack into byte 27. + # The leftover code [1] starts the next byte as 01 00 00 00 = 64. + features = np.array([[0.0, 1.0, 2.0, 3.0, 1.0]]) + + actual = quantize_ndarray( + features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} + ) + + np.testing.assert_array_equal(actual, np.array([[27, 64]], dtype=np.uint8)) + + def test_quantize_ndarray_four_bit_packs_full_byte_ascending_codes(self) -> None: + # With clip range [0, 15], these values equal their 4-bit codes. + # [0, 15] packs as 0000 1111 = 15. + features = np.array([[0.0, 15.0]]) + + actual = quantize_ndarray( + features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} + ) + + np.testing.assert_array_equal(actual, np.array([[15]], dtype=np.uint8)) + + def test_quantize_ndarray_four_bit_packs_full_byte_descending_codes(self) -> None: + # With clip range [0, 15], these values equal their 4-bit codes. + # [15, 0] packs as 1111 0000 = 240. + features = np.array([[15.0, 0.0]]) + + actual = quantize_ndarray( + features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} + ) + + np.testing.assert_array_equal(actual, np.array([[240]], dtype=np.uint8)) + + def test_quantize_ndarray_four_bit_pads_final_byte(self) -> None: + # The first two codes [0, 15] pack into byte 15. + # The leftover code [8] starts the next byte as 1000 0000 = 128. + features = np.array([[0.0, 15.0, 8.0]]) + + actual = quantize_ndarray( + features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} + ) + + np.testing.assert_array_equal(actual, np.array([[15, 128]], dtype=np.uint8)) + + def test_quantize_ndarray_eight_bit_stores_one_code_per_column(self) -> None: + # 8-bit quantization has one code per uint8 column, so no bit packing changes the order. + features = np.array([[0.0, 128.0, 255.0]]) + + actual = quantize_ndarray( + features, bits=8, stats={"clip_min": 0.0, "clip_max": 255.0} + ) + + np.testing.assert_array_equal(actual, np.array([[0, 128, 255]], dtype=np.uint8)) + + def test_quantize_ndarray_clips_multi_bit_values_before_packing(self) -> None: + # Values are clipped before scaling: [-1, 0, 3, 4] over [0, 3] + # becomes codes [0, 0, 3, 3], which packs as 00 00 11 11 = 15. + features = np.array([[-1.0, 0.0, 3.0, 4.0]]) + + actual = quantize_ndarray( + features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} + ) + + np.testing.assert_array_equal(actual, np.array([[15]], dtype=np.uint8)) + + def test_quantize_ndarray_rejects_invalid_bit_width(self) -> None: + with self.assertRaises(ValueError): + quantize_ndarray( + np.zeros((1, 1)), bits=3, stats={"clip_min": 0.0, "clip_max": 1.0} + ) + + def test_quantize_ndarray_rejects_non_2d_features(self) -> None: + with self.assertRaises(ValueError): + quantize_ndarray( + np.zeros((1, 1, 1)), bits=2, stats={"clip_min": 0.0, "clip_max": 1.0} + ) diff --git a/tests/unit/common/utils/feature_quantization/torch_ops_test.py b/tests/unit/common/utils/feature_quantization/torch_ops_test.py new file mode 100644 index 000000000..58dab6d62 --- /dev/null +++ b/tests/unit/common/utils/feature_quantization/torch_ops_test.py @@ -0,0 +1,230 @@ +import torch + +from gigl.common.utils.feature_quantization.torch_ops import dequantize_torch_tensor +from gigl.types.graph import FeatureQuantizationMetadata +from tests.test_assets.test_case import TestCase + + +class TorchFeatureQuantizationOpsTest(TestCase): + def test_dequantize_torch_tensor_single_bit_unpacks_full_byte(self) -> None: + # 0b10101010 = 170 unpacks high-bits-first to [1, 0, 1, 0, 1, 0, 1, 0]. + # Code 1 maps to pos_mean and code 0 maps to neg_mean. + metadata = FeatureQuantizationMetadata( + bits=1, + feature_dim=8, + quantized_feature_indices=tuple(range(8)), + neg_mean=-1.5, + pos_mean=2.5, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[170]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close( + actual, + torch.tensor([[2.5, -1.5, 2.5, -1.5, 2.5, -1.5, 2.5, -1.5]]), + ) + + def test_dequantize_torch_tensor_single_bit_trims_padded_codes(self) -> None: + # 0b10101000 = 168 unpacks to [1, 0, 1, 0, 1, 0, 0, 0]. + # The final three zeros are padding and are trimmed to five logical features. + metadata = FeatureQuantizationMetadata( + bits=1, + feature_dim=5, + quantized_feature_indices=tuple(range(5)), + neg_mean=-1.5, + pos_mean=2.5, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[168]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[2.5, -1.5, 2.5, -1.5, 2.5]])) + + def test_dequantize_torch_tensor_two_bit_unpacks_ascending_codes(self) -> None: + # 27 = 0b00011011 unpacks high-bits-first into 2-bit codes [0, 1, 2, 3]. + # With clip range [0, 3], those codes dequantize exactly to the same values. + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=tuple(range(4)), + clip_min=0.0, + clip_max=3.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[27]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[0.0, 1.0, 2.0, 3.0]])) + + def test_dequantize_torch_tensor_two_bit_unpacks_descending_codes(self) -> None: + # 228 = 0b11100100 unpacks high-bits-first into 2-bit codes [3, 2, 1, 0]. + # With clip range [0, 3], those codes dequantize exactly to the same values. + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=tuple(range(4)), + clip_min=0.0, + clip_max=3.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[228]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[3.0, 2.0, 1.0, 0.0]])) + + def test_dequantize_torch_tensor_two_bit_trims_padded_codes(self) -> None: + # [27, 64] unpacks to [0, 1, 2, 3, 1, 0, 0, 0]. + # The final three zeros are padding and are trimmed to five logical features. + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=5, + quantized_feature_indices=tuple(range(5)), + clip_min=0.0, + clip_max=3.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[27, 64]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[0.0, 1.0, 2.0, 3.0, 1.0]])) + + def test_dequantize_torch_tensor_four_bit_unpacks_ascending_codes(self) -> None: + # 15 = 0x0F unpacks high-bits-first into 4-bit codes [0, 15]. + # With clip range [0, 15], those codes dequantize exactly to the same values. + metadata = FeatureQuantizationMetadata( + bits=4, + feature_dim=2, + quantized_feature_indices=tuple(range(2)), + clip_min=0.0, + clip_max=15.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[15]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[0.0, 15.0]])) + + def test_dequantize_torch_tensor_four_bit_unpacks_descending_codes(self) -> None: + # 240 = 0xF0 unpacks high-bits-first into 4-bit codes [15, 0]. + # With clip range [0, 15], those codes dequantize exactly to the same values. + metadata = FeatureQuantizationMetadata( + bits=4, + feature_dim=2, + quantized_feature_indices=tuple(range(2)), + clip_min=0.0, + clip_max=15.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[240]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[15.0, 0.0]])) + + def test_dequantize_torch_tensor_four_bit_trims_padded_codes(self) -> None: + # [15, 128] unpacks to 4-bit codes [0, 15, 8, 0]. + # The final zero is padding and is trimmed to three logical features. + metadata = FeatureQuantizationMetadata( + bits=4, + feature_dim=3, + quantized_feature_indices=tuple(range(3)), + clip_min=0.0, + clip_max=15.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[15, 128]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[0.0, 15.0, 8.0]])) + + def test_dequantize_torch_tensor_eight_bit_reads_one_code_per_column(self) -> None: + # 8-bit quantization stores one code per uint8 column, so unpacking preserves order. + metadata = FeatureQuantizationMetadata( + bits=8, + feature_dim=3, + quantized_feature_indices=tuple(range(3)), + clip_min=0.0, + clip_max=255.0, + ) + + actual = dequantize_torch_tensor( + torch.tensor([[0, 128, 255]], dtype=torch.uint8), metadata=metadata + ) + + torch.testing.assert_close(actual, torch.tensor([[0.0, 128.0, 255.0]])) + + def test_dequantize_torch_tensor_rejects_wrong_packed_feature_dim(self) -> None: + # Five 2-bit features need two packed bytes: four codes in the first byte, + # then one code plus padding in the second byte. One input byte is too short. + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=5, + quantized_feature_indices=tuple(range(5)), + clip_min=0.0, + clip_max=3.0, + ) + + with self.assertRaises(ValueError): + dequantize_torch_tensor( + torch.tensor([[27]], dtype=torch.uint8), metadata=metadata + ) + + def test_dequantize_torch_tensor_requires_single_bit_neg_mean(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=1, + feature_dim=2, + quantized_feature_indices=(0, 1), + pos_mean=1.0, + ) + + with self.assertRaises(ValueError): + dequantize_torch_tensor( + torch.tensor([[128]], dtype=torch.uint8), metadata=metadata + ) + + def test_dequantize_torch_tensor_requires_single_bit_pos_mean(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=1, + feature_dim=2, + quantized_feature_indices=(0, 1), + neg_mean=-1.0, + ) + + with self.assertRaises(ValueError): + dequantize_torch_tensor( + torch.tensor([[128]], dtype=torch.uint8), metadata=metadata + ) + + def test_dequantize_torch_tensor_requires_multi_bit_clip_min(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=4, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_max=1.0, + ) + + with self.assertRaises(ValueError): + dequantize_torch_tensor( + torch.tensor([[15]], dtype=torch.uint8), metadata=metadata + ) + + def test_dequantize_torch_tensor_requires_multi_bit_clip_max(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=4, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_min=0.0, + ) + + with self.assertRaises(ValueError): + dequantize_torch_tensor( + torch.tensor([[15]], dtype=torch.uint8), metadata=metadata + ) diff --git a/tests/unit/types_tests/graph_test.py b/tests/unit/types_tests/graph_test.py index 3abe43b93..b1b8b54d2 100644 --- a/tests/unit/types_tests/graph_test.py +++ b/tests/unit/types_tests/graph_test.py @@ -8,6 +8,8 @@ from gigl.types.graph import ( DEFAULT_HOMOGENEOUS_EDGE_TYPE, DEFAULT_HOMOGENEOUS_NODE_TYPE, + FeatureQuantizationIndexTensors, + FeatureQuantizationMetadata, LoadedGraphTensors, is_label_edge_type, label_edge_type_to_message_passing_edge_type, @@ -84,6 +86,95 @@ def test_from_heterogeneous_invalid(self, _, input_value): with self.assertRaises(ValueError): to_homogeneous(input_value) + def test_feature_quantization_metadata_counts_quantized_features(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=2, feature_dim=6, quantized_feature_indices=(1, 3, 4) + ) + + self.assertEqual(metadata.quantized_feature_dim, 3) + + def test_feature_quantization_metadata_defaults_to_no_quantized_features( + self, + ) -> None: + metadata = FeatureQuantizationMetadata(bits=8, feature_dim=3) + + self.assertEqual(metadata.quantized_feature_indices, ()) + self.assertEqual(metadata.quantized_feature_dim, 0) + self.assertEqual(metadata.packed_feature_dim, 0) + self.assertEqual(metadata.raw_feature_indices, (0, 1, 2)) + self.assertEqual(metadata.raw_feature_dim, 3) + + def test_feature_quantization_metadata_packed_dim_without_padding(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=2, feature_dim=4, quantized_feature_indices=(0, 1, 2, 3) + ) + + self.assertEqual(metadata.packed_feature_dim, 1) + + def test_feature_quantization_metadata_packed_dim_with_padding(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=2, feature_dim=5, quantized_feature_indices=(0, 1, 2, 3, 4) + ) + + self.assertEqual(metadata.packed_feature_dim, 2) + + def test_feature_quantization_metadata_raw_feature_indices(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=4, feature_dim=6, quantized_feature_indices=(0, 2, 5) + ) + + self.assertEqual(metadata.raw_feature_indices, (1, 3, 4)) + + def test_feature_quantization_metadata_raw_feature_dim(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=4, feature_dim=6, quantized_feature_indices=(0, 2, 5) + ) + + self.assertEqual(metadata.raw_feature_dim, 3) + + def test_feature_quantization_metadata_scatter_index_tensors(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=4, feature_dim=6, quantized_feature_indices=(0, 2, 5) + ) + + index_tensors = metadata.scatter_index_tensors(torch.device("cpu")) + + self.assertIsInstance(index_tensors, FeatureQuantizationIndexTensors) + self.assertEqual(index_tensors.quantized.device, torch.device("cpu")) + self.assertEqual(index_tensors.raw.device, torch.device("cpu")) + self.assertEqual(index_tensors.quantized.dtype, torch.long) + self.assertEqual(index_tensors.raw.dtype, torch.long) + torch.testing.assert_close( + index_tensors.quantized, torch.tensor([0, 2, 5], dtype=torch.long) + ) + torch.testing.assert_close( + index_tensors.raw, torch.tensor([1, 3, 4], dtype=torch.long) + ) + + def test_feature_quantization_metadata_rejects_invalid_bit_width(self) -> None: + with self.assertRaises(ValueError): + FeatureQuantizationMetadata( + bits=3, feature_dim=2, quantized_feature_indices=(0, 1) + ) + + def test_feature_quantization_metadata_rejects_negative_index(self) -> None: + with self.assertRaises(ValueError): + FeatureQuantizationMetadata( + bits=2, feature_dim=2, quantized_feature_indices=(-1, 1) + ) + + def test_feature_quantization_metadata_rejects_index_at_feature_dim(self) -> None: + with self.assertRaises(ValueError): + FeatureQuantizationMetadata( + bits=2, feature_dim=2, quantized_feature_indices=(0, 2) + ) + + def test_feature_quantization_metadata_rejects_duplicate_indices(self) -> None: + with self.assertRaises(ValueError): + FeatureQuantizationMetadata( + bits=2, feature_dim=3, quantized_feature_indices=(0, 1, 1) + ) + @parameterized.expand( [ param( From 4376ea76d63798c72a3877c12be757d900f5da13 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 22:09:16 +0000 Subject: [PATCH 05/78] Whitespace --- .../feature_quantization/numpy_ops_test.py | 20 ---------------- .../feature_quantization/torch_ops_test.py | 23 ------------------- tests/unit/types_tests/graph_test.py | 6 ----- 3 files changed, 49 deletions(-) diff --git a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py index 23fa6ff26..e481bc305 100644 --- a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py @@ -10,9 +10,7 @@ def test_quantize_ndarray_single_bit_packs_full_byte(self) -> None: # [-1, 0, 0.5, 2, -0.5, 3, 4, -4] -> [0, 0, 1, 1, 0, 1, 1, 0]. # High-bits-first packing gives 0b00110110 = 54. features = np.array([[-1.0, 0.0, 0.5, 2.0, -0.5, 3.0, 4.0, -4.0]]) - actual = quantize_ndarray(features, bits=1, stats={}) - np.testing.assert_array_equal(actual, np.array([[54]], dtype=np.uint8)) def test_quantize_ndarray_single_bit_pads_final_byte(self) -> None: @@ -20,96 +18,78 @@ def test_quantize_ndarray_single_bit_pads_final_byte(self) -> None: # [1, -1, 2, -2, 3] becomes [1, 0, 1, 0, 1]. # Padding fills the remaining bit slots with zeros: 0b10101000 = 168. features = np.array([[1.0, -1.0, 2.0, -2.0, 3.0]]) - actual = quantize_ndarray(features, bits=1, stats={}) - np.testing.assert_array_equal(actual, np.array([[168]], dtype=np.uint8)) def test_quantize_ndarray_two_bit_packs_full_byte_ascending_codes(self) -> None: # With clip range [0, 3], these values equal their 2-bit codes. # [0, 1, 2, 3] packs as 00 01 10 11 = 0b00011011 = 27. features = np.array([[0.0, 1.0, 2.0, 3.0]]) - actual = quantize_ndarray( features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} ) - np.testing.assert_array_equal(actual, np.array([[27]], dtype=np.uint8)) def test_quantize_ndarray_two_bit_packs_full_byte_descending_codes(self) -> None: # With clip range [0, 3], these values equal their 2-bit codes. # [3, 2, 1, 0] packs as 11 10 01 00 = 0b11100100 = 228. features = np.array([[3.0, 2.0, 1.0, 0.0]]) - actual = quantize_ndarray( features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} ) - np.testing.assert_array_equal(actual, np.array([[228]], dtype=np.uint8)) def test_quantize_ndarray_two_bit_pads_final_byte(self) -> None: # The first four codes [0, 1, 2, 3] pack into byte 27. # The leftover code [1] starts the next byte as 01 00 00 00 = 64. features = np.array([[0.0, 1.0, 2.0, 3.0, 1.0]]) - actual = quantize_ndarray( features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} ) - np.testing.assert_array_equal(actual, np.array([[27, 64]], dtype=np.uint8)) def test_quantize_ndarray_four_bit_packs_full_byte_ascending_codes(self) -> None: # With clip range [0, 15], these values equal their 4-bit codes. # [0, 15] packs as 0000 1111 = 15. features = np.array([[0.0, 15.0]]) - actual = quantize_ndarray( features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} ) - np.testing.assert_array_equal(actual, np.array([[15]], dtype=np.uint8)) def test_quantize_ndarray_four_bit_packs_full_byte_descending_codes(self) -> None: # With clip range [0, 15], these values equal their 4-bit codes. # [15, 0] packs as 1111 0000 = 240. features = np.array([[15.0, 0.0]]) - actual = quantize_ndarray( features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} ) - np.testing.assert_array_equal(actual, np.array([[240]], dtype=np.uint8)) def test_quantize_ndarray_four_bit_pads_final_byte(self) -> None: # The first two codes [0, 15] pack into byte 15. # The leftover code [8] starts the next byte as 1000 0000 = 128. features = np.array([[0.0, 15.0, 8.0]]) - actual = quantize_ndarray( features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} ) - np.testing.assert_array_equal(actual, np.array([[15, 128]], dtype=np.uint8)) def test_quantize_ndarray_eight_bit_stores_one_code_per_column(self) -> None: # 8-bit quantization has one code per uint8 column, so no bit packing changes the order. features = np.array([[0.0, 128.0, 255.0]]) - actual = quantize_ndarray( features, bits=8, stats={"clip_min": 0.0, "clip_max": 255.0} ) - np.testing.assert_array_equal(actual, np.array([[0, 128, 255]], dtype=np.uint8)) def test_quantize_ndarray_clips_multi_bit_values_before_packing(self) -> None: # Values are clipped before scaling: [-1, 0, 3, 4] over [0, 3] # becomes codes [0, 0, 3, 3], which packs as 00 00 11 11 = 15. features = np.array([[-1.0, 0.0, 3.0, 4.0]]) - actual = quantize_ndarray( features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} ) - np.testing.assert_array_equal(actual, np.array([[15]], dtype=np.uint8)) def test_quantize_ndarray_rejects_invalid_bit_width(self) -> None: diff --git a/tests/unit/common/utils/feature_quantization/torch_ops_test.py b/tests/unit/common/utils/feature_quantization/torch_ops_test.py index 58dab6d62..6f941c548 100644 --- a/tests/unit/common/utils/feature_quantization/torch_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/torch_ops_test.py @@ -16,11 +16,9 @@ def test_dequantize_torch_tensor_single_bit_unpacks_full_byte(self) -> None: neg_mean=-1.5, pos_mean=2.5, ) - actual = dequantize_torch_tensor( torch.tensor([[170]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close( actual, torch.tensor([[2.5, -1.5, 2.5, -1.5, 2.5, -1.5, 2.5, -1.5]]), @@ -36,11 +34,9 @@ def test_dequantize_torch_tensor_single_bit_trims_padded_codes(self) -> None: neg_mean=-1.5, pos_mean=2.5, ) - actual = dequantize_torch_tensor( torch.tensor([[168]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[2.5, -1.5, 2.5, -1.5, 2.5]])) def test_dequantize_torch_tensor_two_bit_unpacks_ascending_codes(self) -> None: @@ -53,11 +49,9 @@ def test_dequantize_torch_tensor_two_bit_unpacks_ascending_codes(self) -> None: clip_min=0.0, clip_max=3.0, ) - actual = dequantize_torch_tensor( torch.tensor([[27]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[0.0, 1.0, 2.0, 3.0]])) def test_dequantize_torch_tensor_two_bit_unpacks_descending_codes(self) -> None: @@ -70,11 +64,9 @@ def test_dequantize_torch_tensor_two_bit_unpacks_descending_codes(self) -> None: clip_min=0.0, clip_max=3.0, ) - actual = dequantize_torch_tensor( torch.tensor([[228]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[3.0, 2.0, 1.0, 0.0]])) def test_dequantize_torch_tensor_two_bit_trims_padded_codes(self) -> None: @@ -87,11 +79,9 @@ def test_dequantize_torch_tensor_two_bit_trims_padded_codes(self) -> None: clip_min=0.0, clip_max=3.0, ) - actual = dequantize_torch_tensor( torch.tensor([[27, 64]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[0.0, 1.0, 2.0, 3.0, 1.0]])) def test_dequantize_torch_tensor_four_bit_unpacks_ascending_codes(self) -> None: @@ -104,11 +94,9 @@ def test_dequantize_torch_tensor_four_bit_unpacks_ascending_codes(self) -> None: clip_min=0.0, clip_max=15.0, ) - actual = dequantize_torch_tensor( torch.tensor([[15]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[0.0, 15.0]])) def test_dequantize_torch_tensor_four_bit_unpacks_descending_codes(self) -> None: @@ -121,11 +109,9 @@ def test_dequantize_torch_tensor_four_bit_unpacks_descending_codes(self) -> None clip_min=0.0, clip_max=15.0, ) - actual = dequantize_torch_tensor( torch.tensor([[240]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[15.0, 0.0]])) def test_dequantize_torch_tensor_four_bit_trims_padded_codes(self) -> None: @@ -138,11 +124,9 @@ def test_dequantize_torch_tensor_four_bit_trims_padded_codes(self) -> None: clip_min=0.0, clip_max=15.0, ) - actual = dequantize_torch_tensor( torch.tensor([[15, 128]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[0.0, 15.0, 8.0]])) def test_dequantize_torch_tensor_eight_bit_reads_one_code_per_column(self) -> None: @@ -154,11 +138,9 @@ def test_dequantize_torch_tensor_eight_bit_reads_one_code_per_column(self) -> No clip_min=0.0, clip_max=255.0, ) - actual = dequantize_torch_tensor( torch.tensor([[0, 128, 255]], dtype=torch.uint8), metadata=metadata ) - torch.testing.assert_close(actual, torch.tensor([[0.0, 128.0, 255.0]])) def test_dequantize_torch_tensor_rejects_wrong_packed_feature_dim(self) -> None: @@ -171,7 +153,6 @@ def test_dequantize_torch_tensor_rejects_wrong_packed_feature_dim(self) -> None: clip_min=0.0, clip_max=3.0, ) - with self.assertRaises(ValueError): dequantize_torch_tensor( torch.tensor([[27]], dtype=torch.uint8), metadata=metadata @@ -184,7 +165,6 @@ def test_dequantize_torch_tensor_requires_single_bit_neg_mean(self) -> None: quantized_feature_indices=(0, 1), pos_mean=1.0, ) - with self.assertRaises(ValueError): dequantize_torch_tensor( torch.tensor([[128]], dtype=torch.uint8), metadata=metadata @@ -197,7 +177,6 @@ def test_dequantize_torch_tensor_requires_single_bit_pos_mean(self) -> None: quantized_feature_indices=(0, 1), neg_mean=-1.0, ) - with self.assertRaises(ValueError): dequantize_torch_tensor( torch.tensor([[128]], dtype=torch.uint8), metadata=metadata @@ -210,7 +189,6 @@ def test_dequantize_torch_tensor_requires_multi_bit_clip_min(self) -> None: quantized_feature_indices=(0, 1), clip_max=1.0, ) - with self.assertRaises(ValueError): dequantize_torch_tensor( torch.tensor([[15]], dtype=torch.uint8), metadata=metadata @@ -223,7 +201,6 @@ def test_dequantize_torch_tensor_requires_multi_bit_clip_max(self) -> None: quantized_feature_indices=(0, 1), clip_min=0.0, ) - with self.assertRaises(ValueError): dequantize_torch_tensor( torch.tensor([[15]], dtype=torch.uint8), metadata=metadata diff --git a/tests/unit/types_tests/graph_test.py b/tests/unit/types_tests/graph_test.py index b1b8b54d2..78493d3d1 100644 --- a/tests/unit/types_tests/graph_test.py +++ b/tests/unit/types_tests/graph_test.py @@ -90,7 +90,6 @@ def test_feature_quantization_metadata_counts_quantized_features(self) -> None: metadata = FeatureQuantizationMetadata( bits=2, feature_dim=6, quantized_feature_indices=(1, 3, 4) ) - self.assertEqual(metadata.quantized_feature_dim, 3) def test_feature_quantization_metadata_defaults_to_no_quantized_features( @@ -108,35 +107,30 @@ def test_feature_quantization_metadata_packed_dim_without_padding(self) -> None: metadata = FeatureQuantizationMetadata( bits=2, feature_dim=4, quantized_feature_indices=(0, 1, 2, 3) ) - self.assertEqual(metadata.packed_feature_dim, 1) def test_feature_quantization_metadata_packed_dim_with_padding(self) -> None: metadata = FeatureQuantizationMetadata( bits=2, feature_dim=5, quantized_feature_indices=(0, 1, 2, 3, 4) ) - self.assertEqual(metadata.packed_feature_dim, 2) def test_feature_quantization_metadata_raw_feature_indices(self) -> None: metadata = FeatureQuantizationMetadata( bits=4, feature_dim=6, quantized_feature_indices=(0, 2, 5) ) - self.assertEqual(metadata.raw_feature_indices, (1, 3, 4)) def test_feature_quantization_metadata_raw_feature_dim(self) -> None: metadata = FeatureQuantizationMetadata( bits=4, feature_dim=6, quantized_feature_indices=(0, 2, 5) ) - self.assertEqual(metadata.raw_feature_dim, 3) def test_feature_quantization_metadata_scatter_index_tensors(self) -> None: metadata = FeatureQuantizationMetadata( bits=4, feature_dim=6, quantized_feature_indices=(0, 2, 5) ) - index_tensors = metadata.scatter_index_tensors(torch.device("cpu")) self.assertIsInstance(index_tensors, FeatureQuantizationIndexTensors) From 838f2edd15e389e58ea69473cb988636eec18972 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 22:11:06 +0000 Subject: [PATCH 06/78] WIP --- gigl/common/utils/feature_quantization/numpy_ops.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index ccc355a7b..f7eb5db75 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -23,7 +23,7 @@ def quantize_ndarray( # 1-bit quantization keeps only sign; values restore from neg/pos means. codes = (features > 0).astype(np.uint8) else: - # Linearly map clipped values into integer buckets. + # Min-max scale using clipped values and map to integer buckets. levels = (1 << bits) - 1 lo, hi = stats["clip_min"], stats["clip_max"] clipped = np.clip(features, lo, hi) From d357421b45c9eb1dc5a349773743997cb0353631 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 22:25:05 +0000 Subject: [PATCH 07/78] Add preprocessor diff --- .../data_preprocessor/data_preprocessor.py | 29 ++ .../lib/transform/feature_quantization.py | 274 ++++++++++++++++++ .../transform/transformed_features_info.py | 4 + .../data_preprocessor/lib/transform/utils.py | 22 ++ gigl/src/data_preprocessor/lib/types.py | 9 +- 5 files changed, 337 insertions(+), 1 deletion(-) create mode 100644 gigl/src/data_preprocessor/lib/transform/feature_quantization.py diff --git a/gigl/src/data_preprocessor/data_preprocessor.py b/gigl/src/data_preprocessor/data_preprocessor.py index 543954753..c2e9c3541 100644 --- a/gigl/src/data_preprocessor/data_preprocessor.py +++ b/gigl/src/data_preprocessor/data_preprocessor.py @@ -1,5 +1,6 @@ import argparse import concurrent.futures +import json import sys import threading from collections import defaultdict @@ -481,6 +482,33 @@ def generate_preprocessed_metadata_pb( feature_dim=feature_dim_output, transform_fn_assets_uri=node_transformed_features_info.transformed_features_transform_fn_assets_path.uri, ) + metadata_path = ( + node_transformed_features_info.feature_quantization_metadata_path.uri + ) + if tf.io.gfile.exists(metadata_path): + logger.info( + f"Loading feature quantization metadata from: {metadata_path}" + ) + with tf.io.gfile.GFile(metadata_path) as f: + metadata = json.loads(f.read()) + logger.info(f"Loaded feature 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"], + ) + 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 + ) preprocessed_metadata_pb.condensed_node_type_to_preprocessed_metadata[ int(condensed_node_type) ].CopyFrom(node_metadata_output_pb) @@ -698,6 +726,7 @@ def inner() -> FeatureSpecDict: pretrained_tft_model_uri=input_node_preprocessing_spec.pretrained_tft_model_uri, features_outputs=input_node_preprocessing_spec.features_outputs, labels_outputs=input_node_preprocessing_spec.labels_outputs, + feature_quantization_spec=input_node_preprocessing_spec.feature_quantization_spec, ) enumerated_node_refs_to_preprocessing_specs[ enumerated_node_metadata.enumerated_node_data_reference diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py new file mode 100644 index 000000000..3a0de2ec3 --- /dev/null +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -0,0 +1,274 @@ +import json +from typing import Iterable + +import apache_beam as beam +import numpy as np +import pyarrow as pa +from apache_beam.transforms.stats import ApproximateQuantiles +from tensorflow_metadata.proto.v0 import schema_pb2 +from tensorflow_transform.tf_metadata import schema_utils +from tensorflow_transform.tf_metadata.dataset_metadata import DatasetMetadata + +from gigl.common.logger import Logger +from gigl.common.utils.feature_quantization.numpy_ops import quantize_ndarray +from gigl.common.utils.tensorflow_schema import feature_spec_to_feature_index_map +from gigl.src.data_preprocessor.lib.types import FeatureQuantizationSpec + +logger = Logger() +_NODE_PACKED_FEATURE_KEY = "node_packed_features" +_SingleBitAcc = tuple[float, int, float, int] + + +def apply_feature_quantization_transform( + transformed_features: beam.PCollection[pa.RecordBatch], + transformed_metadata: DatasetMetadata, + analyzed_metadata: beam.PCollection[DatasetMetadata] | None, + spec: FeatureQuantizationSpec, + feature_keys: list[str], + metadata_path: str, +): + logger.info(f"Applying Beam feature quantization with spec: {spec}") + stats = _build_feature_quantization_stats(transformed_features, spec) + logical_metadata = ( + transformed_metadata + if analyzed_metadata is None + else beam.pvalue.AsSingleton(analyzed_metadata) + ) + _ = ( + stats + | "Build feature quantization metadata JSON" + >> beam.Map( + _feature_quantization_metadata_json, + spec=spec, + feature_keys=feature_keys, + dataset_metadata=logical_metadata, + ) + | "Write feature quantization metadata" + >> beam.io.WriteToText( + metadata_path, + num_shards=1, + shard_name_template="", + ) + ) + transformed_features = transformed_features | ( + "Quantize transformed feature RecordBatches" + >> beam.Map( + _quantize_record_batch, + spec=spec, + stats=beam.pvalue.AsSingleton(stats), + ) + ) + # Encode TFRecords with the compact physical schema. The persisted schema + # remains the original logical TFT schema because dequantization scatters + # features back. + if analyzed_metadata is None: + physical_metadata = DatasetMetadata( + _apply_feature_quantization_schema(transformed_metadata.schema, spec) + ) + else: + physical_metadata = analyzed_metadata | ( + "Apply feature quantization schema" + >> beam.Map( + lambda metadata, spec: DatasetMetadata( + _apply_feature_quantization_schema(metadata.schema, spec) + ), + spec=spec, + ) + ) + physical_metadata = beam.pvalue.AsSingleton(physical_metadata) + return transformed_features, physical_metadata + + +def _build_feature_quantization_stats( + record_batches: beam.PCollection[pa.RecordBatch], + spec: FeatureQuantizationSpec, +) -> beam.PCollection[dict[str, float]]: + if spec.bits not in (1, 2, 4, 8): + raise ValueError(f"bits must be one of 1, 2, 4, or 8, got {spec.bits}.") + if not spec.feature_keys: + raise ValueError("Feature quantization expects at least one feature key.") + logger.info( + f"Building Beam feature quantization stats for {len(spec.feature_keys)} " + f"features with bits={spec.bits}: {spec.feature_keys}" + ) + if spec.bits == 1: + return ( + record_batches + | "Compute single bit quantization stats" + >> beam.CombineGlobally(_SingleBitStatsFn(spec.feature_keys)) + ) + return ( + record_batches + | "Build multi-bit quantization value batches" + >> beam.Map(_build_feature_values, feature_keys=spec.feature_keys) + | "Compute multi-bit quantization quantiles" + >> ApproximateQuantiles.Globally(num_quantiles=1000, input_batched=True) + | "Build multi-bit quantization stats" + >> beam.Map(_multi_bit_stats_from_quantiles) + ) + + +def _quantize_record_batch( + batch: pa.RecordBatch, + spec: FeatureQuantizationSpec, + stats: dict[str, float], +) -> pa.RecordBatch: + features = _build_feature_matrix(batch, spec.feature_keys) + packed = quantize_ndarray(features, bits=spec.bits, stats=stats) + schema_names = batch.schema.names + keep_indices = [ + i for i, name in enumerate(schema_names) if name not in set(spec.feature_keys) + ] + arrays = [batch.column(i) for i in keep_indices] + names = [schema_names[i] for i in keep_indices] + arrays.append( + pa.array([[row.tobytes()] for row in packed], type=pa.list_(pa.binary())) + ) + names.append(_NODE_PACKED_FEATURE_KEY) + return pa.RecordBatch.from_arrays(arrays, names=names) + + +def _feature_quantization_metadata_json( + stats: dict[str, float], + spec: FeatureQuantizationSpec, + feature_keys: list[str], + dataset_metadata: DatasetMetadata, +) -> str: + raw_feature_spec = schema_utils.schema_as_feature_spec( + dataset_metadata.schema + ).feature_spec + feature_key_set = set(feature_keys) + missing = [ + key + for key in spec.feature_keys + if key not in raw_feature_spec or key not in feature_key_set + ] + if missing: + raise ValueError( + f"Quantized feature keys missing from feature outputs: {missing}" + ) + feature_spec = {key: raw_feature_spec[key] for key in feature_keys} + feature_index = feature_spec_to_feature_index_map(feature_spec) + quantized_feature_indices = [] + for key in spec.feature_keys: + start, end = feature_index[key] + if end - start != 1: + raise ValueError( + f"Feature quantization expects scalar features, got {key}." + ) + quantized_feature_indices.append(start) + metadata = { + "packed_feature_key": _NODE_PACKED_FEATURE_KEY, + "quantized_feature_indices": quantized_feature_indices, + "bits": spec.bits, + **stats, + } + logger.info(f"Writing feature quantization metadata: {metadata}") + return json.dumps(metadata) + + +def _apply_feature_quantization_schema( + schema: schema_pb2.Schema, spec: FeatureQuantizationSpec +) -> schema_pb2.Schema: + drop_keys = set(spec.feature_keys) | {_NODE_PACKED_FEATURE_KEY} + quantized_schema = schema_pb2.Schema() + quantized_schema.CopyFrom(schema) + del quantized_schema.feature[:] + quantized_schema.feature.extend( + 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.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(spec.feature_keys)} features and added bytes feature " + f"{_NODE_PACKED_FEATURE_KEY}." + ) + return quantized_schema + + +def _build_feature_values( + batch: pa.RecordBatch, feature_keys: list[str] +) -> list[float]: + values = _build_feature_matrix(batch, feature_keys).reshape(-1) + return values[np.isfinite(values)].astype(float).tolist() + + +def _build_feature_matrix(batch: pa.RecordBatch, feature_keys: list[str]) -> np.ndarray: + key_to_idx = {name: i for i, name in enumerate(batch.schema.names)} + cols: list[np.ndarray] = [] + for key in feature_keys: + if key not in key_to_idx: + raise ValueError(f"Feature key {key} not found in RecordBatch.") + col = batch.column(key_to_idx[key]) + values = np.asarray(col.to_numpy(zero_copy_only=False), dtype=np.float32) + if values.ndim != 1: + raise ValueError( + f"Feature quantization expects scalar features, got {key} with shape {values.shape}." + ) + cols.append(values) + return np.stack(cols, axis=1) + + +def _multi_bit_stats_from_quantiles(quantiles: list[float]) -> dict[str, float]: + if not quantiles: + raise ValueError("Cannot compute quantization stats from no values.") + quantile_count = len(quantiles) - 1 + clip_min = float(quantiles[round(0.005 * quantile_count)]) + clip_max = float(quantiles[round(0.995 * quantile_count)]) + if clip_max <= clip_min: + clip_max = clip_min + 1e-5 + stats = {"clip_min": clip_min, "clip_max": clip_max} + logger.info(f"Computed feature quantization stats: {stats}") + return stats + + +class _SingleBitStatsFn(beam.CombineFn): + """Beam CombineFn that accumulates sums and counts across batches. + + Used to derive the mean of positive and negative feature values for 1-bit quantization. + """ + + def __init__(self, feature_keys: list[str]): + self._feature_keys = feature_keys + + def create_accumulator(self) -> _SingleBitAcc: + return 0.0, 0, 0.0, 0 + + def add_input( + self, accumulator: _SingleBitAcc, batch: pa.RecordBatch + ) -> _SingleBitAcc: + neg_sum, neg_count, pos_sum, pos_count = accumulator + values = _build_feature_matrix(batch, self._feature_keys).reshape(-1) + values = values[np.isfinite(values)] + neg = values <= 0 + pos = values > 0 + return ( + neg_sum + float(values[neg].sum()), + neg_count + int(neg.sum()), + pos_sum + float(values[pos].sum()), + pos_count + int(pos.sum()), + ) + + def merge_accumulators( + self, accumulators: Iterable[_SingleBitAcc] + ) -> _SingleBitAcc: + neg_sum = neg_count = pos_sum = pos_count = 0 + for n_sum, n_count, p_sum, p_count in accumulators: + neg_sum += n_sum + neg_count += n_count + pos_sum += p_sum + pos_count += p_count + return neg_sum, neg_count, pos_sum, pos_count + + def extract_output(self, accumulator: _SingleBitAcc) -> dict[str, float]: + neg_sum, neg_count, pos_sum, pos_count = accumulator + stats = { + "neg_mean": neg_sum / neg_count if neg_count else 0.0, + "pos_mean": pos_sum / pos_count if pos_count else 0.0, + } + logger.info(f"Computed Beam feature quantization stats: {stats}") + return stats diff --git a/gigl/src/data_preprocessor/lib/transform/transformed_features_info.py b/gigl/src/data_preprocessor/lib/transform/transformed_features_info.py index c7c0669eb..903a02057 100644 --- a/gigl/src/data_preprocessor/lib/transform/transformed_features_info.py +++ b/gigl/src/data_preprocessor/lib/transform/transformed_features_info.py @@ -22,6 +22,7 @@ class TransformedFeaturesInfo: raw_data_schema_file_path: GcsUri tft_temp_directory_path: GcsUri transformed_features_file_prefix: GcsUri + feature_quantization_metadata_path: GcsUri transformed_features_schema_path: GcsUri transform_directory_path: GcsUri dataflow_console_uri: Optional[HttpUri] = None @@ -92,6 +93,9 @@ def __init__( custom_identifier=custom_identifier, ) ) + self.feature_quantization_metadata_path = GcsUri.join( + self.transform_directory_path, "feature_quantization_metadata.json" + ) self.transformed_features_schema_path = ( gcs_constants.get_tf_transformed_features_schema_path( diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index 7e9be7081..5dbf79a90 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -27,6 +27,9 @@ EdgeDataReference, NodeDataReference, ) +from gigl.src.data_preprocessor.lib.transform.feature_quantization import ( + apply_feature_quantization_transform, +) from gigl.src.data_preprocessor.lib.transform.tf_value_encoder import TFValueEncoder from gigl.src.data_preprocessor.lib.transform.transformed_features_info import ( TransformedFeaturesInfo, @@ -362,6 +365,25 @@ def get_load_data_and_transform_pipeline_component( if should_use_existing_transform_fn else beam.pvalue.AsSingleton(analyzed_transform_fn[1].deferred_metadata) # type: ignore ) + q_spec = None + if isinstance(preprocessing_spec, NodeDataPreprocessingSpec): + q_spec = preprocessing_spec.feature_quantization_spec + if q_spec is not None: + if should_use_existing_transform_fn: + analyzed_metadata = None + else: + analyzed_metadata = analyzed_transform_fn[1].deferred_metadata # type: ignore + + transformed_features, resolved_transformed_metadata = ( + apply_feature_quantization_transform( + transformed_features=transformed_features, + transformed_metadata=transformed_metadata, + analyzed_metadata=analyzed_metadata, + spec=q_spec, + feature_keys=list(preprocessing_spec.features_outputs or []), + metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, + ) + ) transformed_features | "Write tf record files" >> BetterWriteToTFRecord( file_path_prefix=transformed_features_info.transformed_features_file_prefix.uri, diff --git a/gigl/src/data_preprocessor/lib/types.py b/gigl/src/data_preprocessor/lib/types.py index 491cc9486..a3308b00d 100644 --- a/gigl/src/data_preprocessor/lib/types.py +++ b/gigl/src/data_preprocessor/lib/types.py @@ -48,6 +48,11 @@ class NodeOutputIdentifier(str): """ +class FeatureQuantizationSpec(NamedTuple): + feature_keys: list[str] + bits: int + + class EdgeOutputIdentifier(NamedTuple): """ References the TFTransform output fields / column names for src and dst node ids of an edge. @@ -72,6 +77,7 @@ class NodeDataPreprocessingSpec(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"""NodeDataPreprocessingSpec( @@ -80,7 +86,8 @@ def __repr__(self) -> str: preprocessing_fn={self.preprocessing_fn}, pretrained_tft_model_uri={self.pretrained_tft_model_uri}, features_outputs={self.features_outputs}, - labels_outputs={self.labels_outputs}) + labels_outputs={self.labels_outputs}, + feature_quantization_spec={self.feature_quantization_spec}) """ From 459b8cbd17143475bea4ebbb85bc013677949b48 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 23:15:35 +0000 Subject: [PATCH 08/78] Add test --- .../feature_quantization_transform_test.py | 121 ++++++++++++++++++ 1 file changed, 121 insertions(+) create mode 100644 tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py diff --git a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py new file mode 100644 index 000000000..a852a62a0 --- /dev/null +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -0,0 +1,121 @@ +import json +import os +import tempfile + +import apache_beam as beam +import pyarrow as pa +import tensorflow as tf +from apache_beam.testing.test_pipeline import TestPipeline +from apache_beam.testing.util import assert_that, equal_to +from tensorflow_metadata.proto.v0 import schema_pb2 +from tensorflow_transform.tf_metadata.dataset_metadata import DatasetMetadata + +from gigl.src.data_preprocessor.lib.transform.feature_quantization import ( + apply_feature_quantization_transform, +) +from gigl.src.data_preprocessor.lib.types import FeatureQuantizationSpec +from tests.test_assets.test_case import TestCase + + +def _column_pylist(batch: pa.RecordBatch, name: str) -> list: + return batch.column(batch.schema.names.index(name)).to_pylist() + + +def _record_batch_summary(batch: pa.RecordBatch) -> dict[str, object]: + return { + "names": batch.schema.names, + "node_id": _column_pylist(batch, "node_id"), + "label": _column_pylist(batch, "label"), + "node_packed_features": _column_pylist(batch, "node_packed_features"), + } + + +class FeatureQuantizationTransformTest(TestCase): + def test_apply_feature_quantization_transform_quantizes_multi_bit_features( + self, + ) -> None: + with tempfile.TemporaryDirectory() as temp_dir: + metadata_path = os.path.join(temp_dir, "feature_quantization_metadata.json") + batch = pa.RecordBatch.from_arrays( + [ + pa.array([10, 11], type=pa.int64()), + pa.array([-2.0, 8.0], type=pa.float32()), + pa.array([-2.0, 8.0], type=pa.float32()), + pa.array([0, 1], type=pa.int64()), + ], + names=["node_id", "f0", "f1", "label"], + ) + transformed_metadata = DatasetMetadata.from_feature_spec( + { + "node_id": tf.io.FixedLenFeature(shape=[], dtype=tf.int64), + "f0": tf.io.FixedLenFeature(shape=[], dtype=tf.float32), + "f1": tf.io.FixedLenFeature(shape=[], dtype=tf.float32), + "label": tf.io.FixedLenFeature(shape=[], dtype=tf.int64), + } + ) + + with TestPipeline() as pipeline: + transformed_batches, physical_metadata = ( + apply_feature_quantization_transform( + transformed_features=pipeline + | "Create RecordBatch" >> beam.Create([batch]), + transformed_metadata=transformed_metadata, + analyzed_metadata=None, + spec=FeatureQuantizationSpec(feature_keys=["f0", "f1"], bits=2), + feature_keys=["f0", "f1"], + metadata_path=metadata_path, + ) + ) + physical_features = { + feature.name: feature + for feature in physical_metadata.schema.feature + } + self.assertEqual( + set(physical_features), + {"node_id", "label", "node_packed_features"}, + ) + self.assertEqual( + physical_features["node_packed_features"].type, + schema_pb2.BYTES, + ) + self.assertEqual( + physical_features["node_packed_features"].value_count.min, + 1, + ) + self.assertEqual( + physical_features["node_packed_features"].value_count.max, + 1, + ) + + # These values sit exactly at the learned clip bounds, so this + # does not depend on mid-bucket rounding: min/min maps to + # 00000000 and max/max maps to 11110000 with two padded codes. + assert_that( + transformed_batches + | "Summarize RecordBatch" >> beam.Map(_record_batch_summary), + equal_to( + [ + { + "names": [ + "node_id", + "label", + "node_packed_features", + ], + "node_id": [10, 11], + "label": [0, 1], + "node_packed_features": [ + [bytes([0])], + [bytes([240])], + ], + } + ] + ), + ) + + with open(metadata_path) as metadata_file: + metadata = json.load(metadata_file) + self.assertEqual(metadata["packed_feature_key"], "node_packed_features") + self.assertEqual(metadata["quantized_feature_indices"], [0, 1]) + self.assertEqual(metadata["bits"], 2) + self.assertEqual(metadata["clip_min"], -2.0) + self.assertEqual(metadata["clip_max"], 8.0) From be34beb0178b9fd552c330a558920a5ef433d013 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 3 Aug 2026 23:21:34 +0000 Subject: [PATCH 09/78] Update read side --- gigl/common/data/dataloaders.py | 51 +++++++++- gigl/common/data/load_torch_tensors.py | 41 ++++++++ .../serialized_graph_metadata_translator.py | 93 ++++++++++++++++++- gigl/types/graph.py | 7 ++ tests/unit/common/data/dataloaders_test.py | 24 ++--- 5 files changed, 198 insertions(+), 18 deletions(-) diff --git a/gigl/common/data/dataloaders.py b/gigl/common/data/dataloaders.py index 0ac83f0e7..e0eab92e0 100644 --- a/gigl/common/data/dataloaders.py +++ b/gigl/common/data/dataloaders.py @@ -21,6 +21,7 @@ class LoadedEntityTensors(NamedTuple): ids: torch.Tensor features: Optional[torch.Tensor] + quantized_features: Optional[torch.Tensor] labels: Optional[torch.Tensor] @@ -44,6 +45,10 @@ class SerializedTFRecordInfo: feature_dim: int # Entity ID Key for current entity. If this is a Node Entity, this must be a string. If this is an edge entity, this must be a Tuple[str, str] for the source and destination ids. entity_key: Union[str, Tuple[str, str]] + # Packed uint8 feature name to load for the current node entity. + packed_feature_key: Optional[str] = None + # Number of packed uint8 columns for the current node entity. + packed_feature_dim: int = 0 # Name of the label columns for the current entity, defaults to an empty list. label_keys: Sequence[str] = field(default_factory=list) # The regex pattern to match the TFRecord files at the specified prefix @@ -367,10 +372,11 @@ def load_as_torch_tensors( serialized_tf_record_info (SerializedTFRecordInfo): Information for how TFRecord files are serialized on disk. tf_dataset_options (TFDatasetOptions): The options to use when building the dataset. Returns: - LoadedEntityTensors: The (id_tensor, feature_tensor, label_tensor) for the loaded entities. + LoadedEntityTensors: The (id_tensor, feature_tensor, quantized_feature_tensor, label_tensor) for the loaded entities. """ entity_key = serialized_tf_record_info.entity_key feature_keys = serialized_tf_record_info.feature_keys + packed_feature_key = serialized_tf_record_info.packed_feature_key label_keys = serialized_tf_record_info.label_keys # We make a deep copy of the feature spec dict so that future modifications don't redirect to the input @@ -392,6 +398,16 @@ 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( @@ -435,13 +451,24 @@ def load_as_torch_tensors( else: empty_feature = None + if packed_feature_key is not None: + empty_quantized_feature = torch.empty( + (0, serialized_tf_record_info.packed_feature_dim), + dtype=torch.uint8, + ) + else: + empty_quantized_feature = None + if label_keys: empty_label = torch.empty(0, len(label_keys)) else: empty_label = None return LoadedEntityTensors( - ids=empty_entity, features=empty_feature, labels=empty_label + ids=empty_entity, + features=empty_feature, + quantized_features=empty_quantized_feature, + labels=empty_label, ) dataset = TFRecordDataLoader._build_dataset_for_uris( @@ -454,6 +481,7 @@ def load_as_torch_tensors( num_entities_processed = 0 id_tensors: list[torch.Tensor] = [] feature_tensors: list[torch.Tensor] = [] + quantized_feature_tensors: list[tf.Tensor] = [] label_tensors: list[torch.Tensor] = [] for idx, batch in enumerate(dataset): id_tensors.append(proccess_id_tensor(batch)) @@ -465,6 +493,15 @@ def load_as_torch_tensors( feature_tensors.append(feature_tensor) if label_tensor is not None: label_tensors.append(label_tensor) + if packed_feature_key is not None: + quantized_feature_tensor = tf.io.decode_raw( + batch[packed_feature_key], tf.uint8 + ) + quantized_feature_tensor = tf.reshape( + quantized_feature_tensor, + [-1, serialized_tf_record_info.packed_feature_dim], + ) + quantized_feature_tensors.append(quantized_feature_tensor) num_entities_processed += ( id_tensors[-1].shape[0] if entity_type == FeatureTypes.NODE @@ -483,11 +520,16 @@ def load_as_torch_tensors( tf.concat(id_tensors, axis=id_concat_axis) ) output_feature_tensor: Optional[torch.Tensor] = None + output_quantized_feature_tensor: Optional[torch.Tensor] = None output_label_tensor: Optional[torch.Tensor] = None if feature_tensors: output_feature_tensor = _tf_tensor_to_torch_tensor( tf.concat(feature_tensors, axis=0) ) + if quantized_feature_tensors: + output_quantized_feature_tensor = _tf_tensor_to_torch_tensor( + tf.concat(quantized_feature_tensors, axis=0) + ).to(torch.uint8) if label_tensors: output_label_tensor = _tf_tensor_to_torch_tensor( tf.concat(label_tensors, axis=0) @@ -503,5 +545,8 @@ def load_as_torch_tensors( f"Converted {num_entities_processed:,} {entity_type.name} to torch tensors in {end - start:.2f} seconds" ) return LoadedEntityTensors( - ids=id_tensor, features=output_feature_tensor, labels=output_label_tensor + ids=id_tensor, + features=output_feature_tensor, + quantized_features=output_quantized_feature_tensor, + labels=output_label_tensor, ) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 2ba43ce10..7b3545e33 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -18,6 +18,7 @@ from gigl.types.graph import ( DEFAULT_HOMOGENEOUS_EDGE_TYPE, DEFAULT_HOMOGENEOUS_NODE_TYPE, + FeatureQuantizationMetadata, LoadedGraphTensors, ) from gigl.utils.share_memory import share_memory @@ -26,6 +27,7 @@ _ID_FMT = "{entity}_ids" _FEATURE_FMT = "{entity}_features" +_PACKED_FEATURE_FMT = "{entity}_packed_features" _LABEL_FMT = "{entity}_labels" _EDGE_WEIGHTS_KEY = "edge_weights" _NODE_KEY = "node" @@ -113,6 +115,10 @@ class SerializedGraphMetadata: negative_label_entity_info: Optional[ Union[SerializedTFRecordInfo, dict[EdgeType, SerializedTFRecordInfo]] ] = None + # Optional node quantization metadata. + node_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] + ] = None def _data_loading_process( @@ -178,6 +184,7 @@ def _data_loading_process( ids: dict[Union[NodeType, EdgeType], torch.Tensor] = {} features: dict[Union[NodeType, EdgeType], torch.Tensor] = {} + quantized_features: dict[Union[NodeType, EdgeType], torch.Tensor] = {} labels: dict[Union[NodeType, EdgeType], torch.Tensor] = {} weights: dict[Union[NodeType, EdgeType], torch.Tensor] = {} for ( @@ -192,12 +199,20 @@ 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 + ): + 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, ) entity_ids = loaded_entity.ids entity_features = loaded_entity.features + entity_quantized_features = loaded_entity.quantized_features entity_labels = loaded_entity.labels ids[graph_type] = entity_ids logger.info( @@ -213,6 +228,16 @@ def _data_loading_process( f"Rank {rank} did not detect {entity_type} features for graph type {graph_type} from {serialized_entity_tf_record_info.tfrecord_uri_prefix.uri}" ) + if entity_quantized_features is not None: + quantized_features[graph_type] = entity_quantized_features + logger.info( + f"Rank {rank} finished loading {entity_type} quantized features of shape {entity_quantized_features.shape} for graph type {graph_type} from {serialized_entity_tf_record_info.tfrecord_uri_prefix.uri}" + ) + else: + logger.info( + f"Rank {rank} did not detect {entity_type} quantized features for graph type {graph_type} from {serialized_entity_tf_record_info.tfrecord_uri_prefix.uri}" + ) + if entity_labels is not None: labels[graph_type] = entity_labels logger.info( @@ -291,6 +316,12 @@ def _data_loading_process( share_memory(features) # We convert the features back to homogeneous from the default heterogeneous setup if our provided input was homogeneous + if quantized_features: + logger.info( + f"Rank {rank} is attempting to share {entity_type} quantized feature memory for tfrecord directories: {all_tf_record_uris}" + ) + share_memory(quantized_features) + if labels: logger.info( f"Rank {rank} is attempting to share {entity_type} label memory for tfrecord directories: {all_tf_record_uris}" @@ -310,6 +341,12 @@ def _data_loading_process( output_dict[_FEATURE_FMT.format(entity=entity_type)] = ( list(features.values())[0] if is_input_homogeneous else features ) + if quantized_features: + output_dict[_PACKED_FEATURE_FMT.format(entity=entity_type)] = ( + list(quantized_features.values())[0] + if is_input_homogeneous + else quantized_features + ) if labels: output_dict[_LABEL_FMT.format(entity=entity_type)] = ( list(labels.values())[0] if is_input_homogeneous else labels @@ -480,6 +517,9 @@ def load_torch_tensors_from_tf_record( node_ids = node_output_dict[_ID_FMT.format(entity=_NODE_KEY)] node_features = node_output_dict.get(_FEATURE_FMT.format(entity=_NODE_KEY), None) + node_quantized_features = node_output_dict.get( + _PACKED_FEATURE_FMT.format(entity=_NODE_KEY), None + ) node_labels = node_output_dict.get(_LABEL_FMT.format(entity=_NODE_KEY), None) edge_index = edge_output_dict[_ID_FMT.format(entity=_EDGE_KEY)] @@ -507,6 +547,7 @@ def load_torch_tensors_from_tf_record( return LoadedGraphTensors( node_ids=node_ids, node_features=node_features, + node_quantized_features=node_quantized_features, node_labels=node_labels, edge_index=edge_index, edge_features=edge_features, diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index c774a57e6..35321e8b6 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -3,13 +3,14 @@ from gigl.common import UriFactory from gigl.common.data.dataloaders import SerializedTFRecordInfo from gigl.common.data.load_torch_tensors import SerializedGraphMetadata +from gigl.common.utils.tensorflow_schema import feature_spec_to_feature_index_map from gigl.src.common.types.graph_data import EdgeType, NodeType from gigl.src.common.types.pb_wrappers.graph_metadata import GraphMetadataPbWrapper from gigl.src.common.types.pb_wrappers.preprocessed_metadata import ( PreprocessedMetadataPbWrapper, ) from gigl.src.data_preprocessor.lib.types import FeatureSpecDict -from gigl.types.graph import to_homogeneous +from gigl.types.graph import FeatureQuantizationMetadata, to_homogeneous from snapchat.research.gbml.preprocessed_metadata_pb2 import PreprocessedMetadata @@ -33,19 +34,91 @@ def _build_serialized_tfrecord_entity_info( Returns: SerializedTFRecordInfo: Stored metadata for current entity """ + packed_feature_key = None + packed_feature_dim = 0 + physical_feature_keys = list(preprocessed_metadata.feature_keys) + feature_dim = preprocessed_metadata.feature_dim + + if isinstance( + preprocessed_metadata, PreprocessedMetadata.NodeMetadataOutput + ) and preprocessed_metadata.HasField("quantized_feature_metadata"): + quantization_metadata = _build_feature_quantization_metadata( + quantized_metadata=preprocessed_metadata.quantized_feature_metadata, + feature_dim=preprocessed_metadata.feature_dim, + ) + packed_feature_key = ( + preprocessed_metadata.quantized_feature_metadata.packed_feature_key + ) + packed_feature_dim = quantization_metadata.packed_feature_dim + quantized_indices = set(quantization_metadata.quantized_feature_indices) + feature_index = feature_spec_to_feature_index_map( + {key: feature_spec_dict[key] for key in preprocessed_metadata.feature_keys} + ) + + physical_feature_keys = [] + for key in preprocessed_metadata.feature_keys: + key_indices = set(range(*feature_index[key])) + quantized_key_indices = key_indices.intersection(quantized_indices) + if not quantized_key_indices: + physical_feature_keys.append(key) + elif quantized_key_indices != key_indices: + raise ValueError( + f"Partial feature quantization is not supported for {key}." + ) + feature_dim = quantization_metadata.raw_feature_dim + + physical_keys = set(physical_feature_keys) + physical_keys.update(preprocessed_metadata.label_keys) + if packed_feature_key is not None: + physical_keys.add(packed_feature_key) + feature_spec_dict = { + key: spec for key, spec in feature_spec_dict.items() if key in physical_keys + } + return SerializedTFRecordInfo( tfrecord_uri_prefix=UriFactory.create_uri( preprocessed_metadata.tfrecord_uri_prefix ), - feature_keys=list(preprocessed_metadata.feature_keys), + feature_keys=physical_feature_keys, feature_spec=feature_spec_dict, - feature_dim=preprocessed_metadata.feature_dim, + feature_dim=feature_dim, entity_key=entity_key, + packed_feature_key=packed_feature_key, + packed_feature_dim=packed_feature_dim, label_keys=list(preprocessed_metadata.label_keys), tfrecord_uri_pattern=tfrecord_uri_pattern, ) +def _build_feature_quantization_metadata( + quantized_metadata: PreprocessedMetadata.FeatureQuantizationMetadata, + feature_dim: int, +) -> FeatureQuantizationMetadata: + state = quantized_metadata.WhichOneof("state") + + neg_mean = pos_mean = clip_min = clip_max = None + if state == "single_bit_state": + bits = 1 + neg_mean = quantized_metadata.single_bit_state.neg_mean + pos_mean = quantized_metadata.single_bit_state.pos_mean + elif state == "multi_bit_state": + bits = quantized_metadata.multi_bit_state.bits + clip_min = quantized_metadata.multi_bit_state.clip_min + clip_max = quantized_metadata.multi_bit_state.clip_max + else: + raise ValueError("Expected quantization state to be set.") + + return FeatureQuantizationMetadata( + bits=bits, + feature_dim=feature_dim, + quantized_feature_indices=tuple(quantized_metadata.quantized_feature_indices), + clip_min=clip_min, + clip_max=clip_max, + neg_mean=neg_mean, + pos_mean=pos_mean, + ) + + def convert_pb_to_serialized_graph_metadata( preprocessed_metadata_pb_wrapper: PreprocessedMetadataPbWrapper, graph_metadata_pb_wrapper: GraphMetadataPbWrapper, @@ -65,6 +138,7 @@ def convert_pb_to_serialized_graph_metadata( edge_entity_info: dict[EdgeType, SerializedTFRecordInfo] = {} positive_label_entity_info: dict[EdgeType, SerializedTFRecordInfo] = {} negative_label_entity_info: dict[EdgeType, SerializedTFRecordInfo] = {} + node_quantization_metadata: dict[NodeType, FeatureQuantizationMetadata] = {} preprocessed_metadata_pb = preprocessed_metadata_pb_wrapper.preprocessed_metadata_pb @@ -92,6 +166,13 @@ def convert_pb_to_serialized_graph_metadata( entity_key=node_key, tfrecord_uri_pattern=tfrecord_uri_pattern, ) + if node_metadata.HasField("quantized_feature_metadata"): + node_quantization_metadata[node_type] = ( + _build_feature_quantization_metadata( + quantized_metadata=node_metadata.quantized_feature_metadata, + feature_dim=node_metadata.feature_dim, + ) + ) for edge_type in graph_metadata_pb_wrapper.edge_types: condensed_edge_type = ( @@ -159,6 +240,9 @@ def convert_pb_to_serialized_graph_metadata( negative_label_entity_info=to_homogeneous(negative_label_entity_info) if len(negative_label_entity_info) > 0 else None, + node_quantization_metadata=to_homogeneous(node_quantization_metadata) + if len(node_quantization_metadata) > 0 + else None, ) else: return SerializedGraphMetadata( @@ -170,4 +254,7 @@ def convert_pb_to_serialized_graph_metadata( negative_label_entity_info=negative_label_entity_info if len(negative_label_entity_info) > 0 else None, + node_quantization_metadata=node_quantization_metadata + if len(node_quantization_metadata) > 0 + else None, ) diff --git a/gigl/types/graph.py b/gigl/types/graph.py index 1d4d208a0..fec65b9d8 100644 --- a/gigl/types/graph.py +++ b/gigl/types/graph.py @@ -221,6 +221,10 @@ class LoadedGraphTensors: negative_label: Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]] # Unpartitioned Edge Weights (per-edge sampling weights, one scalar per edge) edge_weights: Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]] = None + # Unpartitioned packed uint8 node features. + node_quantized_features: Optional[ + Union[torch.Tensor, dict[NodeType, torch.Tensor]] + ] = None def treat_labels_as_edges(self, edge_dir: Literal["in", "out"]) -> None: """ @@ -319,6 +323,9 @@ def treat_labels_as_edges(self, edge_dir: Literal["in", "out"]) -> None: self.node_ids = to_heterogeneous_node(self.node_ids) self.node_features = to_heterogeneous_node(self.node_features) + self.node_quantized_features = to_heterogeneous_node( + self.node_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/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index bb7bad002..8911ef982 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -291,7 +291,7 @@ def test_load_as_torch_tensors( ): """Test TFRecordDataLoader's ability to load features and optionally labels.""" loader = TFRecordDataLoader(rank=0, world_size=1) - node_ids, feature_tensor, label_tensor = loader.load_as_torch_tensors( + loaded = loader.load_as_torch_tensors( serialized_tf_record_info=SerializedTFRecordInfo( tfrecord_uri_prefix=UriFactory.create_uri(self.data_dir), feature_spec=feature_spec, @@ -305,11 +305,11 @@ def test_load_as_torch_tensors( ) # Verify entity IDs are loaded correctly - assert_close(node_ids, expected_id_tensor) + assert_close(loaded.ids, expected_id_tensor) - assert_close(feature_tensor, expected_feature_tensor) + assert_close(loaded.features, expected_feature_tensor) - assert_close(label_tensor, expected_label_tensor) + assert_close(loaded.labels, expected_label_tensor) def test_build_dataset_for_uris(self): dataset = TFRecordDataLoader._build_dataset_for_uris( @@ -396,7 +396,7 @@ def test_load_empty_directory( self.addCleanup(temp_dir.cleanup) loader = TFRecordDataLoader(rank=0, world_size=1) - node_ids, feature_tensor, label_tensor = loader.load_as_torch_tensors( + loaded = loader.load_as_torch_tensors( serialized_tf_record_info=SerializedTFRecordInfo( tfrecord_uri_prefix=UriFactory.create_uri(temp_dir.name), feature_spec={}, # Doesn't matter what this is. @@ -408,9 +408,9 @@ def test_load_empty_directory( tf_dataset_options=TFDatasetOptions(deterministic=True), ) - assert_close(node_ids, expected_node_ids) - assert_close(feature_tensor, expected_features) - assert_close(label_tensor, expected_label_tensor) + assert_close(loaded.ids, expected_node_ids) + assert_close(loaded.features, expected_features) + assert_close(loaded.labels, expected_label_tensor) @parameterized.expand( [ @@ -470,7 +470,7 @@ def test_load_labels_from_pb(self): condensed_node_type ] loader = TFRecordDataLoader(rank=0, world_size=1) - _, feature_tensor, label_tensor = loader.load_as_torch_tensors( + loaded = loader.load_as_torch_tensors( serialized_tf_record_info=SerializedTFRecordInfo( tfrecord_uri_prefix=UriFactory.create_uri( node_metadata.tfrecord_uri_prefix @@ -487,9 +487,9 @@ def test_load_labels_from_pb(self): tf_dataset_options=TFDatasetOptions(deterministic=True), ) # Ensure we have loaded data - assert feature_tensor is not None and label_tensor is not None - self.assertEqual(feature_tensor.size(1), node_metadata.feature_dim) - self.assertEqual(label_tensor.size(1), len(node_metadata.label_keys)) + assert loaded.features is not None and loaded.labels is not None + self.assertEqual(loaded.features.size(1), node_metadata.feature_dim) + self.assertEqual(loaded.labels.size(1), len(node_metadata.label_keys)) def test_load_edge_weights_from_tf_record(self): """Edge weight column is extracted from edge features and returned separately. From be4a9fef332e107b479543e0fd2ee1bbf0b3baeb Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 16:13:52 +0000 Subject: [PATCH 10/78] Add integration code --- gigl/distributed/base_dist_loader.py | 1 + gigl/distributed/base_sampler.py | 44 +++++ gigl/distributed/dataset_factory.py | 12 +- gigl/distributed/dist_ablp_neighborloader.py | 21 +++ gigl/distributed/dist_dataset.py | 148 +++++++++++++++ gigl/distributed/dist_partitioner.py | 176 +++++++++++++++--- gigl/distributed/dist_range_partitioner.py | 29 ++- .../distributed/distributed_neighborloader.py | 20 ++ gigl/distributed/graph_store/dist_server.py | 15 +- .../graph_store/remote_dist_dataset.py | 24 +++ gigl/distributed/sampler.py | 1 + gigl/distributed/utils/neighborloader.py | 90 ++++++++- gigl/types/graph.py | 6 + 13 files changed, 558 insertions(+), 29 deletions(-) diff --git a/gigl/distributed/base_dist_loader.py b/gigl/distributed/base_dist_loader.py index 79e5c8b98..60522cc0c 100644 --- a/gigl/distributed/base_dist_loader.py +++ b/gigl/distributed/base_dist_loader.py @@ -245,6 +245,7 @@ def __init__( ) self._node_feature_info = dataset_schema.node_feature_info self._edge_feature_info = dataset_schema.edge_feature_info + self._node_quantization_metadata = dataset_schema.node_quantization_metadata self._sampler_options = sampler_options self._non_blocking_transfers = non_blocking_transfers diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index a76944f55..5ff21b1e3 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -20,6 +20,7 @@ from gigl.common.logger import Logger from gigl.distributed.sampler import ( NEGATIVE_LABEL_METADATA_KEY, + NODE_PACKED_FEATURES_METADATA_KEY, POSITIVE_LABEL_METADATA_KEY, ABLPNodeSamplerInput, ) @@ -110,9 +111,28 @@ def __init__(self, *args, **kwargs) -> None: which GLT's event loop would swallow the same way it swallows the original sampling exception. """ + data = kwargs.get("data") super().__init__(*args, **kwargs) self._sampling_error_sent: bool = False + self.dist_node_quantized_feature: Optional[DistFeature] = None + if ( + self.collect_features + and data is not None + and getattr(data, "node_quantized_features", None) is not None + ): + # Mirrors GLT's dist_node_feature initialization: + # https://github.com/alibaba/graphlearn-for-pytorch/blob/88ff111ac0d9e45c6c9d2d18cfc5883dca07e9f9/graphlearn_torch/python/distributed/dist_neighbor_sampler.py#L162-L167 + self.dist_node_quantized_feature = DistFeature( + data.num_partitions, + data.partition_idx, + data.node_quantized_features, + data.node_pb, + local_only=False, + rpc_router=self.rpc_router, + device=self.device, + ) + def _prepare_sample_loop_inputs( self, inputs: NodeSamplerInput, @@ -369,6 +389,26 @@ async def _collate_fn( futs[f"{as_str(ntype)}.nfeats"] = wrap_torch_future( self.dist_node_feature.async_get(nodes, ntype) ) + if self.dist_node_quantized_feature is not None: + if self.use_all2all: + sorted_ntype = sorted( + self.dist_node_quantized_feature.feature_pb.keys() + ) + quantized_nfeat_dict = self.dist_node_quantized_feature.get_all2all( + output, sorted_ntype + ) + for ntype, quantized_nfeats in quantized_nfeat_dict.items(): + result_map[ + f"#META.{NODE_PACKED_FEATURES_METADATA_KEY}.{as_str(ntype)}" + ] = quantized_nfeats + else: + for ntype, nodes in output.node.items(): + nodes = nodes.to(torch.long) + futs[ + f"#META.{NODE_PACKED_FEATURES_METADATA_KEY}.{as_str(ntype)}" + ] = wrap_torch_future( + self.dist_node_quantized_feature.async_get(nodes, ntype) + ) if self.dist_edge_feature is not None and self.with_edge: for etype in self.edge_types: if self.edge_dir == "in": @@ -416,6 +456,10 @@ async def _collate_fn( futs["nfeats"] = wrap_torch_future( self.dist_node_feature.async_get(output.node) ) + if self.dist_node_quantized_feature is not None: + futs[f"#META.{NODE_PACKED_FEATURES_METADATA_KEY}"] = wrap_torch_future( + self.dist_node_quantized_feature.async_get(output.node) + ) if self.dist_edge_feature is not None: eids = result_map["eids"] futs["efeats"] = wrap_torch_future( diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index cbd7006f8..1a5f859b5 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -180,6 +180,10 @@ def _load_and_build_partitioned_dataset( partitioner.register_node_features( node_features=loaded_graph_tensors.node_features ) + if loaded_graph_tensors.node_quantized_features is not None: + partitioner.register_node_quantized_features( + node_quantized_features=loaded_graph_tensors.node_quantized_features + ) if loaded_graph_tensors.node_labels is not None: partitioner.register_node_labels(node_labels=loaded_graph_tensors.node_labels) if loaded_graph_tensors.edge_weights is not None: @@ -205,6 +209,7 @@ def _load_and_build_partitioned_dataset( del ( loaded_graph_tensors.node_ids, loaded_graph_tensors.node_features, + loaded_graph_tensors.node_quantized_features, loaded_graph_tensors.edge_index, loaded_graph_tensors.edge_features, loaded_graph_tensors.edge_weights, @@ -217,7 +222,12 @@ def _load_and_build_partitioned_dataset( partition_output = partitioner.partition() logger.info(f"Initializing DistDataset instance with edge direction {edge_dir}") - dataset = DistDataset(rank=rank, world_size=world_size, edge_dir=edge_dir) + dataset = DistDataset( + rank=rank, + world_size=world_size, + edge_dir=edge_dir, + node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, + ) dataset.build( partition_output=partition_output, diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index be1856b7c..0787d0691 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -1,3 +1,4 @@ +import time from collections import abc, defaultdict from itertools import count from typing import Optional, Union @@ -33,6 +34,7 @@ extract_edge_type_metadata, extract_metadata, labeled_to_homogeneous, + materialize_quantized_node_features, set_missing_features, shard_nodes_by_process, strip_label_edges, @@ -579,6 +581,7 @@ def _setup_for_colocated( edge_types=edge_types, node_feature_info=dataset.node_feature_info, edge_feature_info=dataset.edge_feature_info, + node_quantization_metadata=dataset.node_quantization_metadata, edge_dir=dataset.edge_dir, ), ) @@ -768,6 +771,7 @@ def _setup_for_graph_store( edge_types=edge_types, node_feature_info=node_feature_info, edge_feature_info=edge_feature_info, + node_quantization_metadata=dataset.node_quantization_metadata, edge_dir=edge_dir, ), backend_key, @@ -866,9 +870,15 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: # around a GLT bug in to_hetero_data. extract_edge_type_metadata then # pulls out labels by prefix. # TODO (mkolodner-sc): Remove the need to extract metadata once GLT's `to_hetero_data` function is fixed + collate_start_time = time.perf_counter() metadata, stripped_msg = extract_metadata(msg, self.to_device) + base_collate_start_time = time.perf_counter() data = super()._collate_fn(stripped_msg) + base_collate_time = time.perf_counter() - base_collate_start_time + logger.debug( + f"Distributed ABLPNeighborLoader GLT base collate time: {base_collate_time:.3f}s" + ) data = set_missing_features( data=data, @@ -908,8 +918,19 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: data, metadata = self._apply_ppr_outputs(data, metadata) + data, metadata = materialize_quantized_node_features( + data=data, + metadata=metadata, + node_quantization_metadata=self._node_quantization_metadata, + ) + # Attach any remaining metadata (e.g. custom user-defined keys) directly onto the # data object so downstream code can access them via attribute lookup. for key, value in metadata.items(): data[key] = value + + collate_time = time.perf_counter() - collate_start_time + logger.debug( + f"Distributed ABLPNeighborLoader end-to-end collate time: {collate_time:.3f}s" + ) return data diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index 181c2c7d9..93ddf2ced 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -23,6 +23,7 @@ from gigl.types.graph import ( FeatureInfo, FeaturePartitionData, + FeatureQuantizationMetadata, GraphPartitionData, PartitionOutput, ) @@ -54,6 +55,9 @@ def __init__( node_feature_partition: Optional[ Union[Feature, dict[NodeType, Feature]] ] = None, + node_quantized_feature_partition: Optional[ + Union[Feature, dict[NodeType, Feature]] + ] = None, edge_feature_partition: Optional[ Union[Feature, dict[EdgeType, Feature]] ] = None, @@ -77,6 +81,15 @@ def __init__( node_feature_info: Optional[ Union[FeatureInfo, dict[NodeType, FeatureInfo]] ] = None, + node_quantized_feature_info: Optional[ + Union[FeatureInfo, dict[NodeType, FeatureInfo]] + ] = None, + node_quantization_metadata: Optional[ + Union[ + FeatureQuantizationMetadata, + dict[NodeType, FeatureQuantizationMetadata], + ] + ] = None, edge_feature_info: Optional[ Union[FeatureInfo, dict[EdgeType, FeatureInfo]] ] = None, @@ -94,9 +107,12 @@ def __init__( rank (int): Rank of the current process world_size (int): World size of the current process edge_dir (Literal["in", "out"]): Edge direction of the provied graph + node_quantization_metadata: Metadata for packed node features. + May be provided during initial construction or IPC rebuild. The below arguments are only expected to be provided when re-serializing an instance of the DistDataset class after build() has been called graph_partition (Optional[Union[Graph, dict[EdgeType, Graph]]]): Partitioned Graph Data node_feature_partition (Optional[Union[Feature, dict[NodeType, Feature]]]): Partitioned Node Feature Data + node_quantized_feature_partition (Optional[Union[Feature, dict[NodeType, Feature]]]): Partitioned packed uint8 node feature data edge_feature_partition (Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]]): Partitioned Edge Feature Data node_labels (Optional[Union[Feature, dict[NodeType, Feature]]]): The labels of each node on the current machine node_partition_book (Optional[Union[PartitionBook, dict[NodeType, PartitionBook]]]): Node Partition Book @@ -109,6 +125,7 @@ def __init__( num_test: (Optional[Union[int, dict[NodeType, int]]]): Number of test nodes on the current machine. Will be a dict if heterogeneous. node_feature_info: Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Dimension of node features and its data type, will be a dict if heterogeneous. Note this will be None in the homogeneous case if the data has no node features, or will only contain node types with node features in the heterogeneous case. + node_quantized_feature_info: Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Dimension and dtype for packed uint8 node features. edge_feature_info: Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]]: Dimension of edge features and its data type, will be a dict if heterogeneous. Note this will be None in the homogeneous case if the data has no edge features, or will only contain edge types with edge features in the heterogeneous case. degree_tensor: Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Pre-computed degree tensor. Lazily computed on first access via the degree_tensor property. @@ -151,6 +168,10 @@ def __init__( self._node_feature_info = node_feature_info self._edge_feature_info = edge_feature_info + self._node_quantized_feature_info = node_quantized_feature_info + self._node_quantized_features = node_quantized_feature_partition + self._node_quantization_metadata = node_quantization_metadata + self._degree_tensor: Optional[ Union[torch.Tensor, dict[NodeType, torch.Tensor]] ] = degree_tensor @@ -209,6 +230,19 @@ def node_features( ): self._node_features = new_node_features + @property + def node_quantized_features( + self, + ) -> Optional[Union[Feature, dict[NodeType, Feature]]]: + return self._node_quantized_features + + @node_quantized_features.setter + def node_quantized_features( + self, + new_node_quantized_features: Optional[Union[Feature, dict[NodeType, Feature]]], + ): + self._node_quantized_features = new_node_quantized_features + @property def edge_features(self) -> Optional[Union[Feature, dict[EdgeType, Feature]]]: """ @@ -303,6 +337,20 @@ def node_feature_info( """ return self._node_feature_info + @property + def node_quantized_feature_info( + self, + ) -> Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: + return self._node_quantized_feature_info + + @property + def node_quantization_metadata( + self, + ) -> Optional[ + Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] + ]: + return self._node_quantization_metadata + @property def edge_feature_info( self, @@ -723,6 +771,73 @@ def _initialize_node_features( ) logger.info("Initialized node features for homogeneous graph to dataset") + def _initialize_node_quantized_features( + self, + node_partition_book: Union[PartitionBook, dict[NodeType, PartitionBook]], + partitioned_node_quantized_features: Optional[ + Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]] + ], + ) -> None: + """ + Initializes packed uint8 node features in a separate feature store. + """ + + node_quantized_features, node_quantized_feature_id_to_index = ( + _prepare_feature_data( + partition_book=node_partition_book, + partitioned_data=partitioned_node_quantized_features, + ) + ) + + if ( + node_quantized_features is None + or node_quantized_feature_id_to_index is None + ): + logger.info("Found no node quantized features to initialize") + return + + # GLT only exposes init_node_features for the standard node feature + # store. Calling it here would overwrite self._node_features, so build + # this sidecar Feature store directly, mirroring GLT's construction: + # https://github.com/alibaba/graphlearn-for-pytorch/blob/88ff111ac0d9e45c6c9d2d18cfc5883dca07e9f9/graphlearn_torch/python/data/dataset.py#L236 + + if isinstance(node_quantized_features, Mapping): + assert isinstance(node_quantized_feature_id_to_index, Mapping) + self._node_quantized_features = { + node_type: Feature( + feature_tensor=features_per_node_type, + id2index=node_quantized_feature_id_to_index[node_type], # ty: ignore[invalid-argument-type] TODO(ty-torch-keyed-access): fix ty false positives for torch-backed keyed container access. + with_gpu=False, + dtype=torch.uint8, + ) + for node_type, features_per_node_type in node_quantized_features.items() + } + self._node_quantized_feature_info = {} + for node_type, features_per_node_type in node_quantized_features.items(): + assert not isinstance(node_type, EdgeType) + self._node_quantized_feature_info[node_type] = FeatureInfo( + dim=features_per_node_type.size(1), # ty: ignore[unresolved-attribute] TODO(ty-torch-keyed-access): fix ty false positives for torch-backed keyed container access. + dtype=features_per_node_type.dtype, # ty: ignore[unresolved-attribute] TODO(ty-torch-keyed-access): fix ty false positives for torch-backed keyed container access. + ) + logger.info( + f"Initialized node quantized features for heterogeneous graph to dataset with node types: {node_quantized_features.keys()}" + ) + else: + assert not isinstance(node_quantized_feature_id_to_index, Mapping) + self._node_quantized_features = Feature( + feature_tensor=node_quantized_features, + id2index=node_quantized_feature_id_to_index, + with_gpu=False, + dtype=torch.uint8, + ) + self._node_quantized_feature_info = FeatureInfo( + dim=node_quantized_features.size(1), + dtype=node_quantized_features.dtype, + ) + logger.info( + "Initialized node quantized features for homogeneous graph to dataset" + ) + def _initialize_node_labels( self, node_partition_book: Union[PartitionBook, dict[NodeType, PartitionBook]], @@ -897,6 +1012,13 @@ def build( partition_output.partitioned_node_features = None gc.collect() + self._initialize_node_quantized_features( + node_partition_book=partition_output.node_partition_book, + partitioned_node_quantized_features=partition_output.partitioned_node_quantized_features, + ) + partition_output.partitioned_node_quantized_features = None + gc.collect() + self._initialize_node_labels( node_partition_book=partition_output.node_partition_book, partitioned_node_labels=partition_output.partitioned_node_labels, @@ -935,6 +1057,7 @@ def share_ipc( Literal["in", "out"], Optional[Union[Graph, dict[EdgeType, Graph]]], Optional[Union[Feature, dict[NodeType, Feature]]], + Optional[Union[Feature, dict[NodeType, Feature]]], Optional[Union[Feature, dict[EdgeType, Feature]]], Optional[Union[Feature, dict[NodeType, Feature]]], Optional[Union[PartitionBook, dict[NodeType, PartitionBook]]], @@ -946,6 +1069,13 @@ def share_ipc( Optional[Union[int, dict[NodeType, int]]], Optional[Union[int, dict[NodeType, int]]], Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]], + Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]], + Optional[ + Union[ + FeatureQuantizationMetadata, + dict[NodeType, FeatureQuantizationMetadata], + ] + ], Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]], Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]], Optional[int], @@ -959,6 +1089,7 @@ def share_ipc( Literal["in", "out"]: Graph Edge Direction Optional[Union[Graph, dict[EdgeType, Graph]]]: Partitioned Graph Data Optional[Union[Feature, dict[NodeType, Feature]]]: Partitioned Node Feature Data + Optional[Union[Feature, dict[NodeType, Feature]]]: Partitioned packed uint8 node feature data Optional[Union[Feature, dict[EdgeType, Feature]]]: Partitioned Edge Feature Data Optional[Union[Feature, dict[NodeType, Feature]]]: Node labels on the current machine. Will be a dict if heterogeneous. Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Node Partition Book Tensor @@ -970,6 +1101,8 @@ def share_ipc( Optional[Union[int, dict[NodeType, int]]]: Number of validation nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[int, dict[NodeType, int]]]: Number of test nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Node feature dim and its data type, will be a dict if heterogeneous + Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Packed uint8 node feature dim and dtype + Optional node quantization metadata. Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]]: Edge feature dim and its data type, will be a dict if heterogeneous Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Degree tensors Optional[int]: Optional per-anchor label cap for ABLP label fetching @@ -990,6 +1123,7 @@ def share_ipc( self._edge_dir, self._graph, self._node_features, + self._node_quantized_features, self._edge_features, self._node_labels, self._node_partition_book, @@ -1001,6 +1135,8 @@ def share_ipc( self._num_val, # Additional field unique to DistDataset class self._num_test, # Additional field unique to DistDataset class self._node_feature_info, # Additional field unique to DistDataset class + self._node_quantized_feature_info, # Additional field unique to DistDataset class + self._node_quantization_metadata, # Additional field unique to DistDataset class self._edge_feature_info, # Additional field unique to DistDataset class self._degree_tensor, # Additional field unique to DistDataset class self._max_labels_per_anchor_node, # Additional field unique to DistDataset class @@ -1234,6 +1370,9 @@ def _rebuild_distributed_dataset( Optional[ Union[Feature, dict[NodeType, Feature]] ], # Partitioned Node Feature Data + Optional[ + Union[Feature, dict[NodeType, Feature]] + ], # Partitioned packed uint8 node feature data Optional[ Union[Feature, dict[EdgeType, Feature]] ], # Partitioned Edge Feature Data @@ -1257,6 +1396,15 @@ def _rebuild_distributed_dataset( Optional[ Union[FeatureInfo, dict[NodeType, FeatureInfo]] ], # Node feature dim and its data type + Optional[ + Union[FeatureInfo, dict[NodeType, FeatureInfo]] + ], # Packed uint8 node feature dim and dtype + Optional[ + Union[ + FeatureQuantizationMetadata, + dict[NodeType, FeatureQuantizationMetadata], + ] + ], # Node quantization metadata Optional[ Union[FeatureInfo, dict[EdgeType, FeatureInfo]] ], # Edge feature dim and its data type diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 2198557ba..04de8ce72 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -150,6 +150,9 @@ def __init__( node_features: Optional[ Union[torch.Tensor, dict[NodeType, torch.Tensor]] ] = None, + node_quantized_features: Optional[ + Union[torch.Tensor, dict[NodeType, torch.Tensor]] + ] = None, edge_index: Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]] = None, edge_features: Optional[ Union[torch.Tensor, dict[EdgeType, torch.Tensor]] @@ -171,6 +174,7 @@ def __init__( should_assign_edges_by_src_node (bool): Whether edges should be assigned to the machine of the source nodes during partitioning node_ids (Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]): Optionally registered node ids from input. Tensors should be of shape [num_nodes_on_current_rank] node_features (Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]): Optionally registered node feats from input. Tensors should be of shope [num_nodes_on_current_rank, node_feat_dim] + node_quantized_features (Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]): Optionally registered packed uint8 node features from input. Tensors should be of shape [num_nodes_on_current_rank, packed_node_feat_dim] edge_index (Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]]): Optionally registered edge indexes from input. Tensors should be of shape [2, num_edges_on_current_rank] edge_features (Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]]): Optionally registered edge features from input. Tensors should be of shape [num_edges_on_current_rank, edge_feat_dim] positive_labels (Optional[Union[torch.Tensor, dict[EdgeType, torch.Tensor]]]): Optionally registered positive labels from input. Tensors should be of shape [2, num_pos_labels_on_current_rank] @@ -195,6 +199,8 @@ def __init__( self._max_node_ids: Optional[dict[NodeType, int]] = None self._node_feat: Optional[dict[NodeType, torch.Tensor]] = None self._node_feat_dim: Optional[dict[NodeType, int]] = None + self._node_quantized_feat: Optional[dict[NodeType, torch.Tensor]] = None + self._node_quantized_feat_dim: Optional[dict[NodeType, int]] = None self._node_labels: Optional[dict[NodeType, torch.Tensor]] = None self._node_labels_dim: Optional[dict[NodeType, int]] = None @@ -218,6 +224,11 @@ def __init__( if node_features is not None: self.register_node_features(node_features=node_features) + if node_quantized_features is not None: + self.register_node_quantized_features( + node_quantized_features=node_quantized_features + ) + if edge_features is not None: self.register_edge_features(edge_features=edge_features) @@ -541,6 +552,40 @@ def register_node_features( for node_type in input_node_features: self._node_feat_dim[node_type] = input_node_features[node_type].shape[1] + def register_node_quantized_features( + self, + node_quantized_features: Union[torch.Tensor, dict[NodeType, torch.Tensor]], + ) -> None: + """Registers packed uint8 node features to the partitioner.""" + + self._assert_and_get_rpc_setup() + + if self._node_quantized_feat is not None: + raise ValueError( + "Node quantized features have already been registered. Cannot re-register node quantized feature data." + ) + + logger.info("Registering Node Quantized Features ...") + + input_node_quantized_features = ( + self._convert_node_entity_to_heterogeneous_format( + input_node_entity=node_quantized_features + ) + ) + + assert input_node_quantized_features, ( + "Node quantized features is an empty dictionary. Please provide node quantized features to register." + ) + + self._node_quantized_feat = convert_to_tensor( + input_node_quantized_features, dtype=torch.uint8 + ) + self._node_quantized_feat_dim = {} + for node_type in input_node_quantized_features: + self._node_quantized_feat_dim[node_type] = input_node_quantized_features[ + node_type + ].shape[1] + def register_node_labels( self, node_labels: Union[torch.Tensor, dict[NodeType, torch.Tensor]] ) -> None: @@ -928,7 +973,11 @@ def _partition_node_features_and_labels( self, node_partition_book: dict[NodeType, PartitionBook], node_type: NodeType, - ) -> tuple[Optional[FeaturePartitionData], Optional[FeaturePartitionData]]: + ) -> tuple[ + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + ]: """ Partitions node features and labels according to the node partition book. @@ -941,6 +990,7 @@ def _partition_node_features_and_labels( Returns: Optional[FeaturePartitionData]: Partitioned data of node features for current node type. + Optional[FeaturePartitionData]: Partitioned data of packed uint8 node features for current node type. Optional[FeaturePartitionData]: Partitioned data of node labels for current node type. """ @@ -980,6 +1030,18 @@ def _partition_node_features_and_labels( if self._node_feat_dim is not None and node_type in self._node_feat_dim else None ) + node_quantized_features = ( + self._node_quantized_feat[node_type] + if self._node_quantized_feat is not None + and node_type in self._node_quantized_feat + else None + ) + node_quantized_feat_dim = ( + self._node_quantized_feat_dim[node_type] + if self._node_quantized_feat_dim is not None + and node_type in self._node_quantized_feat_dim + else None + ) node_labels = ( self._node_labels[node_type] @@ -992,22 +1054,29 @@ def _partition_node_features_and_labels( else None ) - if node_features is not None and node_labels is not None: - input_data: Tuple[torch.Tensor, ...] = ( - node_features, - node_labels, - node_ids, - ) - elif node_features is not None: - input_data = (node_features, node_ids) - elif node_labels is not None: - input_data = (node_labels, node_ids) - else: + input_parts: list[torch.Tensor] = [] + node_feature_ind: Optional[int] = None + node_quantized_feature_ind: Optional[int] = None + node_label_ind: Optional[int] = None + + if node_features is not None: + node_feature_ind = len(input_parts) + input_parts.append(node_features) + if node_quantized_features is not None: + node_quantized_feature_ind = len(input_parts) + input_parts.append(node_quantized_features) + if node_labels is not None: + node_label_ind = len(input_parts) + input_parts.append(node_labels) + if not input_parts: raise ValueError( - f"Found no node features or node labels to partition for node type {node_type}" + f"Found no node features, quantized node features, or node labels to partition for node type {node_type}" ) + input_parts.append(node_ids) + input_data: Tuple[torch.Tensor, ...] = tuple(input_parts) has_node_features = node_features is not None + has_node_quantized_features = node_quantized_features is not None has_node_labels = node_labels is not None def _node_feature_partition_fn(node_feature_ids, _): @@ -1028,12 +1097,21 @@ def _node_feature_partition_fn(node_feature_ids, _): self._remove_key_from_member_dict("_max_node_ids", node_type) self._remove_key_from_member_dict("_node_feat", node_type) self._remove_key_from_member_dict("_node_feat_dim", node_type) + self._remove_key_from_member_dict("_node_quantized_feat", node_type) + self._remove_key_from_member_dict("_node_quantized_feat_dim", node_type) self._remove_key_from_member_dict("_node_labels", node_type) self._remove_key_from_member_dict("_node_labels_dim", node_type) # Since the unpartitioned node ids, features and labels are large, we would like to delete them when # they are no longer needed to free memory. - del node_ids, num_nodes, max_node_ids, node_features, node_labels + del ( + node_ids, + num_nodes, + max_node_ids, + node_features, + node_quantized_features, + node_labels, + ) gc.collect() @@ -1049,18 +1127,29 @@ def _node_feature_partition_fn(node_feature_ids, _): # Partitioned node features are stored at the 0th index ineach tuple of the partitioned results. if has_node_features: + assert node_feature_ind is not None node_feature_partition_data = FeaturePartitionData( - feats=torch.cat([r[0] for r in partitioned_results]), + feats=torch.cat([r[node_feature_ind] for r in partitioned_results]), ids=partitioned_ids, ) else: node_feature_partition_data = None - # Partitioned node labels are stored at the 1st index in each tuple of the partitioned results if we have - # both node features and node labels. Otherwise, it is stored at the 0th index + if has_node_quantized_features: + assert node_quantized_feature_ind is not None + node_quantized_feature_partition_data = FeaturePartitionData( + feats=torch.cat( + [r[node_quantized_feature_ind] for r in partitioned_results] + ), + ids=partitioned_ids, + ) + + else: + node_quantized_feature_partition_data = None + if has_node_labels: - node_label_ind = 1 if has_node_features else 0 + assert node_label_ind is not None node_label_partition_data = FeaturePartitionData( feats=torch.cat([r[node_label_ind] for r in partitioned_results]), ids=partitioned_ids, @@ -1077,6 +1166,14 @@ def _node_feature_partition_fn(node_feature_ids, _): else: node_feature_partition_data = None + if node_quantized_feat_dim is not None: + node_quantized_feature_partition_data = FeaturePartitionData( + feats=torch.empty((0, node_quantized_feat_dim), dtype=torch.uint8), + ids=torch.empty(0), + ) + else: + node_quantized_feature_partition_data = None + if node_labels_dim is not None: node_label_partition_data = FeaturePartitionData( feats=torch.empty((0, node_labels_dim)), @@ -1093,7 +1190,11 @@ def _node_feature_partition_fn(node_feature_ids, _): f"Node feature and label partitioning for node type {node_type} finished, took {time.time() - start_time:.3f}s" ) - return node_feature_partition_data, node_label_partition_data + return ( + node_feature_partition_data, + node_quantized_feature_partition_data, + node_label_partition_data, + ) def _partition_edge_index_and_edge_features( self, @@ -1459,6 +1560,7 @@ def partition_node_features_and_labels( ) -> tuple[ Optional[Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]]], Optional[Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]]], + Optional[Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]]], ]: """ Partitions node features and labels of a graph. If heterogeneous, partitions features and labels for all node type. @@ -1476,6 +1578,7 @@ def partition_node_features_and_labels( node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The Computed Node Partition Book Returns: Optional[Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]]]: Partitioned data of node features. + Optional[Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]]]: Partitioned data of packed uint8 node features. Optional[Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]]]: Partitioned data of node labels. """ assert self._num_nodes is not None and self._node_ids is not None, ( @@ -1509,22 +1612,32 @@ def partition_node_features_and_labels( ) for node_type in self._node_feat.keys(): node_feature_types.add(node_type) - elif self._node_labels is not None: + if self._node_quantized_feat is not None: + self._assert_data_type_consistency( + input_entity=self._node_quantized_feat, + is_node_entity=True, + is_subset=True, + ) + for node_type in self._node_quantized_feat.keys(): + node_feature_types.add(node_type) + if self._node_labels is not None: self._assert_data_type_consistency( input_entity=self._node_labels, is_node_entity=True, is_subset=True ) for node_type in self._node_labels.keys(): node_feature_types.add(node_type) - else: + if not node_feature_types: raise ValueError( - "Node features or labels must be registered prior to partitioning." + "Node features, quantized node features, or labels must be registered prior to partitioning." ) partitioned_node_features: dict[NodeType, FeaturePartitionData] = {} + partitioned_node_quantized_features: dict[NodeType, FeaturePartitionData] = {} partitioned_node_labels: dict[NodeType, FeaturePartitionData] = {} for node_type in sorted(node_feature_types): ( partitioned_node_features_for_node_type, + partitioned_node_quantized_features_for_node_type, partitioned_node_labels_for_node_type, ) = self._partition_node_features_and_labels( node_partition_book=transformed_node_partition_book, node_type=node_type @@ -1533,6 +1646,10 @@ def partition_node_features_and_labels( partitioned_node_features[node_type] = ( partitioned_node_features_for_node_type ) + if partitioned_node_quantized_features_for_node_type is not None: + partitioned_node_quantized_features[node_type] = ( + partitioned_node_quantized_features_for_node_type + ) if partitioned_node_labels_for_node_type is not None: partitioned_node_labels[node_type] = ( partitioned_node_labels_for_node_type @@ -1546,6 +1663,9 @@ def partition_node_features_and_labels( partitioned_node_features[DEFAULT_HOMOGENEOUS_NODE_TYPE] if partitioned_node_features else None, + partitioned_node_quantized_features[DEFAULT_HOMOGENEOUS_NODE_TYPE] + if partitioned_node_quantized_features + else None, partitioned_node_labels[DEFAULT_HOMOGENEOUS_NODE_TYPE] if partitioned_node_labels else None, @@ -1553,6 +1673,9 @@ def partition_node_features_and_labels( else: return ( partitioned_node_features if partitioned_node_features else None, + partitioned_node_quantized_features + if partitioned_node_quantized_features + else None, partitioned_node_labels if partitioned_node_labels else None, ) @@ -1771,15 +1894,21 @@ def partition( node_partition_book=node_partition_book ) - if self._node_feat is not None or self._node_labels is not None: + if ( + self._node_feat is not None + or self._node_quantized_feat is not None + or self._node_labels is not None + ): ( partitioned_node_features, + partitioned_node_quantized_features, partitioned_node_labels, ) = self.partition_node_features_and_labels( node_partition_book=node_partition_book ) else: partitioned_node_features = None + partitioned_node_quantized_features = None partitioned_node_labels = None if self._positive_label_edge_index is not None: @@ -1805,6 +1934,7 @@ def partition( edge_partition_book=edge_partition_book, partitioned_edge_index=partitioned_edge_index, partitioned_node_features=partitioned_node_features, + partitioned_node_quantized_features=partitioned_node_quantized_features, partitioned_edge_features=partitioned_edge_features, partitioned_positive_labels=partitioned_positive_edge_index, partitioned_negative_labels=partitioned_negative_edge_index, diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index faa1cfb57..b7b0754f7 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -128,7 +128,11 @@ def _partition_node_features_and_labels( self, node_partition_book: dict[NodeType, PartitionBook], node_type: NodeType, - ) -> tuple[Optional[FeaturePartitionData], Optional[FeaturePartitionData]]: + ) -> tuple[ + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + Optional[FeaturePartitionData], + ]: """ Partitions node features according to the node partition book. We rely on the functionality from the parent tensor-based partitioner here, and add logic to sort the node features by node indices which is specific to range-based partitioning. This is done so that the range-based @@ -143,6 +147,7 @@ def _partition_node_features_and_labels( """ ( feature_partition_data, + quantized_feature_partition_data, labels_partition_data, ) = super()._partition_node_features_and_labels( node_partition_book=node_partition_book, node_type=node_type @@ -167,6 +172,22 @@ def _partition_node_features_and_labels( else: partitioned_node_feature_data = None + if quantized_feature_partition_data is not None: + ids = quantized_feature_partition_data.ids + assert ids is not None + sorted_node_ids_indices = torch.argsort(ids) + partitioned_node_quantized_features = ( + quantized_feature_partition_data.feats[sorted_node_ids_indices] + ) + partitioned_node_quantized_feature_data = FeaturePartitionData( + feats=partitioned_node_quantized_features, ids=None + ) + + del sorted_node_ids_indices + gc.collect() + else: + partitioned_node_quantized_feature_data = None + if labels_partition_data is not None: ids = labels_partition_data.ids assert ids is not None @@ -183,7 +204,11 @@ def _partition_node_features_and_labels( else: partitioned_node_label_data = None - return partitioned_node_feature_data, partitioned_node_label_data + return ( + partitioned_node_feature_data, + partitioned_node_quantized_feature_data, + partitioned_node_label_data, + ) def _partition_edge_index_and_edge_features( self, diff --git a/gigl/distributed/distributed_neighborloader.py b/gigl/distributed/distributed_neighborloader.py index e3f9b478a..d6a8d9753 100644 --- a/gigl/distributed/distributed_neighborloader.py +++ b/gigl/distributed/distributed_neighborloader.py @@ -1,4 +1,5 @@ import sys +import time from collections import abc from itertools import count from typing import Optional, Tuple, Union @@ -29,6 +30,7 @@ SamplingClusterSetup, extract_metadata, labeled_to_homogeneous, + materialize_quantized_node_features, set_missing_features, shard_nodes_by_process, strip_label_edges, @@ -411,6 +413,7 @@ def _setup_for_graph_store( edge_types=edge_types, node_feature_info=node_feature_info, edge_feature_info=edge_feature_info, + node_quantization_metadata=dataset.node_quantization_metadata, edge_dir=dataset.fetch_edge_dir(), ), backend_key, @@ -528,6 +531,7 @@ def _setup_for_colocated( edge_types=edge_types, node_feature_info=dataset.node_feature_info, edge_feature_info=dataset.edge_feature_info, + node_quantization_metadata=dataset.node_quantization_metadata, edge_dir=dataset.edge_dir, ), ) @@ -542,8 +546,14 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: # as edge types and fails when edge_dir="out" (tries to call # reverse_edge_type on them). We strip them here and re-apply after. # TODO (mkolodner-sc): Remove once GLT's to_hetero_data is fixed. + collate_start_time = time.perf_counter() metadata, stripped_msg = extract_metadata(msg, self.to_device) + base_collate_start_time = time.perf_counter() data = super()._collate_fn(stripped_msg) + base_collate_time = time.perf_counter() - base_collate_start_time + logger.debug( + f"Distributed NeighborLoader GLT base collate time: {base_collate_time:.3f}s" + ) data = set_missing_features( data=data, node_feature_info=self._node_feature_info, @@ -556,9 +566,19 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: data = labeled_to_homogeneous(DEFAULT_HOMOGENEOUS_EDGE_TYPE, data) data, metadata = self._apply_ppr_outputs(data, metadata) + data, metadata = materialize_quantized_node_features( + data=data, + metadata=metadata, + node_quantization_metadata=self._node_quantization_metadata, + ) # Attach any remaining metadata (e.g. custom user-defined keys) directly onto the # data object so downstream code can access them via attribute lookup. for key, value in metadata.items(): data[key] = value + + collate_time = time.perf_counter() - collate_start_time + logger.debug( + f"Distributed NeighborLoader end-to-end collate time: {collate_time:.3f}s" + ) return data diff --git a/gigl/distributed/graph_store/dist_server.py b/gigl/distributed/graph_store/dist_server.py index eb10dad0c..0a92d959d 100644 --- a/gigl/distributed/graph_store/dist_server.py +++ b/gigl/distributed/graph_store/dist_server.py @@ -94,7 +94,12 @@ def compute_process(): ) from gigl.distributed.sampler_options import PPRSamplerOptions from gigl.src.common.types.graph_data import EdgeType, NodeType -from gigl.types.graph import FeatureInfo, reverse_edge_type, select_label_edge_types +from gigl.types.graph import ( + FeatureInfo, + FeatureQuantizationMetadata, + reverse_edge_type, + select_label_edge_types, +) from gigl.utils.data_splitters import get_labels_for_anchor_nodes from gigl.utils.share_memory import share_memory @@ -405,6 +410,14 @@ def get_node_feature_info( """ return self.dataset.node_feature_info + def get_node_quantization_metadata( + self, + ) -> Union[ + FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata], None + ]: + """Get node feature quantization metadata from the dataset.""" + return self.dataset.node_quantization_metadata + def get_edge_feature_info( self, ) -> Union[FeatureInfo, dict[EdgeType, FeatureInfo], None]: diff --git a/gigl/distributed/graph_store/remote_dist_dataset.py b/gigl/distributed/graph_store/remote_dist_dataset.py index e58784dc0..ef3fc24b0 100644 --- a/gigl/distributed/graph_store/remote_dist_dataset.py +++ b/gigl/distributed/graph_store/remote_dist_dataset.py @@ -21,6 +21,7 @@ DEFAULT_HOMOGENEOUS_EDGE_TYPE, DEFAULT_HOMOGENEOUS_NODE_TYPE, FeatureInfo, + FeatureQuantizationMetadata, ) from gigl.utils.sampling import ABLPInputNodes @@ -66,6 +67,29 @@ def fetch_node_feature_info( DistServer.get_node_feature_info, ) + def fetch_node_quantization_metadata( + self, + ) -> Union[ + FeatureQuantizationMetadata, + dict[NodeType, FeatureQuantizationMetadata], + None, + ]: + """Fetch node feature quantization metadata from the registered dataset.""" + return request_server( + 0, + DistServer.get_node_quantization_metadata, + ) + + @property + def node_quantization_metadata( + self, + ) -> Union[ + FeatureQuantizationMetadata, + dict[NodeType, FeatureQuantizationMetadata], + None, + ]: + return self.fetch_node_quantization_metadata() + def fetch_edge_feature_info( self, ) -> Union[FeatureInfo, dict[EdgeType, FeatureInfo], None]: diff --git a/gigl/distributed/sampler.py b/gigl/distributed/sampler.py index e99dd65dc..7789c6731 100644 --- a/gigl/distributed/sampler.py +++ b/gigl/distributed/sampler.py @@ -8,6 +8,7 @@ POSITIVE_LABEL_METADATA_KEY: Final[str] = "gigl_positive_labels." NEGATIVE_LABEL_METADATA_KEY: Final[str] = "gigl_negative_labels." +NODE_PACKED_FEATURES_METADATA_KEY: Final[str] = "node_packed_features" class ABLPNodeSamplerInput(NodeSamplerInput): diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 876582208..b40eb471a 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -1,11 +1,12 @@ """Utils for Neighbor loaders.""" import ast +import time from collections import abc from copy import deepcopy from dataclasses import dataclass from enum import Enum -from typing import Literal, Optional, TypeVar, Union +from typing import Literal, Optional, TypeVar, Union, cast import torch from graphlearn_torch.channel import SampleMessage @@ -13,7 +14,14 @@ from torch_geometric.typing import EdgeType, NodeType from gigl.common.logger import Logger -from gigl.types.graph import FeatureInfo, is_label_edge_type +from gigl.common.utils.feature_quantization.torch_ops import dequantize_torch_tensor +from gigl.distributed.sampler import NODE_PACKED_FEATURES_METADATA_KEY +from gigl.types.graph import ( + DEFAULT_HOMOGENEOUS_NODE_TYPE, + FeatureInfo, + FeatureQuantizationMetadata, + is_label_edge_type, +) logger = Logger() @@ -46,6 +54,10 @@ class DatasetSchema: edge_feature_info: Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]] # Edge direction. edge_dir: Union[str, Literal["in", "out"]] + # Quantization metadata for packed node features. + node_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] + ] = None def patch_fanout_for_sampling( @@ -324,6 +336,80 @@ def set_missing_features( return data +def materialize_quantized_node_features( + data: _GraphType, + metadata: dict[str, torch.Tensor], + node_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] + ], +) -> tuple[_GraphType, dict[str, torch.Tensor]]: + """Materialize packed quantized node features into PyG node feature tensors.""" + if node_quantization_metadata is None: + return data, metadata + materialize_start_time = time.perf_counter() + + def materialize( + store, packed_features: torch.Tensor, q: FeatureQuantizationMetadata + ) -> None: + dequantized = dequantize_torch_tensor(packed_features, metadata=q) + x = getattr(store, "x", None) + out = dequantized.new_empty((dequantized.size(0), q.feature_dim)) + scatter_indices = q.scatter_index_tensors(out.device) + out[:, scatter_indices.quantized] = dequantized + + if x is None and q.raw_feature_dim: + raise ValueError(f"Missing {q.raw_feature_dim} unquantized features") + if x is not None: + if x.size(1) != q.raw_feature_dim: + raise ValueError( + f"Expected {q.raw_feature_dim} raw node feature columns before " + f"dequantization, got {x.size(1)}." + ) + out[:, scatter_indices.raw] = x + store.x = out + + if isinstance(data, Data): + if isinstance(node_quantization_metadata, dict): + homogeneous_quantization_metadata = cast( + dict[NodeType, FeatureQuantizationMetadata], + node_quantization_metadata, + ) + quantization_metadata = homogeneous_quantization_metadata[ + DEFAULT_HOMOGENEOUS_NODE_TYPE + ] + metadata_key = ( + f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}" + ) + else: + quantization_metadata = node_quantization_metadata + metadata_key = NODE_PACKED_FEATURES_METADATA_KEY + packed_features = metadata.pop(metadata_key, None) + if packed_features is None: + raise ValueError( + f"Missing packed quantized node features in metadata key {metadata_key}." + ) + materialize(data, packed_features, quantization_metadata) + else: + assert isinstance(node_quantization_metadata, dict), ( + "Expected per-node-type quantization metadata for heterogeneous data." + ) + node_quantization_metadata = cast( + dict[NodeType, FeatureQuantizationMetadata], node_quantization_metadata + ) + for node_type, quantization_metadata in node_quantization_metadata.items(): + metadata_key = f"{NODE_PACKED_FEATURES_METADATA_KEY}.{node_type}" + packed_features = metadata.pop(metadata_key, None) + if packed_features is None: + continue + materialize(data[node_type], packed_features, quantization_metadata) + + materialize_time = time.perf_counter() - materialize_start_time + logger.debug( + f"Quantized node feature materialization time: {materialize_time:.3f}s" + ) + return data, metadata + + def extract_metadata( msg: SampleMessage, device: torch.device ) -> tuple[dict[str, torch.Tensor], SampleMessage]: diff --git a/gigl/types/graph.py b/gigl/types/graph.py index fec65b9d8..776a1a658 100644 --- a/gigl/types/graph.py +++ b/gigl/types/graph.py @@ -99,6 +99,12 @@ class PartitionOutput: Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]] ] + # Quantized node features on current rank. These are packed uint8 features + # aligned by node id and dequantized/scattered in the sampler collate path. + partitioned_node_quantized_features: Optional[ + Union[FeaturePartitionData, dict[NodeType, FeaturePartitionData]] + ] = None + @dataclass(frozen=True) class FeatureInfo: From 0400961f76256f51ce60723fc3ca3e986db80d55 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 17:43:05 +0000 Subject: [PATCH 11/78] Add docs --- .../utils/feature_quantization/README.md | 75 +++++++++++++++++++ .../utils/feature_quantization/__init__.py | 11 +++ 2 files changed, 86 insertions(+) create mode 100644 gigl/common/utils/feature_quantization/README.md diff --git a/gigl/common/utils/feature_quantization/README.md b/gigl/common/utils/feature_quantization/README.md new file mode 100644 index 000000000..4fe6bf508 --- /dev/null +++ b/gigl/common/utils/feature_quantization/README.md @@ -0,0 +1,75 @@ +# Feature Quantization + +This package contains the low-level NumPy and Torch helpers for node feature +quantization in GiGL. + +Feature quantization is lossy compression: high-precision feature values such as +fp32 are mapped into a lower-precision representation. The current built-in +scheme stores low-bit codes as packed `uint8` bytes, then reconstructs +approximate feature values when a sampled subgraph is materialized for training +or inference. + +The motivation is practical scaling. Large-scale GNN training is often +memory-bound: + +- feature hydration uses irregular memory access, often over the network; +- hydrated features still need to move to the accelerator; +- large feature stores limit the workloads that fit on a given machine. + +Reducing feature size can improve feature-store footprint, network bandwidth, +and PCIe transfer volume. GNNs have also been shown to be relatively tolerant of +input feature quantization in [Degree-Quant: Quantization-Aware Training for +Graph Neural Networks](https://arxiv.org/abs/2207.14696), which motivates this +as a useful tradeoff for GiGL. + +## Current Built-In Flow + +The built-in flow is: + +1. The data preprocessor computes feature summary statistics offline. +2. The preprocessor quantizes selected scalar node feature columns with NumPy. +3. The packed `uint8` feature sidecar is written to TFRecords. +4. Distributed dataset construction partitions and samples the packed bytes. +5. The dataloader collate path dequantizes sampled packed features with Torch. +6. Dequantized columns are scattered back into the logical `x` feature matrix. + +The NumPy/Torch split is intentional: + +- `numpy_ops.py` runs in preprocessing, where data is on CPU and Torch may not + be available. +- `torch_ops.py` runs during dataloader collation, where sampled feature data is + already represented as Torch tensors and may already be on GPU. + +`FeatureQuantizationMetadata` is the contract between those two steps. It records +the bit width, packed feature dimension, logical feature positions, and the +statistics needed to invert the compression step. + +## Current Built-In Scheme + +The current implementation supports `1`, `2`, `4`, and `8` bit quantization. +Codes are packed high-bits-first into bytes. + +For `1` bit, values are represented by sign and reconstructed from the positive +and non-positive means. + +For `2`, `4`, and `8` bits, values are clipped to pre-computed bounds and mapped into +uniform integer buckets between those bounds. + +## TODO: Pluggable Schemes + +The current metadata/proto shape is tied to the built-in quantization scheme. A +useful follow-up is to make the quantization scheme itself pluggable. + +One possible design is: + +- define a quantizer object tied to `FeatureQuantizationSpec`; +- serialize a stable fully qualified name or registry key for the quantizer; +- serialize the quantizer arguments in the proto; +- rebuild the quantizer from that metadata on the read side; +- require each quantizer to provide matching NumPy `quantize` and Torch + `dequantize` implementations. + +That would let developers add new schemes without threading one-off fields +through every proto and loader path. A registry key is likely safer than +arbitrary imports, but either way the important contract is that the serialized +scheme identifies both the compression step and its inverse. diff --git a/gigl/common/utils/feature_quantization/__init__.py b/gigl/common/utils/feature_quantization/__init__.py index e69de29bb..3a1d91791 100644 --- a/gigl/common/utils/feature_quantization/__init__.py +++ b/gigl/common/utils/feature_quantization/__init__.py @@ -0,0 +1,11 @@ +"""Utilities for lossy node feature quantization in GiGL. + +Feature quantization compresses high-precision node feature columns into a +lower-precision representation. GiGL uses this to reduce feature-store size on +disk and in memory, and to reduce feature bandwidth across storage, sampling, +and device-transfer paths. + +The current built-in workflow computes summary statistics offline in the +preprocessor, stores packed feature bytes, partitions and samples those bytes, +then dequantizes sampled subgraph features during dataloader collation. +""" From ec88d0f3a6895f7fd836e5ed7b0a478b8d272830 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 17:47:17 +0000 Subject: [PATCH 12/78] Update docs --- gigl/common/utils/feature_quantization/README.md | 6 +++--- gigl/common/utils/feature_quantization/__init__.py | 12 +----------- 2 files changed, 4 insertions(+), 14 deletions(-) diff --git a/gigl/common/utils/feature_quantization/README.md b/gigl/common/utils/feature_quantization/README.md index 4fe6bf508..02b413d2e 100644 --- a/gigl/common/utils/feature_quantization/README.md +++ b/gigl/common/utils/feature_quantization/README.md @@ -18,9 +18,9 @@ memory-bound: Reducing feature size can improve feature-store footprint, network bandwidth, and PCIe transfer volume. GNNs have also been shown to be relatively tolerant of -input feature quantization in [Degree-Quant: Quantization-Aware Training for -Graph Neural Networks](https://arxiv.org/abs/2207.14696), which motivates this -as a useful tradeoff for GiGL. +input feature quantization in [BiFeat: Supercharge GNN Training via Graph +Feature Quantization](https://arxiv.org/abs/2207.14696), which motivates this as +a useful tradeoff for GiGL. ## Current Built-In Flow diff --git a/gigl/common/utils/feature_quantization/__init__.py b/gigl/common/utils/feature_quantization/__init__.py index 3a1d91791..6f1229294 100644 --- a/gigl/common/utils/feature_quantization/__init__.py +++ b/gigl/common/utils/feature_quantization/__init__.py @@ -1,11 +1 @@ -"""Utilities for lossy node feature quantization in GiGL. - -Feature quantization compresses high-precision node feature columns into a -lower-precision representation. GiGL uses this to reduce feature-store size on -disk and in memory, and to reduce feature bandwidth across storage, sampling, -and device-transfer paths. - -The current built-in workflow computes summary statistics offline in the -preprocessor, stores packed feature bytes, partitions and samples those bytes, -then dequantizes sampled subgraph features during dataloader collation. -""" +"""Utilities for node feature quantization in GiGL.""" From 3e76a5220a9f5e95cd84a8279ad9420dda4ad4b6 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 17:56:05 +0000 Subject: [PATCH 13/78] Add inline comment explaining uint16 cast --- gigl/common/utils/feature_quantization/numpy_ops.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index f7eb5db75..9b4e1ebe7 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -40,6 +40,9 @@ def _pack_codes(codes: np.ndarray, bits: int) -> np.ndarray: # Pad only the feature dimension of this 2D [row, feature] array. codes = np.pad(codes, ((0, 0), (0, pad)), constant_values=0) # Group the padded feature dimension into chunks that each form one byte. + # Valid bit widths pack exactly one byte per group, so the final sum is at + # most 255. uint16 is a conservative arithmetic dtype that avoids relying on + # NumPy's uint8 accumulator behavior before the final uint8 cast. codes = codes.reshape(codes.shape[0], -1, per_byte).astype(np.uint16) shifts = bits * np.arange(per_byte - 1, -1, -1, dtype=np.uint16) weights = (1 << shifts).astype(np.uint16) From c0b4847125d9626e8acaf14dc3301ad7faf47907 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 17:58:07 +0000 Subject: [PATCH 14/78] Add quantize dequantize roundtrip test --- .../feature_quantization/torch_ops_test.py | 82 +++++++++++++++++++ 1 file changed, 82 insertions(+) diff --git a/tests/unit/common/utils/feature_quantization/torch_ops_test.py b/tests/unit/common/utils/feature_quantization/torch_ops_test.py index 6f941c548..79ca95db2 100644 --- a/tests/unit/common/utils/feature_quantization/torch_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/torch_ops_test.py @@ -1,11 +1,93 @@ +import numpy as np import torch +from gigl.common.utils.feature_quantization.numpy_ops import quantize_ndarray from gigl.common.utils.feature_quantization.torch_ops import dequantize_torch_tensor from gigl.types.graph import FeatureQuantizationMetadata from tests.test_assets.test_case import TestCase class TorchFeatureQuantizationOpsTest(TestCase): + def test_quantize_numpy_dequantize_torch_round_trip_single_bit_without_padding( + self, + ) -> None: + features = np.array( + [[-2.0, -0.5, 0.5, 3.0, -3.0, 2.0, -1.0, 1.0]], dtype=np.float32 + ) + metadata = FeatureQuantizationMetadata( + bits=1, + feature_dim=features.shape[1], + quantized_feature_indices=tuple(range(features.shape[1])), + neg_mean=-1.25, + pos_mean=1.75, + ) + + packed = quantize_ndarray(features, bits=metadata.bits, stats={}) + actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) + + self.assertEqual(actual.shape, torch.Size(features.shape)) + self.assertEqual(actual.dtype, torch.float32) + self.assertEqual(set(actual.flatten().tolist()), {-1.25, 1.75}) + + def test_quantize_numpy_dequantize_torch_round_trip_single_bit_with_padding( + self, + ) -> None: + features = np.array([[-2.0, -0.5, 0.5, 3.0, -3.0]], dtype=np.float32) + metadata = FeatureQuantizationMetadata( + bits=1, + feature_dim=features.shape[1], + quantized_feature_indices=tuple(range(features.shape[1])), + neg_mean=-1.25, + pos_mean=1.75, + ) + + packed = quantize_ndarray(features, bits=metadata.bits, stats={}) + actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) + + self.assertEqual(actual.shape, torch.Size(features.shape)) + self.assertEqual(actual.dtype, torch.float32) + self.assertEqual(set(actual.flatten().tolist()), {-1.25, 1.75}) + + def test_quantize_numpy_dequantize_torch_round_trip_multi_bit_without_padding( + self, + ) -> None: + features = np.array([[-1.0, 0.0, 0.5, 1.0]], dtype=np.float32) + stats = {"clip_min": 0.0, "clip_max": 1.0} + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=features.shape[1], + quantized_feature_indices=tuple(range(features.shape[1])), + **stats, + ) + + packed = quantize_ndarray(features, bits=metadata.bits, stats=stats) + actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) + + self.assertEqual(actual.shape, torch.Size(features.shape)) + self.assertEqual(actual.dtype, torch.float32) + self.assertTrue(torch.all(actual >= stats["clip_min"]).item()) + self.assertTrue(torch.all(actual <= stats["clip_max"]).item()) + + def test_quantize_numpy_dequantize_torch_round_trip_multi_bit_with_padding( + self, + ) -> None: + features = np.array([[-1.0, 0.0, 0.5, 1.0, 2.0]], dtype=np.float32) + stats = {"clip_min": 0.0, "clip_max": 1.0} + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=features.shape[1], + quantized_feature_indices=tuple(range(features.shape[1])), + **stats, + ) + + packed = quantize_ndarray(features, bits=metadata.bits, stats=stats) + actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) + + self.assertEqual(actual.shape, torch.Size(features.shape)) + self.assertEqual(actual.dtype, torch.float32) + self.assertTrue(torch.all(actual >= stats["clip_min"]).item()) + self.assertTrue(torch.all(actual <= stats["clip_max"]).item()) + def test_dequantize_torch_tensor_single_bit_unpacks_full_byte(self) -> None: # 0b10101010 = 170 unpacks high-bits-first to [1, 0, 1, 0, 1, 0, 1, 0]. # Code 1 maps to pos_mean and code 0 maps to neg_mean. From 7c70b021b0927a9deef806552bfb17e8c6483a36 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 18:00:30 +0000 Subject: [PATCH 15/78] Simplify test logic --- .../feature_quantization/torch_ops_test.py | 62 +++---------------- 1 file changed, 7 insertions(+), 55 deletions(-) diff --git a/tests/unit/common/utils/feature_quantization/torch_ops_test.py b/tests/unit/common/utils/feature_quantization/torch_ops_test.py index 79ca95db2..fa370a994 100644 --- a/tests/unit/common/utils/feature_quantization/torch_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/torch_ops_test.py @@ -8,30 +8,7 @@ class TorchFeatureQuantizationOpsTest(TestCase): - def test_quantize_numpy_dequantize_torch_round_trip_single_bit_without_padding( - self, - ) -> None: - features = np.array( - [[-2.0, -0.5, 0.5, 3.0, -3.0, 2.0, -1.0, 1.0]], dtype=np.float32 - ) - metadata = FeatureQuantizationMetadata( - bits=1, - feature_dim=features.shape[1], - quantized_feature_indices=tuple(range(features.shape[1])), - neg_mean=-1.25, - pos_mean=1.75, - ) - - packed = quantize_ndarray(features, bits=metadata.bits, stats={}) - actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) - - self.assertEqual(actual.shape, torch.Size(features.shape)) - self.assertEqual(actual.dtype, torch.float32) - self.assertEqual(set(actual.flatten().tolist()), {-1.25, 1.75}) - - def test_quantize_numpy_dequantize_torch_round_trip_single_bit_with_padding( - self, - ) -> None: + def test_quantize_numpy_dequantize_torch_round_trip_single_bit(self) -> None: features = np.array([[-2.0, -0.5, 0.5, 3.0, -3.0]], dtype=np.float32) metadata = FeatureQuantizationMetadata( bits=1, @@ -44,35 +21,13 @@ def test_quantize_numpy_dequantize_torch_round_trip_single_bit_with_padding( packed = quantize_ndarray(features, bits=metadata.bits, stats={}) actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) - self.assertEqual(actual.shape, torch.Size(features.shape)) - self.assertEqual(actual.dtype, torch.float32) - self.assertEqual(set(actual.flatten().tolist()), {-1.25, 1.75}) - - def test_quantize_numpy_dequantize_torch_round_trip_multi_bit_without_padding( - self, - ) -> None: - features = np.array([[-1.0, 0.0, 0.5, 1.0]], dtype=np.float32) - stats = {"clip_min": 0.0, "clip_max": 1.0} - metadata = FeatureQuantizationMetadata( - bits=2, - feature_dim=features.shape[1], - quantized_feature_indices=tuple(range(features.shape[1])), - **stats, + torch.testing.assert_close( + actual, torch.tensor([[-1.25, -1.25, 1.75, 1.75, -1.25]]) ) - packed = quantize_ndarray(features, bits=metadata.bits, stats=stats) - actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) - - self.assertEqual(actual.shape, torch.Size(features.shape)) - self.assertEqual(actual.dtype, torch.float32) - self.assertTrue(torch.all(actual >= stats["clip_min"]).item()) - self.assertTrue(torch.all(actual <= stats["clip_max"]).item()) - - def test_quantize_numpy_dequantize_torch_round_trip_multi_bit_with_padding( - self, - ) -> None: - features = np.array([[-1.0, 0.0, 0.5, 1.0, 2.0]], dtype=np.float32) - stats = {"clip_min": 0.0, "clip_max": 1.0} + def test_quantize_numpy_dequantize_torch_round_trip_multi_bit(self) -> None: + features = np.array([[-1.0, 0.0, 1.0, 2.0, 4.0]], dtype=np.float32) + stats = {"clip_min": 0.0, "clip_max": 3.0} metadata = FeatureQuantizationMetadata( bits=2, feature_dim=features.shape[1], @@ -83,10 +38,7 @@ def test_quantize_numpy_dequantize_torch_round_trip_multi_bit_with_padding( packed = quantize_ndarray(features, bits=metadata.bits, stats=stats) actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) - self.assertEqual(actual.shape, torch.Size(features.shape)) - self.assertEqual(actual.dtype, torch.float32) - self.assertTrue(torch.all(actual >= stats["clip_min"]).item()) - self.assertTrue(torch.all(actual <= stats["clip_max"]).item()) + torch.testing.assert_close(actual, torch.tensor([[0.0, 0.0, 1.0, 2.0, 3.0]])) def test_dequantize_torch_tensor_single_bit_unpacks_full_byte(self) -> None: # 0b10101010 = 170 unpacks high-bits-first to [1, 0, 1, 0, 1, 0, 1, 0]. From 1510b4361020100745e290c47839bdf064106270 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 18:02:17 +0000 Subject: [PATCH 16/78] Defensive check against nan or inf --- .../utils/feature_quantization/numpy_ops.py | 2 ++ .../utils/feature_quantization/numpy_ops_test.py | 16 ++++++++++++++++ 2 files changed, 18 insertions(+) diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index 9b4e1ebe7..e2143ccee 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -19,6 +19,8 @@ def quantize_ndarray( raise ValueError(f"bits must be one of {valid_bits}, got {bits}") if features.ndim != 2: raise ValueError(f"Expected a 2D feature array, got shape {features.shape}.") + if not np.isfinite(features).all(): + raise ValueError("features must be finite; got NaN or Inf") if bits == 1: # 1-bit quantization keeps only sign; values restore from neg/pos means. codes = (features > 0).astype(np.uint8) diff --git a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py index e481bc305..0a4325d2a 100644 --- a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py @@ -103,3 +103,19 @@ def test_quantize_ndarray_rejects_non_2d_features(self) -> None: quantize_ndarray( np.zeros((1, 1, 1)), bits=2, stats={"clip_min": 0.0, "clip_max": 1.0} ) + + def test_quantize_ndarray_rejects_nan_features(self) -> None: + with self.assertRaisesRegex(ValueError, "features must be finite"): + quantize_ndarray( + np.array([[np.nan, 1.0]]), + bits=4, + stats={"clip_min": 0.0, "clip_max": 1.0}, + ) + + def test_quantize_ndarray_rejects_inf_features(self) -> None: + with self.assertRaisesRegex(ValueError, "features must be finite"): + quantize_ndarray( + np.array([[np.inf, 1.0]]), + bits=4, + stats={"clip_min": 0.0, "clip_max": 1.0}, + ) From 5116725a6951de8fc6a90a546858e31a7f254671 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 18:08:39 +0000 Subject: [PATCH 17/78] Pass clip args directly to quantize --- .../utils/feature_quantization/numpy_ops.py | 22 +++++--- .../feature_quantization/numpy_ops_test.py | 54 +++++++------------ .../feature_quantization/torch_ops_test.py | 12 +++-- 3 files changed, 43 insertions(+), 45 deletions(-) diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index e2143ccee..ce15cc008 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -5,15 +5,22 @@ the dataloader collate path operates on torch tensors that may already be on GPU. """ -from collections.abc import Mapping - import numpy as np def quantize_ndarray( - features: np.ndarray, *, bits: int, stats: Mapping[str, float] + features: np.ndarray, + *, + bits: int, + clip_min: float | None = None, + clip_max: float | None = None, ) -> np.ndarray: - """Quantize a 2D float array into packed uint8 codes.""" + """Quantize a 2D float array into packed uint8 codes. + + For multi-bit quantization, `clip_min` and `clip_max` are required and + define the min-max scaling range: values are clipped to that range, scaled + to `[0, 2**bits - 1]`, rounded to integer codes, then packed into bytes. + """ valid_bits = (1, 2, 4, 8) if bits not in valid_bits: raise ValueError(f"bits must be one of {valid_bits}, got {bits}") @@ -26,10 +33,11 @@ def quantize_ndarray( codes = (features > 0).astype(np.uint8) else: # Min-max scale using clipped values and map to integer buckets. + if clip_min is None or clip_max is None: + raise ValueError(f"{bits}-bit quantization requires clip_min/clip_max") levels = (1 << bits) - 1 - lo, hi = stats["clip_min"], stats["clip_max"] - clipped = np.clip(features, lo, hi) - scaled = (clipped - lo) / (hi - lo) + clipped = np.clip(features, clip_min, clip_max) + scaled = (clipped - clip_min) / (clip_max - clip_min) codes = np.rint(scaled * levels).astype(np.uint8) return _pack_codes(codes, bits) diff --git a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py index 0a4325d2a..a6e6f4136 100644 --- a/tests/unit/common/utils/feature_quantization/numpy_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/numpy_ops_test.py @@ -10,7 +10,7 @@ def test_quantize_ndarray_single_bit_packs_full_byte(self) -> None: # [-1, 0, 0.5, 2, -0.5, 3, 4, -4] -> [0, 0, 1, 1, 0, 1, 1, 0]. # High-bits-first packing gives 0b00110110 = 54. features = np.array([[-1.0, 0.0, 0.5, 2.0, -0.5, 3.0, 4.0, -4.0]]) - actual = quantize_ndarray(features, bits=1, stats={}) + actual = quantize_ndarray(features, bits=1) np.testing.assert_array_equal(actual, np.array([[54]], dtype=np.uint8)) def test_quantize_ndarray_single_bit_pads_final_byte(self) -> None: @@ -18,98 +18,83 @@ def test_quantize_ndarray_single_bit_pads_final_byte(self) -> None: # [1, -1, 2, -2, 3] becomes [1, 0, 1, 0, 1]. # Padding fills the remaining bit slots with zeros: 0b10101000 = 168. features = np.array([[1.0, -1.0, 2.0, -2.0, 3.0]]) - actual = quantize_ndarray(features, bits=1, stats={}) + actual = quantize_ndarray(features, bits=1) np.testing.assert_array_equal(actual, np.array([[168]], dtype=np.uint8)) def test_quantize_ndarray_two_bit_packs_full_byte_ascending_codes(self) -> None: # With clip range [0, 3], these values equal their 2-bit codes. # [0, 1, 2, 3] packs as 00 01 10 11 = 0b00011011 = 27. features = np.array([[0.0, 1.0, 2.0, 3.0]]) - actual = quantize_ndarray( - features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} - ) + actual = quantize_ndarray(features, bits=2, clip_min=0.0, clip_max=3.0) np.testing.assert_array_equal(actual, np.array([[27]], dtype=np.uint8)) def test_quantize_ndarray_two_bit_packs_full_byte_descending_codes(self) -> None: # With clip range [0, 3], these values equal their 2-bit codes. # [3, 2, 1, 0] packs as 11 10 01 00 = 0b11100100 = 228. features = np.array([[3.0, 2.0, 1.0, 0.0]]) - actual = quantize_ndarray( - features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} - ) + actual = quantize_ndarray(features, bits=2, clip_min=0.0, clip_max=3.0) np.testing.assert_array_equal(actual, np.array([[228]], dtype=np.uint8)) def test_quantize_ndarray_two_bit_pads_final_byte(self) -> None: # The first four codes [0, 1, 2, 3] pack into byte 27. # The leftover code [1] starts the next byte as 01 00 00 00 = 64. features = np.array([[0.0, 1.0, 2.0, 3.0, 1.0]]) - actual = quantize_ndarray( - features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} - ) + actual = quantize_ndarray(features, bits=2, clip_min=0.0, clip_max=3.0) np.testing.assert_array_equal(actual, np.array([[27, 64]], dtype=np.uint8)) def test_quantize_ndarray_four_bit_packs_full_byte_ascending_codes(self) -> None: # With clip range [0, 15], these values equal their 4-bit codes. # [0, 15] packs as 0000 1111 = 15. features = np.array([[0.0, 15.0]]) - actual = quantize_ndarray( - features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} - ) + actual = quantize_ndarray(features, bits=4, clip_min=0.0, clip_max=15.0) np.testing.assert_array_equal(actual, np.array([[15]], dtype=np.uint8)) def test_quantize_ndarray_four_bit_packs_full_byte_descending_codes(self) -> None: # With clip range [0, 15], these values equal their 4-bit codes. # [15, 0] packs as 1111 0000 = 240. features = np.array([[15.0, 0.0]]) - actual = quantize_ndarray( - features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} - ) + actual = quantize_ndarray(features, bits=4, clip_min=0.0, clip_max=15.0) np.testing.assert_array_equal(actual, np.array([[240]], dtype=np.uint8)) def test_quantize_ndarray_four_bit_pads_final_byte(self) -> None: # The first two codes [0, 15] pack into byte 15. # The leftover code [8] starts the next byte as 1000 0000 = 128. features = np.array([[0.0, 15.0, 8.0]]) - actual = quantize_ndarray( - features, bits=4, stats={"clip_min": 0.0, "clip_max": 15.0} - ) + actual = quantize_ndarray(features, bits=4, clip_min=0.0, clip_max=15.0) np.testing.assert_array_equal(actual, np.array([[15, 128]], dtype=np.uint8)) def test_quantize_ndarray_eight_bit_stores_one_code_per_column(self) -> None: # 8-bit quantization has one code per uint8 column, so no bit packing changes the order. features = np.array([[0.0, 128.0, 255.0]]) - actual = quantize_ndarray( - features, bits=8, stats={"clip_min": 0.0, "clip_max": 255.0} - ) + actual = quantize_ndarray(features, bits=8, clip_min=0.0, clip_max=255.0) np.testing.assert_array_equal(actual, np.array([[0, 128, 255]], dtype=np.uint8)) def test_quantize_ndarray_clips_multi_bit_values_before_packing(self) -> None: # Values are clipped before scaling: [-1, 0, 3, 4] over [0, 3] # becomes codes [0, 0, 3, 3], which packs as 00 00 11 11 = 15. features = np.array([[-1.0, 0.0, 3.0, 4.0]]) - actual = quantize_ndarray( - features, bits=2, stats={"clip_min": 0.0, "clip_max": 3.0} - ) + actual = quantize_ndarray(features, bits=2, clip_min=0.0, clip_max=3.0) np.testing.assert_array_equal(actual, np.array([[15]], dtype=np.uint8)) def test_quantize_ndarray_rejects_invalid_bit_width(self) -> None: with self.assertRaises(ValueError): - quantize_ndarray( - np.zeros((1, 1)), bits=3, stats={"clip_min": 0.0, "clip_max": 1.0} - ) + quantize_ndarray(np.zeros((1, 1)), bits=3, clip_min=0.0, clip_max=1.0) def test_quantize_ndarray_rejects_non_2d_features(self) -> None: with self.assertRaises(ValueError): - quantize_ndarray( - np.zeros((1, 1, 1)), bits=2, stats={"clip_min": 0.0, "clip_max": 1.0} - ) + quantize_ndarray(np.zeros((1, 1, 1)), bits=2, clip_min=0.0, clip_max=1.0) + + def test_quantize_ndarray_requires_multi_bit_clip_bounds(self) -> None: + with self.assertRaisesRegex(ValueError, "requires clip_min/clip_max"): + quantize_ndarray(np.zeros((1, 1)), bits=2) def test_quantize_ndarray_rejects_nan_features(self) -> None: with self.assertRaisesRegex(ValueError, "features must be finite"): quantize_ndarray( np.array([[np.nan, 1.0]]), bits=4, - stats={"clip_min": 0.0, "clip_max": 1.0}, + clip_min=0.0, + clip_max=1.0, ) def test_quantize_ndarray_rejects_inf_features(self) -> None: @@ -117,5 +102,6 @@ def test_quantize_ndarray_rejects_inf_features(self) -> None: quantize_ndarray( np.array([[np.inf, 1.0]]), bits=4, - stats={"clip_min": 0.0, "clip_max": 1.0}, + clip_min=0.0, + clip_max=1.0, ) diff --git a/tests/unit/common/utils/feature_quantization/torch_ops_test.py b/tests/unit/common/utils/feature_quantization/torch_ops_test.py index fa370a994..a720891df 100644 --- a/tests/unit/common/utils/feature_quantization/torch_ops_test.py +++ b/tests/unit/common/utils/feature_quantization/torch_ops_test.py @@ -18,7 +18,7 @@ def test_quantize_numpy_dequantize_torch_round_trip_single_bit(self) -> None: pos_mean=1.75, ) - packed = quantize_ndarray(features, bits=metadata.bits, stats={}) + packed = quantize_ndarray(features, bits=metadata.bits) actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) torch.testing.assert_close( @@ -27,15 +27,19 @@ def test_quantize_numpy_dequantize_torch_round_trip_single_bit(self) -> None: def test_quantize_numpy_dequantize_torch_round_trip_multi_bit(self) -> None: features = np.array([[-1.0, 0.0, 1.0, 2.0, 4.0]], dtype=np.float32) - stats = {"clip_min": 0.0, "clip_max": 3.0} + clip_min = 0.0 + clip_max = 3.0 metadata = FeatureQuantizationMetadata( bits=2, feature_dim=features.shape[1], quantized_feature_indices=tuple(range(features.shape[1])), - **stats, + clip_min=clip_min, + clip_max=clip_max, ) - packed = quantize_ndarray(features, bits=metadata.bits, stats=stats) + packed = quantize_ndarray( + features, bits=metadata.bits, clip_min=clip_min, clip_max=clip_max + ) actual = dequantize_torch_tensor(torch.from_numpy(packed), metadata=metadata) torch.testing.assert_close(actual, torch.tensor([[0.0, 0.0, 1.0, 2.0, 3.0]])) From d74590c2926c74e242b00f24ff586b5ab57fb854 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 4 Aug 2026 18:12:55 +0000 Subject: [PATCH 18/78] Upd --- .../lib/transform/feature_quantization.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index 3a0de2ec3..b8cd27f7c 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -114,7 +114,15 @@ def _quantize_record_batch( stats: dict[str, float], ) -> pa.RecordBatch: features = _build_feature_matrix(batch, spec.feature_keys) - packed = quantize_ndarray(features, bits=spec.bits, stats=stats) + if spec.bits == 1: + packed = quantize_ndarray(features, bits=spec.bits) + else: + packed = quantize_ndarray( + features, + bits=spec.bits, + clip_min=stats["clip_min"], + clip_max=stats["clip_max"], + ) schema_names = batch.schema.names keep_indices = [ i for i, name in enumerate(schema_names) if name not in set(spec.feature_keys) From 1ebbd19dc22891b3b95c86f589eff09865d1e2bf Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 17:24:27 +0000 Subject: [PATCH 19/78] Upd --- .../lib/transform/feature_quantization.py | 212 +++++++++--------- .../data_preprocessor/lib/transform/utils.py | 44 ++-- gigl/src/data_preprocessor/lib/types.py | 10 +- .../feature_quantization_transform_test.py | 14 +- 4 files changed, 148 insertions(+), 132 deletions(-) diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index b8cd27f7c..8a1afd6f0 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -1,5 +1,5 @@ import json -from typing import Iterable +from typing import Final, Iterable, List, TypeAlias import apache_beam as beam import numpy as np @@ -15,92 +15,96 @@ from gigl.src.data_preprocessor.lib.types import FeatureQuantizationSpec logger = Logger() -_NODE_PACKED_FEATURE_KEY = "node_packed_features" -_SingleBitAcc = tuple[float, int, float, int] +_NODE_PACKED_FEATURE_KEY: Final[str] = "node_packed_features" +_SingleBitAcc: TypeAlias = tuple[float, int, float, int] def apply_feature_quantization_transform( - transformed_features: beam.PCollection[pa.RecordBatch], - transformed_metadata: DatasetMetadata, - analyzed_metadata: beam.PCollection[DatasetMetadata] | None, - spec: FeatureQuantizationSpec, - feature_keys: list[str], + logical_features: beam.PCollection[pa.RecordBatch], + transform_output_metadata: DatasetMetadata, + analyzed_logical_metadata: beam.PCollection[DatasetMetadata] | None, + quantization_spec: FeatureQuantizationSpec, + logical_feature_keys: list[str], metadata_path: str, ): - logger.info(f"Applying Beam feature quantization with spec: {spec}") - stats = _build_feature_quantization_stats(transformed_features, spec) - logical_metadata = ( - transformed_metadata - if analyzed_metadata is None - else beam.pvalue.AsSingleton(analyzed_metadata) + if analyzed_logical_metadata is None: + logical_metadata = transform_output_metadata + else: + logical_metadata = beam.pvalue.AsSingleton(analyzed_logical_metadata) + + logger.info(f"Applying Beam feature quantization with spec: {quantization_spec}") + quantization_stats = _build_feature_quantization_stats( + logical_features, quantization_spec ) _ = ( - stats - | "Build feature quantization metadata JSON" + quantization_stats + | "Build feature quantization stats JSON" >> beam.Map( - _feature_quantization_metadata_json, - spec=spec, - feature_keys=feature_keys, - dataset_metadata=logical_metadata, + _feature_quantization_stats_to_json, + quantization_spec=quantization_spec, + logical_feature_keys=logical_feature_keys, + logical_metadata=logical_metadata, ) - | "Write feature quantization metadata" + | "Write feature quantization stats" >> beam.io.WriteToText( - metadata_path, - num_shards=1, - shard_name_template="", + metadata_path, num_shards=1, shard_name_template="" ) ) - transformed_features = transformed_features | ( - "Quantize transformed feature RecordBatches" + + quantized_features = logical_features | ( + "Quantize feature RecordBatches" >> beam.Map( _quantize_record_batch, - spec=spec, - stats=beam.pvalue.AsSingleton(stats), + quantization_spec=quantization_spec, + quantization_stats=beam.pvalue.AsSingleton(quantization_stats), ) ) - # Encode TFRecords with the compact physical schema. The persisted schema - # remains the original logical TFT schema because dequantization scatters - # features back. - if analyzed_metadata is None: - physical_metadata = DatasetMetadata( - _apply_feature_quantization_schema(transformed_metadata.schema, spec) + + # Encode TFRecords with the compact physical schema. The persisted schema remains + # the original logical TFT schema because dequantization scatters features back. + if analyzed_logical_metadata is None: + physical_feature_metadata = DatasetMetadata( + _apply_feature_quantization_schema( + transform_output_metadata.schema, quantization_spec + ) ) else: - physical_metadata = analyzed_metadata | ( + physical_feature_metadata = analyzed_logical_metadata | ( "Apply feature quantization schema" >> beam.Map( - lambda metadata, spec: DatasetMetadata( - _apply_feature_quantization_schema(metadata.schema, spec) + lambda metadata, quantization_spec: DatasetMetadata( + _apply_feature_quantization_schema( + metadata.schema, quantization_spec + ) ), - spec=spec, + quantization_spec=quantization_spec, ) ) - physical_metadata = beam.pvalue.AsSingleton(physical_metadata) - return transformed_features, physical_metadata + physical_feature_metadata = beam.pvalue.AsSingleton(physical_feature_metadata) + return quantized_features, physical_feature_metadata def _build_feature_quantization_stats( - record_batches: beam.PCollection[pa.RecordBatch], - spec: FeatureQuantizationSpec, + logical_features: beam.PCollection[pa.RecordBatch], + quantization_spec: FeatureQuantizationSpec, ) -> beam.PCollection[dict[str, float]]: - if spec.bits not in (1, 2, 4, 8): - raise ValueError(f"bits must be one of 1, 2, 4, or 8, got {spec.bits}.") - if not spec.feature_keys: - raise ValueError("Feature quantization expects at least one feature key.") logger.info( - f"Building Beam feature quantization stats for {len(spec.feature_keys)} " - f"features with bits={spec.bits}: {spec.feature_keys}" + f"Building Beam feature quantization stats for {len(quantization_spec.feature_keys)} " + f"features with bits={quantization_spec.bits}: {quantization_spec.feature_keys}" ) - if spec.bits == 1: + if quantization_spec.bits == 1: return ( - record_batches + logical_features | "Compute single bit quantization stats" - >> beam.CombineGlobally(_SingleBitStatsFn(spec.feature_keys)) + >> beam.CombineGlobally(_SingleBitStatsFn(quantization_spec.feature_keys)) ) return ( - record_batches + logical_features | "Build multi-bit quantization value batches" - >> beam.Map(_build_feature_values, feature_keys=spec.feature_keys) + >> beam.Map( + _flatten_feature_values, + quantized_feature_keys=quantization_spec.feature_keys, + ) | "Compute multi-bit quantization quantiles" >> ApproximateQuantiles.Globally(num_quantiles=1000, input_batched=True) | "Build multi-bit quantization stats" @@ -110,25 +114,27 @@ def _build_feature_quantization_stats( def _quantize_record_batch( batch: pa.RecordBatch, - spec: FeatureQuantizationSpec, - stats: dict[str, float], + quantization_spec: FeatureQuantizationSpec, + quantization_stats: dict[str, float], ) -> pa.RecordBatch: - features = _build_feature_matrix(batch, spec.feature_keys) - if spec.bits == 1: - packed = quantize_ndarray(features, bits=spec.bits) + feature_matrix = _build_feature_matrix(batch, quantization_spec.feature_keys) + if quantization_spec.bits == 1: + packed = quantize_ndarray(feature_matrix, bits=quantization_spec.bits) else: packed = quantize_ndarray( - features, - bits=spec.bits, - clip_min=stats["clip_min"], - clip_max=stats["clip_max"], + feature_matrix, + bits=quantization_spec.bits, + clip_min=quantization_stats["clip_min"], + clip_max=quantization_stats["clip_max"], ) - schema_names = batch.schema.names - keep_indices = [ - i for i, name in enumerate(schema_names) if name not in set(spec.feature_keys) + + quantized_feature_keys = set(quantization_spec.feature_keys) + arrays = [ + batch.column(i) + for i, name in enumerate(batch.schema.names) + if name not in quantized_feature_keys ] - arrays = [batch.column(i) for i in keep_indices] - names = [schema_names[i] for i in keep_indices] + names = [name for name in batch.schema.names if name not in quantized_feature_keys] arrays.append( pa.array([[row.tobytes()] for row in packed], type=pa.list_(pa.binary())) ) @@ -136,49 +142,44 @@ def _quantize_record_batch( return pa.RecordBatch.from_arrays(arrays, names=names) -def _feature_quantization_metadata_json( - stats: dict[str, float], - spec: FeatureQuantizationSpec, - feature_keys: list[str], - dataset_metadata: DatasetMetadata, +def _feature_quantization_stats_to_json( + quantization_stats: dict[str, float], + quantization_spec: FeatureQuantizationSpec, + logical_feature_keys: list[str], + logical_metadata: DatasetMetadata, ) -> str: raw_feature_spec = schema_utils.schema_as_feature_spec( - dataset_metadata.schema + logical_metadata.schema ).feature_spec - feature_key_set = set(feature_keys) - missing = [ - key - for key in spec.feature_keys - if key not in raw_feature_spec or key not in feature_key_set - ] - if missing: + missing_feature_keys = set(quantization_spec.feature_keys) - set( + logical_feature_keys + ) + if missing_feature_keys: raise ValueError( - f"Quantized feature keys missing from feature outputs: {missing}" + f"Quantized features missing from feature outputs: {missing_feature_keys}" ) - feature_spec = {key: raw_feature_spec[key] for key in feature_keys} - feature_index = feature_spec_to_feature_index_map(feature_spec) - quantized_feature_indices = [] - for key in spec.feature_keys: - start, end = feature_index[key] + logical_feature_spec = {key: raw_feature_spec[key] for key in logical_feature_keys} + logical_feature_index = feature_spec_to_feature_index_map(logical_feature_spec) + quantized_feature_indices: List[int] = [] + for key in quantization_spec.feature_keys: + start, end = logical_feature_index[key] if end - start != 1: - raise ValueError( - f"Feature quantization expects scalar features, got {key}." - ) + raise ValueError(f"Quantization expects scalar features, got {key}.") quantized_feature_indices.append(start) metadata = { "packed_feature_key": _NODE_PACKED_FEATURE_KEY, "quantized_feature_indices": quantized_feature_indices, - "bits": spec.bits, - **stats, + "bits": quantization_spec.bits, + **quantization_stats, } logger.info(f"Writing feature quantization metadata: {metadata}") return json.dumps(metadata) def _apply_feature_quantization_schema( - schema: schema_pb2.Schema, spec: FeatureQuantizationSpec + schema: schema_pb2.Schema, quantization_spec: FeatureQuantizationSpec ) -> schema_pb2.Schema: - drop_keys = set(spec.feature_keys) | {_NODE_PACKED_FEATURE_KEY} + drop_keys = set(quantization_spec.feature_keys) | {_NODE_PACKED_FEATURE_KEY} quantized_schema = schema_pb2.Schema() quantized_schema.CopyFrom(schema) del quantized_schema.feature[:] @@ -192,23 +193,24 @@ def _apply_feature_quantization_schema( packed_feature.value_count.max = 1 logger.info( f"Updated transformed schema for feature quantization: dropped " - f"{len(spec.feature_keys)} features and added bytes feature " + f"{len(quantization_spec.feature_keys)} features and added bytes feature " f"{_NODE_PACKED_FEATURE_KEY}." ) return quantized_schema -def _build_feature_values( - batch: pa.RecordBatch, feature_keys: list[str] +def _flatten_feature_values( + batch: pa.RecordBatch, quantized_feature_keys: list[str] ) -> list[float]: - values = _build_feature_matrix(batch, feature_keys).reshape(-1) - return values[np.isfinite(values)].astype(float).tolist() + return _build_feature_matrix(batch, quantized_feature_keys).ravel().tolist() -def _build_feature_matrix(batch: pa.RecordBatch, feature_keys: list[str]) -> np.ndarray: +def _build_feature_matrix( + batch: pa.RecordBatch, quantized_feature_keys: list[str] +) -> np.ndarray: key_to_idx = {name: i for i, name in enumerate(batch.schema.names)} cols: list[np.ndarray] = [] - for key in feature_keys: + for key in quantized_feature_keys: if key not in key_to_idx: raise ValueError(f"Feature key {key} not found in RecordBatch.") col = batch.column(key_to_idx[key]) @@ -218,7 +220,10 @@ def _build_feature_matrix(batch: pa.RecordBatch, feature_keys: list[str]) -> np. f"Feature quantization expects scalar features, got {key} with shape {values.shape}." ) cols.append(values) - return np.stack(cols, axis=1) + feature_matrix = np.stack(cols, axis=1) + if not np.isfinite(feature_matrix).all(): + raise ValueError("Feature quantization expects finite feature values.") + return feature_matrix def _multi_bit_stats_from_quantiles(quantiles: list[float]) -> dict[str, float]: @@ -240,7 +245,7 @@ class _SingleBitStatsFn(beam.CombineFn): Used to derive the mean of positive and negative feature values for 1-bit quantization. """ - def __init__(self, feature_keys: list[str]): + def __init__(self, feature_keys: list[str]) -> None: self._feature_keys = feature_keys def create_accumulator(self) -> _SingleBitAcc: @@ -250,8 +255,7 @@ def add_input( self, accumulator: _SingleBitAcc, batch: pa.RecordBatch ) -> _SingleBitAcc: neg_sum, neg_count, pos_sum, pos_count = accumulator - values = _build_feature_matrix(batch, self._feature_keys).reshape(-1) - values = values[np.isfinite(values)] + values = _build_feature_matrix(batch, self._feature_keys).ravel() neg = values <= 0 pos = values > 0 return ( diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index 5dbf79a90..87ef311b8 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -338,57 +338,59 @@ def get_load_data_and_transform_pipeline_component( ) # Apply TransformFn over raw features - transformed_features, transformed_metadata = ( + logical_features, transform_output_metadata = ( (raw_features, raw_tensor_adapter_config), resolved_transform_fn, ) | "Transform raw features dataset" >> tft_beam.TransformDataset( output_record_batches=True ) - # The transformed_features returned by tft_beam.TransformDataset is a + # The feature batches returned by tft_beam.TransformDataset are a # PCollection of Tuple[pa.RecordBatch, dict[str, pa.Array]]. The first - # one are the transformed features. The second one are the passthrough + # item contains logical features. The second one contains passthrough # features, which doesn't apply here since we do not specify passthrough_keys # in tft_beam.Context. Hence we drop the second one in the tuple. - transformed_features = transformed_features | "Extract RecordBatch" >> beam.Map( + logical_features = logical_features | "Extract RecordBatch" >> beam.Map( lambda element: element[0] ) - # The transformed_metadata returned by tft_beam.TransformDataset can only + # The transform output metadata returned by tft_beam.TransformDataset can only # be relied on for encoding purposes when reusing a pretrained transform_fn, # yet it could be inaccurate when using a new transform_fn built by - # tft_beam.AnalyzeDataset. For the later case, we do not use transformed_metadata + # tft_beam.AnalyzeDataset. For the later case, we do not use transform output metadata # returned by tft_beam.TransformDataset, but use deferred_metadata from # transform_fn instead. - resolved_transformed_metadata = ( - transformed_metadata + tfrecord_metadata = ( + transform_output_metadata if should_use_existing_transform_fn else beam.pvalue.AsSingleton(analyzed_transform_fn[1].deferred_metadata) # type: ignore ) - q_spec = None + quantization_spec = None if isinstance(preprocessing_spec, NodeDataPreprocessingSpec): - q_spec = preprocessing_spec.feature_quantization_spec - if q_spec is not None: + quantization_spec = preprocessing_spec.feature_quantization_spec + if quantization_spec is not None: if should_use_existing_transform_fn: - analyzed_metadata = None + analyzed_logical_metadata = None else: - analyzed_metadata = analyzed_transform_fn[1].deferred_metadata # type: ignore + analyzed_logical_metadata = analyzed_transform_fn[1].deferred_metadata # type: ignore - transformed_features, resolved_transformed_metadata = ( + quantized_features, tfrecord_metadata = ( apply_feature_quantization_transform( - transformed_features=transformed_features, - transformed_metadata=transformed_metadata, - analyzed_metadata=analyzed_metadata, - spec=q_spec, - feature_keys=list(preprocessing_spec.features_outputs or []), + logical_features=logical_features, + transform_output_metadata=transform_output_metadata, + analyzed_logical_metadata=analyzed_logical_metadata, + quantization_spec=quantization_spec, + logical_feature_keys=list(preprocessing_spec.features_outputs or []), metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, ) ) + else: + quantized_features = logical_features - transformed_features | "Write tf record files" >> BetterWriteToTFRecord( + quantized_features | "Write tf record files" >> BetterWriteToTFRecord( file_path_prefix=transformed_features_info.transformed_features_file_prefix.uri, max_bytes_per_shard=int(2e8), # 200mb, - transformed_metadata=resolved_transformed_metadata, + transformed_metadata=tfrecord_metadata, # TODO(mkolodner-sc): Right now, a non-zero value for num_shards overrides the max_bytes_per_shard condition. We need to implement # a solution where num_shards specified is just a minimum, causing the max_bytes_per_shard rule taking precedent over the num_shards rule. This will require # dynamically determining the number of shards produced by max_bytes_per_shard and setting it to be equal to min_num_shards if the value is less than it. diff --git a/gigl/src/data_preprocessor/lib/types.py b/gigl/src/data_preprocessor/lib/types.py index a3308b00d..c805e7bf1 100644 --- a/gigl/src/data_preprocessor/lib/types.py +++ b/gigl/src/data_preprocessor/lib/types.py @@ -1,4 +1,5 @@ from abc import ABC, abstractmethod +from dataclasses import dataclass from typing import Any, Callable, NamedTuple, Optional, Tuple import apache_beam as beam @@ -48,10 +49,17 @@ class NodeOutputIdentifier(str): """ -class FeatureQuantizationSpec(NamedTuple): +@dataclass(frozen=True) +class FeatureQuantizationSpec: feature_keys: list[str] bits: int + def __post_init__(self) -> None: + if not self.feature_keys: + raise ValueError("Feature quantization expects at least one feature key.") + if self.bits not in (1, 2, 4, 8): + raise ValueError(f"bits must be one of 1, 2, 4, or 8, got {self.bits}.") + class EdgeOutputIdentifier(NamedTuple): """ 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 a852a62a0..458337be6 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -45,7 +45,7 @@ def test_apply_feature_quantization_transform_quantizes_multi_bit_features( ], names=["node_id", "f0", "f1", "label"], ) - transformed_metadata = DatasetMetadata.from_feature_spec( + transform_output_metadata = DatasetMetadata.from_feature_spec( { "node_id": tf.io.FixedLenFeature(shape=[], dtype=tf.int64), "f0": tf.io.FixedLenFeature(shape=[], dtype=tf.float32), @@ -57,12 +57,14 @@ def test_apply_feature_quantization_transform_quantizes_multi_bit_features( with TestPipeline() as pipeline: transformed_batches, physical_metadata = ( apply_feature_quantization_transform( - transformed_features=pipeline + logical_features=pipeline | "Create RecordBatch" >> beam.Create([batch]), - transformed_metadata=transformed_metadata, - analyzed_metadata=None, - spec=FeatureQuantizationSpec(feature_keys=["f0", "f1"], bits=2), - feature_keys=["f0", "f1"], + transform_output_metadata=transform_output_metadata, + analyzed_logical_metadata=None, + quantization_spec=FeatureQuantizationSpec( + feature_keys=["f0", "f1"], bits=2 + ), + logical_feature_keys=["f0", "f1"], metadata_path=metadata_path, ) ) From 1fd3e285f68846d8af703abcaa012bd2708e54fb Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 17:31:59 +0000 Subject: [PATCH 20/78] upd --- .../data_preprocessor/data_preprocessor.py | 6 +-- .../data_preprocessor/lib/transform/utils.py | 46 +++++++++---------- 2 files changed, 24 insertions(+), 28 deletions(-) diff --git a/gigl/src/data_preprocessor/data_preprocessor.py b/gigl/src/data_preprocessor/data_preprocessor.py index c2e9c3541..6144cce48 100644 --- a/gigl/src/data_preprocessor/data_preprocessor.py +++ b/gigl/src/data_preprocessor/data_preprocessor.py @@ -486,12 +486,10 @@ def generate_preprocessed_metadata_pb( node_transformed_features_info.feature_quantization_metadata_path.uri ) if tf.io.gfile.exists(metadata_path): - logger.info( - f"Loading feature quantization metadata from: {metadata_path}" - ) + logger.info(f"Loading quantization metadata from: {metadata_path}") with tf.io.gfile.GFile(metadata_path) as f: metadata = json.loads(f.read()) - logger.info(f"Loaded feature quantization metadata: {metadata}") + logger.info(f"Loaded quantization metadata: {metadata}") bits = metadata["bits"] quantized_feature_metadata_pb = preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( packed_feature_key=metadata["packed_feature_key"], diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index 87ef311b8..09e9992c6 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -338,59 +338,57 @@ def get_load_data_and_transform_pipeline_component( ) # Apply TransformFn over raw features - logical_features, transform_output_metadata = ( + transformed_features, transformed_metadata = ( (raw_features, raw_tensor_adapter_config), resolved_transform_fn, ) | "Transform raw features dataset" >> tft_beam.TransformDataset( output_record_batches=True ) - # The feature batches returned by tft_beam.TransformDataset are a + # The transformed_features returned by tft_beam.TransformDataset is a # PCollection of Tuple[pa.RecordBatch, dict[str, pa.Array]]. The first - # item contains logical features. The second one contains passthrough + # one are the transformed features. The second one are the passthrough # features, which doesn't apply here since we do not specify passthrough_keys # in tft_beam.Context. Hence we drop the second one in the tuple. - logical_features = logical_features | "Extract RecordBatch" >> beam.Map( + transformed_features = transformed_features | "Extract RecordBatch" >> beam.Map( lambda element: element[0] ) - # The transform output metadata returned by tft_beam.TransformDataset can only + # The transformed_metadata returned by tft_beam.TransformDataset can only # be relied on for encoding purposes when reusing a pretrained transform_fn, # yet it could be inaccurate when using a new transform_fn built by - # tft_beam.AnalyzeDataset. For the later case, we do not use transform output metadata + # tft_beam.AnalyzeDataset. For the later case, we do not use transformed_metadata # returned by tft_beam.TransformDataset, but use deferred_metadata from # transform_fn instead. - tfrecord_metadata = ( - transform_output_metadata + resolved_transformed_metadata = ( + transformed_metadata if should_use_existing_transform_fn else beam.pvalue.AsSingleton(analyzed_transform_fn[1].deferred_metadata) # type: ignore ) - quantization_spec = None + q_spec = None if isinstance(preprocessing_spec, NodeDataPreprocessingSpec): - quantization_spec = preprocessing_spec.feature_quantization_spec - if quantization_spec is not None: + q_spec = preprocessing_spec.feature_quantization_spec + if q_spec is not None: if should_use_existing_transform_fn: - analyzed_logical_metadata = None + analyzed_metadata = None else: - analyzed_logical_metadata = analyzed_transform_fn[1].deferred_metadata # type: ignore + analyzed_metadata = analyzed_transform_fn[1].deferred_metadata # type: ignore - quantized_features, tfrecord_metadata = ( + transformed_features, resolved_transformed_metadata = ( apply_feature_quantization_transform( - logical_features=logical_features, - transform_output_metadata=transform_output_metadata, - analyzed_logical_metadata=analyzed_logical_metadata, - quantization_spec=quantization_spec, - logical_feature_keys=list(preprocessing_spec.features_outputs or []), - metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, + transformed_features, + transformed_metadata, + analyzed_metadata, + q_spec, + list(preprocessing_spec.features_outputs or []), + transformed_features_info.feature_quantization_metadata_path.uri, ) ) - else: - quantized_features = logical_features - quantized_features | "Write tf record files" >> BetterWriteToTFRecord( + transformed_features | "Write tf record files" >> BetterWriteToTFRecord( file_path_prefix=transformed_features_info.transformed_features_file_prefix.uri, max_bytes_per_shard=int(2e8), # 200mb, - transformed_metadata=tfrecord_metadata, + transformed_metadata=resolved_transformed_metadata, # TODO(mkolodner-sc): Right now, a non-zero value for num_shards overrides the max_bytes_per_shard condition. We need to implement # a solution where num_shards specified is just a minimum, causing the max_bytes_per_shard rule taking precedent over the num_shards rule. This will require # dynamically determining the number of shards produced by max_bytes_per_shard and setting it to be equal to min_num_shards if the value is less than it. From 48bab635ef0b6deb3d43a9ff7900828f068e9aaa Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 18:57:31 +0000 Subject: [PATCH 21/78] WIP --- .../data_preprocessor/data_preprocessor.py | 4 +- .../lib/transform/feature_quantization.py | 70 +++++----- .../data_preprocessor/lib/transform/utils.py | 30 +++-- gigl/src/data_preprocessor/lib/types.py | 2 + .../feature_quantization_transform_test.py | 125 +++++++++--------- 5 files changed, 111 insertions(+), 120 deletions(-) diff --git a/gigl/src/data_preprocessor/data_preprocessor.py b/gigl/src/data_preprocessor/data_preprocessor.py index 6144cce48..4f678669b 100644 --- a/gigl/src/data_preprocessor/data_preprocessor.py +++ b/gigl/src/data_preprocessor/data_preprocessor.py @@ -486,10 +486,10 @@ def generate_preprocessed_metadata_pb( node_transformed_features_info.feature_quantization_metadata_path.uri ) if tf.io.gfile.exists(metadata_path): - logger.info(f"Loading quantization metadata from: {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 quantization metadata: {metadata}") + 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"], diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index 8a1afd6f0..b36ffaa2a 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -16,38 +16,38 @@ logger = Logger() _NODE_PACKED_FEATURE_KEY: Final[str] = "node_packed_features" -_SingleBitAcc: TypeAlias = tuple[float, int, float, int] +_SignStats: TypeAlias = tuple[float, int, float, int] def apply_feature_quantization_transform( logical_features: beam.PCollection[pa.RecordBatch], - transform_output_metadata: DatasetMetadata, - analyzed_logical_metadata: beam.PCollection[DatasetMetadata] | None, + logical_metadata: DatasetMetadata | beam.PCollection[DatasetMetadata], quantization_spec: FeatureQuantizationSpec, logical_feature_keys: list[str], - metadata_path: str, -): - if analyzed_logical_metadata is None: - logical_metadata = transform_output_metadata + quantization_metadata_path: str, +) -> tuple[beam.PCollection[pa.RecordBatch], DatasetMetadata | beam.pvalue.AsSingleton]: + logical_metadata_is_eager = isinstance(logical_metadata, DatasetMetadata) + if logical_metadata_is_eager: + metadata_for_json = logical_metadata else: - logical_metadata = beam.pvalue.AsSingleton(analyzed_logical_metadata) + metadata_for_json = beam.pvalue.AsSingleton(logical_metadata) logger.info(f"Applying Beam feature quantization with spec: {quantization_spec}") quantization_stats = _build_feature_quantization_stats( logical_features, quantization_spec ) - _ = ( + ( quantization_stats | "Build feature quantization stats JSON" >> beam.Map( _feature_quantization_stats_to_json, quantization_spec=quantization_spec, logical_feature_keys=logical_feature_keys, - logical_metadata=logical_metadata, + logical_metadata=metadata_for_json, ) | "Write feature quantization stats" >> beam.io.WriteToText( - metadata_path, num_shards=1, shard_name_template="" + quantization_metadata_path, num_shards=1, shard_name_template="" ) ) @@ -60,16 +60,14 @@ def apply_feature_quantization_transform( ) ) - # Encode TFRecords with the compact physical schema. The persisted schema remains - # the original logical TFT schema because dequantization scatters features back. - if analyzed_logical_metadata is None: + if logical_metadata_is_eager: physical_feature_metadata = DatasetMetadata( _apply_feature_quantization_schema( - transform_output_metadata.schema, quantization_spec + logical_metadata.schema, quantization_spec ) ) else: - physical_feature_metadata = analyzed_logical_metadata | ( + physical_feature_metadata = logical_metadata | ( "Apply feature quantization schema" >> beam.Map( lambda metadata, quantization_spec: DatasetMetadata( @@ -96,7 +94,7 @@ def _build_feature_quantization_stats( return ( logical_features | "Compute single bit quantization stats" - >> beam.CombineGlobally(_SingleBitStatsFn(quantization_spec.feature_keys)) + >> beam.CombineGlobally(_PosNegMeanFn(quantization_spec.feature_keys)) ) return ( logical_features @@ -148,24 +146,23 @@ def _feature_quantization_stats_to_json( logical_feature_keys: list[str], logical_metadata: DatasetMetadata, ) -> str: + missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) + if missing: + raise ValueError(f"Quantized features missing: {missing}") + raw_feature_spec = schema_utils.schema_as_feature_spec( logical_metadata.schema ).feature_spec - missing_feature_keys = set(quantization_spec.feature_keys) - set( - logical_feature_keys - ) - if missing_feature_keys: - raise ValueError( - f"Quantized features missing from feature outputs: {missing_feature_keys}" - ) logical_feature_spec = {key: raw_feature_spec[key] for key in logical_feature_keys} logical_feature_index = feature_spec_to_feature_index_map(logical_feature_spec) + quantized_feature_indices: List[int] = [] for key in quantization_spec.feature_keys: start, end = logical_feature_index[key] if end - start != 1: raise ValueError(f"Quantization expects scalar features, got {key}.") quantized_feature_indices.append(start) + metadata = { "packed_feature_key": _NODE_PACKED_FEATURE_KEY, "quantized_feature_indices": quantized_feature_indices, @@ -208,18 +205,20 @@ def _flatten_feature_values( def _build_feature_matrix( batch: pa.RecordBatch, quantized_feature_keys: list[str] ) -> np.ndarray: - key_to_idx = {name: i for i, name in enumerate(batch.schema.names)} + key_to_idx: dict[str, int] = {name: i for i, name in enumerate(batch.schema.names)} cols: list[np.ndarray] = [] for key in quantized_feature_keys: if key not in key_to_idx: raise ValueError(f"Feature key {key} not found in RecordBatch.") + col = batch.column(key_to_idx[key]) values = np.asarray(col.to_numpy(zero_copy_only=False), dtype=np.float32) if values.ndim != 1: raise ValueError( - f"Feature quantization expects scalar features, got {key} with shape {values.shape}." + f"Quantization expects scalar features, got {key} with shape {values.shape}." ) cols.append(values) + feature_matrix = np.stack(cols, axis=1) if not np.isfinite(feature_matrix).all(): raise ValueError("Feature quantization expects finite feature values.") @@ -239,21 +238,16 @@ def _multi_bit_stats_from_quantiles(quantiles: list[float]) -> dict[str, float]: return stats -class _SingleBitStatsFn(beam.CombineFn): - """Beam CombineFn that accumulates sums and counts across batches. - - Used to derive the mean of positive and negative feature values for 1-bit quantization. - """ +class _PosNegMeanFn(beam.CombineFn): + """Accumulates mean positive and negative feature values for 1-bit quantization.""" def __init__(self, feature_keys: list[str]) -> None: self._feature_keys = feature_keys - def create_accumulator(self) -> _SingleBitAcc: + def create_accumulator(self) -> _SignStats: return 0.0, 0, 0.0, 0 - def add_input( - self, accumulator: _SingleBitAcc, batch: pa.RecordBatch - ) -> _SingleBitAcc: + def add_input(self, accumulator: _SignStats, batch: pa.RecordBatch) -> _SignStats: neg_sum, neg_count, pos_sum, pos_count = accumulator values = _build_feature_matrix(batch, self._feature_keys).ravel() neg = values <= 0 @@ -265,9 +259,7 @@ def add_input( pos_count + int(pos.sum()), ) - def merge_accumulators( - self, accumulators: Iterable[_SingleBitAcc] - ) -> _SingleBitAcc: + def merge_accumulators(self, accumulators: Iterable[_SignStats]) -> _SignStats: neg_sum = neg_count = pos_sum = pos_count = 0 for n_sum, n_count, p_sum, p_count in accumulators: neg_sum += n_sum @@ -276,7 +268,7 @@ def merge_accumulators( pos_count += p_count return neg_sum, neg_count, pos_sum, pos_count - def extract_output(self, accumulator: _SingleBitAcc) -> dict[str, float]: + def extract_output(self, accumulator: _SignStats) -> dict[str, float]: neg_sum, neg_count, pos_sum, pos_count = accumulator stats = { "neg_mean": neg_sum / neg_count if neg_count else 0.0, diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index 09e9992c6..e0c344e21 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -36,6 +36,7 @@ ) from gigl.src.data_preprocessor.lib.types import ( EdgeDataPreprocessingSpec, + FeatureQuantizationSpec, FeatureSpecDict, InstanceDict, NodeDataPreprocessingSpec, @@ -365,23 +366,24 @@ def get_load_data_and_transform_pipeline_component( if should_use_existing_transform_fn else beam.pvalue.AsSingleton(analyzed_transform_fn[1].deferred_metadata) # type: ignore ) - q_spec = None + logical_metadata = ( + transformed_metadata + if should_use_existing_transform_fn + else analyzed_transform_fn[1].deferred_metadata # type: ignore + ) + quantization_spec: FeatureQuantizationSpec | None = None if isinstance(preprocessing_spec, NodeDataPreprocessingSpec): - q_spec = preprocessing_spec.feature_quantization_spec - if q_spec is not None: - if should_use_existing_transform_fn: - analyzed_metadata = None - else: - analyzed_metadata = analyzed_transform_fn[1].deferred_metadata # type: ignore - + quantization_spec = preprocessing_spec.feature_quantization_spec + if quantization_spec is not None: transformed_features, resolved_transformed_metadata = ( apply_feature_quantization_transform( - transformed_features, - transformed_metadata, - analyzed_metadata, - q_spec, - list(preprocessing_spec.features_outputs or []), - transformed_features_info.feature_quantization_metadata_path.uri, + logical_features=transformed_features, + logical_metadata=logical_metadata, + quantization_spec=quantization_spec, + logical_feature_keys=list( + preprocessing_spec.features_outputs or [] + ), + metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, ) ) diff --git a/gigl/src/data_preprocessor/lib/types.py b/gigl/src/data_preprocessor/lib/types.py index c805e7bf1..7012635ff 100644 --- a/gigl/src/data_preprocessor/lib/types.py +++ b/gigl/src/data_preprocessor/lib/types.py @@ -51,6 +51,8 @@ class NodeOutputIdentifier(str): @dataclass(frozen=True) class FeatureQuantizationSpec: + """Selects logical feature fields to pack at a fixed bit width.""" + feature_keys: list[str] bits: int 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 458337be6..a2cc0a686 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -7,6 +7,7 @@ import tensorflow as tf from apache_beam.testing.test_pipeline import TestPipeline from apache_beam.testing.util import assert_that, equal_to +from parameterized import parameterized from tensorflow_metadata.proto.v0 import schema_pb2 from tensorflow_transform.tf_metadata.dataset_metadata import DatasetMetadata @@ -17,34 +18,45 @@ from tests.test_assets.test_case import TestCase -def _column_pylist(batch: pa.RecordBatch, name: str) -> list: - return batch.column(batch.schema.names.index(name)).to_pylist() - - -def _record_batch_summary(batch: pa.RecordBatch) -> dict[str, object]: - return { - "names": batch.schema.names, - "node_id": _column_pylist(batch, "node_id"), - "label": _column_pylist(batch, "label"), - "node_packed_features": _column_pylist(batch, "node_packed_features"), - } - - class FeatureQuantizationTransformTest(TestCase): - def test_apply_feature_quantization_transform_quantizes_multi_bit_features( + @parameterized.expand( + [ + ( + "multibit", + 2, + [(-2.0, -2.0), (8.0, 8.0)], + {"clip_min": -2.0, "clip_max": 8.0}, + ), + ( + "single_bit", + 1, + [(-4.0, -2.0), (4.0, 8.0)], + {"neg_mean": -3.0, "pos_mean": 6.0}, + ), + ] + ) + def test_apply_feature_quantization_transform_writes_metadata( self, + _: str, + bits: int, + feature_values: list[tuple[float, float]], + expected_stats: dict[str, float], ) -> None: with tempfile.TemporaryDirectory() as temp_dir: metadata_path = os.path.join(temp_dir, "feature_quantization_metadata.json") - batch = pa.RecordBatch.from_arrays( - [ - pa.array([10, 11], type=pa.int64()), - pa.array([-2.0, 8.0], type=pa.float32()), - pa.array([-2.0, 8.0], type=pa.float32()), - pa.array([0, 1], type=pa.int64()), - ], - names=["node_id", "f0", "f1", "label"], - ) + logical_feature_keys = ["f0", "f1"] + batches = [ + pa.RecordBatch.from_arrays( + [ + pa.array([node_id], type=pa.int64()), + pa.array([f0], type=pa.float32()), + pa.array([f1], type=pa.float32()), + pa.array([node_id], type=pa.int64()), + ], + names=["node_id", "f0", "f1", "label"], + ) + for node_id, (f0, f1) in enumerate(feature_values) + ] transform_output_metadata = DatasetMetadata.from_feature_spec( { "node_id": tf.io.FixedLenFeature(shape=[], dtype=tf.int64), @@ -58,14 +70,13 @@ def test_apply_feature_quantization_transform_quantizes_multi_bit_features( transformed_batches, physical_metadata = ( apply_feature_quantization_transform( logical_features=pipeline - | "Create RecordBatch" >> beam.Create([batch]), - transform_output_metadata=transform_output_metadata, - analyzed_logical_metadata=None, + | "Create RecordBatches" >> beam.Create(batches), + logical_metadata=transform_output_metadata, quantization_spec=FeatureQuantizationSpec( - feature_keys=["f0", "f1"], bits=2 + feature_keys=logical_feature_keys, bits=bits ), - logical_feature_keys=["f0", "f1"], - metadata_path=metadata_path, + logical_feature_keys=logical_feature_keys, + quantization_metadata_path=metadata_path, ) ) physical_features = { @@ -73,51 +84,35 @@ def test_apply_feature_quantization_transform_quantizes_multi_bit_features( for feature in physical_metadata.schema.feature } self.assertEqual( - set(physical_features), - {"node_id", "label", "node_packed_features"}, + set(physical_features), {"node_id", "label", "node_packed_features"} ) + packed_feature = physical_features["node_packed_features"] self.assertEqual( - physical_features["node_packed_features"].type, - schema_pb2.BYTES, - ) - self.assertEqual( - physical_features["node_packed_features"].value_count.min, - 1, - ) - self.assertEqual( - physical_features["node_packed_features"].value_count.max, - 1, + ( + packed_feature.type, + packed_feature.value_count.min, + packed_feature.value_count.max, + ), + (schema_pb2.BYTES, 1, 1), ) - # These values sit exactly at the learned clip bounds, so this - # does not depend on mid-bucket rounding: min/min maps to - # 00000000 and max/max maps to 11110000 with two padded codes. assert_that( transformed_batches - | "Summarize RecordBatch" >> beam.Map(_record_batch_summary), + | "Extract quantized feature names" + >> beam.Map(lambda batch: batch.schema.names), equal_to( - [ - { - "names": [ - "node_id", - "label", - "node_packed_features", - ], - "node_id": [10, 11], - "label": [0, 1], - "node_packed_features": [ - [bytes([0])], - [bytes([240])], - ], - } - ] + [["node_id", "label", "node_packed_features"]] * len(batches) ), ) with open(metadata_path) as metadata_file: metadata = json.load(metadata_file) - self.assertEqual(metadata["packed_feature_key"], "node_packed_features") - self.assertEqual(metadata["quantized_feature_indices"], [0, 1]) - self.assertEqual(metadata["bits"], 2) - self.assertEqual(metadata["clip_min"], -2.0) - self.assertEqual(metadata["clip_max"], 8.0) + self.assertEqual( + metadata, + { + "packed_feature_key": "node_packed_features", + "quantized_feature_indices": [0, 1], + "bits": bits, + **expected_stats, + }, + ) From 1ecf397bb4d5008178000240c8fe529477076b21 Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 18:59:36 +0000 Subject: [PATCH 22/78] WIP --- .../data_preprocessor/lib/transform/feature_quantization.py | 2 +- gigl/src/data_preprocessor/lib/transform/utils.py | 4 ++-- .../data_preprocessor/feature_quantization_transform_test.py | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index b36ffaa2a..d54f9c3f2 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -22,8 +22,8 @@ def apply_feature_quantization_transform( logical_features: beam.PCollection[pa.RecordBatch], logical_metadata: DatasetMetadata | beam.PCollection[DatasetMetadata], - quantization_spec: FeatureQuantizationSpec, logical_feature_keys: list[str], + quantization_spec: FeatureQuantizationSpec, quantization_metadata_path: str, ) -> tuple[beam.PCollection[pa.RecordBatch], DatasetMetadata | beam.pvalue.AsSingleton]: logical_metadata_is_eager = isinstance(logical_metadata, DatasetMetadata) diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index e0c344e21..07bfeaf7c 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -379,11 +379,11 @@ def get_load_data_and_transform_pipeline_component( apply_feature_quantization_transform( logical_features=transformed_features, logical_metadata=logical_metadata, - quantization_spec=quantization_spec, logical_feature_keys=list( preprocessing_spec.features_outputs or [] ), - metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, + quantization_spec=quantization_spec, + quantization_metadata_path=transformed_features_info.feature_quantization_metadata_path.uri, ) ) 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 a2cc0a686..cc146f8e9 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -72,10 +72,10 @@ def test_apply_feature_quantization_transform_writes_metadata( logical_features=pipeline | "Create RecordBatches" >> beam.Create(batches), logical_metadata=transform_output_metadata, + logical_feature_keys=logical_feature_keys, quantization_spec=FeatureQuantizationSpec( feature_keys=logical_feature_keys, bits=bits ), - logical_feature_keys=logical_feature_keys, quantization_metadata_path=metadata_path, ) ) From 2fd2c64037818ec1c6e2ef37a8e76e49cf5ef00c Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 19:13:11 +0000 Subject: [PATCH 23/78] WIP --- .../lib/transform/feature_quantization.py | 69 ++++++++++--------- 1 file changed, 36 insertions(+), 33 deletions(-) diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index d54f9c3f2..e80e30aca 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -1,5 +1,5 @@ import json -from typing import Final, Iterable, List, TypeAlias +from typing import Final, Iterable, TypeAlias import apache_beam as beam import numpy as np @@ -26,26 +26,28 @@ def apply_feature_quantization_transform( quantization_spec: FeatureQuantizationSpec, quantization_metadata_path: str, ) -> tuple[beam.PCollection[pa.RecordBatch], DatasetMetadata | beam.pvalue.AsSingleton]: + missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) + if missing: + raise ValueError(f"Quantized features missing: {missing}") + logical_metadata_is_eager = isinstance(logical_metadata, DatasetMetadata) if logical_metadata_is_eager: metadata_for_json = logical_metadata else: metadata_for_json = beam.pvalue.AsSingleton(logical_metadata) - logger.info(f"Applying Beam feature quantization with spec: {quantization_spec}") - quantization_stats = _build_feature_quantization_stats( - logical_features, quantization_spec - ) + logger.info(f"Applying feature quantization with spec: {quantization_spec}") + quantization_stats = _build_quantization_stats(logical_features, quantization_spec) ( quantization_stats - | "Build feature quantization stats JSON" + | "Build quantization stats JSON" >> beam.Map( - _feature_quantization_stats_to_json, + _quantization_stats_to_json, quantization_spec=quantization_spec, logical_feature_keys=logical_feature_keys, logical_metadata=metadata_for_json, ) - | "Write feature quantization stats" + | "Write quantization stats" >> beam.io.WriteToText( quantization_metadata_path, num_shards=1, shard_name_template="" ) @@ -62,18 +64,14 @@ def apply_feature_quantization_transform( if logical_metadata_is_eager: physical_feature_metadata = DatasetMetadata( - _apply_feature_quantization_schema( - logical_metadata.schema, quantization_spec - ) + _apply_quantization_schema(logical_metadata.schema, quantization_spec) ) else: physical_feature_metadata = logical_metadata | ( "Apply feature quantization schema" >> beam.Map( lambda metadata, quantization_spec: DatasetMetadata( - _apply_feature_quantization_schema( - metadata.schema, quantization_spec - ) + _apply_quantization_schema(metadata.schema, quantization_spec) ), quantization_spec=quantization_spec, ) @@ -82,7 +80,7 @@ def apply_feature_quantization_transform( return quantized_features, physical_feature_metadata -def _build_feature_quantization_stats( +def _build_quantization_stats( logical_features: beam.PCollection[pa.RecordBatch], quantization_spec: FeatureQuantizationSpec, ) -> beam.PCollection[dict[str, float]]: @@ -140,40 +138,45 @@ def _quantize_record_batch( return pa.RecordBatch.from_arrays(arrays, names=names) -def _feature_quantization_stats_to_json( +def _quantization_stats_to_json( quantization_stats: dict[str, float], quantization_spec: FeatureQuantizationSpec, logical_feature_keys: list[str], logical_metadata: DatasetMetadata, ) -> str: - missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) - if missing: - raise ValueError(f"Quantized features missing: {missing}") + metadata = { + "packed_feature_key": _NODE_PACKED_FEATURE_KEY, + "quantized_feature_indices": _quantized_feature_indices( + logical_metadata, logical_feature_keys, quantization_spec.feature_keys + ), + "bits": quantization_spec.bits, + **quantization_stats, + } + logger.info(f"Writing feature quantization metadata: {metadata}") + return json.dumps(metadata) + +def _quantized_feature_indices( + logical_metadata: DatasetMetadata, + logical_feature_keys: list[str], + quantized_feature_keys: list[str], +) -> list[int]: raw_feature_spec = schema_utils.schema_as_feature_spec( logical_metadata.schema ).feature_spec logical_feature_spec = {key: raw_feature_spec[key] for key in logical_feature_keys} logical_feature_index = feature_spec_to_feature_index_map(logical_feature_spec) - quantized_feature_indices: List[int] = [] - for key in quantization_spec.feature_keys: + feature_indices: list[int] = [] + for key in quantized_feature_keys: start, end = logical_feature_index[key] if end - start != 1: - raise ValueError(f"Quantization expects scalar features, got {key}.") - quantized_feature_indices.append(start) - - metadata = { - "packed_feature_key": _NODE_PACKED_FEATURE_KEY, - "quantized_feature_indices": quantized_feature_indices, - "bits": quantization_spec.bits, - **quantization_stats, - } - logger.info(f"Writing feature quantization metadata: {metadata}") - return json.dumps(metadata) + raise ValueError(f"Quantization expects scalar features, got {key}") + feature_indices.append(start) + return feature_indices -def _apply_feature_quantization_schema( +def _apply_quantization_schema( schema: schema_pb2.Schema, quantization_spec: FeatureQuantizationSpec ) -> schema_pb2.Schema: drop_keys = set(quantization_spec.feature_keys) | {_NODE_PACKED_FEATURE_KEY} From 3fbe66a60205523063fa65990a616d945aa55ecd Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 21:51:18 +0000 Subject: [PATCH 24/78] Add SUPPORTED_QUANTIZATION_BITS const --- gigl/common/utils/feature_quantization/__init__.py | 4 ++++ gigl/common/utils/feature_quantization/numpy_ops.py | 9 ++++++--- gigl/src/data_preprocessor/lib/types.py | 7 +++++-- gigl/types/graph.py | 8 +++++--- 4 files changed, 20 insertions(+), 8 deletions(-) diff --git a/gigl/common/utils/feature_quantization/__init__.py b/gigl/common/utils/feature_quantization/__init__.py index 6f1229294..ed7d37bd9 100644 --- a/gigl/common/utils/feature_quantization/__init__.py +++ b/gigl/common/utils/feature_quantization/__init__.py @@ -1 +1,5 @@ """Utilities for node feature quantization in GiGL.""" + +from typing import Final + +SUPPORTED_QUANTIZATION_BITS: Final[tuple[int, ...]] = (1, 2, 4, 8) diff --git a/gigl/common/utils/feature_quantization/numpy_ops.py b/gigl/common/utils/feature_quantization/numpy_ops.py index ce15cc008..e2fbd393a 100644 --- a/gigl/common/utils/feature_quantization/numpy_ops.py +++ b/gigl/common/utils/feature_quantization/numpy_ops.py @@ -7,6 +7,8 @@ import numpy as np +from gigl.common.utils.feature_quantization import SUPPORTED_QUANTIZATION_BITS + def quantize_ndarray( features: np.ndarray, @@ -21,9 +23,10 @@ def quantize_ndarray( define the min-max scaling range: values are clipped to that range, scaled to `[0, 2**bits - 1]`, rounded to integer codes, then packed into bytes. """ - valid_bits = (1, 2, 4, 8) - if bits not in valid_bits: - raise ValueError(f"bits must be one of {valid_bits}, got {bits}") + if bits not in SUPPORTED_QUANTIZATION_BITS: + raise ValueError( + f"bits must be one of {SUPPORTED_QUANTIZATION_BITS}, got {bits}" + ) if features.ndim != 2: raise ValueError(f"Expected a 2D feature array, got shape {features.shape}.") if not np.isfinite(features).all(): diff --git a/gigl/src/data_preprocessor/lib/types.py b/gigl/src/data_preprocessor/lib/types.py index 7012635ff..3a660933d 100644 --- a/gigl/src/data_preprocessor/lib/types.py +++ b/gigl/src/data_preprocessor/lib/types.py @@ -9,6 +9,7 @@ from tensorflow_transform import common_types from gigl.common import Uri +from gigl.common.utils.feature_quantization import SUPPORTED_QUANTIZATION_BITS # TODO (mkolodner-sc): Move these variables to a more general location, as they are used even outside of context of data preprocessor @@ -59,8 +60,10 @@ class FeatureQuantizationSpec: def __post_init__(self) -> None: if not self.feature_keys: raise ValueError("Feature quantization expects at least one feature key.") - if self.bits not in (1, 2, 4, 8): - raise ValueError(f"bits must be one of 1, 2, 4, or 8, got {self.bits}.") + if self.bits not in SUPPORTED_QUANTIZATION_BITS: + raise ValueError( + f"bits must be one of {SUPPORTED_QUANTIZATION_BITS}, got {self.bits}." + ) class EdgeOutputIdentifier(NamedTuple): diff --git a/gigl/types/graph.py b/gigl/types/graph.py index 1d4d208a0..3c3622b9a 100644 --- a/gigl/types/graph.py +++ b/gigl/types/graph.py @@ -9,6 +9,7 @@ from gigl.common.data.dataloaders import SerializedTFRecordInfo from gigl.common.logger import Logger +from gigl.common.utils.feature_quantization import SUPPORTED_QUANTIZATION_BITS # TODO(kmonte) - we should move gigl.src.common.types.graph_data to this file. from gigl.src.common.types.graph_data import EdgeType, NodeType, Relation @@ -129,9 +130,10 @@ class FeatureQuantizationMetadata: pos_mean: Optional[float] = None def __post_init__(self) -> None: - valid_bits = (1, 2, 4, 8) - if self.bits not in valid_bits: - raise ValueError(f"bits must be one of {valid_bits}, got {self.bits}") + if self.bits not in SUPPORTED_QUANTIZATION_BITS: + raise ValueError( + f"bits must be one of {SUPPORTED_QUANTIZATION_BITS}, got {self.bits}" + ) if any(i < 0 or i >= self.feature_dim for i in self.quantized_feature_indices): raise ValueError( f"quantized_feature_indices must be in [0, {self.feature_dim}), got {self.quantized_feature_indices}" From 109e1e5976da34ccc2556785fae7a58b1038f23a Mon Sep 17 00:00:00 2001 From: jchmura Date: Fri, 7 Aug 2026 21:57:12 +0000 Subject: [PATCH 25/78] Add docstring to quantization transform --- .../lib/transform/feature_quantization.py | 22 +++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index e80e30aca..5ae5dbdac 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -26,6 +26,28 @@ def apply_feature_quantization_transform( quantization_spec: FeatureQuantizationSpec, quantization_metadata_path: 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. + + Side Effects: + Writes the quantization statistics JSON that ``data_preprocessor.py`` + reads and serializes into the preprocessing metadata protobuf. + + Args: + logical_features: RecordBatches containing the logical feature columns. + logical_metadata: Eager or deferred metadata for the logical schema. + 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. + + 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 + logical model schema remains unchanged. + """ missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) if missing: raise ValueError(f"Quantized features missing: {missing}") From 7ed4b29f30f1bc717057f5f33d484d669908726c Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 15:41:08 +0000 Subject: [PATCH 26/78] Format docs --- .../utils/feature_quantization/README.md | 60 ++++++++----------- 1 file changed, 25 insertions(+), 35 deletions(-) diff --git a/gigl/common/utils/feature_quantization/README.md b/gigl/common/utils/feature_quantization/README.md index 02b413d2e..0453c9e3c 100644 --- a/gigl/common/utils/feature_quantization/README.md +++ b/gigl/common/utils/feature_quantization/README.md @@ -1,26 +1,21 @@ # Feature Quantization -This package contains the low-level NumPy and Torch helpers for node feature -quantization in GiGL. +This package contains the low-level NumPy and Torch helpers for node feature quantization in GiGL. -Feature quantization is lossy compression: high-precision feature values such as -fp32 are mapped into a lower-precision representation. The current built-in -scheme stores low-bit codes as packed `uint8` bytes, then reconstructs -approximate feature values when a sampled subgraph is materialized for training -or inference. +Feature quantization is lossy compression: high-precision feature values such as fp32 are mapped into a lower-precision +representation. The current built-in scheme stores low-bit codes as packed `uint8` bytes, then reconstructs approximate +feature values when a sampled subgraph is materialized for training or inference. -The motivation is practical scaling. Large-scale GNN training is often -memory-bound: +The motivation is practical scaling. Large-scale GNN training is often memory-bound: - feature hydration uses irregular memory access, often over the network; - hydrated features still need to move to the accelerator; - large feature stores limit the workloads that fit on a given machine. -Reducing feature size can improve feature-store footprint, network bandwidth, -and PCIe transfer volume. GNNs have also been shown to be relatively tolerant of -input feature quantization in [BiFeat: Supercharge GNN Training via Graph -Feature Quantization](https://arxiv.org/abs/2207.14696), which motivates this as -a useful tradeoff for GiGL. +Reducing feature size can improve feature-store footprint, network bandwidth, and PCIe transfer volume. GNNs have also +been shown to be relatively tolerant of input feature quantization in +[BiFeat: Supercharge GNN Training via Graph Feature Quantization](https://arxiv.org/abs/2207.14696), which motivates +this as a useful tradeoff for GiGL. ## Current Built-In Flow @@ -35,30 +30,27 @@ The built-in flow is: The NumPy/Torch split is intentional: -- `numpy_ops.py` runs in preprocessing, where data is on CPU and Torch may not - be available. -- `torch_ops.py` runs during dataloader collation, where sampled feature data is - already represented as Torch tensors and may already be on GPU. +- `numpy_ops.py` runs in preprocessing, where data is on CPU and Torch may not be available. +- `torch_ops.py` runs during dataloader collation, where sampled feature data is already represented as Torch tensors + and may already be on GPU. -`FeatureQuantizationMetadata` is the contract between those two steps. It records -the bit width, packed feature dimension, logical feature positions, and the -statistics needed to invert the compression step. +`FeatureQuantizationMetadata` is the contract between those two steps. It records the bit width, packed feature +dimension, logical feature positions, and the statistics needed to invert the compression step. ## Current Built-In Scheme -The current implementation supports `1`, `2`, `4`, and `8` bit quantization. -Codes are packed high-bits-first into bytes. +The current implementation supports `1`, `2`, `4`, and `8` bit quantization. Codes are packed high-bits-first into +bytes. -For `1` bit, values are represented by sign and reconstructed from the positive -and non-positive means. +For `1` bit, values are represented by sign and reconstructed from the positive and non-positive means. -For `2`, `4`, and `8` bits, values are clipped to pre-computed bounds and mapped into -uniform integer buckets between those bounds. +For `2`, `4`, and `8` bits, values are clipped to pre-computed bounds and mapped into uniform integer buckets between +those bounds. ## TODO: Pluggable Schemes -The current metadata/proto shape is tied to the built-in quantization scheme. A -useful follow-up is to make the quantization scheme itself pluggable. +The current metadata/proto shape is tied to the built-in quantization scheme. A useful follow-up is to make the +quantization scheme itself pluggable. One possible design is: @@ -66,10 +58,8 @@ One possible design is: - serialize a stable fully qualified name or registry key for the quantizer; - serialize the quantizer arguments in the proto; - rebuild the quantizer from that metadata on the read side; -- require each quantizer to provide matching NumPy `quantize` and Torch - `dequantize` implementations. +- require each quantizer to provide matching NumPy `quantize` and Torch `dequantize` implementations. -That would let developers add new schemes without threading one-off fields -through every proto and loader path. A registry key is likely safer than -arbitrary imports, but either way the important contract is that the serialized -scheme identifies both the compression step and its inverse. +That would let developers add new schemes without threading one-off fields through every proto and loader path. A +registry key is likely safer than arbitrary imports, but either way the important contract is that the serialized scheme +identifies both the compression step and its inverse. From b490889d22bc6a3fec3b37f36970372ffc35f75b Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 15:43:10 +0000 Subject: [PATCH 27/78] Revert diff --- gigl/common/utils/feature_quantization/README.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/gigl/common/utils/feature_quantization/README.md b/gigl/common/utils/feature_quantization/README.md index 0453c9e3c..181548b31 100644 --- a/gigl/common/utils/feature_quantization/README.md +++ b/gigl/common/utils/feature_quantization/README.md @@ -1,6 +1,6 @@ # Feature Quantization -This package contains the low-level NumPy and Torch helpers for node feature quantization in GiGL. +This package contains the low-level NumPy and Torch helpers for feature quantization in GiGL. Feature quantization is lossy compression: high-precision feature values such as fp32 are mapped into a lower-precision representation. The current built-in scheme stores low-bit codes as packed `uint8` bytes, then reconstructs approximate @@ -22,7 +22,7 @@ this as a useful tradeoff for GiGL. The built-in flow is: 1. The data preprocessor computes feature summary statistics offline. -2. The preprocessor quantizes selected scalar node feature columns with NumPy. +2. The preprocessor quantizes selected scalar feature columns with NumPy. 3. The packed `uint8` feature sidecar is written to TFRecords. 4. Distributed dataset construction partitions and samples the packed bytes. 5. The dataloader collate path dequantizes sampled packed features with Torch. From 42676f5f8707243222ab9a93438995e73eba2db7 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 15:46:17 +0000 Subject: [PATCH 28/78] Revert diff in graph types --- gigl/types/graph.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/gigl/types/graph.py b/gigl/types/graph.py index 3c3622b9a..0b3a4c729 100644 --- a/gigl/types/graph.py +++ b/gigl/types/graph.py @@ -167,11 +167,14 @@ def raw_feature_dim(self) -> int: """Number of logical features that remain in raw form.""" return len(self.raw_feature_indices) - @lru_cache(maxsize=2) + # One cache entry matches the expected single-device dataloader call path. + @lru_cache(maxsize=1) def scatter_index_tensors( self, device: torch.device ) -> FeatureQuantizationIndexTensors: - """Device-local indices for scattering quantized and raw features.""" + """Device-local scatter indices for the single-device hot path.""" + # The logical indices never change across batches, so cache their + # device-local tensors for repeated quantized/raw feature scatter writes. quantized = torch.tensor( self.quantized_feature_indices, dtype=torch.long, device=device ) From 6c8c3c43899b84e59c2c2c0a14a13b44c12d7365 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 16:37:05 +0000 Subject: [PATCH 29/78] Cleanup --- gigl/common/data/dataloaders.py | 2 +- .../serialized_graph_metadata_translator.py | 60 +++++++++---------- 2 files changed, 31 insertions(+), 31 deletions(-) diff --git a/gigl/common/data/dataloaders.py b/gigl/common/data/dataloaders.py index e0eab92e0..824e5225d 100644 --- a/gigl/common/data/dataloaders.py +++ b/gigl/common/data/dataloaders.py @@ -529,7 +529,7 @@ def load_as_torch_tensors( if quantized_feature_tensors: output_quantized_feature_tensor = _tf_tensor_to_torch_tensor( tf.concat(quantized_feature_tensors, axis=0) - ).to(torch.uint8) + ) if label_tensors: output_label_tensor = _tf_tensor_to_torch_tensor( tf.concat(label_tensors, axis=0) diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index 35321e8b6..ac2bb447c 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -1,4 +1,4 @@ -from typing import Tuple, Union +from typing import Optional, Tuple, Union from gigl.common import UriFactory from gigl.common.data.dataloaders import SerializedTFRecordInfo @@ -21,6 +21,7 @@ def _build_serialized_tfrecord_entity_info( feature_spec_dict: FeatureSpecDict, entity_key: Union[str, Tuple[str, str]], tfrecord_uri_pattern: str, + quantization_metadata: Optional[FeatureQuantizationMetadata] = None, ) -> SerializedTFRecordInfo: """ Populates a SerializedTFRecordInfo field from provided arguments for either a node or edge entity of a single node/edge type. @@ -31,21 +32,16 @@ def _build_serialized_tfrecord_entity_info( feature_spec_dict (FeatureSpecDict): Feature spec to register to SerializedTFRecordInfo entity_key (Union[str, Tuple[str, str]]): Entity key to register to SerializedTFRecordInfo, is a str if Node entity or Tuple[str, str] if Edge entity tfrecord_uri_pattern (str): Regex pattern for loading serialized tf records + quantization_metadata (Optional[FeatureQuantizationMetadata]): Quantization + metadata for a node entity, when its features are quantized. Returns: SerializedTFRecordInfo: Stored metadata for current entity """ - packed_feature_key = None - packed_feature_dim = 0 - physical_feature_keys = list(preprocessed_metadata.feature_keys) - feature_dim = preprocessed_metadata.feature_dim - - if isinstance( - preprocessed_metadata, PreprocessedMetadata.NodeMetadataOutput - ) and preprocessed_metadata.HasField("quantized_feature_metadata"): - quantization_metadata = _build_feature_quantization_metadata( - quantized_metadata=preprocessed_metadata.quantized_feature_metadata, - feature_dim=preprocessed_metadata.feature_dim, - ) + if quantization_metadata is not None: + if not isinstance( + preprocessed_metadata, PreprocessedMetadata.NodeMetadataOutput + ): + raise ValueError("Quantization is supported only for node entities.") packed_feature_key = ( preprocessed_metadata.quantized_feature_metadata.packed_feature_key ) @@ -55,31 +51,36 @@ def _build_serialized_tfrecord_entity_info( {key: feature_spec_dict[key] for key in preprocessed_metadata.feature_keys} ) - physical_feature_keys = [] + feature_keys_to_load: list[str] = [] + # Keep only wholly unquantized fields: TFRecord decoding cannot split one + # field between float and packed storage. for key in preprocessed_metadata.feature_keys: key_indices = set(range(*feature_index[key])) quantized_key_indices = key_indices.intersection(quantized_indices) if not quantized_key_indices: - physical_feature_keys.append(key) + feature_keys_to_load.append(key) elif quantized_key_indices != key_indices: - raise ValueError( - f"Partial feature quantization is not supported for {key}." - ) + raise ValueError(f"Partial quantization not supported for {key}") feature_dim = quantization_metadata.raw_feature_dim + else: + packed_feature_key = None + packed_feature_dim = 0 + feature_keys_to_load = list(preprocessed_metadata.feature_keys) + feature_dim = preprocessed_metadata.feature_dim - physical_keys = set(physical_feature_keys) - physical_keys.update(preprocessed_metadata.label_keys) + serialized_keys = set(feature_keys_to_load) + serialized_keys.update(preprocessed_metadata.label_keys) if packed_feature_key is not None: - physical_keys.add(packed_feature_key) + serialized_keys.add(packed_feature_key) feature_spec_dict = { - key: spec for key, spec in feature_spec_dict.items() if key in physical_keys + key: spec for key, spec in feature_spec_dict.items() if key in serialized_keys } return SerializedTFRecordInfo( tfrecord_uri_prefix=UriFactory.create_uri( preprocessed_metadata.tfrecord_uri_prefix ), - feature_keys=physical_feature_keys, + feature_keys=feature_keys_to_load, feature_spec=feature_spec_dict, feature_dim=feature_dim, entity_key=entity_key, @@ -159,20 +160,19 @@ def convert_pb_to_serialized_graph_metadata( ) node_key = node_metadata.node_id_key + if node_metadata.HasField("quantized_feature_metadata"): + node_quantization_metadata[node_type] = _build_feature_quantization_metadata( + quantized_metadata=node_metadata.quantized_feature_metadata, + feature_dim=node_metadata.feature_dim, + ) node_entity_info[node_type] = _build_serialized_tfrecord_entity_info( preprocessed_metadata=node_metadata, feature_spec_dict=node_feature_spec_dict, entity_key=node_key, tfrecord_uri_pattern=tfrecord_uri_pattern, + quantization_metadata=node_quantization_metadata.get(node_type), ) - if node_metadata.HasField("quantized_feature_metadata"): - node_quantization_metadata[node_type] = ( - _build_feature_quantization_metadata( - quantized_metadata=node_metadata.quantized_feature_metadata, - feature_dim=node_metadata.feature_dim, - ) - ) for edge_type in graph_metadata_pb_wrapper.edge_types: condensed_edge_type = ( From 8e59f967ec1cbe8f87f4838878e31dcc98a6ebf4 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 17:06:12 +0000 Subject: [PATCH 30/78] Update unit tests --- tests/unit/common/data/dataloaders_test.py | 46 ++++++++++++++ .../dataset_input_metadata_translator_test.py | 63 +++++++++++++++++++ tests/unit/types_tests/graph_test.py | 26 +++++++- 3 files changed, 134 insertions(+), 1 deletion(-) diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 8911ef982..3bfaff851 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -309,8 +309,53 @@ def test_load_as_torch_tensors( assert_close(loaded.features, expected_feature_tensor) + self.assertIsNone(loaded.quantized_features) + assert_close(loaded.labels, expected_label_tensor) + def test_load_as_torch_tensors_decodes_packed_node_features(self) -> None: + packed_feature_values = [[1, 2], [254, 255]] + with tf.io.TFRecordWriter(str(self.data_dir / "packed.tfrecord")) as writer: + for node_id, packed_features in enumerate(packed_feature_values): + writer.write( + tf.train.Example( + features=tf.train.Features( + feature={ + "node_id": tf.train.Feature( + int64_list=tf.train.Int64List(value=[node_id]) + ), + "packed_features": tf.train.Feature( + bytes_list=tf.train.BytesList( + value=[bytes(packed_features)] + ) + ), + } + ) + ).SerializeToString() + ) + + loaded = TFRecordDataLoader(rank=0, world_size=1).load_as_torch_tensors( + serialized_tf_record_info=SerializedTFRecordInfo( + tfrecord_uri_prefix=UriFactory.create_uri(self.data_dir), + feature_spec={"node_id": tf.io.FixedLenFeature([], tf.int64)}, + feature_keys=[], + feature_dim=0, + entity_key="node_id", + packed_feature_key="packed_features", + packed_feature_dim=2, + tfrecord_uri_pattern="packed.tfrecord", + ), + tf_dataset_options=TFDatasetOptions(deterministic=True), + ) + + assert_close(loaded.ids, torch.tensor([0, 1])) + self.assertIsNone(loaded.features) + assert_close( + loaded.quantized_features, + torch.tensor(packed_feature_values, dtype=torch.uint8), + ) + self.assertIsNone(loaded.labels) + def test_build_dataset_for_uris(self): dataset = TFRecordDataLoader._build_dataset_for_uris( uris=[UriFactory.create_uri(self.data_dir / "100.tfrecord")], @@ -410,6 +455,7 @@ def test_load_empty_directory( assert_close(loaded.ids, expected_node_ids) assert_close(loaded.features, expected_features) + self.assertIsNone(loaded.quantized_features) assert_close(loaded.labels, expected_label_tensor) @parameterized.expand( diff --git a/tests/unit/distributed/dataset_input_metadata_translator_test.py b/tests/unit/distributed/dataset_input_metadata_translator_test.py index 49156166f..8e37f62c4 100644 --- a/tests/unit/distributed/dataset_input_metadata_translator_test.py +++ b/tests/unit/distributed/dataset_input_metadata_translator_test.py @@ -9,6 +9,9 @@ ) from gigl.src.common.types.graph_data import EdgeType, NodeType from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper +from gigl.src.common.types.pb_wrappers.preprocessed_metadata import ( + PreprocessedMetadataPbWrapper, +) from gigl.src.mocking.lib.mocked_dataset_resources import MockedDatasetInfo from gigl.src.mocking.lib.versioning import ( MockedDatasetArtifactMetadata, @@ -20,6 +23,8 @@ CORA_USER_DEFINED_NODE_ANCHOR_MOCKED_DATASET_INFO, DBLP_GRAPH_NODE_ANCHOR_MOCKED_DATASET_INFO, ) +from gigl.types.graph import FeatureQuantizationMetadata +from snapchat.research.gbml import preprocessed_metadata_pb2 from tests.test_assets.test_case import TestCase @@ -65,6 +70,64 @@ def _assert_data_type_correctness( else: self.assertNotIsInstance(entity_info, abc.Mapping) + def test_translates_quantized_node_metadata(self) -> None: + mocked_dataset_artifact_metadata = self._name_to_mocked_dataset_map[ + CORA_NODE_CLASSIFICATION_MOCKED_DATASET_INFO.name + ] + gbml_config_pb_wrapper = ( + GbmlConfigPbWrapper.get_gbml_config_pb_wrapper_from_uri( + gbml_config_uri=mocked_dataset_artifact_metadata.frozen_gbml_config_uri + ) + ) + preprocessed_metadata_pb = preprocessed_metadata_pb2.PreprocessedMetadata() + preprocessed_metadata_pb.CopyFrom( + gbml_config_pb_wrapper.preprocessed_metadata_pb_wrapper.preprocessed_metadata_pb + ) + condensed_node_type = gbml_config_pb_wrapper.graph_metadata_pb_wrapper.homogeneous_condensed_node_type + node_metadata = ( + preprocessed_metadata_pb.condensed_node_type_to_preprocessed_metadata[ + condensed_node_type + ] + ) + node_metadata.quantized_feature_metadata.CopyFrom( + preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( + packed_feature_key="packed_features", + quantized_feature_indices=[0, 1], + multi_bit_state=preprocessed_metadata_pb2.PreprocessedMetadata.MultiBitQuantizationState( + bits=2, + clip_min=-1.0, + clip_max=1.0, + ), + ) + ) + + serialized_metadata = convert_pb_to_serialized_graph_metadata( + preprocessed_metadata_pb_wrapper=PreprocessedMetadataPbWrapper( + preprocessed_metadata_pb + ), + graph_metadata_pb_wrapper=gbml_config_pb_wrapper.graph_metadata_pb_wrapper, + ) + + self.assertEqual( + serialized_metadata.node_quantization_metadata, + FeatureQuantizationMetadata( + bits=2, + feature_dim=node_metadata.feature_dim, + quantized_feature_indices=(0, 1), + clip_min=-1.0, + clip_max=1.0, + ), + ) + node_entity_info = serialized_metadata.node_entity_info + self.assertIsInstance(node_entity_info, SerializedTFRecordInfo) + assert isinstance(node_entity_info, SerializedTFRecordInfo) + self.assertEqual( + node_entity_info.feature_keys, list(node_metadata.feature_keys)[2:] + ) + self.assertEqual(node_entity_info.feature_dim, node_metadata.feature_dim - 2) + self.assertEqual(node_entity_info.packed_feature_key, "packed_features") + self.assertEqual(node_entity_info.packed_feature_dim, 1) + @parameterized.expand( [ param( diff --git a/tests/unit/types_tests/graph_test.py b/tests/unit/types_tests/graph_test.py index 78493d3d1..9a5cca5ec 100644 --- a/tests/unit/types_tests/graph_test.py +++ b/tests/unit/types_tests/graph_test.py @@ -1,4 +1,4 @@ -from typing import Literal, Union +from typing import Literal, Union, cast import torch from absl.testing import absltest @@ -422,6 +422,30 @@ def test_treat_labels_as_edges_converts_homogeneous_edge_weights(self) -> None: self.assertIn(DEFAULT_HOMOGENEOUS_EDGE_TYPE, edge_weights) torch.testing.assert_close(edge_weights[DEFAULT_HOMOGENEOUS_EDGE_TYPE], weights) + def test_treat_labels_as_edges_converts_quantized_node_features(self) -> None: + quantized_features = torch.tensor([[1], [2], [3]], dtype=torch.uint8) + graph_tensors = LoadedGraphTensors( + node_ids=torch.tensor([0, 1, 2]), + node_features=None, + node_labels=None, + edge_index=torch.tensor([[0, 1], [1, 2]]), + edge_features=None, + positive_label=torch.tensor([[0], [2]]), + negative_label=None, + node_quantized_features=quantized_features, + ) + + graph_tensors.treat_labels_as_edges(edge_dir="out") + + heterogeneous_quantized_features = cast( + dict[NodeType, torch.Tensor], graph_tensors.node_quantized_features + ) + self.assertIsInstance(heterogeneous_quantized_features, dict) + torch.testing.assert_close( + heterogeneous_quantized_features[DEFAULT_HOMOGENEOUS_NODE_TYPE], + quantized_features, + ) + def test_select_label_edge_types(self): message_passing_edge_type = DEFAULT_HOMOGENEOUS_EDGE_TYPE edge_types = [ From 6ce592255e4c52377be8aa1fa05390f882848033 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 17:12:10 +0000 Subject: [PATCH 31/78] Better comments --- .../distributed/utils/serialized_graph_metadata_translator.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index ac2bb447c..813a2c49c 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -52,14 +52,14 @@ def _build_serialized_tfrecord_entity_info( ) feature_keys_to_load: list[str] = [] - # Keep only wholly unquantized fields: TFRecord decoding cannot split one - # field between float and packed storage. + # Identify non-quantized feature fields to load from float storage. for key in preprocessed_metadata.feature_keys: key_indices = set(range(*feature_index[key])) quantized_key_indices = key_indices.intersection(quantized_indices) if not quantized_key_indices: feature_keys_to_load.append(key) elif quantized_key_indices != key_indices: + # TFRecord decoding cannot split a field between float and packed storage. raise ValueError(f"Partial quantization not supported for {key}") feature_dim = quantization_metadata.raw_feature_dim else: From be16daef2622eaa0f3f7e46a43c3acb6b0ac2e68 Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 17:47:18 +0000 Subject: [PATCH 32/78] Expand multi-line decleration with typedefs --- .../utils/serialized_graph_metadata_translator.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index 813a2c49c..b5164d1f9 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -41,6 +41,7 @@ def _build_serialized_tfrecord_entity_info( if not isinstance( preprocessed_metadata, PreprocessedMetadata.NodeMetadataOutput ): + # TODO(quantization): Support edge feature quantization. raise ValueError("Quantization is supported only for node entities.") packed_feature_key = ( preprocessed_metadata.quantized_feature_metadata.packed_feature_key @@ -97,7 +98,10 @@ def _build_feature_quantization_metadata( ) -> FeatureQuantizationMetadata: state = quantized_metadata.WhichOneof("state") - neg_mean = pos_mean = clip_min = clip_max = None + neg_mean: Optional[float] = None + pos_mean: Optional[float] = None + clip_min: Optional[float] = None + clip_max: Optional[float] = None if state == "single_bit_state": bits = 1 neg_mean = quantized_metadata.single_bit_state.neg_mean From 5bf105b3f81abfc2ade84c0f18f3ffa72bd466fc Mon Sep 17 00:00:00 2001 From: jchmura Date: Mon, 10 Aug 2026 17:52:42 +0000 Subject: [PATCH 33/78] Use len() check instead of full set equality for partial quantization safe guard --- gigl/distributed/utils/serialized_graph_metadata_translator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index b5164d1f9..7a59f3140 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -59,7 +59,7 @@ def _build_serialized_tfrecord_entity_info( quantized_key_indices = key_indices.intersection(quantized_indices) if not quantized_key_indices: feature_keys_to_load.append(key) - elif quantized_key_indices != key_indices: + elif len(quantized_key_indices) != len(key_indices): # TFRecord decoding cannot split a field between float and packed storage. raise ValueError(f"Partial quantization not supported for {key}") feature_dim = quantization_metadata.raw_feature_dim From 29aabbbe885a04534820245fd3696c3aebae992a Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 15:09:05 +0000 Subject: [PATCH 34/78] Run format --- gigl/common/data/load_torch_tensors.py | 1 + .../utils/serialized_graph_metadata_translator.py | 12 ++++++++---- 2 files changed, 9 insertions(+), 4 deletions(-) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 7b3545e33..3a1174888 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -203,6 +203,7 @@ def _data_loading_process( 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" ) diff --git a/gigl/distributed/utils/serialized_graph_metadata_translator.py b/gigl/distributed/utils/serialized_graph_metadata_translator.py index 7a59f3140..36ad31c52 100644 --- a/gigl/distributed/utils/serialized_graph_metadata_translator.py +++ b/gigl/distributed/utils/serialized_graph_metadata_translator.py @@ -42,7 +42,9 @@ def _build_serialized_tfrecord_entity_info( preprocessed_metadata, PreprocessedMetadata.NodeMetadataOutput ): # TODO(quantization): Support edge feature quantization. - raise ValueError("Quantization is supported only for node entities.") + raise NotImplementedError( + "Feature quantization is not supported for edge entities." + ) packed_feature_key = ( preprocessed_metadata.quantized_feature_metadata.packed_feature_key ) @@ -165,9 +167,11 @@ def convert_pb_to_serialized_graph_metadata( node_key = node_metadata.node_id_key if node_metadata.HasField("quantized_feature_metadata"): - node_quantization_metadata[node_type] = _build_feature_quantization_metadata( - quantized_metadata=node_metadata.quantized_feature_metadata, - feature_dim=node_metadata.feature_dim, + node_quantization_metadata[node_type] = ( + _build_feature_quantization_metadata( + quantized_metadata=node_metadata.quantized_feature_metadata, + feature_dim=node_metadata.feature_dim, + ) ) node_entity_info[node_type] = _build_serialized_tfrecord_entity_info( From 6270a0b9bcc360c68fd2220cd3e927afd22b4b81 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 15:36:13 +0000 Subject: [PATCH 35/78] Lazy log format --- gigl/distributed/dist_ablp_neighborloader.py | 6 ++++-- gigl/distributed/distributed_neighborloader.py | 6 ++++-- gigl/distributed/utils/neighborloader.py | 4 +--- 3 files changed, 9 insertions(+), 7 deletions(-) diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index 1b1f06738..6b35f0e03 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -933,7 +933,8 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: data = super()._collate_fn(stripped_msg) base_collate_time = time.perf_counter() - base_collate_start_time logger.debug( - f"Distributed ABLPNeighborLoader GLT base collate time: {base_collate_time:.3f}s" + "Distributed ABLPNeighborLoader GLT base collate time: %.3fs", + base_collate_time, ) data = set_missing_features( @@ -987,6 +988,7 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: collate_time = time.perf_counter() - collate_start_time logger.debug( - f"Distributed ABLPNeighborLoader end-to-end collate time: {collate_time:.3f}s" + "Distributed ABLPNeighborLoader end-to-end collate time: %.3fs", + collate_time, ) return data diff --git a/gigl/distributed/distributed_neighborloader.py b/gigl/distributed/distributed_neighborloader.py index d6a8d9753..7ed34de25 100644 --- a/gigl/distributed/distributed_neighborloader.py +++ b/gigl/distributed/distributed_neighborloader.py @@ -552,7 +552,8 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: data = super()._collate_fn(stripped_msg) base_collate_time = time.perf_counter() - base_collate_start_time logger.debug( - f"Distributed NeighborLoader GLT base collate time: {base_collate_time:.3f}s" + "Distributed NeighborLoader GLT base collate time: %.3fs", + base_collate_time, ) data = set_missing_features( data=data, @@ -579,6 +580,7 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: collate_time = time.perf_counter() - collate_start_time logger.debug( - f"Distributed NeighborLoader end-to-end collate time: {collate_time:.3f}s" + "Distributed NeighborLoader end-to-end collate time: %.3fs", + collate_time, ) return data diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index b40eb471a..576b5564b 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -404,9 +404,7 @@ def materialize( materialize(data[node_type], packed_features, quantization_metadata) materialize_time = time.perf_counter() - materialize_start_time - logger.debug( - f"Quantized node feature materialization time: {materialize_time:.3f}s" - ) + logger.debug("Quantized node feature materialization time: %.3fs", materialize_time) return data, metadata From 99ef03d394c9919832cbecee54fbebf633415cdf Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 15:50:37 +0000 Subject: [PATCH 36/78] Update graph store dist server contract to match distdataset --- gigl/distributed/graph_store/dist_server.py | 6 ++++++ .../graph_store/remote_dist_dataset.py | 16 ++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/gigl/distributed/graph_store/dist_server.py b/gigl/distributed/graph_store/dist_server.py index 0a92d959d..104c37750 100644 --- a/gigl/distributed/graph_store/dist_server.py +++ b/gigl/distributed/graph_store/dist_server.py @@ -410,6 +410,12 @@ def get_node_feature_info( """ return self.dataset.node_feature_info + def get_node_quantized_feature_info( + self, + ) -> Union[FeatureInfo, dict[NodeType, FeatureInfo], None]: + """Get packed node feature information from the dataset.""" + return self.dataset.node_quantized_feature_info + def get_node_quantization_metadata( self, ) -> Union[ diff --git a/gigl/distributed/graph_store/remote_dist_dataset.py b/gigl/distributed/graph_store/remote_dist_dataset.py index ef3fc24b0..dbf67c86c 100644 --- a/gigl/distributed/graph_store/remote_dist_dataset.py +++ b/gigl/distributed/graph_store/remote_dist_dataset.py @@ -67,6 +67,22 @@ def fetch_node_feature_info( DistServer.get_node_feature_info, ) + def fetch_node_quantized_feature_info( + self, + ) -> Union[FeatureInfo, dict[NodeType, FeatureInfo], None]: + """Fetch packed node feature information from the registered dataset.""" + return request_server( + 0, + DistServer.get_node_quantized_feature_info, + ) + + @property + def node_quantized_feature_info( + self, + ) -> Union[FeatureInfo, dict[NodeType, FeatureInfo], None]: + """Return packed node feature information from the storage cluster.""" + return self.fetch_node_quantized_feature_info() + def fetch_node_quantization_metadata( self, ) -> Union[ From 2b4204299709f7ef07c6201106bd05d7defb144a Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 15:51:41 +0000 Subject: [PATCH 37/78] Add unit tests --- .../distributed/distributed_dataset_test.py | 44 ++++++++++++ .../distributed/utils/neighborloader_test.py | 70 ++++++++++++++++++- 2 files changed, 113 insertions(+), 1 deletion(-) diff --git a/tests/unit/distributed/distributed_dataset_test.py b/tests/unit/distributed/distributed_dataset_test.py index 692f7816b..1f109aa67 100644 --- a/tests/unit/distributed/distributed_dataset_test.py +++ b/tests/unit/distributed/distributed_dataset_test.py @@ -39,6 +39,7 @@ DEFAULT_HOMOGENEOUS_EDGE_TYPE, FeatureInfo, FeaturePartitionData, + FeatureQuantizationMetadata, GraphPartitionData, PartitionOutput, ) @@ -569,6 +570,49 @@ def test_building_homogeneous_dataset_preserves_node_features_and_labels(self): dataset.node_features.feature_tensor, torch.zeros(10, 2) ) + def test_building_dataset_preserves_packed_node_features_and_metadata(self) -> None: + packed_features = torch.tensor([[48], [144]], dtype=torch.uint8) + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + partition_output = PartitionOutput( + node_partition_book=torch.zeros(2), + edge_partition_book=torch.zeros(1), + partitioned_edge_index=GraphPartitionData( + edge_index=torch.tensor([[0], [1]]), edge_ids=None + ), + partitioned_node_features=None, + partitioned_edge_features=None, + partitioned_positive_labels=None, + partitioned_negative_labels=None, + partitioned_node_labels=None, + partitioned_node_quantized_features=FeaturePartitionData( + feats=packed_features, ids=torch.arange(2) + ), + ) + + dataset = DistDataset( + rank=0, + world_size=1, + edge_dir="out", + node_quantization_metadata=quantization_metadata, + ) + dataset.build(partition_output=partition_output) + + assert isinstance(dataset.node_quantized_features, Feature) + self.assert_tensor_equality( + dataset.node_quantized_features.feature_tensor, packed_features + ) + self.assertEqual( + dataset.node_quantized_feature_info, + FeatureInfo(dim=1, dtype=torch.uint8), + ) + self.assertEqual(dataset.node_quantization_metadata, quantization_metadata) + def test_building_heterogeneous_dataset_preserves_node_features_and_labels(self): partition_output = PartitionOutput( node_partition_book={ diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index 00c49fb55..ca7faf9e2 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -8,6 +8,7 @@ from gigl.distributed.sampler import ( NEGATIVE_LABEL_METADATA_KEY, + NODE_PACKED_FEATURES_METADATA_KEY, POSITIVE_LABEL_METADATA_KEY, ) from gigl.distributed.utils.neighborloader import ( @@ -15,13 +16,18 @@ extract_edge_type_metadata, extract_metadata, labeled_to_homogeneous, + materialize_quantized_node_features, patch_fanout_for_sampling, set_missing_features, shard_nodes_by_process, strip_label_edges, strip_non_ppr_edge_types, ) -from gigl.types.graph import FeatureInfo, message_passing_to_positive_label +from gigl.types.graph import ( + FeatureInfo, + FeatureQuantizationMetadata, + message_passing_to_positive_label, +) from tests.test_assets.test_case import TestCase _U2U_EDGE_TYPE = ("user", "to", "user") @@ -65,6 +71,68 @@ def test_shard_nodes_by_process( ) self.assert_tensor_equality(sharded_tensor, expected_sharded_tensor) + def test_materialize_quantized_node_features_reconstructs_feature_order( + self, + ) -> None: + data = Data(x=torch.tensor([[10.0, 20.0], [30.0, 40.0]])) + metadata = { + NODE_PACKED_FEATURES_METADATA_KEY: torch.tensor( + [[48], [144]], dtype=torch.uint8 + ), + "request_id": torch.tensor([7]), + } + + materialized_data, remaining_metadata = materialize_quantized_node_features( + data=data, + metadata=metadata, + node_quantization_metadata=FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ), + ) + + self.assert_tensor_equality( + materialized_data.x, + torch.tensor([[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]]), + ) + self.assertEqual(set(remaining_metadata), {"request_id"}) + + def test_materialize_quantized_node_features_uses_per_node_type_metadata( + self, + ) -> None: + data = HeteroData() + metadata = { + f"{NODE_PACKED_FEATURES_METADATA_KEY}.user": torch.tensor( + [[48]], dtype=torch.uint8 + ), + "request_id": torch.tensor([7]), + } + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_min=0.0, + clip_max=3.0, + ) + + materialized_data, remaining_metadata = materialize_quantized_node_features( + data=data, + metadata=metadata, + node_quantization_metadata={ + "user": quantization_metadata, + "item": quantization_metadata, + }, + ) + + self.assert_tensor_equality( + materialized_data["user"].x, torch.tensor([[0.0, 3.0]]) + ) + self.assertFalse(hasattr(materialized_data["item"], "x")) + self.assertEqual(set(remaining_metadata), {"request_id"}) + @parameterized.expand( [ param( From 779a947e170c7112e5f4fa52d83d5a7a1dcfc0ae Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 16:12:33 +0000 Subject: [PATCH 38/78] Update tests --- .../run_distributed_partitioner.py | 26 +++++++- tests/unit/distributed/dist_server_test.py | 42 +++++++++++++ .../distributed/distributed_dataset_test.py | 6 +- .../distributed_neighborloader_test.py | 59 +++++++++++++++++++ .../distributed_partitioner_test.py | 32 +++++++++- .../distributed/utils/neighborloader_test.py | 2 + 6 files changed, 163 insertions(+), 4 deletions(-) diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 3bbc3406c..d67121d72 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -1,5 +1,5 @@ from enum import Enum -from typing import Optional, Type, Union +from typing import Optional, Type, Union, cast import torch from graphlearn_torch.distributed import init_rpc, init_worker_group @@ -34,6 +34,7 @@ def run_distributed_partitioner( master_port: int, input_data_strategy: InputDataStrategy, partitioner_class: Type[DistPartitioner], + include_quantized_features: bool = False, rank_to_edge_weights: Optional[ dict[int, Union[torch.Tensor, dict[EdgeType, torch.Tensor]]] ] = None, @@ -50,6 +51,7 @@ def run_distributed_partitioner( master_port (int): Master port for initializing rpc for partitioning input_data_strategy (InputDataStrategy): Strategy for registering inputs to the partitioner partitioner_class (Type[DistPartitioner]): The class to use for partitioning + include_quantized_features: Whether to register one packed uint8 feature per node. rank_to_edge_weights (Optional[dict[int, Union[torch.Tensor, dict[EdgeType, torch.Tensor]]]]): Optional mapping of rank to 1D edge weight tensor (or dict per EdgeType for heterogeneous). Only supported with REGISTER_ALL_ENTITIES_SEPARATELY strategy. """ @@ -61,6 +63,9 @@ def run_distributed_partitioner( positive_labels: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] negative_labels: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] node_labels: Union[torch.Tensor, dict[NodeType, torch.Tensor]] + node_quantized_features: Optional[ + Union[torch.Tensor, dict[NodeType, torch.Tensor]] + ] = None if not is_heterogeneous: node_ids = input_graph.node_ids[USER_NODE_TYPE] @@ -79,6 +84,16 @@ def run_distributed_partitioner( negative_labels = input_graph.negative_labels node_labels = input_graph.node_labels + if include_quantized_features: + if isinstance(node_ids, dict): + node_ids_by_type = cast(dict[NodeType, torch.Tensor], node_ids) + node_quantized_features = { + node_type: node_type_ids.to(torch.uint8).unsqueeze(1) + for node_type, node_type_ids in node_ids_by_type.items() + } + else: + node_quantized_features = node_ids.to(torch.uint8).unsqueeze(1) + partition_output: PartitionOutput init_worker_group(world_size=MOCKED_NUM_PARTITIONS, rank=rank) @@ -115,11 +130,17 @@ def run_distributed_partitioner( ) dist_partitioner.register_node_features(node_features=node_features) + if node_quantized_features is not None: + dist_partitioner.register_node_quantized_features( + node_quantized_features=node_quantized_features + ) dist_partitioner.register_node_labels(node_labels=node_labels) del node_labels del node_features + del node_quantized_features ( output_node_features, + output_node_quantized_features, output_node_labels, ) = dist_partitioner.partition_node_features_and_labels( node_partition_book=output_node_partition_book @@ -146,6 +167,7 @@ def run_distributed_partitioner( edge_partition_book=output_edge_partition_book, partitioned_edge_index=output_edge_index, partitioned_node_features=output_node_features, + partitioned_node_quantized_features=output_node_quantized_features, partitioned_node_labels=output_node_labels, partitioned_edge_features=output_edge_features, partitioned_positive_labels=output_positive_labels, @@ -186,6 +208,7 @@ def run_distributed_partitioner( should_assign_edges_by_src_node=should_assign_edges_by_src_node, node_ids=node_ids, node_features=node_features, + node_quantized_features=node_quantized_features, edge_index=edge_index, edge_features=edge_features, positive_labels=positive_labels, @@ -196,6 +219,7 @@ def run_distributed_partitioner( del ( node_ids, node_features, + node_quantized_features, edge_index, edge_features, positive_labels, diff --git a/tests/unit/distributed/dist_server_test.py b/tests/unit/distributed/dist_server_test.py index 3512ba276..64204ea83 100644 --- a/tests/unit/distributed/dist_server_test.py +++ b/tests/unit/distributed/dist_server_test.py @@ -5,6 +5,7 @@ from absl.testing import absltest from graphlearn_torch.sampler import NodeSamplerInput, SamplingConfig, SamplingType +from gigl.distributed.dist_dataset import DistDataset from gigl.distributed.graph_store import dist_server from gigl.distributed.graph_store.messages import ( FetchABLPInputRequest, @@ -12,8 +13,10 @@ InitSamplingBackendRequest, RegisterBackendRequest, ) +from gigl.distributed.graph_store.remote_dist_dataset import RemoteDistDataset from gigl.distributed.graph_store.sharding import ServerSlice from gigl.src.common.types.graph_data import Relation +from gigl.types.graph import FeatureQuantizationMetadata from tests.test_assets.distributed.test_dataset import ( DEFAULT_HETEROGENEOUS_EDGE_INDICES, DEFAULT_HOMOGENEOUS_EDGE_INDEX, @@ -62,6 +65,45 @@ def test_get_node_feature_info_with_homogeneous_dataset(self) -> None: # Verify it returns the correct feature info self.assertIsNone(node_feature_info) + def test_get_node_quantization_metadata(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_min=0.0, + clip_max=3.0, + ) + dataset = DistDataset( + rank=0, + world_size=1, + edge_dir="out", + node_quantization_metadata=metadata, + ) + + server = dist_server.DistServer(dataset) + + self.assertEqual(server.get_node_quantization_metadata(), metadata) + + def test_remote_dataset_fetches_node_quantization_metadata(self) -> None: + metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_min=0.0, + clip_max=3.0, + ) + with patch( + "gigl.distributed.graph_store.remote_dist_dataset.request_server", + return_value=metadata, + ) as request_server: + remote_dataset = RemoteDistDataset(cluster_info=MagicMock(), local_rank=0) + + self.assertEqual(remote_dataset.node_quantization_metadata, metadata) + + request_server.assert_called_once_with( + 0, dist_server.DistServer.get_node_quantization_metadata + ) + def test_get_edge_feature_info_with_heterogeneous_dataset(self) -> None: """Test get_edge_feature_info with a heterogeneous dataset.""" dataset = create_heterogeneous_dataset( diff --git a/tests/unit/distributed/distributed_dataset_test.py b/tests/unit/distributed/distributed_dataset_test.py index 1f109aa67..7c98d34dd 100644 --- a/tests/unit/distributed/distributed_dataset_test.py +++ b/tests/unit/distributed/distributed_dataset_test.py @@ -591,7 +591,7 @@ def test_building_dataset_preserves_packed_node_features_and_metadata(self) -> N partitioned_negative_labels=None, partitioned_node_labels=None, partitioned_node_quantized_features=FeaturePartitionData( - feats=packed_features, ids=torch.arange(2) + feats=packed_features, ids=torch.tensor([3, 7]) ), ) @@ -607,6 +607,10 @@ def test_building_dataset_preserves_packed_node_features_and_metadata(self) -> N self.assert_tensor_equality( dataset.node_quantized_features.feature_tensor, packed_features ) + self.assert_tensor_equality( + dataset.node_quantized_features[torch.tensor([7, 3])], + packed_features.flip(0), + ) self.assertEqual( dataset.node_quantized_feature_info, FeatureInfo(dim=1, dtype=torch.uint8), diff --git a/tests/unit/distributed/distributed_neighborloader_test.py b/tests/unit/distributed/distributed_neighborloader_test.py index 02cc41a11..c9caee50f 100644 --- a/tests/unit/distributed/distributed_neighborloader_test.py +++ b/tests/unit/distributed/distributed_neighborloader_test.py @@ -30,6 +30,7 @@ DEFAULT_HOMOGENEOUS_EDGE_TYPE, FeatureInfo, FeaturePartitionData, + FeatureQuantizationMetadata, GraphPartitionData, PartitionOutput, message_passing_to_negative_label, @@ -426,6 +427,26 @@ def _run_featureless_edge_ids_absent( shutdown_rpc() +def _run_quantized_feature_neighbor_loader(_: int, dataset: DistDataset) -> None: + create_test_process_group() + loader = DistNeighborLoader( + dataset=dataset, + input_nodes=torch.tensor([0, 1]), + num_neighbors=[0], + batch_size=2, + pin_memory_device=torch.device("cpu"), + ) + + expected_features = torch.tensor([[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]]) + batch_count = 0 + for batch in loader: + assert isinstance(batch, Data) + assert_tensor_equality(batch.x, expected_features[batch.node]) + batch_count += 1 + assert batch_count == 1 + shutdown_rpc() + + class DistributedNeighborLoaderTest(TestCase): def setUp(self): super().setUp() @@ -777,6 +798,44 @@ def test_isolated_homogeneous_neighbor_loader( args=(dataset, 18), ) + def test_distributed_neighbor_loader_materializes_quantized_node_features( + self, + ) -> None: + partition_output = PartitionOutput( + node_partition_book=torch.zeros(2), + edge_partition_book=torch.zeros(2), + partitioned_edge_index=GraphPartitionData( + edge_index=torch.tensor([[0, 1], [1, 0]]), edge_ids=None + ), + partitioned_node_features=FeaturePartitionData( + feats=torch.tensor([[10.0, 20.0], [30.0, 40.0]]), + ids=torch.arange(2), + ), + partitioned_node_quantized_features=FeaturePartitionData( + feats=torch.tensor([[48], [144]], dtype=torch.uint8), + ids=torch.arange(2), + ), + partitioned_edge_features=None, + partitioned_positive_labels=None, + partitioned_negative_labels=None, + partitioned_node_labels=None, + ) + dataset = DistDataset( + rank=0, + world_size=1, + edge_dir="out", + node_quantization_metadata=FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ), + ) + dataset.build(partition_output=partition_output) + + mp.spawn(fn=_run_quantized_feature_neighbor_loader, args=(dataset,)) + @parameterized.expand( [ param( diff --git a/tests/unit/distributed/distributed_partitioner_test.py b/tests/unit/distributed/distributed_partitioner_test.py index 0b02b8e2b..6860d0bdd 100644 --- a/tests/unit/distributed/distributed_partitioner_test.py +++ b/tests/unit/distributed/distributed_partitioner_test.py @@ -280,7 +280,7 @@ def _assert_node_data_outputs( ], expected_node_types: list[NodeType], expected_edge_types: list[EdgeType], - entity_name: Literal["features", "labels"], + entity_name: Literal["features", "quantized_features", "labels"], ) -> None: """ Checks correctness for node feature or label outputs of partitioning @@ -384,7 +384,7 @@ def _assert_node_data_outputs( tensor_a=node_data.feats[idx], tensor_b=torch.tensor([n_id], dtype=torch.int64), ) - else: + elif entity_name == "features": # We expect the shape of the node features to be equal to the expected node feature dimension self.assertEqual( node_data.feats.size(1), @@ -401,6 +401,14 @@ def _assert_node_data_outputs( * n_id * 0.1, ) + else: + self.assertEqual(node_data.feats.dtype, torch.uint8) + self.assertEqual(node_data.feats.size(1), 1) + for idx, n_id in enumerate(node_data_ids): + self.assert_tensor_equality( + tensor_a=node_data.feats[idx], + tensor_b=torch.tensor([n_id], dtype=torch.uint8), + ) def _assert_edge_feature_outputs( self, @@ -621,6 +629,7 @@ def _assert_label_outputs( should_assign_edges_by_src_node=True, partitioner_class=DistPartitioner, expected_pb_dtype=torch.uint8, + include_quantized_features=True, ), param( "Homogeneous Partitioning By Dest Node- Register All Entites together through Constructor", @@ -661,6 +670,7 @@ def _assert_label_outputs( should_assign_edges_by_src_node=True, partitioner_class=DistRangePartitioner, expected_pb_dtype=torch.int64, + include_quantized_features=True, ), param( "Heterogeneous Partitioning By Source Node - Range Partitioning", @@ -688,6 +698,7 @@ def test_partitioning_correctness( should_assign_edges_by_src_node: bool, partitioner_class: Type[DistPartitioner], expected_pb_dtype: torch.dtype, + include_quantized_features: bool = False, ) -> None: """ Tests partitioning functionality and correctness on mocked inputs @@ -717,6 +728,7 @@ def test_partitioning_correctness( master_port, input_data_strategy, partitioner_class, + include_quantized_features, ), nprocs=MOCKED_NUM_PARTITIONS, join=True, @@ -822,6 +834,22 @@ def test_partitioning_correctness( entity_name="features", ) + if include_quantized_features: + assert ( + partition_output.partitioned_node_quantized_features is not None + ) + self._assert_node_data_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_node_data=partition_output.partitioned_node_quantized_features, + expected_node_types=MOCKED_HETEROGENEOUS_NODE_TYPES, + expected_edge_types=MOCKED_HETEROGENEOUS_EDGE_TYPES, + entity_name="quantized_features", + ) + if isinstance(partition_output.partitioned_node_features, abc.Mapping): for ( node_type, diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index ca7faf9e2..083dd6fd8 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -99,6 +99,7 @@ def test_materialize_quantized_node_features_reconstructs_feature_order( torch.tensor([[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]]), ) self.assertEqual(set(remaining_metadata), {"request_id"}) + self.assert_tensor_equality(remaining_metadata["request_id"], torch.tensor([7])) def test_materialize_quantized_node_features_uses_per_node_type_metadata( self, @@ -132,6 +133,7 @@ def test_materialize_quantized_node_features_uses_per_node_type_metadata( ) self.assertFalse(hasattr(materialized_data["item"], "x")) self.assertEqual(set(remaining_metadata), {"request_id"}) + self.assert_tensor_equality(remaining_metadata["request_id"], torch.tensor([7])) @parameterized.expand( [ From bdc3e0fca566ade08ad44e6a2b3fafa183b67944 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 16:57:47 +0000 Subject: [PATCH 39/78] Fix implicit default homogenous key mismatch in quant metadata when ablp treat labels as edges --- gigl/distributed/dataset_factory.py | 51 +++++++- .../dist_ablp_neighborloader_test.py | 115 +++++++++++++++++- 2 files changed, 160 insertions(+), 6 deletions(-) diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 1a5f859b5..a49bc94dc 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -41,7 +41,7 @@ from gigl.distributed.utils.serialized_graph_metadata_translator import ( convert_pb_to_serialized_graph_metadata, ) -from gigl.src.common.types.graph_data import EdgeType +from gigl.src.common.types.graph_data import EdgeType, NodeType from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper from gigl.src.common.types.pb_wrappers.task_metadata import TaskMetadataType from gigl.utils.data_splitters import ( @@ -52,10 +52,46 @@ get_max_labels_per_anchor_node_from_runtime_args, select_ssl_positive_label_edges, ) +from gigl.types.graph import ( + DEFAULT_HOMOGENEOUS_NODE_TYPE, + FeatureQuantizationMetadata, +) logger = Logger() +def _normalize_node_quantization_metadata( + node_quantized_features: Optional[ + Union[torch.Tensor, dict[NodeType, torch.Tensor]] + ], + node_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] + ], + labels_converted_to_edges: bool, +) -> Optional[Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]]]: + """Align node quantization metadata with packed feature representation.""" + if ( + labels_converted_to_edges + and node_quantization_metadata is not None + and not isinstance(node_quantization_metadata, Mapping) + ): + node_quantization_metadata = { + DEFAULT_HOMOGENEOUS_NODE_TYPE: node_quantization_metadata + } + + if node_quantized_features is None or node_quantization_metadata is None: + return node_quantization_metadata + + packed_features_are_typed = isinstance(node_quantized_features, Mapping) + metadata_is_typed = isinstance(node_quantization_metadata, Mapping) + if packed_features_are_typed != metadata_is_typed: + raise ValueError( + "Packed node features and node quantization metadata must both be " + "scalar or both be keyed by node type." + ) + return node_quantization_metadata + + @tf_on_cpu def _load_and_build_partitioned_dataset( serialized_graph_metadata: SerializedGraphMetadata, @@ -151,12 +187,19 @@ def _load_and_build_partitioned_dataset( loaded_graph_tensors.positive_label = positive_label_edges - if ( + labels_converted_to_edges = ( isinstance(splitter, NodeAnchorLinkSplitter) and splitter.should_convert_labels_to_edges - ): + ) + if labels_converted_to_edges: loaded_graph_tensors.treat_labels_as_edges(edge_dir=edge_dir) + node_quantization_metadata = _normalize_node_quantization_metadata( + node_quantized_features=loaded_graph_tensors.node_quantized_features, + node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, + labels_converted_to_edges=labels_converted_to_edges, + ) + should_assign_edges_by_src_node: bool = False if edge_dir == "in" else True if partitioner_class is None: @@ -226,7 +269,7 @@ def _load_and_build_partitioned_dataset( rank=rank, world_size=world_size, edge_dir=edge_dir, - node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, + node_quantization_metadata=node_quantization_metadata, ) dataset.build( diff --git a/tests/unit/distributed/dist_ablp_neighborloader_test.py b/tests/unit/distributed/dist_ablp_neighborloader_test.py index 35607301f..df26164d2 100644 --- a/tests/unit/distributed/dist_ablp_neighborloader_test.py +++ b/tests/unit/distributed/dist_ablp_neighborloader_test.py @@ -1,6 +1,6 @@ import unittest from collections import defaultdict -from typing import Literal, Optional, Union +from typing import Literal, Optional, Union, cast import torch import torch.multiprocessing as mp @@ -10,7 +10,10 @@ from parameterized import param, parameterized from torch_geometric.data import Data, HeteroData -from gigl.distributed.dataset_factory import build_dataset +from gigl.distributed.dataset_factory import ( + _validate_node_quantization_metadata_representation, + build_dataset, +) from gigl.distributed.dist_ablp_neighborloader import DistABLPLoader from gigl.distributed.dist_dataset import DistDataset from gigl.distributed.dist_partitioner import DistPartitioner @@ -28,7 +31,11 @@ ) from gigl.types.graph import ( DEFAULT_HOMOGENEOUS_EDGE_TYPE, + DEFAULT_HOMOGENEOUS_NODE_TYPE, + FeaturePartitionData, + FeatureQuantizationMetadata, GraphPartitionData, + LoadedGraphTensors, PartitionOutput, is_label_edge_type, message_passing_to_negative_label, @@ -196,6 +203,32 @@ def _run_cora_supervised( shutdown_rpc() +def _run_quantized_homogeneous_ablp_loader(_: int, dataset: DistDataset) -> None: + """Assert homogeneous ABLP materializes partial packed features.""" + create_test_process_group() + loader = DistABLPLoader( + dataset=dataset, + num_neighbors=[2], + input_nodes=torch.tensor([0]), + batch_size=1, + pin_memory_device=torch.device("cpu"), + ) + + expected_features = torch.tensor( + [[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0], [1.0, 50.0, 2.0, 60.0]] + ) + batch_count = 0 + for batch in loader: + assert isinstance(batch, Data) + assert_tensor_equality(batch.x, expected_features[batch.node]) + assert _global_pair_set(batch.node, batch.node, batch.y_positive) == [(0, 1)] + assert _global_pair_set(batch.node, batch.node, batch.y_negative) == [(0, 2)] + batch_count += 1 + + assert batch_count == 1 + shutdown_rpc() + + def _run_dblp_supervised( _, dataset: DistDataset, @@ -735,6 +768,84 @@ def test_ablp_dataloader( ), ) + def test_homogeneous_ablp_materializes_quantized_features(self) -> None: + """Promoted metadata supports quantized homogeneous ABLP batches.""" + loaded_graph_tensors = LoadedGraphTensors( + node_ids=torch.arange(3), + node_features=torch.tensor([[10.0, 20.0], [30.0, 40.0], [50.0, 60.0]]), + node_quantized_features=torch.tensor( + [[48], [144], [96]], dtype=torch.uint8 + ), + node_labels=None, + edge_index=torch.tensor([[0, 1], [1, 0]]), + edge_features=None, + positive_label=torch.tensor([[0], [1]]), + negative_label=torch.tensor([[0], [2]]), + ) + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + + loaded_graph_tensors.treat_labels_as_edges(edge_dir="out") + promoted_metadata = {DEFAULT_HOMOGENEOUS_NODE_TYPE: quantization_metadata} + _validate_node_quantization_metadata_representation( + node_quantized_features=loaded_graph_tensors.node_quantized_features, + node_quantization_metadata=promoted_metadata, + ) + + assert isinstance(loaded_graph_tensors.edge_index, dict) + assert isinstance(loaded_graph_tensors.node_features, dict) + assert isinstance(loaded_graph_tensors.node_quantized_features, dict) + edge_index = cast(dict[EdgeType, torch.Tensor], loaded_graph_tensors.edge_index) + node_features = cast( + dict[NodeType, torch.Tensor], loaded_graph_tensors.node_features + ) + node_quantized_features = cast( + dict[NodeType, torch.Tensor], loaded_graph_tensors.node_quantized_features + ) + partition_output = PartitionOutput( + node_partition_book={DEFAULT_HOMOGENEOUS_NODE_TYPE: torch.zeros(3)}, + edge_partition_book={ + edge_type: torch.zeros(3) + for edge_type in edge_index + }, + partitioned_edge_index={ + edge_type: GraphPartitionData( + edge_tensor, torch.arange(edge_tensor.size(1)) + ) + for edge_type, edge_tensor in edge_index.items() + }, + partitioned_node_features={ + DEFAULT_HOMOGENEOUS_NODE_TYPE: FeaturePartitionData( + feats=node_features[DEFAULT_HOMOGENEOUS_NODE_TYPE], + ids=torch.arange(3), + ) + }, + partitioned_node_quantized_features={ + DEFAULT_HOMOGENEOUS_NODE_TYPE: FeaturePartitionData( + feats=node_quantized_features[DEFAULT_HOMOGENEOUS_NODE_TYPE], + ids=torch.arange(3), + ) + }, + partitioned_edge_features=None, + partitioned_negative_labels=None, + partitioned_positive_labels=None, + partitioned_node_labels=None, + ) + dataset = DistDataset( + rank=0, + world_size=1, + edge_dir="out", + node_quantization_metadata=promoted_metadata, + ) + dataset.build(partition_output=partition_output) + + mp.spawn(fn=_run_quantized_homogeneous_ablp_loader, args=(dataset,)) + @parameterized.expand( [ param( From 7bcafc4efc740a7a127c44ec75e93d689b38c664 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 17:45:05 +0000 Subject: [PATCH 40/78] Fix test --- gigl/distributed/dataset_factory.py | 12 +++++++----- .../distributed/dist_ablp_neighborloader_test.py | 14 ++------------ 2 files changed, 9 insertions(+), 17 deletions(-) diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index a49bc94dc..87971425a 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -44,6 +44,10 @@ from gigl.src.common.types.graph_data import EdgeType, NodeType from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper from gigl.src.common.types.pb_wrappers.task_metadata import TaskMetadataType +from gigl.types.graph import ( + DEFAULT_HOMOGENEOUS_NODE_TYPE, + FeatureQuantizationMetadata, +) from gigl.utils.data_splitters import ( DistNodeAnchorLinkSplitter, DistNodeSplitter, @@ -52,10 +56,6 @@ get_max_labels_per_anchor_node_from_runtime_args, select_ssl_positive_label_edges, ) -from gigl.types.graph import ( - DEFAULT_HOMOGENEOUS_NODE_TYPE, - FeatureQuantizationMetadata, -) logger = Logger() @@ -68,7 +68,9 @@ def _normalize_node_quantization_metadata( Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] ], labels_converted_to_edges: bool, -) -> Optional[Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]]]: +) -> Optional[ + Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] +]: """Align node quantization metadata with packed feature representation.""" if ( labels_converted_to_edges diff --git a/tests/unit/distributed/dist_ablp_neighborloader_test.py b/tests/unit/distributed/dist_ablp_neighborloader_test.py index df26164d2..abe69b04b 100644 --- a/tests/unit/distributed/dist_ablp_neighborloader_test.py +++ b/tests/unit/distributed/dist_ablp_neighborloader_test.py @@ -10,10 +10,7 @@ from parameterized import param, parameterized from torch_geometric.data import Data, HeteroData -from gigl.distributed.dataset_factory import ( - _validate_node_quantization_metadata_representation, - build_dataset, -) +from gigl.distributed.dataset_factory import build_dataset from gigl.distributed.dist_ablp_neighborloader import DistABLPLoader from gigl.distributed.dist_dataset import DistDataset from gigl.distributed.dist_partitioner import DistPartitioner @@ -792,10 +789,6 @@ def test_homogeneous_ablp_materializes_quantized_features(self) -> None: loaded_graph_tensors.treat_labels_as_edges(edge_dir="out") promoted_metadata = {DEFAULT_HOMOGENEOUS_NODE_TYPE: quantization_metadata} - _validate_node_quantization_metadata_representation( - node_quantized_features=loaded_graph_tensors.node_quantized_features, - node_quantization_metadata=promoted_metadata, - ) assert isinstance(loaded_graph_tensors.edge_index, dict) assert isinstance(loaded_graph_tensors.node_features, dict) @@ -809,10 +802,7 @@ def test_homogeneous_ablp_materializes_quantized_features(self) -> None: ) partition_output = PartitionOutput( node_partition_book={DEFAULT_HOMOGENEOUS_NODE_TYPE: torch.zeros(3)}, - edge_partition_book={ - edge_type: torch.zeros(3) - for edge_type in edge_index - }, + edge_partition_book={edge_type: torch.zeros(3) for edge_type in edge_index}, partitioned_edge_index={ edge_type: GraphPartitionData( edge_tensor, torch.arange(edge_tensor.size(1)) From 37291adbb1dd940da53de1652791846841a50383 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 17:49:28 +0000 Subject: [PATCH 41/78] Inline quant metadata normalization --- gigl/distributed/dataset_factory.py | 65 ++++++++++------------------- 1 file changed, 21 insertions(+), 44 deletions(-) diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 87971425a..2aaeb2328 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -41,13 +41,10 @@ from gigl.distributed.utils.serialized_graph_metadata_translator import ( convert_pb_to_serialized_graph_metadata, ) -from gigl.src.common.types.graph_data import EdgeType, NodeType +from gigl.src.common.types.graph_data import EdgeType from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper from gigl.src.common.types.pb_wrappers.task_metadata import TaskMetadataType -from gigl.types.graph import ( - DEFAULT_HOMOGENEOUS_NODE_TYPE, - FeatureQuantizationMetadata, -) +from gigl.types.graph import DEFAULT_HOMOGENEOUS_NODE_TYPE from gigl.utils.data_splitters import ( DistNodeAnchorLinkSplitter, DistNodeSplitter, @@ -60,40 +57,6 @@ logger = Logger() -def _normalize_node_quantization_metadata( - node_quantized_features: Optional[ - Union[torch.Tensor, dict[NodeType, torch.Tensor]] - ], - node_quantization_metadata: Optional[ - Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] - ], - labels_converted_to_edges: bool, -) -> Optional[ - Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] -]: - """Align node quantization metadata with packed feature representation.""" - if ( - labels_converted_to_edges - and node_quantization_metadata is not None - and not isinstance(node_quantization_metadata, Mapping) - ): - node_quantization_metadata = { - DEFAULT_HOMOGENEOUS_NODE_TYPE: node_quantization_metadata - } - - if node_quantized_features is None or node_quantization_metadata is None: - return node_quantization_metadata - - packed_features_are_typed = isinstance(node_quantized_features, Mapping) - metadata_is_typed = isinstance(node_quantization_metadata, Mapping) - if packed_features_are_typed != metadata_is_typed: - raise ValueError( - "Packed node features and node quantization metadata must both be " - "scalar or both be keyed by node type." - ) - return node_quantization_metadata - - @tf_on_cpu def _load_and_build_partitioned_dataset( serialized_graph_metadata: SerializedGraphMetadata, @@ -196,11 +159,25 @@ def _load_and_build_partitioned_dataset( if labels_converted_to_edges: loaded_graph_tensors.treat_labels_as_edges(edge_dir=edge_dir) - node_quantization_metadata = _normalize_node_quantization_metadata( - node_quantized_features=loaded_graph_tensors.node_quantized_features, - node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, - labels_converted_to_edges=labels_converted_to_edges, - ) + node_quantization_metadata = serialized_graph_metadata.node_quantization_metadata + if ( + labels_converted_to_edges + and node_quantization_metadata is not None + and not isinstance(node_quantization_metadata, Mapping) + ): + node_quantization_metadata = { + DEFAULT_HOMOGENEOUS_NODE_TYPE: node_quantization_metadata + } + if ( + loaded_graph_tensors.node_quantized_features is not None + and node_quantization_metadata is not None + and isinstance(loaded_graph_tensors.node_quantized_features, Mapping) + != isinstance(node_quantization_metadata, Mapping) + ): + raise ValueError( + "Packed node features and node quantization metadata must both be " + "scalar or both be keyed by node type." + ) should_assign_edges_by_src_node: bool = False if edge_dir == "in" else True From 6e25193b62b484a0cec230e4e3902ac0aac08ad1 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 18:01:53 +0000 Subject: [PATCH 42/78] Remove collate timers --- gigl/distributed/dist_ablp_neighborloader.py | 13 ------------- gigl/distributed/distributed_neighborloader.py | 13 ------------- 2 files changed, 26 deletions(-) diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index 6b35f0e03..56a887c10 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -1,4 +1,3 @@ -import time import warnings from collections import abc from itertools import count @@ -926,16 +925,9 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: # around a GLT bug in to_hetero_data. extract_edge_type_metadata then # pulls out labels by prefix. # TODO (mkolodner-sc): Remove the need to extract metadata once GLT's `to_hetero_data` function is fixed - collate_start_time = time.perf_counter() metadata, stripped_msg = extract_metadata(msg, self.to_device) - base_collate_start_time = time.perf_counter() data = super()._collate_fn(stripped_msg) - base_collate_time = time.perf_counter() - base_collate_start_time - logger.debug( - "Distributed ABLPNeighborLoader GLT base collate time: %.3fs", - base_collate_time, - ) data = set_missing_features( data=data, @@ -986,9 +978,4 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: for key, value in metadata.items(): data[key] = value - collate_time = time.perf_counter() - collate_start_time - logger.debug( - "Distributed ABLPNeighborLoader end-to-end collate time: %.3fs", - collate_time, - ) return data diff --git a/gigl/distributed/distributed_neighborloader.py b/gigl/distributed/distributed_neighborloader.py index 7ed34de25..b96082fd0 100644 --- a/gigl/distributed/distributed_neighborloader.py +++ b/gigl/distributed/distributed_neighborloader.py @@ -1,5 +1,4 @@ import sys -import time from collections import abc from itertools import count from typing import Optional, Tuple, Union @@ -546,15 +545,8 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: # as edge types and fails when edge_dir="out" (tries to call # reverse_edge_type on them). We strip them here and re-apply after. # TODO (mkolodner-sc): Remove once GLT's to_hetero_data is fixed. - collate_start_time = time.perf_counter() metadata, stripped_msg = extract_metadata(msg, self.to_device) - base_collate_start_time = time.perf_counter() data = super()._collate_fn(stripped_msg) - base_collate_time = time.perf_counter() - base_collate_start_time - logger.debug( - "Distributed NeighborLoader GLT base collate time: %.3fs", - base_collate_time, - ) data = set_missing_features( data=data, node_feature_info=self._node_feature_info, @@ -578,9 +570,4 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: for key, value in metadata.items(): data[key] = value - collate_time = time.perf_counter() - collate_start_time - logger.debug( - "Distributed NeighborLoader end-to-end collate time: %.3fs", - collate_time, - ) return data From 5d8d14899dc42c4b1959634a7449d139994f240b Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 18:06:31 +0000 Subject: [PATCH 43/78] Whitesapce --- gigl/distributed/dist_ablp_neighborloader.py | 1 - gigl/distributed/distributed_neighborloader.py | 1 - 2 files changed, 2 deletions(-) diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index 56a887c10..78147754b 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -977,5 +977,4 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: # data object so downstream code can access them via attribute lookup. for key, value in metadata.items(): data[key] = value - return data diff --git a/gigl/distributed/distributed_neighborloader.py b/gigl/distributed/distributed_neighborloader.py index b96082fd0..92135ea6b 100644 --- a/gigl/distributed/distributed_neighborloader.py +++ b/gigl/distributed/distributed_neighborloader.py @@ -569,5 +569,4 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: # data object so downstream code can access them via attribute lookup. for key, value in metadata.items(): data[key] = value - return data From 438e07c8bd536c5777875c9a5a6aa792ca1b57bf Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 18:22:34 +0000 Subject: [PATCH 44/78] Add descriptive comment for why we need to remap quantization metadat --- gigl/distributed/dataset_factory.py | 24 ++++++++++++++---------- 1 file changed, 14 insertions(+), 10 deletions(-) diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 2aaeb2328..03b3ca5fe 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -159,7 +159,11 @@ def _load_and_build_partitioned_dataset( if labels_converted_to_edges: loaded_graph_tensors.treat_labels_as_edges(edge_dir=edge_dir) + packed_node_features = loaded_graph_tensors.node_quantized_features node_quantization_metadata = serialized_graph_metadata.node_quantization_metadata + # Quantization metadata is serialization configuration, not a loaded graph + # tensor. ``treat_labels_as_edges`` converts the tensors to per-node-type + # form, so convert the metadata separately to keep them aligned. if ( labels_converted_to_edges and node_quantization_metadata is not None @@ -168,16 +172,16 @@ def _load_and_build_partitioned_dataset( node_quantization_metadata = { DEFAULT_HOMOGENEOUS_NODE_TYPE: node_quantization_metadata } - if ( - loaded_graph_tensors.node_quantized_features is not None - and node_quantization_metadata is not None - and isinstance(loaded_graph_tensors.node_quantized_features, Mapping) - != isinstance(node_quantization_metadata, Mapping) - ): - raise ValueError( - "Packed node features and node quantization metadata must both be " - "scalar or both be keyed by node type." - ) + if packed_node_features is not None and node_quantization_metadata is not None: + packed_features_are_mapping = isinstance(packed_node_features, Mapping) + metadata_is_mapping = isinstance(node_quantization_metadata, Mapping) + if packed_features_are_mapping != metadata_is_mapping: + raise ValueError( + "Packed node features and node quantization metadata must both be " + "scalar or keyed by node type. " + f"Got packed node features of type {type(packed_node_features)} " + f"and node quantization metadata {node_quantization_metadata}." + ) should_assign_edges_by_src_node: bool = False if edge_dir == "in" else True From c8d43a58d9ec3d03a9c47825ed0266deb2a18f2d Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 20:25:23 +0000 Subject: [PATCH 45/78] Don't need node quantized feature info --- gigl/distributed/dist_dataset.py | 23 ------------------- gigl/distributed/graph_store/dist_server.py | 6 ----- .../graph_store/remote_dist_dataset.py | 16 ------------- .../distributed/distributed_dataset_test.py | 4 ---- 4 files changed, 49 deletions(-) diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index 93ddf2ced..4fc486d4e 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -81,9 +81,6 @@ def __init__( node_feature_info: Optional[ Union[FeatureInfo, dict[NodeType, FeatureInfo]] ] = None, - node_quantized_feature_info: Optional[ - Union[FeatureInfo, dict[NodeType, FeatureInfo]] - ] = None, node_quantization_metadata: Optional[ Union[ FeatureQuantizationMetadata, @@ -125,7 +122,6 @@ def __init__( num_test: (Optional[Union[int, dict[NodeType, int]]]): Number of test nodes on the current machine. Will be a dict if heterogeneous. node_feature_info: Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Dimension of node features and its data type, will be a dict if heterogeneous. Note this will be None in the homogeneous case if the data has no node features, or will only contain node types with node features in the heterogeneous case. - node_quantized_feature_info: Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Dimension and dtype for packed uint8 node features. edge_feature_info: Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]]: Dimension of edge features and its data type, will be a dict if heterogeneous. Note this will be None in the homogeneous case if the data has no edge features, or will only contain edge types with edge features in the heterogeneous case. degree_tensor: Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Pre-computed degree tensor. Lazily computed on first access via the degree_tensor property. @@ -168,7 +164,6 @@ def __init__( self._node_feature_info = node_feature_info self._edge_feature_info = edge_feature_info - self._node_quantized_feature_info = node_quantized_feature_info self._node_quantized_features = node_quantized_feature_partition self._node_quantization_metadata = node_quantization_metadata @@ -337,12 +332,6 @@ def node_feature_info( """ return self._node_feature_info - @property - def node_quantized_feature_info( - self, - ) -> Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: - return self._node_quantized_feature_info - @property def node_quantization_metadata( self, @@ -812,13 +801,6 @@ def _initialize_node_quantized_features( ) for node_type, features_per_node_type in node_quantized_features.items() } - self._node_quantized_feature_info = {} - for node_type, features_per_node_type in node_quantized_features.items(): - assert not isinstance(node_type, EdgeType) - self._node_quantized_feature_info[node_type] = FeatureInfo( - dim=features_per_node_type.size(1), # ty: ignore[unresolved-attribute] TODO(ty-torch-keyed-access): fix ty false positives for torch-backed keyed container access. - dtype=features_per_node_type.dtype, # ty: ignore[unresolved-attribute] TODO(ty-torch-keyed-access): fix ty false positives for torch-backed keyed container access. - ) logger.info( f"Initialized node quantized features for heterogeneous graph to dataset with node types: {node_quantized_features.keys()}" ) @@ -830,10 +812,6 @@ def _initialize_node_quantized_features( with_gpu=False, dtype=torch.uint8, ) - self._node_quantized_feature_info = FeatureInfo( - dim=node_quantized_features.size(1), - dtype=node_quantized_features.dtype, - ) logger.info( "Initialized node quantized features for homogeneous graph to dataset" ) @@ -1135,7 +1113,6 @@ def share_ipc( self._num_val, # Additional field unique to DistDataset class self._num_test, # Additional field unique to DistDataset class self._node_feature_info, # Additional field unique to DistDataset class - self._node_quantized_feature_info, # Additional field unique to DistDataset class self._node_quantization_metadata, # Additional field unique to DistDataset class self._edge_feature_info, # Additional field unique to DistDataset class self._degree_tensor, # Additional field unique to DistDataset class diff --git a/gigl/distributed/graph_store/dist_server.py b/gigl/distributed/graph_store/dist_server.py index 104c37750..0a92d959d 100644 --- a/gigl/distributed/graph_store/dist_server.py +++ b/gigl/distributed/graph_store/dist_server.py @@ -410,12 +410,6 @@ def get_node_feature_info( """ return self.dataset.node_feature_info - def get_node_quantized_feature_info( - self, - ) -> Union[FeatureInfo, dict[NodeType, FeatureInfo], None]: - """Get packed node feature information from the dataset.""" - return self.dataset.node_quantized_feature_info - def get_node_quantization_metadata( self, ) -> Union[ diff --git a/gigl/distributed/graph_store/remote_dist_dataset.py b/gigl/distributed/graph_store/remote_dist_dataset.py index dbf67c86c..ef3fc24b0 100644 --- a/gigl/distributed/graph_store/remote_dist_dataset.py +++ b/gigl/distributed/graph_store/remote_dist_dataset.py @@ -67,22 +67,6 @@ def fetch_node_feature_info( DistServer.get_node_feature_info, ) - def fetch_node_quantized_feature_info( - self, - ) -> Union[FeatureInfo, dict[NodeType, FeatureInfo], None]: - """Fetch packed node feature information from the registered dataset.""" - return request_server( - 0, - DistServer.get_node_quantized_feature_info, - ) - - @property - def node_quantized_feature_info( - self, - ) -> Union[FeatureInfo, dict[NodeType, FeatureInfo], None]: - """Return packed node feature information from the storage cluster.""" - return self.fetch_node_quantized_feature_info() - def fetch_node_quantization_metadata( self, ) -> Union[ diff --git a/tests/unit/distributed/distributed_dataset_test.py b/tests/unit/distributed/distributed_dataset_test.py index 7c98d34dd..d26a80d1b 100644 --- a/tests/unit/distributed/distributed_dataset_test.py +++ b/tests/unit/distributed/distributed_dataset_test.py @@ -611,10 +611,6 @@ def test_building_dataset_preserves_packed_node_features_and_metadata(self) -> N dataset.node_quantized_features[torch.tensor([7, 3])], packed_features.flip(0), ) - self.assertEqual( - dataset.node_quantized_feature_info, - FeatureInfo(dim=1, dtype=torch.uint8), - ) self.assertEqual(dataset.node_quantization_metadata, quantization_metadata) def test_building_heterogeneous_dataset_preserves_node_features_and_labels(self): From 0f2f9b9964ebf36c9a92010c47a991b4687ae009 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 20:28:19 +0000 Subject: [PATCH 46/78] Remove debug timing --- gigl/distributed/utils/neighborloader.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 576b5564b..617aba7ce 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -1,7 +1,6 @@ """Utils for Neighbor loaders.""" import ast -import time from collections import abc from copy import deepcopy from dataclasses import dataclass @@ -346,7 +345,6 @@ def materialize_quantized_node_features( """Materialize packed quantized node features into PyG node feature tensors.""" if node_quantization_metadata is None: return data, metadata - materialize_start_time = time.perf_counter() def materialize( store, packed_features: torch.Tensor, q: FeatureQuantizationMetadata @@ -403,8 +401,6 @@ def materialize( continue materialize(data[node_type], packed_features, quantization_metadata) - materialize_time = time.perf_counter() - materialize_start_time - logger.debug("Quantized node feature materialization time: %.3fs", materialize_time) return data, metadata From c9a9912ae32fc22ffac769e6df3c5ff1a03e1fac Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 20:37:16 +0000 Subject: [PATCH 47/78] Value error instead of assertion --- gigl/distributed/utils/neighborloader.py | 5 ++--- .../distributed/utils/neighborloader_test.py | 16 ++++++++++++++++ 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 617aba7ce..07eb84be4 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -388,9 +388,8 @@ def materialize( ) materialize(data, packed_features, quantization_metadata) else: - assert isinstance(node_quantization_metadata, dict), ( - "Expected per-node-type quantization metadata for heterogeneous data." - ) + if not isinstance(node_quantization_metadata, dict): + raise ValueError("Expected per-node-type metadata for heterogeneous data.") node_quantization_metadata = cast( dict[NodeType, FeatureQuantizationMetadata], node_quantization_metadata ) diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index 083dd6fd8..20ad8b710 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -135,6 +135,22 @@ def test_materialize_quantized_node_features_uses_per_node_type_metadata( self.assertEqual(set(remaining_metadata), {"request_id"}) self.assert_tensor_equality(remaining_metadata["request_id"], torch.tensor([7])) + def test_materialize_quantized_node_features_rejects_scalar_metadata_for_heterogeneous_data( + self, + ) -> None: + with self.assertRaises(ValueError): + materialize_quantized_node_features( + data=HeteroData(), + metadata={}, + node_quantization_metadata=FeatureQuantizationMetadata( + bits=2, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_min=0.0, + clip_max=3.0, + ), + ) + @parameterized.expand( [ param( From feb138edd61ce7dd289849349635488bf2a98dbf Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 20:42:03 +0000 Subject: [PATCH 48/78] Remove stale share ipc entry --- gigl/distributed/dist_dataset.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index 4fc486d4e..5e201f8b3 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -1079,7 +1079,6 @@ def share_ipc( Optional[Union[int, dict[NodeType, int]]]: Number of validation nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[int, dict[NodeType, int]]]: Number of test nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Node feature dim and its data type, will be a dict if heterogeneous - Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Packed uint8 node feature dim and dtype Optional node quantization metadata. Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]]: Edge feature dim and its data type, will be a dict if heterogeneous Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Degree tensors @@ -1373,9 +1372,6 @@ def _rebuild_distributed_dataset( Optional[ Union[FeatureInfo, dict[NodeType, FeatureInfo]] ], # Node feature dim and its data type - Optional[ - Union[FeatureInfo, dict[NodeType, FeatureInfo]] - ], # Packed uint8 node feature dim and dtype Optional[ Union[ FeatureQuantizationMetadata, From 67957317a93b6021cff2f0ad6f47f5863a4e0a70 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 20:44:26 +0000 Subject: [PATCH 49/78] Upd --- gigl/distributed/dist_dataset.py | 1 - 1 file changed, 1 deletion(-) diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index 5e201f8b3..9c45c5eb4 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -1047,7 +1047,6 @@ def share_ipc( Optional[Union[int, dict[NodeType, int]]], Optional[Union[int, dict[NodeType, int]]], Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]], - Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]], Optional[ Union[ FeatureQuantizationMetadata, From 41f3a21708ba04627490db5a2c5dffe55c4957d1 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 20:45:53 +0000 Subject: [PATCH 50/78] Upd --- gigl/distributed/dist_dataset.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index 9c45c5eb4..a41c0e12a 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -1078,7 +1078,7 @@ def share_ipc( Optional[Union[int, dict[NodeType, int]]]: Number of validation nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[int, dict[NodeType, int]]]: Number of test nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Node feature dim and its data type, will be a dict if heterogeneous - Optional node quantization metadata. + Optional[Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]]]: Node quantization metadata. Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]]: Edge feature dim and its data type, will be a dict if heterogeneous Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Degree tensors Optional[int]: Optional per-anchor label cap for ABLP label fetching From 1e375ac39141130a684dba9d96187a0ecc29100a Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 21:08:03 +0000 Subject: [PATCH 51/78] Update tests --- .../run_distributed_partitioner.py | 30 +++++++----------- .../distributed_partitioner_test.py | 31 +++++++------------ 2 files changed, 24 insertions(+), 37 deletions(-) diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index d67121d72..046b8bf49 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -34,7 +34,6 @@ def run_distributed_partitioner( master_port: int, input_data_strategy: InputDataStrategy, partitioner_class: Type[DistPartitioner], - include_quantized_features: bool = False, rank_to_edge_weights: Optional[ dict[int, Union[torch.Tensor, dict[EdgeType, torch.Tensor]]] ] = None, @@ -51,7 +50,6 @@ def run_distributed_partitioner( master_port (int): Master port for initializing rpc for partitioning input_data_strategy (InputDataStrategy): Strategy for registering inputs to the partitioner partitioner_class (Type[DistPartitioner]): The class to use for partitioning - include_quantized_features: Whether to register one packed uint8 feature per node. rank_to_edge_weights (Optional[dict[int, Union[torch.Tensor, dict[EdgeType, torch.Tensor]]]]): Optional mapping of rank to 1D edge weight tensor (or dict per EdgeType for heterogeneous). Only supported with REGISTER_ALL_ENTITIES_SEPARATELY strategy. """ @@ -63,9 +61,7 @@ def run_distributed_partitioner( positive_labels: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] negative_labels: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] node_labels: Union[torch.Tensor, dict[NodeType, torch.Tensor]] - node_quantized_features: Optional[ - Union[torch.Tensor, dict[NodeType, torch.Tensor]] - ] = None + node_quantized_features: Union[torch.Tensor, dict[NodeType, torch.Tensor]] if not is_heterogeneous: node_ids = input_graph.node_ids[USER_NODE_TYPE] @@ -84,15 +80,14 @@ def run_distributed_partitioner( negative_labels = input_graph.negative_labels node_labels = input_graph.node_labels - if include_quantized_features: - if isinstance(node_ids, dict): - node_ids_by_type = cast(dict[NodeType, torch.Tensor], node_ids) - node_quantized_features = { - node_type: node_type_ids.to(torch.uint8).unsqueeze(1) - for node_type, node_type_ids in node_ids_by_type.items() - } - else: - node_quantized_features = node_ids.to(torch.uint8).unsqueeze(1) + if isinstance(node_ids, dict): + node_ids_by_type = cast(dict[NodeType, torch.Tensor], node_ids) + node_quantized_features = { + node_type: node_type_ids.to(torch.uint8).unsqueeze(1) + for node_type, node_type_ids in node_ids_by_type.items() + } + else: + node_quantized_features = node_ids.to(torch.uint8).unsqueeze(1) partition_output: PartitionOutput @@ -130,10 +125,9 @@ def run_distributed_partitioner( ) dist_partitioner.register_node_features(node_features=node_features) - if node_quantized_features is not None: - dist_partitioner.register_node_quantized_features( - node_quantized_features=node_quantized_features - ) + dist_partitioner.register_node_quantized_features( + node_quantized_features=node_quantized_features + ) dist_partitioner.register_node_labels(node_labels=node_labels) del node_labels del node_features diff --git a/tests/unit/distributed/distributed_partitioner_test.py b/tests/unit/distributed/distributed_partitioner_test.py index 6860d0bdd..0f817bafa 100644 --- a/tests/unit/distributed/distributed_partitioner_test.py +++ b/tests/unit/distributed/distributed_partitioner_test.py @@ -629,7 +629,6 @@ def _assert_label_outputs( should_assign_edges_by_src_node=True, partitioner_class=DistPartitioner, expected_pb_dtype=torch.uint8, - include_quantized_features=True, ), param( "Homogeneous Partitioning By Dest Node- Register All Entites together through Constructor", @@ -670,7 +669,6 @@ def _assert_label_outputs( should_assign_edges_by_src_node=True, partitioner_class=DistRangePartitioner, expected_pb_dtype=torch.int64, - include_quantized_features=True, ), param( "Heterogeneous Partitioning By Source Node - Range Partitioning", @@ -698,7 +696,6 @@ def test_partitioning_correctness( should_assign_edges_by_src_node: bool, partitioner_class: Type[DistPartitioner], expected_pb_dtype: torch.dtype, - include_quantized_features: bool = False, ) -> None: """ Tests partitioning functionality and correctness on mocked inputs @@ -728,7 +725,6 @@ def test_partitioning_correctness( master_port, input_data_strategy, partitioner_class, - include_quantized_features, ), nprocs=MOCKED_NUM_PARTITIONS, join=True, @@ -834,21 +830,18 @@ def test_partitioning_correctness( entity_name="features", ) - if include_quantized_features: - assert ( - partition_output.partitioned_node_quantized_features is not None - ) - self._assert_node_data_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_node_data=partition_output.partitioned_node_quantized_features, - expected_node_types=MOCKED_HETEROGENEOUS_NODE_TYPES, - expected_edge_types=MOCKED_HETEROGENEOUS_EDGE_TYPES, - entity_name="quantized_features", - ) + assert partition_output.partitioned_node_quantized_features is not None + self._assert_node_data_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_node_data=partition_output.partitioned_node_quantized_features, + expected_node_types=MOCKED_HETEROGENEOUS_NODE_TYPES, + expected_edge_types=MOCKED_HETEROGENEOUS_EDGE_TYPES, + entity_name="quantized_features", + ) if isinstance(partition_output.partitioned_node_features, abc.Mapping): for ( From bec1f10bc1485f6225078301495de01287c56435 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 23:30:46 +0000 Subject: [PATCH 52/78] Remove quantization metadata property to make RPC call explicit --- gigl/distributed/dist_ablp_neighborloader.py | 4 ++-- gigl/distributed/distributed_neighborloader.py | 2 +- gigl/distributed/graph_store/remote_dist_dataset.py | 10 ---------- tests/unit/distributed/dist_server_test.py | 4 +++- 4 files changed, 6 insertions(+), 14 deletions(-) diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index 78147754b..a828f5a48 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -608,7 +608,7 @@ def _setup_for_colocated( edge_types=edge_types, node_feature_info=dataset.node_feature_info, edge_feature_info=dataset.edge_feature_info, - node_quantization_metadata=dataset.node_quantization_metadata, + node_quantization_metadata=dataset.fetch_node_quantization_metadata(), edge_dir=dataset.edge_dir, ), ) @@ -798,7 +798,7 @@ def _setup_for_graph_store( edge_types=edge_types, node_feature_info=node_feature_info, edge_feature_info=edge_feature_info, - node_quantization_metadata=dataset.node_quantization_metadata, + node_quantization_metadata=dataset.fetch_node_quantization_metadata(), edge_dir=edge_dir, ), backend_key, diff --git a/gigl/distributed/distributed_neighborloader.py b/gigl/distributed/distributed_neighborloader.py index 92135ea6b..3effa01c3 100644 --- a/gigl/distributed/distributed_neighborloader.py +++ b/gigl/distributed/distributed_neighborloader.py @@ -412,7 +412,7 @@ def _setup_for_graph_store( edge_types=edge_types, node_feature_info=node_feature_info, edge_feature_info=edge_feature_info, - node_quantization_metadata=dataset.node_quantization_metadata, + node_quantization_metadata=dataset.fetch_node_quantization_metadata(), edge_dir=dataset.fetch_edge_dir(), ), backend_key, diff --git a/gigl/distributed/graph_store/remote_dist_dataset.py b/gigl/distributed/graph_store/remote_dist_dataset.py index ef3fc24b0..81609961b 100644 --- a/gigl/distributed/graph_store/remote_dist_dataset.py +++ b/gigl/distributed/graph_store/remote_dist_dataset.py @@ -80,16 +80,6 @@ def fetch_node_quantization_metadata( DistServer.get_node_quantization_metadata, ) - @property - def node_quantization_metadata( - self, - ) -> Union[ - FeatureQuantizationMetadata, - dict[NodeType, FeatureQuantizationMetadata], - None, - ]: - return self.fetch_node_quantization_metadata() - def fetch_edge_feature_info( self, ) -> Union[FeatureInfo, dict[EdgeType, FeatureInfo], None]: diff --git a/tests/unit/distributed/dist_server_test.py b/tests/unit/distributed/dist_server_test.py index 64204ea83..c876fcef1 100644 --- a/tests/unit/distributed/dist_server_test.py +++ b/tests/unit/distributed/dist_server_test.py @@ -98,7 +98,9 @@ def test_remote_dataset_fetches_node_quantization_metadata(self) -> None: ) as request_server: remote_dataset = RemoteDistDataset(cluster_info=MagicMock(), local_rank=0) - self.assertEqual(remote_dataset.node_quantization_metadata, metadata) + self.assertEqual( + remote_dataset.fetch_node_quantization_metadata(), metadata + ) request_server.assert_called_once_with( 0, dist_server.DistServer.get_node_quantization_metadata From 84e3e58b8a7598592d13dfa0527ae5d83bcaa279 Mon Sep 17 00:00:00 2001 From: jchmura Date: Tue, 11 Aug 2026 23:33:13 +0000 Subject: [PATCH 53/78] Update --- gigl/distributed/dist_ablp_neighborloader.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index a828f5a48..10638c330 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -608,7 +608,7 @@ def _setup_for_colocated( edge_types=edge_types, node_feature_info=dataset.node_feature_info, edge_feature_info=dataset.edge_feature_info, - node_quantization_metadata=dataset.fetch_node_quantization_metadata(), + node_quantization_metadata=dataset.node_quantization_metadata, edge_dir=dataset.edge_dir, ), ) From ca7fdea25a4a0fa2ce5c0d678f387f259e438012 Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 15:37:00 +0000 Subject: [PATCH 54/78] Add type to storage --- gigl/distributed/utils/neighborloader.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 07eb84be4..a13572bdf 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -10,6 +10,7 @@ import torch from graphlearn_torch.channel import SampleMessage from torch_geometric.data import Data, HeteroData +from torch_geometric.data.storage import NodeStorage from torch_geometric.typing import EdgeType, NodeType from gigl.common.logger import Logger @@ -347,7 +348,9 @@ def materialize_quantized_node_features( return data, metadata def materialize( - store, packed_features: torch.Tensor, q: FeatureQuantizationMetadata + store: Union[Data, NodeStorage], + packed_features: torch.Tensor, + q: FeatureQuantizationMetadata, ) -> None: dequantized = dequantize_torch_tensor(packed_features, metadata=q) x = getattr(store, "x", None) From 54e53c54cea5563ef9534c0ec8ea2532468d414b Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 15:50:06 +0000 Subject: [PATCH 55/78] No need for metadata promotion to dict on labeled homogeneous ablp --- gigl/distributed/dataset_factory.py | 32 ++----------------- gigl/distributed/utils/neighborloader.py | 28 +++++++--------- .../dist_ablp_neighborloader_test.py | 5 ++- .../distributed/utils/neighborloader_test.py | 29 +++++++++++++++++ 4 files changed, 46 insertions(+), 48 deletions(-) diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 03b3ca5fe..1a5f859b5 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -44,7 +44,6 @@ from gigl.src.common.types.graph_data import EdgeType from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper from gigl.src.common.types.pb_wrappers.task_metadata import TaskMetadataType -from gigl.types.graph import DEFAULT_HOMOGENEOUS_NODE_TYPE from gigl.utils.data_splitters import ( DistNodeAnchorLinkSplitter, DistNodeSplitter, @@ -152,36 +151,11 @@ def _load_and_build_partitioned_dataset( loaded_graph_tensors.positive_label = positive_label_edges - labels_converted_to_edges = ( + if ( isinstance(splitter, NodeAnchorLinkSplitter) and splitter.should_convert_labels_to_edges - ) - if labels_converted_to_edges: - loaded_graph_tensors.treat_labels_as_edges(edge_dir=edge_dir) - - packed_node_features = loaded_graph_tensors.node_quantized_features - node_quantization_metadata = serialized_graph_metadata.node_quantization_metadata - # Quantization metadata is serialization configuration, not a loaded graph - # tensor. ``treat_labels_as_edges`` converts the tensors to per-node-type - # form, so convert the metadata separately to keep them aligned. - if ( - labels_converted_to_edges - and node_quantization_metadata is not None - and not isinstance(node_quantization_metadata, Mapping) ): - node_quantization_metadata = { - DEFAULT_HOMOGENEOUS_NODE_TYPE: node_quantization_metadata - } - if packed_node_features is not None and node_quantization_metadata is not None: - packed_features_are_mapping = isinstance(packed_node_features, Mapping) - metadata_is_mapping = isinstance(node_quantization_metadata, Mapping) - if packed_features_are_mapping != metadata_is_mapping: - raise ValueError( - "Packed node features and node quantization metadata must both be " - "scalar or keyed by node type. " - f"Got packed node features of type {type(packed_node_features)} " - f"and node quantization metadata {node_quantization_metadata}." - ) + loaded_graph_tensors.treat_labels_as_edges(edge_dir=edge_dir) should_assign_edges_by_src_node: bool = False if edge_dir == "in" else True @@ -252,7 +226,7 @@ def _load_and_build_partitioned_dataset( rank=rank, world_size=world_size, edge_dir=edge_dir, - node_quantization_metadata=node_quantization_metadata, + node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, ) dataset.build( diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index a13572bdf..1ab2f9846 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -371,25 +371,21 @@ def materialize( if isinstance(data, Data): if isinstance(node_quantization_metadata, dict): - homogeneous_quantization_metadata = cast( - dict[NodeType, FeatureQuantizationMetadata], - node_quantization_metadata, - ) - quantization_metadata = homogeneous_quantization_metadata[ - DEFAULT_HOMOGENEOUS_NODE_TYPE - ] - metadata_key = ( - f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}" - ) - else: - quantization_metadata = node_quantization_metadata - metadata_key = NODE_PACKED_FEATURES_METADATA_KEY - packed_features = metadata.pop(metadata_key, None) + raise ValueError("Expect scalar quantization metadata for homogeneous data") + metadata_keys = ( + NODE_PACKED_FEATURES_METADATA_KEY, + f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}", + ) + packed_features = None + for metadata_key in metadata_keys: + packed_features = metadata.pop(metadata_key, None) + if packed_features is not None: + break if packed_features is None: raise ValueError( - f"Missing packed quantized node features in metadata key {metadata_key}." + f"Missing packed quantized features in metadata keys {metadata_keys}" ) - materialize(data, packed_features, quantization_metadata) + materialize(data, packed_features, node_quantization_metadata) else: if not isinstance(node_quantization_metadata, dict): raise ValueError("Expected per-node-type metadata for heterogeneous data.") diff --git a/tests/unit/distributed/dist_ablp_neighborloader_test.py b/tests/unit/distributed/dist_ablp_neighborloader_test.py index abe69b04b..a3d25dfeb 100644 --- a/tests/unit/distributed/dist_ablp_neighborloader_test.py +++ b/tests/unit/distributed/dist_ablp_neighborloader_test.py @@ -766,7 +766,7 @@ def test_ablp_dataloader( ) def test_homogeneous_ablp_materializes_quantized_features(self) -> None: - """Promoted metadata supports quantized homogeneous ABLP batches.""" + """Scalar metadata supports quantized homogeneous ABLP batches.""" loaded_graph_tensors = LoadedGraphTensors( node_ids=torch.arange(3), node_features=torch.tensor([[10.0, 20.0], [30.0, 40.0], [50.0, 60.0]]), @@ -788,7 +788,6 @@ def test_homogeneous_ablp_materializes_quantized_features(self) -> None: ) loaded_graph_tensors.treat_labels_as_edges(edge_dir="out") - promoted_metadata = {DEFAULT_HOMOGENEOUS_NODE_TYPE: quantization_metadata} assert isinstance(loaded_graph_tensors.edge_index, dict) assert isinstance(loaded_graph_tensors.node_features, dict) @@ -830,7 +829,7 @@ def test_homogeneous_ablp_materializes_quantized_features(self) -> None: rank=0, world_size=1, edge_dir="out", - node_quantization_metadata=promoted_metadata, + node_quantization_metadata=quantization_metadata, ) dataset.build(partition_output=partition_output) diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index 20ad8b710..03a4455b3 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -24,6 +24,7 @@ strip_non_ppr_edge_types, ) from gigl.types.graph import ( + DEFAULT_HOMOGENEOUS_NODE_TYPE, FeatureInfo, FeatureQuantizationMetadata, message_passing_to_positive_label, @@ -101,6 +102,34 @@ def test_materialize_quantized_node_features_reconstructs_feature_order( self.assertEqual(set(remaining_metadata), {"request_id"}) self.assert_tensor_equality(remaining_metadata["request_id"], torch.tensor([7])) + def test_materialize_quantized_node_features_accepts_labeled_homogeneous_key( + self, + ) -> None: + data = Data() + metadata = { + f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}": torch.tensor( + [[27]], dtype=torch.uint8 + ) + } + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 1, 2, 3), + clip_min=0.0, + clip_max=3.0, + ) + + materialized_data, remaining_metadata = materialize_quantized_node_features( + data=data, + metadata=metadata, + node_quantization_metadata=quantization_metadata, + ) + + self.assert_tensor_equality( + materialized_data.x, torch.tensor([[0.0, 1.0, 2.0, 3.0]]) + ) + self.assertEmpty(remaining_metadata) + def test_materialize_quantized_node_features_uses_per_node_type_metadata( self, ) -> None: From 1cdeb88bacb0c70533e423887781717e69ed6fbb Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 15:56:52 +0000 Subject: [PATCH 56/78] Simplify test --- gigl/distributed/utils/neighborloader.py | 14 ++------- .../distributed/utils/neighborloader_test.py | 29 ------------------- 2 files changed, 3 insertions(+), 40 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 1ab2f9846..484a8f1c3 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -17,7 +17,6 @@ from gigl.common.utils.feature_quantization.torch_ops import dequantize_torch_tensor from gigl.distributed.sampler import NODE_PACKED_FEATURES_METADATA_KEY from gigl.types.graph import ( - DEFAULT_HOMOGENEOUS_NODE_TYPE, FeatureInfo, FeatureQuantizationMetadata, is_label_edge_type, @@ -372,18 +371,11 @@ def materialize( if isinstance(data, Data): if isinstance(node_quantization_metadata, dict): raise ValueError("Expect scalar quantization metadata for homogeneous data") - metadata_keys = ( - NODE_PACKED_FEATURES_METADATA_KEY, - f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}", - ) - packed_features = None - for metadata_key in metadata_keys: - packed_features = metadata.pop(metadata_key, None) - if packed_features is not None: - break + packed_features = metadata.pop(NODE_PACKED_FEATURES_METADATA_KEY, None) if packed_features is None: raise ValueError( - f"Missing packed quantized features in metadata keys {metadata_keys}" + f"Missing packed quantized features in metadata key " + f"{NODE_PACKED_FEATURES_METADATA_KEY}" ) materialize(data, packed_features, node_quantization_metadata) else: diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index 03a4455b3..20ad8b710 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -24,7 +24,6 @@ strip_non_ppr_edge_types, ) from gigl.types.graph import ( - DEFAULT_HOMOGENEOUS_NODE_TYPE, FeatureInfo, FeatureQuantizationMetadata, message_passing_to_positive_label, @@ -102,34 +101,6 @@ def test_materialize_quantized_node_features_reconstructs_feature_order( self.assertEqual(set(remaining_metadata), {"request_id"}) self.assert_tensor_equality(remaining_metadata["request_id"], torch.tensor([7])) - def test_materialize_quantized_node_features_accepts_labeled_homogeneous_key( - self, - ) -> None: - data = Data() - metadata = { - f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}": torch.tensor( - [[27]], dtype=torch.uint8 - ) - } - quantization_metadata = FeatureQuantizationMetadata( - bits=2, - feature_dim=4, - quantized_feature_indices=(0, 1, 2, 3), - clip_min=0.0, - clip_max=3.0, - ) - - materialized_data, remaining_metadata = materialize_quantized_node_features( - data=data, - metadata=metadata, - node_quantization_metadata=quantization_metadata, - ) - - self.assert_tensor_equality( - materialized_data.x, torch.tensor([[0.0, 1.0, 2.0, 3.0]]) - ) - self.assertEmpty(remaining_metadata) - def test_materialize_quantized_node_features_uses_per_node_type_metadata( self, ) -> None: From 50e6a8d4035c21ddee1842773262fbbdf8464c55 Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 16:25:34 +0000 Subject: [PATCH 57/78] Add type to scatter index --- gigl/distributed/utils/neighborloader.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 484a8f1c3..2e31d23dc 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -18,6 +18,7 @@ from gigl.distributed.sampler import NODE_PACKED_FEATURES_METADATA_KEY from gigl.types.graph import ( FeatureInfo, + FeatureQuantizationIndexTensors, FeatureQuantizationMetadata, is_label_edge_type, ) @@ -354,18 +355,19 @@ def materialize( dequantized = dequantize_torch_tensor(packed_features, metadata=q) x = getattr(store, "x", None) out = dequantized.new_empty((dequantized.size(0), q.feature_dim)) - scatter_indices = q.scatter_index_tensors(out.device) - out[:, scatter_indices.quantized] = dequantized + scatter_idx: FeatureQuantizationIndexTensors = q.scatter_index_tensors( + out.device + ) + out[:, scatter_idx.quantized] = dequantized if x is None and q.raw_feature_dim: raise ValueError(f"Missing {q.raw_feature_dim} unquantized features") if x is not None: if x.size(1) != q.raw_feature_dim: raise ValueError( - f"Expected {q.raw_feature_dim} raw node feature columns before " - f"dequantization, got {x.size(1)}." + f"Expected {q.raw_feature_dim} raw node features before dequantization, got {x.size(1)}" ) - out[:, scatter_indices.raw] = x + out[:, scatter_idx.raw] = x store.x = out if isinstance(data, Data): @@ -374,8 +376,7 @@ def materialize( packed_features = metadata.pop(NODE_PACKED_FEATURES_METADATA_KEY, None) if packed_features is None: raise ValueError( - f"Missing packed quantized features in metadata key " - f"{NODE_PACKED_FEATURES_METADATA_KEY}" + f"Missing packed quantized features in metadata key {NODE_PACKED_FEATURES_METADATA_KEY}" ) materialize(data, packed_features, node_quantization_metadata) else: From 207631ae3ce9a87f335653dd92e1b6c4b5fa4f5d Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 17:46:43 +0000 Subject: [PATCH 58/78] Add edge feature quantization --- gigl/common/data/load_torch_tensors.py | 151 ++++++++++++++- .../utils/feature_quantization/README.md | 4 +- gigl/distributed/base_dist_loader.py | 8 +- gigl/distributed/base_sampler.py | 34 ++++ gigl/distributed/dataset_factory.py | 10 + gigl/distributed/dist_ablp_neighborloader.py | 8 + gigl/distributed/dist_dataset.py | 91 ++++++++- gigl/distributed/dist_partitioner.py | 98 +++++++++- gigl/distributed/dist_range_partitioner.py | 56 +++++- .../distributed/distributed_neighborloader.py | 8 + gigl/distributed/graph_store/dist_server.py | 8 + .../graph_store/remote_dist_dataset.py | 8 + gigl/distributed/sampler.py | 1 + gigl/distributed/utils/neighborloader.py | 181 +++++++++++++++++- .../serialized_graph_metadata_translator.py | 24 ++- .../data_preprocessor/data_preprocessor.py | 34 +++- .../lib/transform/feature_quantization.py | 56 ++++-- .../data_preprocessor/lib/transform/utils.py | 10 +- gigl/src/data_preprocessor/lib/types.py | 1 + gigl/types/graph.py | 9 + .../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 +- .../feature_quantization_transform_test.py | 34 ++++ .../run_distributed_partitioner.py | 27 ++- tests/unit/common/data/dataloaders_test.py | 85 ++++++++ .../distributed_partitioner_test.py | 37 +++- .../distributed/utils/neighborloader_test.py | 170 ++++++++++++++++ 32 files changed, 1234 insertions(+), 106 deletions(-) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 3a1174888..3db6e1dfb 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 typing import MutableMapping, Optional, Union, cast import torch import torch.multiprocessing as mp @@ -119,6 +119,138 @@ 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: + """Validate sampling-weight configuration before TFRecord loading.""" + 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 a string for homogeneous graphs." + ) + configured_weights = [ + ( + DEFAULT_HOMOGENEOUS_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 a dict[EdgeType, str] for " + "heterogeneous graphs 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( + "weight_edge_feat_name contains unknown edge types: " + f"{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 remain an unquantized scalar edge feature. Available raw " + f"features: {serialized_info.feature_keys}" + ) + feature_spec = serialized_info.feature_spec[feature_name] + feature_width = feature_spec.shape[-1] if feature_spec.shape else 1 + if feature_width != 1: + raise ValueError( + f"Sampling-weight field '{feature_name}' for edge type {edge_type} " + f"must be scalar, but has width {feature_width}." + ) + + +def _remove_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 the separately stored sampling-weight column from model metadata.""" + quantization_metadata = serialized_graph_metadata.edge_quantization_metadata + if quantization_metadata is None or weight_edge_feat_name is None: + return quantization_metadata + + edge_info_by_type: dict[EdgeType, SerializedTFRecordInfo] + metadata_by_type: dict[EdgeType, FeatureQuantizationMetadata] + weight_by_type: dict[EdgeType, str] + 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 = { + DEFAULT_HOMOGENEOUS_EDGE_TYPE: serialized_graph_metadata.edge_entity_info + } + metadata_by_type = {DEFAULT_HOMOGENEOUS_EDGE_TYPE: quantization_metadata} + weight_by_type = {DEFAULT_HOMOGENEOUS_EDGE_TYPE: weight_edge_feat_name} + is_homogeneous = True + else: + assert isinstance(quantization_metadata, dict) + edge_info_by_type = serialized_graph_metadata.edge_entity_info + metadata_by_type = cast( + dict[EdgeType, FeatureQuantizationMetadata], quantization_metadata + ) + if isinstance(weight_edge_feat_name, str): + edge_type = next(iter(edge_info_by_type)) + weight_by_type = {edge_type: weight_edge_feat_name} + else: + weight_by_type = 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_metadata[edge_type] = FeatureQuantizationMetadata( + bits=metadata.bits, + feature_dim=metadata.feature_dim - 1, + quantized_feature_indices=tuple( + index - int(index > weight_logical_index) + for index in metadata.quantized_feature_indices + ), + clip_min=metadata.clip_min, + clip_max=metadata.clip_max, + neg_mean=metadata.neg_mean, + pos_mean=metadata.pos_mean, + ) + + if is_homogeneous: + return adjusted_metadata[DEFAULT_HOMOGENEOUS_EDGE_TYPE] + return adjusted_metadata def _data_loading_process( @@ -199,14 +331,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 +520,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 +654,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 +684,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/common/utils/feature_quantization/README.md b/gigl/common/utils/feature_quantization/README.md index 181548b31..b8ece0cdc 100644 --- a/gigl/common/utils/feature_quantization/README.md +++ b/gigl/common/utils/feature_quantization/README.md @@ -22,11 +22,11 @@ this as a useful tradeoff for GiGL. The built-in flow is: 1. The data preprocessor computes feature summary statistics offline. -2. The preprocessor quantizes selected scalar feature columns with NumPy. +2. The preprocessor quantizes selected scalar node or main-edge feature columns with NumPy. 3. The packed `uint8` feature sidecar is written to TFRecords. 4. Distributed dataset construction partitions and samples the packed bytes. 5. The dataloader collate path dequantizes sampled packed features with Torch. -6. Dequantized columns are scattered back into the logical `x` feature matrix. +6. Dequantized columns are scattered back into the logical `x` or `edge_attr` feature matrix. The NumPy/Torch split is intentional: diff --git a/gigl/distributed/base_dist_loader.py b/gigl/distributed/base_dist_loader.py index ea2e91bfc..5c27a648e 100644 --- a/gigl/distributed/base_dist_loader.py +++ b/gigl/distributed/base_dist_loader.py @@ -56,6 +56,7 @@ from gigl.distributed.utils.channel import MonitoredShmChannel from gigl.distributed.utils.neighborloader import ( DatasetSchema, + _map_to_effective_edge_types, attach_ppr_outputs, extract_edge_type_metadata, patch_fanout_for_sampling, @@ -244,8 +245,13 @@ def __init__( dataset_schema.is_homogeneous_with_labeled_edge_type ) self._node_feature_info = dataset_schema.node_feature_info - self._edge_feature_info = dataset_schema.edge_feature_info + self._edge_feature_info = _map_to_effective_edge_types( + dataset_schema.edge_feature_info, dataset_schema.edge_dir + ) self._node_quantization_metadata = dataset_schema.node_quantization_metadata + self._edge_quantization_metadata = _map_to_effective_edge_types( + dataset_schema.edge_quantization_metadata, dataset_schema.edge_dir + ) self._sampler_options = sampler_options self._non_blocking_transfers = non_blocking_transfers diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index 5ff21b1e3..4a27b583c 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -19,6 +19,7 @@ from gigl.common.logger import Logger from gigl.distributed.sampler import ( + EDGE_PACKED_FEATURES_METADATA_KEY, NEGATIVE_LABEL_METADATA_KEY, NODE_PACKED_FEATURES_METADATA_KEY, POSITIVE_LABEL_METADATA_KEY, @@ -116,6 +117,7 @@ def __init__(self, *args, **kwargs) -> None: self._sampling_error_sent: bool = False self.dist_node_quantized_feature: Optional[DistFeature] = None + self.dist_edge_quantized_feature: Optional[DistFeature] = None if ( self.collect_features and data is not None @@ -132,6 +134,20 @@ def __init__(self, *args, **kwargs) -> None: rpc_router=self.rpc_router, device=self.device, ) + if ( + self.collect_features + and data is not None + and getattr(data, "edge_quantized_features", None) is not None + ): + self.dist_edge_quantized_feature = DistFeature( + data.num_partitions, + data.partition_idx, + data.edge_quantized_features, + data.edge_pb, + local_only=False, + rpc_router=self.rpc_router, + device=self.device, + ) def _prepare_sample_loop_inputs( self, @@ -426,6 +442,20 @@ async def _collate_fn( futs[result_key] = wrap_torch_future( self.dist_edge_feature.async_get(eids, etype) ) + if self.dist_edge_quantized_feature is not None and self.with_edge: + for etype in self.edge_types: + result_edge_type = ( + reverse_edge_type(etype) if self.edge_dir == "in" else etype + ) + eids = result_map.get(f"{as_str(result_edge_type)}.eids") + if eids is not None: + futs[ + f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{result_edge_type}" + ] = wrap_torch_future( + self.dist_edge_quantized_feature.async_get( + eids.to(torch.long), etype + ) + ) if output.batch is not None: for ntype, batch in output.batch.items(): result_map[f"{as_str(ntype)}.batch"] = batch @@ -465,6 +495,10 @@ async def _collate_fn( futs["efeats"] = wrap_torch_future( self.dist_edge_feature.async_get(eids) ) + if self.dist_edge_quantized_feature is not None: + futs[f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}"] = wrap_torch_future( + self.dist_edge_quantized_feature.async_get(result_map["eids"]) + ) if output.batch is not None: result_map["batch"] = output.batch diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 1a5f859b5..07dace589 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -24,6 +24,7 @@ from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, TFDatasetOptions, + _remove_weight_from_edge_quantization_metadata, load_torch_tensors_from_tf_record, ) from gigl.common.logger import Logger @@ -194,6 +195,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 +217,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, @@ -227,6 +233,10 @@ def _load_and_build_partitioned_dataset( world_size=world_size, edge_dir=edge_dir, node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, + edge_quantization_metadata=_remove_weight_from_edge_quantization_metadata( + serialized_graph_metadata=serialized_graph_metadata, + weight_edge_feat_name=weight_edge_feat_name, + ), ) dataset.build( diff --git a/gigl/distributed/dist_ablp_neighborloader.py b/gigl/distributed/dist_ablp_neighborloader.py index 10638c330..b15d96dff 100644 --- a/gigl/distributed/dist_ablp_neighborloader.py +++ b/gigl/distributed/dist_ablp_neighborloader.py @@ -38,6 +38,7 @@ extract_edge_type_metadata, extract_metadata, labeled_to_homogeneous, + materialize_quantized_edge_features, materialize_quantized_node_features, set_missing_features, shard_nodes_by_process, @@ -609,6 +610,7 @@ def _setup_for_colocated( node_feature_info=dataset.node_feature_info, edge_feature_info=dataset.edge_feature_info, node_quantization_metadata=dataset.node_quantization_metadata, + edge_quantization_metadata=dataset.edge_quantization_metadata, edge_dir=dataset.edge_dir, ), ) @@ -799,6 +801,7 @@ def _setup_for_graph_store( node_feature_info=node_feature_info, edge_feature_info=edge_feature_info, node_quantization_metadata=dataset.fetch_node_quantization_metadata(), + edge_quantization_metadata=dataset.fetch_edge_quantization_metadata(), edge_dir=edge_dir, ), backend_key, @@ -972,6 +975,11 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: metadata=metadata, node_quantization_metadata=self._node_quantization_metadata, ) + data, metadata = materialize_quantized_edge_features( + data=data, + metadata=metadata, + edge_quantization_metadata=self._edge_quantization_metadata, + ) # Attach any remaining metadata (e.g. custom user-defined keys) directly onto the # data object so downstream code can access them via attribute lookup. diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index a41c0e12a..b1182a86c 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -4,7 +4,7 @@ import time from collections.abc import Mapping from multiprocessing.reduction import ForkingPickler -from typing import Literal, Optional, Tuple, TypeVar, Union, overload +from typing import Literal, Optional, Tuple, TypeVar, Union, cast, overload import graphlearn_torch as glt import torch @@ -58,6 +58,9 @@ def __init__( node_quantized_feature_partition: Optional[ Union[Feature, dict[NodeType, Feature]] ] = None, + edge_quantized_feature_partition: Optional[ + Union[Feature, dict[EdgeType, Feature]] + ] = None, edge_feature_partition: Optional[ Union[Feature, dict[EdgeType, Feature]] ] = None, @@ -87,6 +90,11 @@ def __init__( dict[NodeType, FeatureQuantizationMetadata], ] ] = None, + edge_quantization_metadata: Optional[ + Union[ + FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata] + ] + ] = None, edge_feature_info: Optional[ Union[FeatureInfo, dict[EdgeType, FeatureInfo]] ] = None, @@ -166,6 +174,8 @@ def __init__( self._node_quantized_features = node_quantized_feature_partition self._node_quantization_metadata = node_quantization_metadata + self._edge_quantized_features = edge_quantized_feature_partition + self._edge_quantization_metadata = edge_quantization_metadata self._degree_tensor: Optional[ Union[torch.Tensor, dict[NodeType, torch.Tensor]] @@ -253,6 +263,13 @@ def edge_features( ): self._edge_features = new_edge_features + @property + def edge_quantized_features( + self, + ) -> Optional[Union[Feature, dict[EdgeType, Feature]]]: + """Packed uint8 main-edge feature sidecar.""" + return self._edge_quantized_features + @property def node_pb( self, @@ -340,6 +357,15 @@ def node_quantization_metadata( ]: return self._node_quantization_metadata + @property + def edge_quantization_metadata( + self, + ) -> Optional[ + Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] + ]: + """Metadata required to materialize packed main-edge features.""" + return self._edge_quantization_metadata + @property def edge_feature_info( self, @@ -899,6 +925,42 @@ def _initialize_edge_features( ) logger.info(f"Initialized edge features for homogeneous graph to dataset") + def _initialize_edge_quantized_features( + self, + edge_partition_book: Union[PartitionBook, dict[EdgeType, PartitionBook]], + partitioned_edge_quantized_features: Optional[ + Union[FeaturePartitionData, dict[EdgeType, FeaturePartitionData]] + ], + ) -> None: + """Initialize packed uint8 main-edge feature storage.""" + features, id_to_index = _prepare_feature_data( + partition_book=edge_partition_book, + partitioned_data=partitioned_edge_quantized_features, + ) + if features is None or id_to_index is None: + return + if isinstance(features, Mapping): + assert isinstance(id_to_index, Mapping) + features = cast(dict[EdgeType, torch.Tensor], features) + id_to_index = cast(dict[EdgeType, torch.Tensor], id_to_index) + self._edge_quantized_features = { + edge_type: Feature( + feature_tensor=features_per_edge_type, + id2index=id_to_index[edge_type], + with_gpu=False, + dtype=torch.uint8, + ) + for edge_type, features_per_edge_type in features.items() + } + else: + assert not isinstance(id_to_index, Mapping) + self._edge_quantized_features = Feature( + feature_tensor=features, + id2index=id_to_index, + with_gpu=False, + dtype=torch.uint8, + ) + def build( self, partition_output: PartitionOutput, @@ -1011,6 +1073,13 @@ def build( partition_output.partitioned_edge_features = None gc.collect() + self._initialize_edge_quantized_features( + edge_partition_book=partition_output.edge_partition_book, + partitioned_edge_quantized_features=partition_output.partitioned_edge_quantized_features, + ) + partition_output.partitioned_edge_quantized_features = None + gc.collect() + self._node_partition_book = partition_output.node_partition_book self._edge_partition_book = partition_output.edge_partition_book @@ -1037,6 +1106,7 @@ def share_ipc( Optional[Union[Feature, dict[NodeType, Feature]]], Optional[Union[Feature, dict[NodeType, Feature]]], Optional[Union[Feature, dict[EdgeType, Feature]]], + Optional[Union[Feature, dict[EdgeType, Feature]]], Optional[Union[Feature, dict[NodeType, Feature]]], Optional[Union[PartitionBook, dict[NodeType, PartitionBook]]], Optional[Union[PartitionBook, dict[EdgeType, PartitionBook]]], @@ -1053,6 +1123,12 @@ def share_ipc( dict[NodeType, FeatureQuantizationMetadata], ] ], + Optional[ + Union[ + FeatureQuantizationMetadata, + dict[EdgeType, FeatureQuantizationMetadata], + ] + ], Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]], Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]], Optional[int], @@ -1067,6 +1143,7 @@ def share_ipc( Optional[Union[Graph, dict[EdgeType, Graph]]]: Partitioned Graph Data Optional[Union[Feature, dict[NodeType, Feature]]]: Partitioned Node Feature Data Optional[Union[Feature, dict[NodeType, Feature]]]: Partitioned packed uint8 node feature data + Optional[Union[Feature, dict[EdgeType, Feature]]]: Partitioned packed uint8 edge feature data Optional[Union[Feature, dict[EdgeType, Feature]]]: Partitioned Edge Feature Data Optional[Union[Feature, dict[NodeType, Feature]]]: Node labels on the current machine. Will be a dict if heterogeneous. Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Node Partition Book Tensor @@ -1079,6 +1156,7 @@ def share_ipc( Optional[Union[int, dict[NodeType, int]]]: Number of test nodes on the current machine. Will be a dict if heterogeneous. Optional[Union[FeatureInfo, dict[NodeType, FeatureInfo]]]: Node feature dim and its data type, will be a dict if heterogeneous Optional[Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]]]: Node quantization metadata. + Optional[Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]]]: Edge quantization metadata. Optional[Union[FeatureInfo, dict[EdgeType, FeatureInfo]]]: Edge feature dim and its data type, will be a dict if heterogeneous Optional[Union[torch.Tensor, dict[NodeType, torch.Tensor]]]: Degree tensors Optional[int]: Optional per-anchor label cap for ABLP label fetching @@ -1100,6 +1178,7 @@ def share_ipc( self._graph, self._node_features, self._node_quantized_features, + self._edge_quantized_features, self._edge_features, self._node_labels, self._node_partition_book, @@ -1112,6 +1191,7 @@ def share_ipc( self._num_test, # Additional field unique to DistDataset class self._node_feature_info, # Additional field unique to DistDataset class self._node_quantization_metadata, # Additional field unique to DistDataset class + self._edge_quantization_metadata, # Additional field unique to DistDataset class self._edge_feature_info, # Additional field unique to DistDataset class self._degree_tensor, # Additional field unique to DistDataset class self._max_labels_per_anchor_node, # Additional field unique to DistDataset class @@ -1348,6 +1428,9 @@ def _rebuild_distributed_dataset( Optional[ Union[Feature, dict[NodeType, Feature]] ], # Partitioned packed uint8 node feature data + Optional[ + Union[Feature, dict[EdgeType, Feature]] + ], # Partitioned packed uint8 edge feature data Optional[ Union[Feature, dict[EdgeType, Feature]] ], # Partitioned Edge Feature Data @@ -1377,6 +1460,12 @@ def _rebuild_distributed_dataset( dict[NodeType, FeatureQuantizationMetadata], ] ], # Node quantization metadata + Optional[ + Union[ + FeatureQuantizationMetadata, + dict[EdgeType, FeatureQuantizationMetadata], + ] + ], # Edge quantization metadata Optional[ Union[FeatureInfo, dict[EdgeType, FeatureInfo]] ], # Edge feature dim and its data type diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 04de8ce72..353978191 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -208,6 +208,11 @@ 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._partitioned_edge_quantized_features: dict[ + EdgeType, FeaturePartitionData + ] = {} 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 +674,26 @@ 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.") + packed_features = self._convert_edge_entity_to_heterogeneous_format( + input_edge_entity=edge_quantized_features + ) + if not packed_features: + raise ValueError("Edge quantized features cannot be empty.") + self._edge_quantized_feat = convert_to_tensor( + packed_features, dtype=torch.uint8 + ) + self._edge_quantized_feat_dim = { + edge_type: features.shape[1] + for edge_type, features in packed_features.items() + } + def register_edge_weights( self, edge_weights: Union[torch.Tensor, dict[EdgeType, torch.Tensor]] ) -> None: @@ -1225,11 +1250,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 @@ -1289,6 +1320,7 @@ def _edge_pfn(_, chunk_range): # IDs are always at r[-1]; features at r[0]; weights at r[1] when # features are also present, else r[0]. 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 +1341,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,6 +1350,11 @@ 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] @@ -1323,15 +1362,20 @@ def _edge_pfn(_, chunk_range): input_parts: list[torch.Tensor] = [] if edge_feat is not None: input_parts.append(edge_feat) + if edge_quantized_features is not None: + input_parts.append(edge_quantized_features) if edge_weights_tensor is not None: 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 + quantized_feat_idx: Optional[int] = 1 if has_edge_feats else 0 + if not has_edge_quantized_feats: + quantized_feat_idx = None weight_idx: Optional[int] = None if has_weights_for_edge_type: - weight_idx = 1 if has_edge_feats else 0 + weight_idx = int(has_edge_feats) + int(has_edge_quantized_feats) def _edge_feat_weight_pfn( ids_chunk: torch.Tensor, _: object @@ -1360,6 +1404,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 +1436,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 +1454,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 +1485,14 @@ def _edge_feat_weight_pfn( weights=partitioned_weights, ) - return current_graph_part, current_feat_part, edge_partition_book + if current_quantized_feat_part is not None: + self._partitioned_edge_quantized_features[edge_type] = ( + current_quantized_feat_part + ) + persistent_edge_partition_book = ( + edge_partition_book if has_edge_feats or has_edge_quantized_feats else None + ) + return current_graph_part, current_feat_part, persistent_edge_partition_book def _partition_label_edge_index( self, @@ -1757,9 +1839,9 @@ def 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 ) @@ -1936,6 +2018,12 @@ 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=( + to_homogeneous(self._partitioned_edge_quantized_features) + if self._is_input_homogeneous + and self._partitioned_edge_quantized_features + else self._partitioned_edge_quantized_features or None + ), 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..67f34c466 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -243,6 +243,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,12 +259,19 @@ 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] @@ -273,6 +284,10 @@ def _partition_edge_index_and_edge_features( 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 +316,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 +334,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 +346,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,6 +369,11 @@ 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() @@ -354,22 +391,29 @@ 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: 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 ) + if partitioned_edge_quantized_features is not None: + self._partitioned_edge_quantized_features[edge_type] = ( + FeaturePartitionData( + feats=partitioned_edge_quantized_features, ids=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}" ) @@ -457,9 +501,9 @@ def 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 ) diff --git a/gigl/distributed/distributed_neighborloader.py b/gigl/distributed/distributed_neighborloader.py index 3effa01c3..6764cd582 100644 --- a/gigl/distributed/distributed_neighborloader.py +++ b/gigl/distributed/distributed_neighborloader.py @@ -29,6 +29,7 @@ SamplingClusterSetup, extract_metadata, labeled_to_homogeneous, + materialize_quantized_edge_features, materialize_quantized_node_features, set_missing_features, shard_nodes_by_process, @@ -413,6 +414,7 @@ def _setup_for_graph_store( node_feature_info=node_feature_info, edge_feature_info=edge_feature_info, node_quantization_metadata=dataset.fetch_node_quantization_metadata(), + edge_quantization_metadata=dataset.fetch_edge_quantization_metadata(), edge_dir=dataset.fetch_edge_dir(), ), backend_key, @@ -531,6 +533,7 @@ def _setup_for_colocated( node_feature_info=dataset.node_feature_info, edge_feature_info=dataset.edge_feature_info, node_quantization_metadata=dataset.node_quantization_metadata, + edge_quantization_metadata=dataset.edge_quantization_metadata, edge_dir=dataset.edge_dir, ), ) @@ -564,6 +567,11 @@ def _collate_fn(self, msg: SampleMessage) -> Union[Data, HeteroData]: metadata=metadata, node_quantization_metadata=self._node_quantization_metadata, ) + data, metadata = materialize_quantized_edge_features( + data=data, + metadata=metadata, + edge_quantization_metadata=self._edge_quantization_metadata, + ) # Attach any remaining metadata (e.g. custom user-defined keys) directly onto the # data object so downstream code can access them via attribute lookup. diff --git a/gigl/distributed/graph_store/dist_server.py b/gigl/distributed/graph_store/dist_server.py index 0a92d959d..b608be433 100644 --- a/gigl/distributed/graph_store/dist_server.py +++ b/gigl/distributed/graph_store/dist_server.py @@ -418,6 +418,14 @@ def get_node_quantization_metadata( """Get node feature quantization metadata from the dataset.""" return self.dataset.node_quantization_metadata + def get_edge_quantization_metadata( + self, + ) -> Union[ + FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata], None + ]: + """Get main-edge feature quantization metadata from the dataset.""" + return self.dataset.edge_quantization_metadata + def get_edge_feature_info( self, ) -> Union[FeatureInfo, dict[EdgeType, FeatureInfo], None]: diff --git a/gigl/distributed/graph_store/remote_dist_dataset.py b/gigl/distributed/graph_store/remote_dist_dataset.py index 81609961b..127078fca 100644 --- a/gigl/distributed/graph_store/remote_dist_dataset.py +++ b/gigl/distributed/graph_store/remote_dist_dataset.py @@ -80,6 +80,14 @@ def fetch_node_quantization_metadata( DistServer.get_node_quantization_metadata, ) + def fetch_edge_quantization_metadata( + self, + ) -> Union[ + FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata], None + ]: + """Fetch main-edge feature quantization metadata from storage.""" + return request_server(0, DistServer.get_edge_quantization_metadata) + def fetch_edge_feature_info( self, ) -> Union[FeatureInfo, dict[EdgeType, FeatureInfo], None]: diff --git a/gigl/distributed/sampler.py b/gigl/distributed/sampler.py index 7789c6731..1e01ee85f 100644 --- a/gigl/distributed/sampler.py +++ b/gigl/distributed/sampler.py @@ -9,6 +9,7 @@ POSITIVE_LABEL_METADATA_KEY: Final[str] = "gigl_positive_labels." NEGATIVE_LABEL_METADATA_KEY: Final[str] = "gigl_negative_labels." NODE_PACKED_FEATURES_METADATA_KEY: Final[str] = "node_packed_features" +EDGE_PACKED_FEATURES_METADATA_KEY: Final[str] = "edge_packed_features" class ABLPNodeSamplerInput(NodeSamplerInput): diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 2e31d23dc..3dea282b6 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -9,13 +9,17 @@ import torch from graphlearn_torch.channel import SampleMessage +from graphlearn_torch.utils import reverse_edge_type from torch_geometric.data import Data, HeteroData -from torch_geometric.data.storage import NodeStorage +from torch_geometric.data.storage import EdgeStorage, NodeStorage from torch_geometric.typing import EdgeType, NodeType from gigl.common.logger import Logger from gigl.common.utils.feature_quantization.torch_ops import dequantize_torch_tensor -from gigl.distributed.sampler import NODE_PACKED_FEATURES_METADATA_KEY +from gigl.distributed.sampler import ( + EDGE_PACKED_FEATURES_METADATA_KEY, + NODE_PACKED_FEATURES_METADATA_KEY, +) from gigl.types.graph import ( FeatureInfo, FeatureQuantizationIndexTensors, @@ -26,6 +30,7 @@ logger = Logger() _GraphType = TypeVar("_GraphType", Data, HeteroData) +_EdgeMetadataValue = TypeVar("_EdgeMetadataValue") class SamplingClusterSetup(Enum): @@ -58,6 +63,23 @@ class DatasetSchema: node_quantization_metadata: Optional[ Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] ] = None + edge_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] + ] = None + + +def _map_to_effective_edge_types( + metadata: Optional[Union[_EdgeMetadataValue, dict[EdgeType, _EdgeMetadataValue]]], + edge_dir: Union[str, Literal["in", "out"]], +) -> Optional[Union[_EdgeMetadataValue, dict[EdgeType, _EdgeMetadataValue]]]: + """Map stored heterogeneous metadata to sampled edge-type direction.""" + if edge_dir != "in" or not isinstance(metadata, dict): + return metadata + typed_metadata = cast(dict[EdgeType, _EdgeMetadataValue], metadata) + return { + reverse_edge_type(edge_type): value + for edge_type, value in typed_metadata.items() + } def patch_fanout_for_sampling( @@ -395,6 +417,161 @@ def materialize( return data, metadata +def materialize_quantized_edge_features( + data: _GraphType, + metadata: dict[str, torch.Tensor], + edge_quantization_metadata: Optional[ + Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] + ], +) -> tuple[_GraphType, dict[str, torch.Tensor]]: + """Materialize packed quantized edge features into PyG edge attributes. + + Args: + data: Sampled homogeneous or heterogeneous PyG graph. + metadata: Sample metadata containing packed edge-feature tensors. + edge_quantization_metadata: Reconstruction metadata for each configured + edge store, or scalar metadata for a homogeneous graph. + + Returns: + The updated graph and metadata with packed edge-feature entries removed. + + Raises: + ValueError: If graph and metadata shapes disagree, packed data is missing + for sampled edges, or exact logical feature reconstruction is not + possible. + """ + if edge_quantization_metadata is None: + return data, metadata + + def materialize( + store: Union[Data, EdgeStorage], + packed_features: torch.Tensor, + quantization_metadata: FeatureQuantizationMetadata, + ) -> None: + if packed_features.ndim != 2: + raise ValueError( + "Expected packed edge features to be a 2-D tensor, got shape " + f"{tuple(packed_features.shape)}" + ) + if packed_features.dtype != torch.uint8: + raise ValueError( + "Expected packed edge features to use torch.uint8 storage, got " + f"{packed_features.dtype}" + ) + dequantized = dequantize_torch_tensor( + packed_features, metadata=quantization_metadata + ) + edge_attr = getattr(store, "edge_attr", None) + edge_index = getattr(store, "edge_index", None) + if edge_index is not None and edge_index.size(1) != packed_features.size(0): + raise ValueError( + f"Expected {edge_index.size(1)} packed edge feature rows, got " + f"{packed_features.size(0)}" + ) + if edge_attr is not None and edge_attr.size(0) != packed_features.size(0): + raise ValueError( + f"Expected {packed_features.size(0)} raw edge feature rows, got " + f"{edge_attr.size(0)}" + ) + output = dequantized.new_empty( + (dequantized.size(0), quantization_metadata.feature_dim) + ) + scatter_indices = quantization_metadata.scatter_index_tensors(output.device) + output[:, scatter_indices.quantized] = dequantized + if edge_attr is None and quantization_metadata.raw_feature_dim: + raise ValueError( + f"Missing {quantization_metadata.raw_feature_dim} unquantized edge features" + ) + if edge_attr is not None: + if edge_attr.size(1) != quantization_metadata.raw_feature_dim: + raise ValueError( + "Expected " + f"{quantization_metadata.raw_feature_dim} raw edge features before " + f"dequantization, got {edge_attr.size(1)}" + ) + output[:, scatter_indices.raw] = edge_attr + store.edge_attr = output + + if isinstance(data, Data): + if isinstance(edge_quantization_metadata, dict): + raise ValueError( + "Expected scalar quantization metadata for homogeneous data" + ) + packed_features = metadata.pop(EDGE_PACKED_FEATURES_METADATA_KEY, None) + if packed_features is None: + edge_index = getattr(data, "edge_index", None) + edge_attr = getattr(data, "edge_attr", None) + num_edges = ( + edge_index.size(1) + if edge_index is not None + else edge_attr.size(0) + if edge_attr is not None + else 0 + ) + if num_edges: + raise ValueError( + "Missing packed quantized features in metadata key " + f"{EDGE_PACKED_FEATURES_METADATA_KEY}" + ) + packed_features = torch.empty( + (0, edge_quantization_metadata.packed_feature_dim), + dtype=torch.uint8, + device=( + edge_attr.device + if edge_attr is not None + else edge_index.device + if edge_index is not None + else None + ), + ) + materialize(data, packed_features, edge_quantization_metadata) + else: + if not isinstance(edge_quantization_metadata, dict): + raise ValueError("Expected per-edge-type metadata for heterogeneous data") + edge_quantization_metadata = cast( + dict[EdgeType, FeatureQuantizationMetadata], edge_quantization_metadata + ) + for edge_type, quantization_metadata in edge_quantization_metadata.items(): + metadata_key = f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{edge_type}" + packed_features = metadata.pop(metadata_key, None) + if edge_type not in data.edge_types: + if packed_features is not None: + raise ValueError( + f"Found packed edge features for missing edge store {edge_type}" + ) + continue + + store = data[edge_type] + if packed_features is None: + edge_index = getattr(store, "edge_index", None) + edge_attr = getattr(store, "edge_attr", None) + num_edges = ( + edge_index.size(1) + if edge_index is not None + else edge_attr.size(0) + if edge_attr is not None + else 0 + ) + if num_edges: + raise ValueError( + "Missing packed quantized features in metadata key " + f"{metadata_key}" + ) + packed_features = torch.empty( + (0, quantization_metadata.packed_feature_dim), + dtype=torch.uint8, + device=( + edge_attr.device + if edge_attr is not None + else edge_index.device + if edge_index is not None + else None + ), + ) + materialize(store, packed_features, quantization_metadata) + return data, metadata + + def extract_metadata( msg: SampleMessage, device: torch.device ) -> tuple[dict[str, torch.Tensor], SampleMessage]: 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/src/data_preprocessor/data_preprocessor.py b/gigl/src/data_preprocessor/data_preprocessor.py index 9d84a8b42..4d2f2c328 100644 --- a/gigl/src/data_preprocessor/data_preprocessor.py +++ b/gigl/src/data_preprocessor/data_preprocessor.py @@ -216,13 +216,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 +433,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 +442,24 @@ 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: + with tf.io.gfile.GFile( + transformed_features_info.feature_quantization_metadata_path.uri + ) as metadata_file: + metadata = json.loads(metadata_file.read()) + quantization_metadata = preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( + packed_feature_key=metadata["packed_feature_key"], + quantized_feature_indices=metadata["quantized_feature_indices"], + ) + if metadata["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 = metadata["bits"] + quantization_metadata.multi_bit_state.clip_min = metadata["clip_min"] + quantization_metadata.multi_bit_state.clip_max = metadata["clip_max"] + output.quantized_feature_metadata.CopyFrom(quantization_metadata) + return output def generate_preprocessed_metadata_pb( self, @@ -782,6 +805,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..bd1bafc4d 100644 --- a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py +++ b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py @@ -16,6 +16,7 @@ logger = Logger() _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 = _NODE_PACKED_FEATURE_KEY, ) -> 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,13 +50,25 @@ 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. """ + if isinstance(logical_metadata, DatasetMetadata) and any( + feature.name == packed_feature_key + for feature in logical_metadata.schema.feature + ): + raise ValueError( + f"Reserved packed feature key {packed_feature_key} already exists in the logical schema." + ) missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) if missing: raise ValueError(f"Quantized features missing: {missing}") @@ -73,6 +89,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 +103,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 +161,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 +188,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 +197,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 +232,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 +248,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..0a4485403 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -28,6 +28,7 @@ NodeDataReference, ) from gigl.src.data_preprocessor.lib.transform.feature_quantization import ( + EDGE_PACKED_FEATURE_KEY, apply_feature_quantization_transform, ) from gigl.src.data_preprocessor.lib.transform.tf_value_encoder import TFValueEncoder @@ -372,7 +373,9 @@ 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: transformed_features, resolved_transformed_metadata = ( @@ -384,6 +387,11 @@ 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=( + EDGE_PACKED_FEATURE_KEY + if isinstance(preprocessing_spec, EdgeDataPreprocessingSpec) + else "node_packed_features" + ), ) ) 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/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/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): 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..fd3cfa4ad 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -19,6 +19,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( [ ( diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 046b8bf49..955f78e4c 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_features, dict): + edge_features_by_type = cast(dict[EdgeType, torch.Tensor], edge_features) + edge_quantized_features = { + edge_type: features.to(torch.uint8) + for edge_type, features in edge_features_by_type.items() + } + else: + edge_quantized_features = edge_features.to(torch.uint8) + 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, ): diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 3bfaff851..548c020df 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -18,6 +18,7 @@ ) from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, + _remove_weight_from_edge_quantization_metadata, load_torch_tensors_from_tf_record, ) from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper @@ -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,89 @@ 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.assertRaisesRegex(ValueError, "must remain an unquantized scalar"): + 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_is_removed_from_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_feature": tf.io.FixedLenFeature([], tf.float32), + "weight": tf.io.FixedLenFeature([], tf.float32), + "edge_packed_features": tf.io.FixedLenFeature([], tf.string), + }, + feature_keys=["raw_feature", "weight"], + feature_dim=2, + 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=(0, 2), + clip_min=0.0, + clip_max=3.0, + ), + ) + + adjusted_metadata = _remove_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=(0, 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..71845d444 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,21 @@ 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), + ) + elif ( input_data_strategy == InputDataStrategy.REGISTER_MINIMAL_ENTITIES_SEPARATELY ): diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index 20ad8b710..134fe5da4 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -7,15 +7,18 @@ from torch_geometric.typing import EdgeType from gigl.distributed.sampler import ( + EDGE_PACKED_FEATURES_METADATA_KEY, NEGATIVE_LABEL_METADATA_KEY, NODE_PACKED_FEATURES_METADATA_KEY, POSITIVE_LABEL_METADATA_KEY, ) from gigl.distributed.utils.neighborloader import ( + _map_to_effective_edge_types, attach_ppr_outputs, extract_edge_type_metadata, extract_metadata, labeled_to_homogeneous, + materialize_quantized_edge_features, materialize_quantized_node_features, patch_fanout_for_sampling, set_missing_features, @@ -101,6 +104,173 @@ def test_materialize_quantized_node_features_reconstructs_feature_order( self.assertEqual(set(remaining_metadata), {"request_id"}) self.assert_tensor_equality(remaining_metadata["request_id"], torch.tensor([7])) + def test_materialize_quantized_edge_features_reconstructs_feature_order( + self, + ) -> None: + data = Data(edge_attr=torch.tensor([[10.0, 20.0], [30.0, 40.0]])) + metadata = { + "edge_packed_features": torch.tensor([[48], [144]], dtype=torch.uint8), + "request_id": torch.tensor([7]), + } + + materialized_data, remaining_metadata = materialize_quantized_edge_features( + data=data, + metadata=metadata, + edge_quantization_metadata=FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ), + ) + + self.assert_tensor_equality( + materialized_data.edge_attr, + torch.tensor([[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]]), + ) + self.assertEqual(set(remaining_metadata), {"request_id"}) + + def test_materialize_quantized_edge_features_uses_effective_edge_type( + self, + ) -> None: + edge_type = ("item", "rev_to", "user") + data = HeteroData() + data[edge_type].edge_attr = torch.tensor([[10.0]]) + metadata = { + f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{edge_type}": torch.tensor( + [[48]], dtype=torch.uint8 + ) + } + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + + materialized_data, remaining_metadata = materialize_quantized_edge_features( + data=data, + metadata=metadata, + edge_quantization_metadata={edge_type: quantization_metadata}, + ) + + self.assert_tensor_equality( + materialized_data[edge_type].edge_attr, + torch.tensor([[0.0, 10.0, 3.0]]), + ) + self.assertEqual(remaining_metadata, {}) + + def test_map_to_effective_edge_types_reverses_inbound_metadata(self) -> None: + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=2, + quantized_feature_indices=(0, 1), + clip_min=0.0, + clip_max=3.0, + ) + + effective_metadata = _map_to_effective_edge_types( + {_U2I_EDGE_TYPE: quantization_metadata}, edge_dir="in" + ) + + self.assertEqual( + effective_metadata, + {("item", "rev_to", "user"): quantization_metadata}, + ) + + def test_materialize_quantized_edge_features_rejects_malformed_sidecars( + self, + ) -> None: + edge_type = ("user", "to", "item") + quantization_metadata = FeatureQuantizationMetadata( + bits=2, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + metadata_key = f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{edge_type}" + + def edge_data(raw_features: torch.Tensor, num_edges: int = 1) -> HeteroData: + data = HeteroData() + data[edge_type].edge_index = torch.zeros((2, num_edges), dtype=torch.long) + data[edge_type].edge_attr = raw_features + return data + + cases = [ + ( + edge_data(torch.tensor([[10.0]])), + {}, + ), + ( + edge_data(torch.tensor([[10.0]])), + {metadata_key: torch.tensor([[48, 0]], dtype=torch.uint8)}, + ), + ( + edge_data(torch.tensor([[10.0], [20.0]]), num_edges=2), + {metadata_key: torch.tensor([[48]], dtype=torch.uint8)}, + ), + ( + edge_data(torch.tensor([[10.0, 20.0]])), + {metadata_key: torch.tensor([[48]], dtype=torch.uint8)}, + ), + ] + + for data, metadata in cases: + with self.subTest(metadata=metadata), self.assertRaises(ValueError): + materialize_quantized_edge_features( + data=data, + metadata=metadata, + edge_quantization_metadata={edge_type: quantization_metadata}, + ) + + def test_materialize_quantized_edge_features_preserves_empty_shape( + self, + ) -> None: + edge_type = ("user", "to", "item") + data = HeteroData() + data[edge_type].edge_index = torch.empty((2, 0), dtype=torch.long) + data[edge_type].edge_attr = torch.empty((0, 1)) + + materialized_data, remaining_metadata = materialize_quantized_edge_features( + data=data, + metadata={}, + edge_quantization_metadata={ + edge_type: FeatureQuantizationMetadata( + bits=2, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + }, + ) + + self.assertEqual( + materialized_data[edge_type].edge_attr.shape, torch.Size([0, 3]) + ) + self.assertEqual(materialized_data[edge_type].edge_attr.dtype, torch.float32) + self.assertEqual(remaining_metadata, {}) + + homogeneous_data, _ = materialize_quantized_edge_features( + data=Data( + edge_index=torch.empty((2, 0), dtype=torch.long), + edge_attr=torch.empty((0, 1)), + ), + metadata={}, + edge_quantization_metadata=FeatureQuantizationMetadata( + bits=2, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ), + ) + self.assertEqual(homogeneous_data.edge_attr.shape, torch.Size([0, 3])) + self.assertEqual(homogeneous_data.edge_attr.dtype, torch.float32) + def test_materialize_quantized_node_features_uses_per_node_type_metadata( self, ) -> None: From 40e1460f8c2921a75144ba3ea5f7ed0ebe8ba3e9 Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 18:31:30 +0000 Subject: [PATCH 59/78] Update --- gigl/distributed/utils/neighborloader.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 2e31d23dc..03b0c4054 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -17,6 +17,7 @@ from gigl.common.utils.feature_quantization.torch_ops import dequantize_torch_tensor from gigl.distributed.sampler import NODE_PACKED_FEATURES_METADATA_KEY from gigl.types.graph import ( + DEFAULT_HOMOGENEOUS_NODE_TYPE, FeatureInfo, FeatureQuantizationIndexTensors, FeatureQuantizationMetadata, @@ -374,9 +375,18 @@ def materialize( if isinstance(node_quantization_metadata, dict): raise ValueError("Expect scalar quantization metadata for homogeneous data") packed_features = metadata.pop(NODE_PACKED_FEATURES_METADATA_KEY, None) + labeled_homogeneous_packed_features_key = ( + f"{NODE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_NODE_TYPE}" + ) + if packed_features is None: + # Labeled homogeneous graphs are sampled as heterogeneous graphs, so + # the packed-feature transport key retains the default node type. + packed_features = metadata.pop( + labeled_homogeneous_packed_features_key, None + ) if packed_features is None: raise ValueError( - f"Missing packed quantized features in metadata key {NODE_PACKED_FEATURES_METADATA_KEY}" + f"Missing packed quantized features in metadata keys {NODE_PACKED_FEATURES_METADATA_KEY} or {labeled_homogeneous_packed_features_key}" ) materialize(data, packed_features, node_quantization_metadata) else: From dcb77222c8308a536004d543d8159837f04a13f4 Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 20:02:41 +0000 Subject: [PATCH 60/78] Improve test readability --- .../distributed_neighborloader_test.py | 27 +++++++++++++++---- 1 file changed, 22 insertions(+), 5 deletions(-) diff --git a/tests/unit/distributed/distributed_neighborloader_test.py b/tests/unit/distributed/distributed_neighborloader_test.py index c9caee50f..9390c4168 100644 --- a/tests/unit/distributed/distributed_neighborloader_test.py +++ b/tests/unit/distributed/distributed_neighborloader_test.py @@ -427,7 +427,11 @@ def _run_featureless_edge_ids_absent( shutdown_rpc() -def _run_quantized_feature_neighbor_loader(_: int, dataset: DistDataset) -> None: +def _run_quantized_feature_neighbor_loader( + _: int, + dataset: DistDataset, + expected_features: torch.Tensor, +) -> None: create_test_process_group() loader = DistNeighborLoader( dataset=dataset, @@ -437,7 +441,6 @@ def _run_quantized_feature_neighbor_loader(_: int, dataset: DistDataset) -> None pin_memory_device=torch.device("cpu"), ) - expected_features = torch.tensor([[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]]) batch_count = 0 for batch in loader: assert isinstance(batch, Data) @@ -801,6 +804,17 @@ def test_isolated_homogeneous_neighbor_loader( def test_distributed_neighbor_loader_materializes_quantized_node_features( self, ) -> None: + raw_node_features = torch.tensor([[10.0, 20.0], [30.0, 40.0]]) + # Each row is one node with one packed uint8. Its two high-order 2-bit + # codes are scattered into feature indices 0 and 2; the remaining two + # codes are padding. 48 (0b00_11_00_00) yields [0, 3], while 144 + # (0b10_01_00_00) yields [2, 1]. + packed_quantized_node_features = torch.tensor( + [[48], [144]], dtype=torch.uint8 + ) + expected_features = torch.tensor( + [[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]] + ) partition_output = PartitionOutput( node_partition_book=torch.zeros(2), edge_partition_book=torch.zeros(2), @@ -808,11 +822,11 @@ def test_distributed_neighbor_loader_materializes_quantized_node_features( edge_index=torch.tensor([[0, 1], [1, 0]]), edge_ids=None ), partitioned_node_features=FeaturePartitionData( - feats=torch.tensor([[10.0, 20.0], [30.0, 40.0]]), + feats=raw_node_features, ids=torch.arange(2), ), partitioned_node_quantized_features=FeaturePartitionData( - feats=torch.tensor([[48], [144]], dtype=torch.uint8), + feats=packed_quantized_node_features, ids=torch.arange(2), ), partitioned_edge_features=None, @@ -834,7 +848,10 @@ def test_distributed_neighbor_loader_materializes_quantized_node_features( ) dataset.build(partition_output=partition_output) - mp.spawn(fn=_run_quantized_feature_neighbor_loader, args=(dataset,)) + mp.spawn( + fn=_run_quantized_feature_neighbor_loader, + args=(dataset, expected_features), + ) @parameterized.expand( [ From 5de7dd9b3ccb24eeb59a9f643802ba08013d042a Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 20:20:46 +0000 Subject: [PATCH 61/78] Fix heterogenous sampler collate with only partially quantized node types --- gigl/distributed/base_sampler.py | 35 ++++++-- .../distributed_neighborloader_test.py | 83 ++++++++++++++++++- 2 files changed, 110 insertions(+), 8 deletions(-) diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index 5ff21b1e3..67ab6d183 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -1,7 +1,7 @@ import asyncio import traceback from collections import defaultdict -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import Optional, Union import torch @@ -377,14 +377,29 @@ async def _collate_fn( ] if self.dist_node_feature is not None: if self.use_all2all: - sorted_ntype = sorted(self.dist_node_feature.feature_pb.keys()) + sorted_ntype = sorted(self.dist_node_feature.local_feature.keys()) + # GLT get_all2all() iterates every type in output.node, not + # just sorted_ntype. feature_pb contains partition books for + # every node type, while local_feature contains only types + # registered in this feature store, such as when only some + # heterogeneous node types have quantized features. + feature_output = replace( + output, + node={ + ntype: nodes + for ntype, nodes in output.node.items() + if ntype in sorted_ntype + }, + ) nfeat_dict = self.dist_node_feature.get_all2all( - output, sorted_ntype + feature_output, sorted_ntype ) for ntype, nfeats in nfeat_dict.items(): result_map[f"{as_str(ntype)}.nfeats"] = nfeats else: for ntype, nodes in output.node.items(): + if ntype not in self.dist_node_feature.local_feature: + continue nodes = nodes.to(torch.long) futs[f"{as_str(ntype)}.nfeats"] = wrap_torch_future( self.dist_node_feature.async_get(nodes, ntype) @@ -392,10 +407,18 @@ async def _collate_fn( if self.dist_node_quantized_feature is not None: if self.use_all2all: sorted_ntype = sorted( - self.dist_node_quantized_feature.feature_pb.keys() + self.dist_node_quantized_feature.local_feature.keys() + ) + feature_output = replace( + output, + node={ + ntype: nodes + for ntype, nodes in output.node.items() + if ntype in sorted_ntype + }, ) quantized_nfeat_dict = self.dist_node_quantized_feature.get_all2all( - output, sorted_ntype + feature_output, sorted_ntype ) for ntype, quantized_nfeats in quantized_nfeat_dict.items(): result_map[ @@ -403,6 +426,8 @@ async def _collate_fn( ] = quantized_nfeats else: for ntype, nodes in output.node.items(): + if ntype not in self.dist_node_quantized_feature.local_feature: + continue nodes = nodes.to(torch.long) futs[ f"#META.{NODE_PACKED_FEATURES_METADATA_KEY}.{as_str(ntype)}" diff --git a/tests/unit/distributed/distributed_neighborloader_test.py b/tests/unit/distributed/distributed_neighborloader_test.py index 9390c4168..e51f10ea0 100644 --- a/tests/unit/distributed/distributed_neighborloader_test.py +++ b/tests/unit/distributed/distributed_neighborloader_test.py @@ -450,6 +450,27 @@ def _run_quantized_feature_neighbor_loader( shutdown_rpc() +def _run_heterogeneous_partially_quantized_neighbor_loader( + _: int, + dataset: DistDataset, + expected_features: dict[NodeType, torch.Tensor], +) -> None: + create_test_process_group() + loader = DistNeighborLoader( + dataset=dataset, + input_nodes=(_USER, torch.tensor([0])), + num_neighbors=[1], + batch_size=1, + pin_memory_device=torch.device("cpu"), + ) + + batch = next(iter(loader)) + assert isinstance(batch, HeteroData) + for node_type, expected_features_for_node_type in expected_features.items(): + assert_tensor_equality(batch[node_type].x, expected_features_for_node_type) + shutdown_rpc() + + class DistributedNeighborLoaderTest(TestCase): def setUp(self): super().setUp() @@ -809,9 +830,7 @@ def test_distributed_neighbor_loader_materializes_quantized_node_features( # codes are scattered into feature indices 0 and 2; the remaining two # codes are padding. 48 (0b00_11_00_00) yields [0, 3], while 144 # (0b10_01_00_00) yields [2, 1]. - packed_quantized_node_features = torch.tensor( - [[48], [144]], dtype=torch.uint8 - ) + packed_quantized_node_features = torch.tensor([[48], [144]], dtype=torch.uint8) expected_features = torch.tensor( [[0.0, 10.0, 3.0, 20.0], [2.0, 30.0, 1.0, 40.0]] ) @@ -853,6 +872,64 @@ def test_distributed_neighbor_loader_materializes_quantized_node_features( args=(dataset, expected_features), ) + def test_heterogeneous_loader_supports_partially_quantized_node_types( + self, + ) -> None: + expected_features = { + _USER: torch.tensor([[0.0, 10.0]]), + _STORY: torch.tensor([[20.0]]), + } + partition_output = PartitionOutput( + node_partition_book={ + _USER: torch.zeros(1), + _STORY: torch.zeros(1), + }, + edge_partition_book={_USER_TO_STORY: torch.zeros(1)}, + partitioned_edge_index={ + _USER_TO_STORY: GraphPartitionData( + edge_index=torch.tensor([[0], [0]]), edge_ids=None + ) + }, + partitioned_node_features={ + _USER: FeaturePartitionData( + feats=torch.tensor([[10.0]]), ids=torch.tensor([0]) + ), + _STORY: FeaturePartitionData( + feats=torch.tensor([[20.0]]), ids=torch.tensor([0]) + ), + }, + partitioned_node_quantized_features={ + _USER: FeaturePartitionData( + feats=torch.tensor([[0]], dtype=torch.uint8), + ids=torch.tensor([0]), + ) + }, + partitioned_edge_features=None, + partitioned_positive_labels=None, + partitioned_negative_labels=None, + partitioned_node_labels=None, + ) + dataset = DistDataset( + rank=0, + world_size=1, + edge_dir="out", + node_quantization_metadata={ + _USER: FeatureQuantizationMetadata( + bits=2, + feature_dim=2, + quantized_feature_indices=(0,), + clip_min=0.0, + clip_max=3.0, + ) + }, + ) + dataset.build(partition_output=partition_output) + + mp.spawn( + fn=_run_heterogeneous_partially_quantized_neighbor_loader, + args=(dataset, expected_features), + ) + @parameterized.expand( [ param( From a5281645028513af9f15ebe6b1fcccf672401cd6 Mon Sep 17 00:00:00 2001 From: jchmura Date: Wed, 12 Aug 2026 22:59:30 +0000 Subject: [PATCH 62/78] Remove planning docs --- CONTEXT.md | 30 ---------- .../node_feature_quantization.md | 60 ------------------- docs/user_guide/index.rst | 1 - 3 files changed, 91 deletions(-) delete mode 100644 CONTEXT.md delete mode 100644 docs/user_guide/config_guides/node_feature_quantization.md diff --git a/CONTEXT.md b/CONTEXT.md deleted file mode 100644 index eaa278572..000000000 --- a/CONTEXT.md +++ /dev/null @@ -1,30 +0,0 @@ -# Graph Feature Quantization - -This context defines the storage optimization used for graph features while preserving the feature tensors consumed by -models. - -## Language - -**Feature quantization**: An opt-in representation that stores selected logical feature fields as packed low-bit values -and reconstructs floating-point tensors before model consumption. - -**Logical feature**: A model-facing feature in its original position and floating-point representation, independent of -its stored representation. - -**Packed feature**: The serialized and distributed-storage representation of one or more quantized logical features. - -**Raw feature sidecar**: The unquantized floating-point columns retained beside a packed feature for the same entities. - -**Quantization metadata**: The bit width, logical feature positions, and dequantization state required to reconstruct a -logical feature vector. - -**Main edge**: An edge in the graph used by neighbor sampling and message passing. _Avoid_: Training edge, regular edge - -**Supervision edge**: A positive or negative source-destination pair used as a link-prediction label rather than as part -of the sampled message-passing feature store. _Avoid_: Main edge, message-passing edge - -**Sampling weight**: A raw scalar edge value consumed before neighbor selection, distinct from model-facing edge -features materialized after sampling. - -**Transparent materialization**: Reconstruction of packed features into their logical floating-point tensor and feature -order before a sampled batch reaches the model. _Avoid_: Model-side dequantization diff --git a/docs/user_guide/config_guides/node_feature_quantization.md b/docs/user_guide/config_guides/node_feature_quantization.md deleted file mode 100644 index 0940fc942..000000000 --- a/docs/user_guide/config_guides/node_feature_quantization.md +++ /dev/null @@ -1,60 +0,0 @@ -# Node Feature Quantization - -Quantization stores selected scalar floating-point node features as packed low-bit values. GiGL reconstructs approximate -floating-point values before the model receives a batch. - -## Enable quantization - -Add `FeatureQuantizationSpec` to the `NodeDataPreprocessingSpec` for each node type you want to quantize, then rerun -data preprocessing. - -```python -from gigl.src.data_preprocessor.lib.types import ( - FeatureQuantizationSpec, - NodeDataPreprocessingSpec, - NodeOutputIdentifier, -) - -node_preprocessing_spec = NodeDataPreprocessingSpec( - identifier_output=NodeOutputIdentifier("node_id"), - features_outputs=["embedding_0", "embedding_1", "embedding_2"], - feature_spec_fn=feature_spec_fn, - preprocessing_fn=preprocessing_fn, - feature_quantization_spec=FeatureQuantizationSpec( - feature_keys=["embedding_0", "embedding_1", "embedding_2"], - bits=4, - ), -) -``` - -`feature_keys` must name distinct scalar output features from `features_outputs`. The supported bit widths are `1`, `2`, -`4`, and `8`. - -No model, trainer, sampler, or inference changes are required. GiGL restores the original feature-vector order and -dimension before the batch reaches the model. - -## Choose features and bit width - -- Select only scalar floating-point outputs. IDs, labels, and non-scalar outputs cannot be quantized. -- Use `8` bits when input precision is more important than payload reduction. Use fewer bits only after evaluating the - task metric with quantization enabled. -- At `2`, `4`, and `8` bits, GiGL clips every selected feature to one shared range and maps it to uniform levels. -- At `1` bit, GiGL stores only the sign and reconstructs values using the global mean of positive or non-positive - values. This is the most lossy option. - -## Expected benefit - -The packed payload for selected features is smaller than 32-bit floats by this factor when the number of selected -features fills whole bytes: - -| Bit width | Packed values per byte | Maximum reduction for selected features | -| --------- | ---------------------- | --------------------------------------- | -| 8 | 1 | 4x | -| 4 | 2 | 8x | -| 2 | 4 | 16x | -| 1 | 8 | 32x | - -The final packed byte is padded when the selected feature count does not fill it, so the actual payload reduction can be -smaller. GiGL does not currently publish an end-to-end storage, throughput, or model-quality guarantee. - -TODO: Add workload benchmarks for storage, transfer, and task-quality impact. diff --git a/docs/user_guide/index.rst b/docs/user_guide/index.rst index e414b73fe..e8b2950fa 100644 --- a/docs/user_guide/index.rst +++ b/docs/user_guide/index.rst @@ -39,7 +39,6 @@ Welcome to the GiGL User Guide. This guide provides detailed documentation to he config_guides/resource_config_guide config_guides/task_config_guide config_guides/data_preprocessor_spec_guide - config_guides/node_feature_quantization .. toctree:: From d14118536c40101795fc0adc3c033a968c6552d1 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 17:18:28 +0000 Subject: [PATCH 63/78] Simple improvements --- gigl/distributed/base_sampler.py | 3 ++- gigl/distributed/dist_dataset.py | 1 + .../lib/transform/feature_quantization.py | 11 ++--------- gigl/src/data_preprocessor/lib/transform/utils.py | 11 ++++++----- .../feature_quantization_transform_test.py | 2 ++ 5 files changed, 13 insertions(+), 15 deletions(-) diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index 0b6cb23a5..ff73dcbcb 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -480,11 +480,12 @@ async def _collate_fn( ) eids = result_map.get(f"{as_str(result_edge_type)}.eids") if eids is not None: + eids = eids.to(torch.long) futs[ f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{result_edge_type}" ] = wrap_torch_future( self.dist_edge_quantized_feature.async_get( - eids.to(torch.long), etype + eids, etype ) ) if output.batch is not None: diff --git a/gigl/distributed/dist_dataset.py b/gigl/distributed/dist_dataset.py index b1182a86c..3ad3a7167 100644 --- a/gigl/distributed/dist_dataset.py +++ b/gigl/distributed/dist_dataset.py @@ -938,6 +938,7 @@ def _initialize_edge_quantized_features( partitioned_data=partitioned_edge_quantized_features, ) if features is None or id_to_index is None: + logger.info("Found no packed quantized edge features to initialize") return if isinstance(features, Mapping): assert isinstance(id_to_index, Mapping) diff --git a/gigl/src/data_preprocessor/lib/transform/feature_quantization.py b/gigl/src/data_preprocessor/lib/transform/feature_quantization.py index bd1bafc4d..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,7 @@ 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] @@ -26,7 +26,7 @@ def apply_feature_quantization_transform( logical_feature_keys: list[str], quantization_spec: FeatureQuantizationSpec, quantization_metadata_path: str, - packed_feature_key: str = _NODE_PACKED_FEATURE_KEY, + packed_feature_key: str, ) -> tuple[beam.PCollection[pa.RecordBatch], DatasetMetadata | beam.pvalue.AsSingleton]: """Quantizes selected feature columns and bit-packs each record's values. @@ -62,13 +62,6 @@ def apply_feature_quantization_transform( ValueError: If the reserved packed key already exists, a selected feature is absent or non-scalar, or feature values cannot be quantized. """ - if isinstance(logical_metadata, DatasetMetadata) and any( - feature.name == packed_feature_key - for feature in logical_metadata.schema.feature - ): - raise ValueError( - f"Reserved packed feature key {packed_feature_key} already exists in the logical schema." - ) missing = set(quantization_spec.feature_keys) - set(logical_feature_keys) if missing: raise ValueError(f"Quantized features missing: {missing}") diff --git a/gigl/src/data_preprocessor/lib/transform/utils.py b/gigl/src/data_preprocessor/lib/transform/utils.py index 0a4485403..39a3ccc7b 100644 --- a/gigl/src/data_preprocessor/lib/transform/utils.py +++ b/gigl/src/data_preprocessor/lib/transform/utils.py @@ -29,6 +29,7 @@ ) 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 @@ -378,6 +379,10 @@ def get_load_data_and_transform_pipeline_component( ): 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, @@ -387,11 +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=( - EDGE_PACKED_FEATURE_KEY - if isinstance(preprocessing_spec, EdgeDataPreprocessingSpec) - else "node_packed_features" - ), + packed_feature_key=packed_feature_key, ) ) 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 fd3cfa4ad..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 @@ -121,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 7dfee9c356f8124fc14ca50e6b4dc29ca70c9b71 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 18:00:00 +0000 Subject: [PATCH 64/78] Clarify sampling weight quantization metadata --- gigl/common/data/load_torch_tensors.py | 31 +++++++++++++++------- gigl/distributed/dataset_factory.py | 4 +-- tests/unit/common/data/dataloaders_test.py | 4 +-- 3 files changed, 26 insertions(+), 13 deletions(-) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 3db6e1dfb..4087aceba 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -184,20 +184,30 @@ def _validate_weight_edge_feature_name( ) -def _remove_weight_from_edge_quantization_metadata( +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 the separately stored sampling-weight column from model metadata.""" + """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 - edge_info_by_type: dict[EdgeType, SerializedTFRecordInfo] - metadata_by_type: dict[EdgeType, FeatureQuantizationMetadata] - weight_by_type: dict[EdgeType, str] if isinstance(serialized_graph_metadata.edge_entity_info, SerializedTFRecordInfo): assert isinstance(quantization_metadata, FeatureQuantizationMetadata) assert isinstance(weight_edge_feat_name, str) @@ -235,13 +245,16 @@ def _remove_weight_from_edge_quantization_metadata( 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] = FeatureQuantizationMetadata( bits=metadata.bits, feature_dim=metadata.feature_dim - 1, - quantized_feature_indices=tuple( - index - int(index > weight_logical_index) - for index in metadata.quantized_feature_indices - ), + quantized_feature_indices=adjusted_quantized_feature_indices, clip_min=metadata.clip_min, clip_max=metadata.clip_max, neg_mean=metadata.neg_mean, diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index 07dace589..ba24a708d 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -24,7 +24,7 @@ from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, TFDatasetOptions, - _remove_weight_from_edge_quantization_metadata, + remove_sampling_weight_from_edge_quantization_metadata, load_torch_tensors_from_tf_record, ) from gigl.common.logger import Logger @@ -233,7 +233,7 @@ def _load_and_build_partitioned_dataset( world_size=world_size, edge_dir=edge_dir, node_quantization_metadata=serialized_graph_metadata.node_quantization_metadata, - edge_quantization_metadata=_remove_weight_from_edge_quantization_metadata( + edge_quantization_metadata=remove_sampling_weight_from_edge_quantization_metadata( serialized_graph_metadata=serialized_graph_metadata, weight_edge_feat_name=weight_edge_feat_name, ), diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 548c020df..4998dec40 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -18,7 +18,7 @@ ) from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, - _remove_weight_from_edge_quantization_metadata, + remove_sampling_weight_from_edge_quantization_metadata, load_torch_tensors_from_tf_record, ) from gigl.src.common.types.pb_wrappers.gbml_config import GbmlConfigPbWrapper @@ -713,7 +713,7 @@ def test_sampling_weight_is_removed_from_edge_quantization_metadata(self) -> Non ), ) - adjusted_metadata = _remove_weight_from_edge_quantization_metadata( + adjusted_metadata = remove_sampling_weight_from_edge_quantization_metadata( serialized_graph_metadata=serialized_graph_metadata, weight_edge_feat_name="weight", ) From 10c6574929c54ede5c7bf9ef7542add64a24d035 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 18:03:58 +0000 Subject: [PATCH 65/78] Simplify sampling weight metadata adjustment --- gigl/common/data/load_torch_tensors.py | 30 ++++++++++++++------------ 1 file changed, 16 insertions(+), 14 deletions(-) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index 4087aceba..bf55a7ba7 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -1,6 +1,6 @@ import time import traceback -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import MutableMapping, Optional, Union, cast import torch @@ -211,23 +211,29 @@ def remove_sampling_weight_from_edge_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 = { + edge_info_by_type: dict[EdgeType, SerializedTFRecordInfo] = { DEFAULT_HOMOGENEOUS_EDGE_TYPE: serialized_graph_metadata.edge_entity_info } - metadata_by_type = {DEFAULT_HOMOGENEOUS_EDGE_TYPE: quantization_metadata} - weight_by_type = {DEFAULT_HOMOGENEOUS_EDGE_TYPE: weight_edge_feat_name} + 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 = serialized_graph_metadata.edge_entity_info - metadata_by_type = cast( + 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 = {edge_type: weight_edge_feat_name} + weight_by_type: dict[EdgeType, str] = {edge_type: weight_edge_feat_name} else: - weight_by_type = weight_edge_feat_name + weight_by_type: dict[EdgeType, str] = weight_edge_feat_name is_homogeneous = False adjusted_metadata: dict[EdgeType, FeatureQuantizationMetadata] = {} @@ -251,14 +257,10 @@ def remove_sampling_weight_from_edge_quantization_metadata( else quantized_feature_index for quantized_feature_index in metadata.quantized_feature_indices ) - adjusted_metadata[edge_type] = FeatureQuantizationMetadata( - bits=metadata.bits, + adjusted_metadata[edge_type] = replace( + metadata, feature_dim=metadata.feature_dim - 1, quantized_feature_indices=adjusted_quantized_feature_indices, - clip_min=metadata.clip_min, - clip_max=metadata.clip_max, - neg_mean=metadata.neg_mean, - pos_mean=metadata.pos_mean, ) if is_homogeneous: From 828ef2dd946bdcbb8f157001a745d3520ba570c6 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 18:16:29 +0000 Subject: [PATCH 66/78] Inline sampled edge metadata mapping --- gigl/distributed/base_dist_loader.py | 29 ++++++++++++++----- gigl/distributed/utils/neighborloader.py | 17 ----------- tests/unit/common/data/dataloaders_test.py | 14 +++++---- .../distributed/utils/neighborloader_test.py | 19 ------------ 4 files changed, 30 insertions(+), 49 deletions(-) diff --git a/gigl/distributed/base_dist_loader.py b/gigl/distributed/base_dist_loader.py index 793107c7a..1c070959d 100644 --- a/gigl/distributed/base_dist_loader.py +++ b/gigl/distributed/base_dist_loader.py @@ -23,6 +23,7 @@ get_context, ) from graphlearn_torch.distributed.rpc import rpc_is_initialized +from graphlearn_torch.utils import reverse_edge_type from graphlearn_torch.sampler import ( EdgeSamplerInput, NodeSamplerInput, @@ -56,7 +57,6 @@ from gigl.distributed.utils.channel import MonitoredShmChannel from gigl.distributed.utils.neighborloader import ( DatasetSchema, - _map_to_effective_edge_types, attach_ppr_outputs, extract_edge_type_metadata, patch_fanout_for_sampling, @@ -245,13 +245,28 @@ def __init__( dataset_schema.is_homogeneous_with_labeled_edge_type ) self._node_feature_info = dataset_schema.node_feature_info - self._edge_feature_info = _map_to_effective_edge_types( - dataset_schema.edge_feature_info, dataset_schema.edge_dir - ) + # GLT returns heterogeneous edge stores in the sampled direction. For + # incoming sampling, that reverses each stored edge type, so feature + # metadata must use the same keys as the returned stores. Otherwise + # empty stores miss their feature shape and packed features cannot be + # reconstructed into their sampled edge attributes. + edge_feature_info = dataset_schema.edge_feature_info + if dataset_schema.edge_dir == "in" and isinstance(edge_feature_info, dict): + edge_feature_info = { + reverse_edge_type(edge_type): feature_info + for edge_type, feature_info in edge_feature_info.items() + } + self._edge_feature_info = edge_feature_info self._node_quantization_metadata = dataset_schema.node_quantization_metadata - self._edge_quantization_metadata = _map_to_effective_edge_types( - dataset_schema.edge_quantization_metadata, dataset_schema.edge_dir - ) + edge_quantization_metadata = dataset_schema.edge_quantization_metadata + if dataset_schema.edge_dir == "in" and isinstance( + edge_quantization_metadata, dict + ): + edge_quantization_metadata = { + reverse_edge_type(edge_type): quantization_metadata + for edge_type, quantization_metadata in edge_quantization_metadata.items() + } + self._edge_quantization_metadata = edge_quantization_metadata self._sampler_options = sampler_options self._non_blocking_transfers = non_blocking_transfers diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 3a968eac1..82120e084 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -9,7 +9,6 @@ import torch from graphlearn_torch.channel import SampleMessage -from graphlearn_torch.utils import reverse_edge_type from torch_geometric.data import Data, HeteroData from torch_geometric.data.storage import EdgeStorage, NodeStorage from torch_geometric.typing import EdgeType, NodeType @@ -32,7 +31,6 @@ logger = Logger() _GraphType = TypeVar("_GraphType", Data, HeteroData) -_EdgeMetadataValue = TypeVar("_EdgeMetadataValue") class SamplingClusterSetup(Enum): @@ -69,21 +67,6 @@ class DatasetSchema: Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] ] = None - -def _map_to_effective_edge_types( - metadata: Optional[Union[_EdgeMetadataValue, dict[EdgeType, _EdgeMetadataValue]]], - edge_dir: Union[str, Literal["in", "out"]], -) -> Optional[Union[_EdgeMetadataValue, dict[EdgeType, _EdgeMetadataValue]]]: - """Map stored heterogeneous metadata to sampled edge-type direction.""" - if edge_dir != "in" or not isinstance(metadata, dict): - return metadata - typed_metadata = cast(dict[EdgeType, _EdgeMetadataValue], metadata) - return { - reverse_edge_type(edge_type): value - for edge_type, value in typed_metadata.items() - } - - def patch_fanout_for_sampling( edge_types: Optional[list[EdgeType]], num_neighbors: Union[list[int], dict[EdgeType, list[int]]], diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 4998dec40..a5b8d8707 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -679,7 +679,9 @@ def test_load_edge_weights_rejects_non_raw_field_before_loading(self) -> None: weight_edge_feat_name="quantized_weight", ) - def test_sampling_weight_is_removed_from_edge_quantization_metadata(self) -> None: + 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( @@ -694,12 +696,12 @@ def test_sampling_weight_is_removed_from_edge_quantization_metadata(self) -> Non feature_spec={ "src_id": tf.io.FixedLenFeature([], tf.int64), "dst_id": tf.io.FixedLenFeature([], tf.int64), - "raw_feature": tf.io.FixedLenFeature([], tf.float32), + "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_feature", "weight"], - feature_dim=2, + feature_keys=["raw_embedding", "weight"], + feature_dim=3, entity_key=("src_id", "dst_id"), packed_feature_key="edge_packed_features", packed_feature_dim=1, @@ -707,7 +709,7 @@ def test_sampling_weight_is_removed_from_edge_quantization_metadata(self) -> Non edge_quantization_metadata=FeatureQuantizationMetadata( bits=2, feature_dim=4, - quantized_feature_indices=(0, 2), + quantized_feature_indices=(3,), clip_min=0.0, clip_max=3.0, ), @@ -723,7 +725,7 @@ def test_sampling_weight_is_removed_from_edge_quantization_metadata(self) -> Non FeatureQuantizationMetadata( bits=2, feature_dim=3, - quantized_feature_indices=(0, 2), + quantized_feature_indices=(2,), clip_min=0.0, clip_max=3.0, ), diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index aa1dde141..9c8c447fa 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -13,7 +13,6 @@ POSITIVE_LABEL_METADATA_KEY, ) from gigl.distributed.utils.neighborloader import ( - _map_to_effective_edge_types, attach_ppr_outputs, extract_edge_type_metadata, extract_metadata, @@ -204,24 +203,6 @@ def test_materialize_quantized_edge_features_uses_effective_edge_type( ) self.assertEqual(remaining_metadata, {}) - def test_map_to_effective_edge_types_reverses_inbound_metadata(self) -> None: - quantization_metadata = FeatureQuantizationMetadata( - bits=2, - feature_dim=2, - quantized_feature_indices=(0, 1), - clip_min=0.0, - clip_max=3.0, - ) - - effective_metadata = _map_to_effective_edge_types( - {_U2I_EDGE_TYPE: quantization_metadata}, edge_dir="in" - ) - - self.assertEqual( - effective_metadata, - {("item", "rev_to", "user"): quantization_metadata}, - ) - def test_materialize_quantized_edge_features_rejects_malformed_sidecars( self, ) -> None: From 06c3b71a1011216477482316e9c90c0a33b1b0f4 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 18:29:25 +0000 Subject: [PATCH 67/78] Format edge quantization changes --- gigl/distributed/base_dist_loader.py | 2 +- gigl/distributed/base_sampler.py | 4 +--- gigl/distributed/dataset_factory.py | 2 +- gigl/distributed/utils/neighborloader.py | 1 + tests/unit/common/data/dataloaders_test.py | 2 +- 5 files changed, 5 insertions(+), 6 deletions(-) diff --git a/gigl/distributed/base_dist_loader.py b/gigl/distributed/base_dist_loader.py index 1c070959d..66f0e0ff4 100644 --- a/gigl/distributed/base_dist_loader.py +++ b/gigl/distributed/base_dist_loader.py @@ -23,7 +23,6 @@ get_context, ) from graphlearn_torch.distributed.rpc import rpc_is_initialized -from graphlearn_torch.utils import reverse_edge_type from graphlearn_torch.sampler import ( EdgeSamplerInput, NodeSamplerInput, @@ -31,6 +30,7 @@ SamplingConfig, SamplingType, ) +from graphlearn_torch.utils import reverse_edge_type from torch_geometric.data import Data, HeteroData from torch_geometric.typing import EdgeType from typing_extensions import Self diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index ff73dcbcb..db47ab811 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -484,9 +484,7 @@ async def _collate_fn( futs[ f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{result_edge_type}" ] = wrap_torch_future( - self.dist_edge_quantized_feature.async_get( - eids, etype - ) + self.dist_edge_quantized_feature.async_get(eids, etype) ) if output.batch is not None: for ntype, batch in output.batch.items(): diff --git a/gigl/distributed/dataset_factory.py b/gigl/distributed/dataset_factory.py index ba24a708d..0ffa3c462 100644 --- a/gigl/distributed/dataset_factory.py +++ b/gigl/distributed/dataset_factory.py @@ -24,8 +24,8 @@ from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, TFDatasetOptions, - remove_sampling_weight_from_edge_quantization_metadata, load_torch_tensors_from_tf_record, + remove_sampling_weight_from_edge_quantization_metadata, ) from gigl.common.logger import Logger from gigl.common.utils.decorator import tf_on_cpu diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 82120e084..d01873077 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -67,6 +67,7 @@ class DatasetSchema: Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] ] = None + def patch_fanout_for_sampling( edge_types: Optional[list[EdgeType]], num_neighbors: Union[list[int], dict[EdgeType, list[int]]], diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index a5b8d8707..3973b898f 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -18,8 +18,8 @@ ) from gigl.common.data.load_torch_tensors import ( SerializedGraphMetadata, - remove_sampling_weight_from_edge_quantization_metadata, 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 From 20bd8b4b908c94f85a5ce4099d0d683ea8981749 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 18:35:54 +0000 Subject: [PATCH 68/78] Improved docs for materialize_quantized_node_features --- gigl/distributed/utils/neighborloader.py | 62 ++++++++++++++++++++---- 1 file changed, 52 insertions(+), 10 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 03b0c4054..6b81306ee 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -344,29 +344,71 @@ def materialize_quantized_node_features( Union[FeatureQuantizationMetadata, dict[NodeType, FeatureQuantizationMetadata]] ], ) -> tuple[_GraphType, dict[str, torch.Tensor]]: - """Materialize packed quantized node features into PyG node feature tensors.""" + """Materialize packed quantized node features into PyG node feature tensors. + + Reconstructs each node feature tensor in its original column order by + dequantizing packed features and combining them with any unquantized + feature columns already present in ``data``. Consumed packed-feature + entries are removed from ``metadata``. + + Args: + data: Homogeneous or heterogeneous sampled graph containing raw node + feature columns. + metadata: Sample metadata containing packed node feature tensors. + node_quantization_metadata: Quantization metadata for the graph's node + features. Homogeneous graphs require a single value; heterogeneous + graphs require metadata for each node type. + + Returns: + A tuple containing the graph with reconstructed node features and the + remaining sample metadata. + + Raises: + ValueError: If the graph and quantization metadata shapes do not match, + required packed features are missing, or raw feature dimensions are + inconsistent. + """ if node_quantization_metadata is None: return data, metadata def materialize( store: Union[Data, NodeStorage], packed_features: torch.Tensor, - q: FeatureQuantizationMetadata, + quantization_metadata: FeatureQuantizationMetadata, ) -> None: - dequantized = dequantize_torch_tensor(packed_features, metadata=q) + """Reconstruct and assign node features for one PyG node store. + + Args: + store: Node store receiving the reconstructed ``x`` tensor. + packed_features: Quantized feature columns for the sampled nodes. + quantization_metadata: Column layout and dequantization metadata. + + Raises: + ValueError: If expected raw feature columns are absent or have an + unexpected dimension. + """ + dequantized = dequantize_torch_tensor( + packed_features, metadata=quantization_metadata + ) x = getattr(store, "x", None) - out = dequantized.new_empty((dequantized.size(0), q.feature_dim)) - scatter_idx: FeatureQuantizationIndexTensors = q.scatter_index_tensors( - out.device + out = dequantized.new_empty( + (dequantized.size(0), quantization_metadata.feature_dim) + ) + scatter_idx: FeatureQuantizationIndexTensors = ( + quantization_metadata.scatter_index_tensors(out.device) ) out[:, scatter_idx.quantized] = dequantized - if x is None and q.raw_feature_dim: - raise ValueError(f"Missing {q.raw_feature_dim} unquantized features") + if x is None and quantization_metadata.raw_feature_dim: + raise ValueError( + f"Missing {quantization_metadata.raw_feature_dim} unquantized features" + ) if x is not None: - if x.size(1) != q.raw_feature_dim: + if x.size(1) != quantization_metadata.raw_feature_dim: raise ValueError( - f"Expected {q.raw_feature_dim} raw node features before dequantization, got {x.size(1)}" + "Expected " + f"{quantization_metadata.raw_feature_dim} raw node features " + f"before dequantization, got {x.size(1)}" ) out[:, scatter_idx.raw] = x store.x = out From 03b0e491d3c176b76974024676969c73893cdf72 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 19:05:57 +0000 Subject: [PATCH 69/78] Return partitioned quantized edge features --- gigl/distributed/dist_partitioner.py | 73 ++++++++++++------- gigl/distributed/dist_range_partitioner.py | 72 +++++++++++------- .../run_distributed_partitioner.py | 2 + 3 files changed, 94 insertions(+), 53 deletions(-) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 76fdd3375..49ad8b3a6 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -210,9 +210,6 @@ def __init__( 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._partitioned_edge_quantized_features: dict[ - EdgeType, FeaturePartitionData - ] = {} 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 @@ -1255,10 +1252,12 @@ 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. + r"""Partition topology and feature sidecars for one edge type. Args: node_partition_book (dict[NodeType, PartitionBook]): The partition books of all graph nodes. @@ -1266,8 +1265,9 @@ 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[PartitionBook]: The partition book of graph edges, will be None if there are no edge features for the current edge type + Optional[FeaturePartitionData]: Standard edge features on the current partition. + Optional[FeaturePartitionData]: Quantized edge features on the current partition. + Optional[PartitionBook]: The partition book of graph edges. """ start_time = time.time() @@ -1514,14 +1514,15 @@ def _edge_feat_weight_pfn( weights=partitioned_weights, ) - if current_quantized_feat_part is not None: - self._partitioned_edge_quantized_features[edge_type] = ( - current_quantized_feat_part - ) persistent_edge_partition_book = ( edge_partition_book if should_generate_partition_book else None ) - return current_graph_part, current_feat_part, persistent_edge_partition_book + return ( + current_graph_part, + current_feat_part, + current_quantized_feat_part, + persistent_edge_partition_book, + ) def _partition_label_edge_index( self, @@ -1794,26 +1795,32 @@ 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]], ], ]: - """ - Partitions edges of a graph, including edge indices and edge features. If there are no edge features, only edge indices are partitioned. - If heterogeneous, partitions edges/features for all edge types. Must call `partition_node` first to get the node partition book as input. + """Partition edges and feature sidecars for all edge types. + + Must call `partition_node` first to get the node partition book as input. + Args: 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]]], - ]: 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. + 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, standard edge features, quantized edge features, + and the edge partition book. Sidecars and partition books are None + when no registered data requires edge lookup. """ self._assert_and_get_rpc_setup() @@ -1859,10 +1866,12 @@ 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 @@ -1874,6 +1883,10 @@ def partition_edge_index_and_edge_features( 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") @@ -1895,6 +1908,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: @@ -1902,6 +1918,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, ) @@ -2000,6 +2021,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 @@ -2047,12 +2069,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=( - to_homogeneous(self._partitioned_edge_quantized_features) - if self._is_input_homogeneous - and self._partitioned_edge_quantized_features - else self._partitioned_edge_quantized_features or None - ), + 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 9e76423dd..af43b7eee 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -215,15 +215,13 @@ 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 - edges and edge features (if they exist) together. Once they have been partitioned across machines, - we build the edge partition book based on the number of edges assigned to each machine. Then, we infer - the edge IDs from the edge partition book's ranges. - - If there are no edge features for the current edge type, both the returned edge feature and edge partition book will be None. + Partition topology and feature sidecars for one edge type. Args: node_partition_book (dict[NodeType, PartitionBook]): The partition books of all graph nodes. @@ -231,8 +229,9 @@ 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[PartitionBook]: The partition book of graph edges, will be None if there are no edge features for the current edge type + Optional[FeaturePartitionData]: Standard edge features on the current partition. + Optional[FeaturePartitionData]: Quantized edge features on the current partition. + Optional[PartitionBook]: The partition book of graph edges. """ start_time = time.time() @@ -412,12 +411,6 @@ def edge_partition_fn(rank_indices, _): if partitioned_edge_features is not None else None ) - if partitioned_edge_quantized_features is not None: - self._partitioned_edge_quantized_features[edge_type] = ( - FeaturePartitionData( - feats=partitioned_edge_quantized_features, ids=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}" ) @@ -434,33 +427,48 @@ 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]], ], ]: - """ - Partitions edges of a graph, including edge indices and edge features. If heterogeneous, partitions edges - for all edge types. You must call `partition_node` first to get the node partition book as input. The difference - between this function and its parent is that we no longer need to check that the `edge_ids` have been - pre-computed as a prerequisite for partitioning edges and edge features. + """Partition edges and feature sidecars using range-based partition books. + + Must call `partition_node` first to get the node partition book as input. 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]]], - ]: Partitioned Graph Data, Feature Data, and corresponding edge partition book, is a dictionary if heterogeneous. + 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, standard edge features, quantized edge features, + and the edge partition book. Sidecars and partition books are None + when no registered data requires edge lookup. """ self._assert_and_get_rpc_setup() @@ -496,10 +504,12 @@ 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 @@ -511,6 +521,10 @@ def partition_edge_index_and_edge_features( 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") @@ -529,6 +543,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: @@ -536,5 +553,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/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 955f78e4c..2aa090f77 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -144,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 @@ -206,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 From d0f62cd413d1acefb11912ba82462378707482ee Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 19:07:33 +0000 Subject: [PATCH 70/78] Minimize range partitioner documentation diff --- gigl/distributed/dist_range_partitioner.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index af43b7eee..5e65f0009 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -221,7 +221,12 @@ def _partition_edge_index_and_edge_features( Optional[PartitionBook], ]: """ - Partition topology and feature sidecars for one edge type. + Partition graph topology of a specific edge type. For range-based partitioning, we partition + edges and edge features (if they exist) together. Once they have been partitioned across machines, + we build the edge partition book based on the number of edges assigned to each machine. Then, we infer + the edge IDs from the edge partition book's ranges. + + If there are no edge features for the current edge type, both the returned edge feature and edge partition book will be None. Args: node_partition_book (dict[NodeType, PartitionBook]): The partition books of all graph nodes. From 3e58bd92d6e2dde9c4dbb8c9ef533a084215fcaa Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 20:13:20 +0000 Subject: [PATCH 71/78] Simplify edge partitioner diffs --- gigl/distributed/dist_partitioner.py | 22 ++++++++-------------- gigl/distributed/dist_range_partitioner.py | 12 ++++++------ 2 files changed, 14 insertions(+), 20 deletions(-) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 49ad8b3a6..e10bf00d8 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -690,8 +690,7 @@ def register_edge_quantized_features( unknown_edge_types = set(packed_features) - set(self._edge_index) if unknown_edge_types: raise ValueError( - "Packed edge features contain unregistered edge types: " - f"{unknown_edge_types}" + f"Packed edge features contain unregistered edge types: {unknown_edge_types}" ) for edge_type, features in packed_features.items(): if not isinstance(features, torch.Tensor): @@ -700,19 +699,16 @@ def register_edge_quantized_features( ) if features.dtype != torch.uint8: raise ValueError( - f"Packed edge features for {edge_type} must use torch.uint8, " - f"got {features.dtype}." + f"Packed edge features for {edge_type} must use torch.uint8, got {features.dtype}." ) if features.ndim != 2: raise ValueError( - f"Packed edge features for {edge_type} must be 2-D, got " - f"shape {tuple(features.shape)}." + f"Packed edge features for {edge_type} must be 2-D, got shape {tuple(features.shape)}." ) expected_rows = self._edge_index[edge_type].size(1) if features.size(0) != expected_rows: raise ValueError( - f"Packed edge features for {edge_type} have {features.size(0)} " - f"rows, expected {expected_rows}." + f"Packed edge features for {edge_type} have {features.size(0)} rows, expected {expected_rows}." ) self._edge_quantized_feat = convert_to_tensor(packed_features) self._edge_quantized_feat_dim = { @@ -1807,9 +1803,9 @@ def partition_edge_index_and_edge_features( Optional[dict[EdgeType, PartitionBook]], ], ]: - """Partition edges and feature sidecars for all edge types. - - Must call `partition_node` first to get the node partition book as input. + """ + Partitions edges of a graph, including edge indices and edge features. If there are no edge features, only edge indices are partitioned. + If heterogeneous, partitions edges/features for all edge types. Must call `partition_node` first to get the node partition book as input. Args: node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book @@ -1818,9 +1814,7 @@ def partition_edge_index_and_edge_features( Union[ 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, standard edge features, quantized edge features, - and the edge partition book. Sidecars and partition books are None - when no registered data requires edge lookup. + ]: Partitioned graph data, standard edge features, quantized edge features, and the corresponding edge partition book. """ self._assert_and_get_rpc_setup() diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index 5e65f0009..ffcb1c5a8 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -460,9 +460,11 @@ def partition_edge_index_and_edge_features( Optional[dict[EdgeType, PartitionBook]], ], ]: - """Partition edges and feature sidecars using range-based partition books. - - Must call `partition_node` first to get the node partition book as input. + """ + Partitions edges of a graph, including edge indices and edge features. If heterogeneous, partitions edges + for all edge types. You must call `partition_node` first to get the node partition book as input. The difference + between this function and its parent is that we no longer need to check that the `edge_ids` have been + pre-computed as a prerequisite for partitioning edges and edge features. Args: node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book @@ -471,9 +473,7 @@ def partition_edge_index_and_edge_features( Union[ 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, standard edge features, quantized edge features, - and the edge partition book. Sidecars and partition books are None - when no registered data requires edge lookup. + ]: Partitioned graph data, standard edge features, quantized edge features, and the corresponding edge partition book. """ self._assert_and_get_rpc_setup() From 4bcab57fcd07f21bb966cc0f0c6bfdd4c794235e Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 20:16:21 +0000 Subject: [PATCH 72/78] Revert diff --- gigl/distributed/dist_partitioner.py | 15 ++++++++------- gigl/distributed/dist_range_partitioner.py | 10 ++++------ 2 files changed, 12 insertions(+), 13 deletions(-) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index e10bf00d8..8a6898bc4 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -1253,7 +1253,8 @@ def _partition_edge_index_and_edge_features( Optional[FeaturePartitionData], Optional[PartitionBook], ]: - r"""Partition topology and feature sidecars for one edge type. + 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. Args: node_partition_book (dict[NodeType, PartitionBook]): The partition books of all graph nodes. @@ -1261,9 +1262,9 @@ def _partition_edge_index_and_edge_features( Returns: GraphPartitionData: The graph data of the current partition. - Optional[FeaturePartitionData]: Standard edge features on the current partition. - Optional[FeaturePartitionData]: Quantized edge features on the current partition. - Optional[PartitionBook]: The partition book of graph edges. + 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 """ start_time = time.time() @@ -1806,15 +1807,15 @@ def partition_edge_index_and_edge_features( """ Partitions edges of a graph, including edge indices and edge features. If there are no edge features, only edge indices are partitioned. If heterogeneous, partitions edges/features for all edge types. Must call `partition_node` first to get the node partition book as input. - Args: node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book - Returns: Union[ 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, standard edge features, quantized edge features, and the corresponding edge partition book. + ]: 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. """ self._assert_and_get_rpc_setup() diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index ffcb1c5a8..000ae1222 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -234,9 +234,9 @@ def _partition_edge_index_and_edge_features( Returns: GraphPartitionData: The graph data of the current partition. - Optional[FeaturePartitionData]: Standard edge features on the current partition. - Optional[FeaturePartitionData]: Quantized edge features on the current partition. - Optional[PartitionBook]: The partition book of graph edges. + 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 """ start_time = time.time() @@ -465,15 +465,13 @@ def partition_edge_index_and_edge_features( for all edge types. You must call `partition_node` first to get the node partition book as input. The difference between this function and its parent is that we no longer need to check that the `edge_ids` have been pre-computed as a prerequisite for partitioning edges and edge features. - Args: node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book - Returns: Union[ 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, standard edge features, quantized edge features, and the corresponding edge partition book. + ]: Partitioned Graph Data, Feature Data, and corresponding edge partition book, is a dictionary if heterogeneous. """ self._assert_and_get_rpc_setup() From a488153a5ba5722bcaa6bc216591553c9649ed2b Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 20:22:07 +0000 Subject: [PATCH 73/78] Remove redundant diff --- gigl/distributed/dist_partitioner.py | 54 +++++++------------ gigl/distributed/dist_range_partitioner.py | 2 + .../distributed_partitioner_test.py | 26 --------- 3 files changed, 20 insertions(+), 62 deletions(-) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index 8a6898bc4..df1ae3ca3 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -675,45 +675,30 @@ 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.") - if self._edge_index is None: raise ValueError( - "Register edge indices before registering packed edge features." + "Edge quantized features have already been registered. Cannot re-register edge quantized feature data." ) - packed_features = self._convert_edge_entity_to_heterogeneous_format( - input_edge_entity=edge_quantized_features - ) - if not packed_features: - raise ValueError("Edge quantized features cannot be empty.") - unknown_edge_types = set(packed_features) - set(self._edge_index) - if unknown_edge_types: - raise ValueError( - f"Packed edge features contain unregistered edge types: {unknown_edge_types}" + logger.info("Registering Edge Quantized Features ...") + + input_edge_quantized_features = ( + self._convert_edge_entity_to_heterogeneous_format( + input_edge_entity=edge_quantized_features ) - for edge_type, features in packed_features.items(): - if not isinstance(features, torch.Tensor): - raise ValueError( - f"Packed edge features for {edge_type} must be a torch.Tensor." - ) - if features.dtype != torch.uint8: - raise ValueError( - f"Packed edge features for {edge_type} must use torch.uint8, got {features.dtype}." - ) - if features.ndim != 2: - raise ValueError( - f"Packed edge features for {edge_type} must be 2-D, got shape {tuple(features.shape)}." - ) - expected_rows = self._edge_index[edge_type].size(1) - if features.size(0) != expected_rows: - raise ValueError( - f"Packed edge features for {edge_type} have {features.size(0)} rows, expected {expected_rows}." - ) - self._edge_quantized_feat = convert_to_tensor(packed_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 packed_features.items() + for edge_type, features in input_edge_quantized_features.items() } def register_edge_weights( @@ -1511,14 +1496,11 @@ def _edge_feat_weight_pfn( weights=partitioned_weights, ) - persistent_edge_partition_book = ( - edge_partition_book if should_generate_partition_book else None - ) return ( current_graph_part, current_feat_part, current_quantized_feat_part, - persistent_edge_partition_book, + edge_partition_book, ) def _partition_label_edge_index( diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index 000ae1222..1cce2f487 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -465,8 +465,10 @@ def partition_edge_index_and_edge_features( for all edge types. You must call `partition_node` first to get the node partition book as input. The difference between this function and its parent is that we no longer need to check that the `edge_ids` have been pre-computed as a prerequisite for partitioning edges and edge features. + Args: node_partition_book (Union[PartitionBook, dict[NodeType, PartitionBook]]): The computed Node Partition Book + Returns: Union[ Tuple[GraphPartitionData, Optional[FeaturePartitionData], Optional[FeaturePartitionData], Optional[PartitionBook]], diff --git a/tests/unit/distributed/distributed_partitioner_test.py b/tests/unit/distributed/distributed_partitioner_test.py index d07859c7f..71845d444 100644 --- a/tests/unit/distributed/distributed_partitioner_test.py +++ b/tests/unit/distributed/distributed_partitioner_test.py @@ -1383,32 +1383,6 @@ def test_edge_features_re_registration(self) -> None: ): partitioner.register_edge_features(edge_features=edge_features) - def test_register_edge_quantized_features_rejects_malformed_sidecars( - self, - ) -> None: - master_port = get_free_port() - init_worker_group(world_size=1, rank=0, group_name=get_process_group_name(0)) - init_rpc( - master_addr=self._master_ip_address, - master_port=master_port, - num_rpc_threads=4, - ) - - edge_index = torch.tensor([[0, 1], [1, 0]]) - malformed_features = [ - torch.ones((2, 1), dtype=torch.float32), - torch.ones(2, dtype=torch.uint8), - torch.ones((1, 1), dtype=torch.uint8), - ] - for features in malformed_features: - with self.subTest(shape=features.shape, dtype=features.dtype): - partitioner = DistPartitioner(should_assign_edges_by_src_node=True) - partitioner.register_edge_index(edge_index=edge_index) - with self.assertRaises(ValueError): - partitioner.register_edge_quantized_features( - edge_quantized_features=features - ) - def test_positive_labels_re_registration(self) -> None: """Test that re-registering labels raises an error.""" master_port = get_free_port() From df98e0fed6cc7f8bf58bf5457ce0e7923fabce6c Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 20:40:04 +0000 Subject: [PATCH 74/78] Merge main --- gigl/distributed/utils/neighborloader.py | 168 +++++------------- .../distributed/utils/neighborloader_test.py | 107 ----------- 2 files changed, 49 insertions(+), 226 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 554740d39..fc10f49af 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -461,34 +461,31 @@ def materialize_quantized_edge_features( Union[FeatureQuantizationMetadata, dict[EdgeType, FeatureQuantizationMetadata]] ], ) -> tuple[_GraphType, dict[str, torch.Tensor]]: - """Materialize packed quantized edge features into PyG edge attributes. + """Materialize packed quantized edge features into PyG edge feature tensors. + + Reconstructs each edge feature tensor in its original column order by + dequantizing packed features and combining them with any unquantized + feature columns already present in ``data``. Consumed packed-feature + entries are removed from ``metadata``. Args: - data: Sampled homogeneous or heterogeneous PyG graph. - metadata: Sample metadata containing packed edge-feature tensors. - edge_quantization_metadata: Reconstruction metadata for each configured - edge store, or scalar metadata for a homogeneous graph. + data: Homogeneous or heterogeneous sampled graph containing raw edge + feature columns. + metadata: Sample metadata containing packed edge feature tensors. + edge_quantization_metadata: Quantization metadata for the graph's edge + features. Homogeneous graphs require a single value; heterogeneous + graphs require metadata for each edge type. Returns: - The updated graph and metadata with packed edge-feature entries removed. + A tuple containing the graph with reconstructed edge features and the + remaining sample metadata. Raises: - ValueError: If graph and metadata shapes disagree, packed data is missing - for sampled edges, or exact logical feature reconstruction is not - possible. + ValueError: If the graph and quantization metadata shapes do not match, + required packed features are missing, or raw feature dimensions are + inconsistent. """ - packed_metadata_keys = [ - key - for key in metadata - if key == EDGE_PACKED_FEATURES_METADATA_KEY - or key.startswith(f"{EDGE_PACKED_FEATURES_METADATA_KEY}.") - ] if edge_quantization_metadata is None: - if packed_metadata_keys: - raise ValueError( - "Found packed edge features without edge quantization metadata: " - f"{packed_metadata_keys}" - ) return data, metadata def materialize( @@ -496,42 +493,29 @@ def materialize( packed_features: torch.Tensor, quantization_metadata: FeatureQuantizationMetadata, ) -> None: - if packed_features.ndim != 2: - raise ValueError( - "Expected packed edge features to be a 2-D tensor, got shape " - f"{tuple(packed_features.shape)}" - ) - if packed_features.dtype != torch.uint8: - raise ValueError( - "Expected packed edge features to use torch.uint8 storage, got " - f"{packed_features.dtype}" - ) + """Reconstruct and assign edge features for one PyG edge store. + + Args: + store: Edge store receiving the reconstructed ``edge_attr`` tensor. + packed_features: Quantized feature columns for the sampled edges. + quantization_metadata: Column layout and dequantization metadata. + + Raises: + ValueError: If expected raw feature columns are absent or have an + unexpected dimension. + """ + dequantized = dequantize_torch_tensor( + packed_features, metadata=quantization_metadata + ) edge_attr = getattr(store, "edge_attr", None) - edge_index = getattr(store, "edge_index", None) - if edge_index is not None and edge_index.size(1) != packed_features.size(0): - raise ValueError( - f"Expected {edge_index.size(1)} packed edge feature rows, got " - f"{packed_features.size(0)}" - ) - if edge_attr is not None and edge_attr.size(0) != packed_features.size(0): - raise ValueError( - f"Expected {packed_features.size(0)} raw edge feature rows, got " - f"{edge_attr.size(0)}" - ) - if packed_features.size(0) == 0: - dequantized = packed_features.new_empty( - (0, quantization_metadata.quantized_feature_dim), - dtype=torch.float32, - ) - else: - dequantized = dequantize_torch_tensor( - packed_features, metadata=quantization_metadata - ) - output = dequantized.new_empty( + out = dequantized.new_empty( (dequantized.size(0), quantization_metadata.feature_dim) ) - scatter_indices = quantization_metadata.scatter_index_tensors(output.device) - output[:, scatter_indices.quantized] = dequantized + scatter_idx: FeatureQuantizationIndexTensors = ( + quantization_metadata.scatter_index_tensors(out.device) + ) + out[:, scatter_idx.quantized] = dequantized + if edge_attr is None and quantization_metadata.raw_feature_dim: raise ValueError( f"Missing {quantization_metadata.raw_feature_dim} unquantized edge features" @@ -540,97 +524,43 @@ def materialize( if edge_attr.size(1) != quantization_metadata.raw_feature_dim: raise ValueError( "Expected " - f"{quantization_metadata.raw_feature_dim} raw edge features before " - f"dequantization, got {edge_attr.size(1)}" + f"{quantization_metadata.raw_feature_dim} raw edge features " + f"before dequantization, got {edge_attr.size(1)}" ) - output[:, scatter_indices.raw] = edge_attr - store.edge_attr = output + out[:, scatter_idx.raw] = edge_attr + store.edge_attr = out if isinstance(data, Data): if isinstance(edge_quantization_metadata, dict): - raise ValueError( - "Expected scalar quantization metadata for homogeneous data" - ) + raise ValueError("Expect scalar quantization metadata for homogeneous data") packed_features = metadata.pop(EDGE_PACKED_FEATURES_METADATA_KEY, None) labeled_homogeneous_packed_features_key = ( f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{DEFAULT_HOMOGENEOUS_EDGE_TYPE}" ) if packed_features is None: + # Labeled homogeneous graphs are sampled as heterogeneous graphs, so + # the packed-feature transport key retains the default edge type. packed_features = metadata.pop( labeled_homogeneous_packed_features_key, None ) if packed_features is None: - edge_index = getattr(data, "edge_index", None) - edge_attr = getattr(data, "edge_attr", None) - num_edges = ( - edge_index.size(1) - if edge_index is not None - else edge_attr.size(0) - if edge_attr is not None - else 0 - ) - if num_edges: - raise ValueError( - "Missing packed quantized features in metadata keys " - f"{EDGE_PACKED_FEATURES_METADATA_KEY} or " - f"{labeled_homogeneous_packed_features_key}" - ) - packed_features = torch.empty( - (0, edge_quantization_metadata.packed_feature_dim), - dtype=torch.uint8, - device=( - edge_attr.device - if edge_attr is not None - else edge_index.device - if edge_index is not None - else None - ), + raise ValueError( + f"Missing packed quantized features in metadata keys {EDGE_PACKED_FEATURES_METADATA_KEY} or {labeled_homogeneous_packed_features_key}" ) materialize(data, packed_features, edge_quantization_metadata) else: if not isinstance(edge_quantization_metadata, dict): - raise ValueError("Expected per-edge-type metadata for heterogeneous data") + raise ValueError("Expected per-edge-type metadata for heterogeneous data.") edge_quantization_metadata = cast( dict[EdgeType, FeatureQuantizationMetadata], edge_quantization_metadata ) for edge_type, quantization_metadata in edge_quantization_metadata.items(): metadata_key = f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{edge_type}" packed_features = metadata.pop(metadata_key, None) - if edge_type not in data.edge_types: - if packed_features is not None: - raise ValueError( - f"Found packed edge features for missing edge store {edge_type}" - ) + if packed_features is None: continue + materialize(data[edge_type], packed_features, quantization_metadata) - store = data[edge_type] - if packed_features is None: - edge_index = getattr(store, "edge_index", None) - edge_attr = getattr(store, "edge_attr", None) - num_edges = ( - edge_index.size(1) - if edge_index is not None - else edge_attr.size(0) - if edge_attr is not None - else 0 - ) - if num_edges: - raise ValueError( - "Missing packed quantized features in metadata key " - f"{metadata_key}" - ) - packed_features = torch.empty( - (0, quantization_metadata.packed_feature_dim), - dtype=torch.uint8, - device=( - edge_attr.device - if edge_attr is not None - else edge_index.device - if edge_index is not None - else None - ), - ) - materialize(store, packed_features, quantization_metadata) return data, metadata diff --git a/tests/unit/distributed/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index 9c8c447fa..de0b5ab31 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -156,22 +156,6 @@ def test_materialize_quantized_edge_features_uses_labeled_homogeneous_key( ) self.assertEqual(remaining_metadata, {}) - def test_materialize_quantized_edge_features_rejects_packed_data_without_metadata( - self, - ) -> None: - data = Data(edge_index=torch.tensor([[0], [1]])) - - with self.assertRaises(ValueError): - materialize_quantized_edge_features( - data=data, - metadata={ - EDGE_PACKED_FEATURES_METADATA_KEY: torch.tensor( - [[48]], dtype=torch.uint8 - ) - }, - edge_quantization_metadata=None, - ) - def test_materialize_quantized_edge_features_uses_effective_edge_type( self, ) -> None: @@ -203,97 +187,6 @@ def test_materialize_quantized_edge_features_uses_effective_edge_type( ) self.assertEqual(remaining_metadata, {}) - def test_materialize_quantized_edge_features_rejects_malformed_sidecars( - self, - ) -> None: - edge_type = ("user", "to", "item") - quantization_metadata = FeatureQuantizationMetadata( - bits=2, - feature_dim=3, - quantized_feature_indices=(0, 2), - clip_min=0.0, - clip_max=3.0, - ) - metadata_key = f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{edge_type}" - - def edge_data(raw_features: torch.Tensor, num_edges: int = 1) -> HeteroData: - data = HeteroData() - data[edge_type].edge_index = torch.zeros((2, num_edges), dtype=torch.long) - data[edge_type].edge_attr = raw_features - return data - - cases = [ - ( - edge_data(torch.tensor([[10.0]])), - {}, - ), - ( - edge_data(torch.tensor([[10.0]])), - {metadata_key: torch.tensor([[48, 0]], dtype=torch.uint8)}, - ), - ( - edge_data(torch.tensor([[10.0], [20.0]]), num_edges=2), - {metadata_key: torch.tensor([[48]], dtype=torch.uint8)}, - ), - ( - edge_data(torch.tensor([[10.0, 20.0]])), - {metadata_key: torch.tensor([[48]], dtype=torch.uint8)}, - ), - ] - - for data, metadata in cases: - with self.subTest(metadata=metadata), self.assertRaises(ValueError): - materialize_quantized_edge_features( - data=data, - metadata=metadata, - edge_quantization_metadata={edge_type: quantization_metadata}, - ) - - def test_materialize_quantized_edge_features_preserves_empty_shape( - self, - ) -> None: - edge_type = ("user", "to", "item") - data = HeteroData() - data[edge_type].edge_index = torch.empty((2, 0), dtype=torch.long) - data[edge_type].edge_attr = torch.empty((0, 1)) - - materialized_data, remaining_metadata = materialize_quantized_edge_features( - data=data, - metadata={}, - edge_quantization_metadata={ - edge_type: FeatureQuantizationMetadata( - bits=2, - feature_dim=3, - quantized_feature_indices=(0, 2), - clip_min=0.0, - clip_max=3.0, - ) - }, - ) - - self.assertEqual( - materialized_data[edge_type].edge_attr.shape, torch.Size([0, 3]) - ) - self.assertEqual(materialized_data[edge_type].edge_attr.dtype, torch.float32) - self.assertEqual(remaining_metadata, {}) - - homogeneous_data, _ = materialize_quantized_edge_features( - data=Data( - edge_index=torch.empty((2, 0), dtype=torch.long), - edge_attr=torch.empty((0, 1)), - ), - metadata={}, - edge_quantization_metadata=FeatureQuantizationMetadata( - bits=2, - feature_dim=3, - quantized_feature_indices=(0, 2), - clip_min=0.0, - clip_max=3.0, - ), - ) - self.assertEqual(homogeneous_data.edge_attr.shape, torch.Size([0, 3])) - self.assertEqual(homogeneous_data.edge_attr.dtype, torch.float32) - def test_materialize_quantized_node_features_uses_per_node_type_metadata( self, ) -> None: From 34ea45988cfff9b69ec14d12bb42985d2fa9bf61 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 20:51:07 +0000 Subject: [PATCH 75/78] Extract common helper function for scatter quantized feats --- gigl/distributed/utils/neighborloader.py | 151 ++++++++++------------- 1 file changed, 63 insertions(+), 88 deletions(-) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index fc10f49af..c35d9da2a 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -344,6 +344,45 @@ def set_missing_features( return data +def _materialize_quantized_features( + store: Union[Data, NodeStorage, EdgeStorage], + packed_features: torch.Tensor, + quantization_metadata: FeatureQuantizationMetadata, + feature_attribute: Literal["x", "edge_attr"], +) -> None: + """Reconstruct and assign quantized features for one PyG feature store. + + Args: + store: PyG data or typed storage receiving the reconstructed features. + packed_features: Quantized feature columns for sampled graph entities. + quantization_metadata: Column layout and dequantization metadata. + feature_attribute: PyG attribute that stores the feature tensor. + + Raises: + ValueError: If expected raw feature columns are absent or have an + unexpected dimension. + """ + dequantized = dequantize_torch_tensor( + packed_features, metadata=quantization_metadata + ) + raw_features = getattr(store, feature_attribute, None) + materialized_features = dequantized.new_empty( + (dequantized.size(0), quantization_metadata.feature_dim) + ) + scatter_idx: FeatureQuantizationIndexTensors = ( + quantization_metadata.scatter_index_tensors(materialized_features.device) + ) + materialized_features[:, scatter_idx.quantized] = dequantized + + if raw_features is None and quantization_metadata.raw_feature_dim: + raise ValueError(f"Missing {quantization_metadata.raw_feature_dim} unquantized features") # fmt: skip + if raw_features is not None: + if raw_features.size(1) != quantization_metadata.raw_feature_dim: + raise ValueError(f"Expected {quantization_metadata.raw_feature_dim} raw features before dequantization, got {raw_features.size(1)}") # fmt: skip + materialized_features[:, scatter_idx.raw] = raw_features + setattr(store, feature_attribute, materialized_features) + + def materialize_quantized_node_features( data: _GraphType, metadata: dict[str, torch.Tensor], @@ -378,48 +417,6 @@ def materialize_quantized_node_features( if node_quantization_metadata is None: return data, metadata - def materialize( - store: Union[Data, NodeStorage], - packed_features: torch.Tensor, - quantization_metadata: FeatureQuantizationMetadata, - ) -> None: - """Reconstruct and assign node features for one PyG node store. - - Args: - store: Node store receiving the reconstructed ``x`` tensor. - packed_features: Quantized feature columns for the sampled nodes. - quantization_metadata: Column layout and dequantization metadata. - - Raises: - ValueError: If expected raw feature columns are absent or have an - unexpected dimension. - """ - dequantized = dequantize_torch_tensor( - packed_features, metadata=quantization_metadata - ) - x = getattr(store, "x", None) - out = dequantized.new_empty( - (dequantized.size(0), quantization_metadata.feature_dim) - ) - scatter_idx: FeatureQuantizationIndexTensors = ( - quantization_metadata.scatter_index_tensors(out.device) - ) - out[:, scatter_idx.quantized] = dequantized - - if x is None and quantization_metadata.raw_feature_dim: - raise ValueError( - f"Missing {quantization_metadata.raw_feature_dim} unquantized features" - ) - if x is not None: - if x.size(1) != quantization_metadata.raw_feature_dim: - raise ValueError( - "Expected " - f"{quantization_metadata.raw_feature_dim} raw node features " - f"before dequantization, got {x.size(1)}" - ) - out[:, scatter_idx.raw] = x - store.x = out - if isinstance(data, Data): if isinstance(node_quantization_metadata, dict): raise ValueError("Expect scalar quantization metadata for homogeneous data") @@ -437,7 +434,12 @@ def materialize( raise ValueError( f"Missing packed quantized features in metadata keys {NODE_PACKED_FEATURES_METADATA_KEY} or {labeled_homogeneous_packed_features_key}" ) - materialize(data, packed_features, node_quantization_metadata) + _materialize_quantized_features( + data, + packed_features, + node_quantization_metadata, + feature_attribute="x", + ) else: if not isinstance(node_quantization_metadata, dict): raise ValueError("Expected per-node-type metadata for heterogeneous data.") @@ -449,7 +451,12 @@ def materialize( packed_features = metadata.pop(metadata_key, None) if packed_features is None: continue - materialize(data[node_type], packed_features, quantization_metadata) + _materialize_quantized_features( + data[node_type], + packed_features, + quantization_metadata, + feature_attribute="x", + ) return data, metadata @@ -488,48 +495,6 @@ def materialize_quantized_edge_features( if edge_quantization_metadata is None: return data, metadata - def materialize( - store: Union[Data, EdgeStorage], - packed_features: torch.Tensor, - quantization_metadata: FeatureQuantizationMetadata, - ) -> None: - """Reconstruct and assign edge features for one PyG edge store. - - Args: - store: Edge store receiving the reconstructed ``edge_attr`` tensor. - packed_features: Quantized feature columns for the sampled edges. - quantization_metadata: Column layout and dequantization metadata. - - Raises: - ValueError: If expected raw feature columns are absent or have an - unexpected dimension. - """ - dequantized = dequantize_torch_tensor( - packed_features, metadata=quantization_metadata - ) - edge_attr = getattr(store, "edge_attr", None) - out = dequantized.new_empty( - (dequantized.size(0), quantization_metadata.feature_dim) - ) - scatter_idx: FeatureQuantizationIndexTensors = ( - quantization_metadata.scatter_index_tensors(out.device) - ) - out[:, scatter_idx.quantized] = dequantized - - if edge_attr is None and quantization_metadata.raw_feature_dim: - raise ValueError( - f"Missing {quantization_metadata.raw_feature_dim} unquantized edge features" - ) - if edge_attr is not None: - if edge_attr.size(1) != quantization_metadata.raw_feature_dim: - raise ValueError( - "Expected " - f"{quantization_metadata.raw_feature_dim} raw edge features " - f"before dequantization, got {edge_attr.size(1)}" - ) - out[:, scatter_idx.raw] = edge_attr - store.edge_attr = out - if isinstance(data, Data): if isinstance(edge_quantization_metadata, dict): raise ValueError("Expect scalar quantization metadata for homogeneous data") @@ -547,7 +512,12 @@ def materialize( raise ValueError( f"Missing packed quantized features in metadata keys {EDGE_PACKED_FEATURES_METADATA_KEY} or {labeled_homogeneous_packed_features_key}" ) - materialize(data, packed_features, edge_quantization_metadata) + _materialize_quantized_features( + data, + packed_features, + edge_quantization_metadata, + feature_attribute="edge_attr", + ) else: if not isinstance(edge_quantization_metadata, dict): raise ValueError("Expected per-edge-type metadata for heterogeneous data.") @@ -559,7 +529,12 @@ def materialize( packed_features = metadata.pop(metadata_key, None) if packed_features is None: continue - materialize(data[edge_type], packed_features, quantization_metadata) + _materialize_quantized_features( + data[edge_type], + packed_features, + quantization_metadata, + feature_attribute="edge_attr", + ) return data, metadata From b1a26f971dd7e3b73c2a561bc552966d0c5d36da Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 21:20:48 +0000 Subject: [PATCH 76/78] Simplify quantized feature materialization --- gigl/common/data/load_torch_tensors.py | 31 +++++----------------- gigl/distributed/utils/neighborloader.py | 8 ++++-- tests/unit/common/data/dataloaders_test.py | 2 +- 3 files changed, 13 insertions(+), 28 deletions(-) diff --git a/gigl/common/data/load_torch_tensors.py b/gigl/common/data/load_torch_tensors.py index bf55a7ba7..667bc348d 100644 --- a/gigl/common/data/load_torch_tensors.py +++ b/gigl/common/data/load_torch_tensors.py @@ -130,29 +130,20 @@ def _validate_weight_edge_feature_name( ], weight_edge_feat_name: Optional[Union[str, dict[EdgeType, str]]], ) -> None: - """Validate sampling-weight configuration before TFRecord loading.""" 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 a string for homogeneous graphs." - ) - configured_weights = [ - ( - DEFAULT_HOMOGENEOUS_EDGE_TYPE, - weight_edge_feat_name, - edge_entity_info, - ) - ] + 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 a dict[EdgeType, str] for " - "heterogeneous graphs with multiple edge types." + "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)] @@ -160,8 +151,7 @@ def _validate_weight_edge_feature_name( unknown_edge_types = set(weight_edge_feat_name) - set(edge_entity_info) if unknown_edge_types: raise ValueError( - "weight_edge_feat_name contains unknown edge types: " - f"{unknown_edge_types}" + f"weight_edge_feat_name contains unknown edge types: {unknown_edge_types}" ) configured_weights = [ (edge_type, feature_name, edge_entity_info[edge_type]) @@ -171,16 +161,7 @@ def _validate_weight_edge_feature_name( 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 remain an unquantized scalar edge feature. Available raw " - f"features: {serialized_info.feature_keys}" - ) - feature_spec = serialized_info.feature_spec[feature_name] - feature_width = feature_spec.shape[-1] if feature_spec.shape else 1 - if feature_width != 1: - raise ValueError( - f"Sampling-weight field '{feature_name}' for edge type {edge_type} " - f"must be scalar, but has width {feature_width}." + f"Sampling-weight field '{feature_name}' for edge type {edge_type} must be an unquantized raw edge feature." ) diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index c35d9da2a..04d787498 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -375,10 +375,14 @@ def _materialize_quantized_features( materialized_features[:, scatter_idx.quantized] = dequantized if raw_features is None and quantization_metadata.raw_feature_dim: - raise ValueError(f"Missing {quantization_metadata.raw_feature_dim} unquantized features") # fmt: skip + raise ValueError( + f"Missing {quantization_metadata.raw_feature_dim} unquantized features" + ) if raw_features is not None: if raw_features.size(1) != quantization_metadata.raw_feature_dim: - raise ValueError(f"Expected {quantization_metadata.raw_feature_dim} raw features before dequantization, got {raw_features.size(1)}") # fmt: skip + raise ValueError( + f"Expected {quantization_metadata.raw_feature_dim} raw features before dequantization, got {raw_features.size(1)}" + ) materialized_features[:, scatter_idx.raw] = raw_features setattr(store, feature_attribute, materialized_features) diff --git a/tests/unit/common/data/dataloaders_test.py b/tests/unit/common/data/dataloaders_test.py index 3973b898f..12b5de2e5 100644 --- a/tests/unit/common/data/dataloaders_test.py +++ b/tests/unit/common/data/dataloaders_test.py @@ -671,7 +671,7 @@ def test_load_edge_weights_rejects_non_raw_field_before_loading(self) -> None: ), ) - with self.assertRaisesRegex(ValueError, "must remain an unquantized scalar"): + 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, From 217ce93b31a0a3549dcc53f003c5d69505935f99 Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 22:14:12 +0000 Subject: [PATCH 77/78] Run review --- gigl/common/data/dataloaders.py | 21 +- gigl/distributed/base_dist_loader.py | 24 +-- gigl/distributed/base_sampler.py | 6 +- gigl/distributed/dist_partitioner.py | 25 +-- gigl/distributed/dist_range_partitioner.py | 4 +- gigl/distributed/utils/neighborloader.py | 6 +- .../data_preprocessor/data_preprocessor.py | 75 +++---- .../feature_quantization_transform_test.py | 139 ++++++++++++ .../run_distributed_partitioner.py | 10 +- tests/unit/distributed/dist_server_test.py | 49 ++++- .../distributed_neighborloader_test.py | 200 +++++++++++++++++- .../distributed_partitioner_test.py | 11 + .../distributed/utils/neighborloader_test.py | 24 +++ 13 files changed, 483 insertions(+), 111 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/distributed/base_dist_loader.py b/gigl/distributed/base_dist_loader.py index 66f0e0ff4..afe8d564f 100644 --- a/gigl/distributed/base_dist_loader.py +++ b/gigl/distributed/base_dist_loader.py @@ -30,7 +30,6 @@ SamplingConfig, SamplingType, ) -from graphlearn_torch.utils import reverse_edge_type from torch_geometric.data import Data, HeteroData from torch_geometric.typing import EdgeType from typing_extensions import Self @@ -245,28 +244,9 @@ def __init__( dataset_schema.is_homogeneous_with_labeled_edge_type ) self._node_feature_info = dataset_schema.node_feature_info - # GLT returns heterogeneous edge stores in the sampled direction. For - # incoming sampling, that reverses each stored edge type, so feature - # metadata must use the same keys as the returned stores. Otherwise - # empty stores miss their feature shape and packed features cannot be - # reconstructed into their sampled edge attributes. - edge_feature_info = dataset_schema.edge_feature_info - if dataset_schema.edge_dir == "in" and isinstance(edge_feature_info, dict): - edge_feature_info = { - reverse_edge_type(edge_type): feature_info - for edge_type, feature_info in edge_feature_info.items() - } - self._edge_feature_info = edge_feature_info + self._edge_feature_info = dataset_schema.edge_feature_info self._node_quantization_metadata = dataset_schema.node_quantization_metadata - edge_quantization_metadata = dataset_schema.edge_quantization_metadata - if dataset_schema.edge_dir == "in" and isinstance( - edge_quantization_metadata, dict - ): - edge_quantization_metadata = { - reverse_edge_type(edge_type): quantization_metadata - for edge_type, quantization_metadata in edge_quantization_metadata.items() - } - self._edge_quantization_metadata = edge_quantization_metadata + self._edge_quantization_metadata = dataset_schema.edge_quantization_metadata self._sampler_options = sampler_options self._non_blocking_transfers = non_blocking_transfers diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index db47ab811..30d3569ec 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -481,8 +481,12 @@ async def _collate_fn( eids = result_map.get(f"{as_str(result_edge_type)}.eids") if eids is not None: eids = eids.to(torch.long) + # GLT maps incoming wire edge types back to the dataset edge + # type during collation. Metadata bypasses that mapping, so its + # transport key must already match the final output store. + metadata_edge_type = etype futs[ - f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{result_edge_type}" + f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{metadata_edge_type}" ] = wrap_torch_future( self.dist_edge_quantized_feature.async_get(eids, etype) ) diff --git a/gigl/distributed/dist_partitioner.py b/gigl/distributed/dist_partitioner.py index df1ae3ca3..9549c9e3f 100644 --- a/gigl/distributed/dist_partitioner.py +++ b/gigl/distributed/dist_partitioner.py @@ -1325,11 +1325,10 @@ 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 @@ -1371,30 +1370,28 @@ def _edge_pfn(_, chunk_range): 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 - quantized_feat_idx: Optional[int] = 1 if has_edge_feats else 0 - if not has_edge_quantized_feats: - quantized_feat_idx = None - weight_idx: Optional[int] = None - if has_weights_for_edge_type: - weight_idx = int(has_edge_feats) + int(has_edge_quantized_feats) - + # 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, diff --git a/gigl/distributed/dist_range_partitioner.py b/gigl/distributed/dist_range_partitioner.py index 1cce2f487..170d5de09 100644 --- a/gigl/distributed/dist_range_partitioner.py +++ b/gigl/distributed/dist_range_partitioner.py @@ -280,8 +280,8 @@ def _partition_edge_index_and_edge_features( 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 diff --git a/gigl/distributed/utils/neighborloader.py b/gigl/distributed/utils/neighborloader.py index 04d787498..4f316298b 100644 --- a/gigl/distributed/utils/neighborloader.py +++ b/gigl/distributed/utils/neighborloader.py @@ -532,7 +532,11 @@ def materialize_quantized_edge_features( metadata_key = f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{edge_type}" packed_features = metadata.pop(metadata_key, None) if packed_features is None: - continue + if edge_type not in data.edge_types or data[edge_type].num_edges == 0: + continue + raise ValueError( + f"Missing packed quantized edge features for sampled edge type {edge_type}" + ) _materialize_quantized_features( data[edge_type], packed_features, diff --git a/gigl/src/data_preprocessor/data_preprocessor.py b/gigl/src/data_preprocessor/data_preprocessor.py index 4d2f2c328..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] @@ -443,21 +475,10 @@ def _generate_edge_metadata_info_pb( transform_fn_assets_uri=transformed_features_info.transformed_features_transform_fn_assets_path.uri, ) if transformed_features_info.feature_quantization_enabled: - with tf.io.gfile.GFile( - transformed_features_info.feature_quantization_metadata_path.uri - ) as metadata_file: - metadata = json.loads(metadata_file.read()) - quantization_metadata = preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( - packed_feature_key=metadata["packed_feature_key"], - quantized_feature_indices=metadata["quantized_feature_indices"], + 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}", ) - if metadata["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 = metadata["bits"] - quantization_metadata.multi_bit_state.clip_min = metadata["clip_min"] - quantization_metadata.multi_bit_state.clip_max = metadata["clip_max"] output.quantized_feature_metadata.CopyFrom(quantization_metadata) return output @@ -515,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 ) 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 540140be9..3bd43425d 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -5,21 +5,160 @@ import apache_beam as beam import pyarrow as pa import tensorflow as tf +import tensorflow_data_validation as tfdv +import torch from apache_beam.testing.test_pipeline import TestPipeline from apache_beam.testing.util import assert_that, equal_to from parameterized import parameterized from tensorflow_metadata.proto.v0 import schema_pb2 from tensorflow_transform.tf_metadata.dataset_metadata import DatasetMetadata +from torch_geometric.data import Data +from gigl.common.beam.better_tfrecordio import BetterWriteToTFRecord +from gigl.common.data.dataloaders import TFDatasetOptions, TFRecordDataLoader +from gigl.distributed.utils.neighborloader import ( + EDGE_PACKED_FEATURES_METADATA_KEY, + materialize_quantized_edge_features, +) +from gigl.distributed.utils.serialized_graph_metadata_translator import ( + convert_pb_to_serialized_graph_metadata, +) +from gigl.src.common.types.pb_wrappers.graph_metadata import GraphMetadataPbWrapper +from gigl.src.common.types.pb_wrappers.preprocessed_metadata import ( + PreprocessedMetadataPbWrapper, +) +from gigl.src.data_preprocessor.data_preprocessor import ( + _load_feature_quantization_metadata_pb, +) 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.types import FeatureQuantizationSpec +from snapchat.research.gbml import graph_schema_pb2, preprocessed_metadata_pb2 from tests.test_assets.test_case import TestCase class FeatureQuantizationTransformTest(TestCase): + def test_edge_quantization_round_trips_through_storage_and_loading(self) -> None: + with tempfile.TemporaryDirectory() as temp_dir: + metadata_path = os.path.join(temp_dir, "feature_quantization_metadata.json") + tfrecord_prefix = os.path.join(temp_dir, "edges") + schema_path = os.path.join(temp_dir, "schema.pbtxt") + logical_metadata = DatasetMetadata.from_feature_spec( + { + "src": tf.io.FixedLenFeature(shape=[], dtype=tf.int64), + "dst": tf.io.FixedLenFeature(shape=[], dtype=tf.int64), + "quantized": tf.io.FixedLenFeature(shape=[], dtype=tf.float32), + "raw": tf.io.FixedLenFeature(shape=[], dtype=tf.float32), + } + ) + logical_batches = [ + pa.RecordBatch.from_arrays( + [ + pa.array([[0], [1]], type=pa.list_(pa.int64())), + pa.array([[1], [0]], type=pa.list_(pa.int64())), + pa.array([[-2.0], [8.0]], type=pa.list_(pa.float32())), + pa.array([[10.0], [20.0]], type=pa.list_(pa.float32())), + ], + names=["src", "dst", "quantized", "raw"], + ) + ] + + with TestPipeline() as pipeline: + transformed_batches, physical_metadata = ( + apply_feature_quantization_transform( + logical_features=pipeline + | "Create edge RecordBatches" >> beam.Create(logical_batches), + logical_metadata=logical_metadata, + logical_feature_keys=["quantized", "raw"], + quantization_spec=FeatureQuantizationSpec( + feature_keys=["quantized"], bits=2 + ), + quantization_metadata_path=metadata_path, + packed_feature_key=EDGE_PACKED_FEATURE_KEY, + ) + ) + transformed_batches | "Write edge TFRecords" >> BetterWriteToTFRecord( + file_path_prefix=tfrecord_prefix, + transformed_metadata=physical_metadata, + num_shards=1, + ) + + tfdv.write_schema_text(logical_metadata.schema, schema_path) + quantization_metadata_pb = _load_feature_quantization_metadata_pb( + metadata_path=metadata_path, + entity_description="test edge type", + ) + self.assertEqual( + quantization_metadata_pb.packed_feature_key, "edge_packed_features" + ) + self.assertEqual( + list(quantization_metadata_pb.quantized_feature_indices), [0] + ) + + preprocessed_metadata_pb = preprocessed_metadata_pb2.PreprocessedMetadata() + preprocessed_metadata_pb.condensed_node_type_to_preprocessed_metadata[ + 0 + ].node_id_key = "node_id" + edge_metadata = ( + preprocessed_metadata_pb.condensed_edge_type_to_preprocessed_metadata[0] + ) + edge_metadata.src_node_id_key = "src" + edge_metadata.dst_node_id_key = "dst" + edge_metadata.main_edge_info.CopyFrom( + preprocessed_metadata_pb2.PreprocessedMetadata.EdgeMetadataInfo( + feature_keys=["quantized", "raw"], + feature_dim=2, + tfrecord_uri_prefix=temp_dir, + schema_uri=schema_path, + quantized_feature_metadata=quantization_metadata_pb, + ) + ) + graph_metadata_pb = graph_schema_pb2.GraphMetadata( + node_types=["node"], + edge_types=[ + graph_schema_pb2.EdgeType( + src_node_type="node", relation="connects", dst_node_type="node" + ) + ], + condensed_node_type_map={0: "node"}, + condensed_edge_type_map={ + 0: graph_schema_pb2.EdgeType( + src_node_type="node", relation="connects", dst_node_type="node" + ) + }, + ) + serialized_metadata = convert_pb_to_serialized_graph_metadata( + preprocessed_metadata_pb_wrapper=PreprocessedMetadataPbWrapper( + preprocessed_metadata_pb + ), + graph_metadata_pb_wrapper=GraphMetadataPbWrapper(graph_metadata_pb), + tfrecord_uri_pattern="edges.*\\.tfrecord", + ) + loaded = TFRecordDataLoader(rank=0, world_size=1).load_as_torch_tensors( + serialized_tf_record_info=serialized_metadata.edge_entity_info, + tf_dataset_options=TFDatasetOptions(deterministic=True), + ) + + assert loaded.features is not None + assert loaded.quantized_features is not None + self.assert_tensor_equality(loaded.ids, torch.tensor([[0, 1], [1, 0]])) + self.assert_tensor_equality(loaded.features, torch.tensor([[10.0], [20.0]])) + self.assert_tensor_equality( + loaded.quantized_features, torch.tensor([[0], [192]], dtype=torch.uint8) + ) + materialized, remaining_metadata = materialize_quantized_edge_features( + data=Data(edge_index=loaded.ids, edge_attr=loaded.features), + metadata={EDGE_PACKED_FEATURES_METADATA_KEY: loaded.quantized_features}, + edge_quantization_metadata=serialized_metadata.edge_quantization_metadata, + ) + self.assert_tensor_equality( + materialized.edge_attr, torch.tensor([[-2.0, 10.0], [8.0, 20.0]]) + ) + self.assertEqual(remaining_metadata, {}) + def test_apply_feature_quantization_transform_rejects_reserved_schema_key( self, ) -> None: diff --git a/tests/test_assets/distributed/run_distributed_partitioner.py b/tests/test_assets/distributed/run_distributed_partitioner.py index 2aa090f77..89c863d07 100644 --- a/tests/test_assets/distributed/run_distributed_partitioner.py +++ b/tests/test_assets/distributed/run_distributed_partitioner.py @@ -108,14 +108,14 @@ def run_distributed_partitioner( 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_features, dict): - edge_features_by_type = cast(dict[EdgeType, torch.Tensor], edge_features) + if isinstance(edge_index, dict): + edge_index_by_type = cast(dict[EdgeType, torch.Tensor], edge_index) edge_quantized_features = { - edge_type: features.to(torch.uint8) - for edge_type, features in edge_features_by_type.items() + edge_type: indices[0].to(torch.uint8).unsqueeze(1) + for edge_type, indices in edge_index_by_type.items() } else: - edge_quantized_features = edge_features.to(torch.uint8) + edge_quantized_features = edge_index[0].to(torch.uint8).unsqueeze(1) dist_partitioner.register_edge_quantized_features( edge_quantized_features=edge_quantized_features ) diff --git a/tests/unit/distributed/dist_server_test.py b/tests/unit/distributed/dist_server_test.py index c876fcef1..1b0f1675e 100644 --- a/tests/unit/distributed/dist_server_test.py +++ b/tests/unit/distributed/dist_server_test.py @@ -1,5 +1,5 @@ import threading -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import torch from absl.testing import absltest @@ -65,45 +65,72 @@ def test_get_node_feature_info_with_homogeneous_dataset(self) -> None: # Verify it returns the correct feature info self.assertIsNone(node_feature_info) - def test_get_node_quantization_metadata(self) -> None: - metadata = FeatureQuantizationMetadata( + def test_get_quantization_metadata(self) -> None: + node_metadata = FeatureQuantizationMetadata( bits=2, feature_dim=2, quantized_feature_indices=(0, 1), clip_min=0.0, clip_max=3.0, ) + edge_metadata = { + USER_TO_STORY: FeatureQuantizationMetadata( + bits=4, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=-1.0, + clip_max=1.0, + ) + } dataset = DistDataset( rank=0, world_size=1, edge_dir="out", - node_quantization_metadata=metadata, + node_quantization_metadata=node_metadata, + edge_quantization_metadata=edge_metadata, ) server = dist_server.DistServer(dataset) - self.assertEqual(server.get_node_quantization_metadata(), metadata) + self.assertEqual(server.get_node_quantization_metadata(), node_metadata) + self.assertEqual(server.get_edge_quantization_metadata(), edge_metadata) - def test_remote_dataset_fetches_node_quantization_metadata(self) -> None: - metadata = FeatureQuantizationMetadata( + def test_remote_dataset_fetches_quantization_metadata(self) -> None: + node_metadata = FeatureQuantizationMetadata( bits=2, feature_dim=2, quantized_feature_indices=(0, 1), clip_min=0.0, clip_max=3.0, ) + edge_metadata = { + USER_TO_STORY: FeatureQuantizationMetadata( + bits=4, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=-1.0, + clip_max=1.0, + ) + } with patch( "gigl.distributed.graph_store.remote_dist_dataset.request_server", - return_value=metadata, + side_effect=[node_metadata, edge_metadata], ) as request_server: remote_dataset = RemoteDistDataset(cluster_info=MagicMock(), local_rank=0) self.assertEqual( - remote_dataset.fetch_node_quantization_metadata(), metadata + remote_dataset.fetch_node_quantization_metadata(), node_metadata + ) + self.assertEqual( + remote_dataset.fetch_edge_quantization_metadata(), edge_metadata ) - request_server.assert_called_once_with( - 0, dist_server.DistServer.get_node_quantization_metadata + self.assertEqual( + request_server.call_args_list, + [ + call(0, dist_server.DistServer.get_node_quantization_metadata), + call(0, dist_server.DistServer.get_edge_quantization_metadata), + ], ) def test_get_edge_feature_info_with_heterogeneous_dataset(self) -> None: diff --git a/tests/unit/distributed/distributed_neighborloader_test.py b/tests/unit/distributed/distributed_neighborloader_test.py index c942792a3..93cbcad2c 100644 --- a/tests/unit/distributed/distributed_neighborloader_test.py +++ b/tests/unit/distributed/distributed_neighborloader_test.py @@ -12,6 +12,7 @@ from gigl.distributed.dataset_factory import build_dataset from gigl.distributed.dist_dataset import DistDataset from gigl.distributed.distributed_neighborloader import DistNeighborLoader +from gigl.distributed.sampler import EDGE_PACKED_FEATURES_METADATA_KEY from gigl.distributed.utils import get_free_port from gigl.distributed.utils.neighborloader import DatasetSchema from gigl.distributed.utils.serialized_graph_metadata_translator import ( @@ -1188,6 +1189,112 @@ def test_packed_edge_metadata_requests_sampled_edge_ids(self) -> None: self.assertTrue(config.with_edge) +# NOTE on the test strategy: GiGL loaders always sample via the multiprocess +# producer, which spawns worker subprocesses with a *fresh* interpreter +# (`mp.get_context("spawn")`, dist_sampling_producer.py). A `mock.patch` applied in the +# loader process therefore never reaches the sampler running in that subprocess, so we +# cannot inject a synthetic failure by mocking the sampler. Instead we reproduce a real +# sampler failure end-to-end: a heterogeneous dataset with an incomplete feature store +# for one message-passing edge type. When the missing edge ID is reached during sampling, +# its feature lookup raises inside the sampling coroutine - the exact swallowed-exception +# case this change surfaces. Without the change this hangs forever, so the test uses a +# bounded join. + + +def _run_partial_edge_feature_coverage_raises( + _, + dataset: DistDataset, + error_holder, +): + create_test_process_group() + assert isinstance(dataset.node_ids, Mapping) + loader = DistNeighborLoader( + dataset=dataset, + input_nodes=(_USER, dataset.node_ids[_USER]), # ty: ignore[invalid-argument-type] + num_neighbors=[2, 2], + pin_memory_device=torch.device("cpu"), + ) + try: + for _datum in loader: + pass + except RuntimeError as e: + error_holder["msg"] = str(e) + finally: + shutdown_rpc() + + +class TestSamplingErrorPropagation(TestCase): + def _build_partial_edge_feature_dataset(self) -> DistDataset: + """Build a hetero dataset with an incomplete edge feature store. + + Both edge types are reachable from ``user`` seeds within a 2-hop fanout. + ``story-to-user`` omits the last edge ID, so its feature lookup raises inside + the sampling coroutine. + """ + n = 5 + edge_index = torch.tensor([[0, 1, 2, 3, 4], [0, 1, 2, 3, 4]]) + partition_output = PartitionOutput( + node_partition_book={_USER: torch.zeros(n), _STORY: torch.zeros(n)}, + edge_partition_book={ + _USER_TO_STORY: torch.zeros(n), + _STORY_TO_USER: torch.zeros(n), + }, + partitioned_edge_index={ + _USER_TO_STORY: GraphPartitionData( + edge_index=edge_index, edge_ids=None + ), + _STORY_TO_USER: GraphPartitionData( + edge_index=edge_index, edge_ids=None + ), + }, + partitioned_node_features={ + _USER: FeaturePartitionData( + feats=torch.zeros(n, 2), ids=torch.arange(n) + ), + _STORY: FeaturePartitionData( + feats=torch.zeros(n, 2), ids=torch.arange(n) + ), + }, + partitioned_edge_features={ + _USER_TO_STORY: FeaturePartitionData( + feats=torch.ones(n, 3), ids=torch.arange(n) + ), + _STORY_TO_USER: FeaturePartitionData( + feats=torch.ones(n - 1, 3), ids=torch.arange(n - 1) + ), + }, + partitioned_positive_labels=None, + partitioned_negative_labels=None, + partitioned_node_labels=None, + ) + dataset = DistDataset(rank=0, world_size=1, edge_dir="out") + dataset.build(partition_output=partition_output) + return dataset + + def test_reachable_sampler_failure_raises_not_hangs(self) -> None: + dataset = self._build_partial_edge_feature_dataset() + manager = mp.Manager() + error_holder = manager.dict() + proc = mp.get_context("spawn").Process( + target=_run_partial_edge_feature_coverage_raises, + args=(0, dataset, error_holder), + ) + proc.start() + proc.join(timeout=180) # bounded: the pre-fix behavior hangs indefinitely + alive = proc.is_alive() + if alive: + proc.terminate() + proc.join(timeout=10) + self.assertFalse( + alive, "loader hung instead of failing fast on a sampler error" + ) + message = error_holder.get("msg", "") + # The training process raised with the worker's real traceback embedded. + self.assertIn("sampling worker failed", message.lower()) + self.assertIn("IndexError", message) + self.assertIn("index 4 is out of bounds", message) + + def _run_heterogeneous_partially_quantized_edge_feature_neighbor_loader( _, dataset: DistDataset, @@ -1210,6 +1317,37 @@ def _run_heterogeneous_partially_quantized_edge_feature_neighbor_loader( shutdown_rpc() +def _run_incoming_heterogeneous_quantized_edge_feature_neighbor_loader( + _, + dataset: DistDataset, + expected_edge_features: torch.Tensor, +) -> None: + create_test_process_group() + loader = DistNeighborLoader( + dataset=dataset, + input_nodes=(_STORY, torch.tensor([0])), + num_neighbors=[1], + batch_size=1, + pin_memory_device=torch.device("cpu"), + ) + + edge_feature_info = loader._edge_feature_info + edge_quantization_metadata = loader._edge_quantization_metadata + assert isinstance(edge_feature_info, dict) + assert isinstance(edge_quantization_metadata, dict) + assert set(edge_feature_info) == {_USER_TO_STORY} + assert set(edge_quantization_metadata) == {_USER_TO_STORY} + assert edge_feature_info[_USER_TO_STORY].dim == 2 + assert edge_quantization_metadata[_USER_TO_STORY].feature_dim == 4 + + batch = next(iter(loader)) + assert isinstance(batch, HeteroData) + assert_tensor_equality(batch[_USER_TO_STORY].edge_attr, expected_edge_features) + assert EDGE_PACKED_FEATURES_METADATA_KEY not in batch + assert f"{EDGE_PACKED_FEATURES_METADATA_KEY}.{_USER_TO_STORY}" not in batch + shutdown_rpc() + + class HeterogeneousEdgeFeatureLookupTest(TestCase): def test_heterogeneous_loader_supports_partially_quantized_edge_types( self, @@ -1219,17 +1357,11 @@ def test_heterogeneous_loader_supports_partially_quantized_edge_types( expected_edge_features = {_USER_TO_STORY: torch.tensor([[0.0, 10.0]])} partition_output = PartitionOutput( node_partition_book={_USER: torch.zeros(1), _STORY: torch.zeros(1)}, - edge_partition_book={ - _USER_TO_STORY: torch.zeros(1), - _STORY_TO_USER: torch.zeros(1), - }, + edge_partition_book={_USER_TO_STORY: torch.zeros(1)}, partitioned_edge_index={ _USER_TO_STORY: GraphPartitionData( edge_index=torch.tensor([[0], [0]]), edge_ids=None - ), - _STORY_TO_USER: GraphPartitionData( - edge_index=torch.tensor([[0], [0]]), edge_ids=None - ), + ) }, partitioned_node_features=None, partitioned_edge_features={ @@ -1268,6 +1400,58 @@ def test_heterogeneous_loader_supports_partially_quantized_edge_types( args=(dataset, expected_edge_features), ) + def test_incoming_edges_reverse_feature_metadata_and_output_stores(self) -> None: + partition_output = PartitionOutput( + node_partition_book={_USER: torch.zeros(1), _STORY: torch.zeros(1)}, + edge_partition_book={ + _USER_TO_STORY: torch.zeros(1), + _STORY_TO_USER: torch.zeros(1), + }, + partitioned_edge_index={ + _USER_TO_STORY: GraphPartitionData( + edge_index=torch.tensor([[0], [0]]), edge_ids=None + ), + _STORY_TO_USER: GraphPartitionData( + edge_index=torch.tensor([[0], [0]]), edge_ids=None + ), + }, + partitioned_node_features=None, + partitioned_edge_features={ + _USER_TO_STORY: FeaturePartitionData( + feats=torch.tensor([[10.0, 20.0]]), ids=torch.tensor([0]) + ) + }, + partitioned_edge_quantized_features={ + _USER_TO_STORY: FeaturePartitionData( + feats=torch.tensor([[48]], dtype=torch.uint8), + ids=torch.tensor([0]), + ) + }, + partitioned_positive_labels=None, + partitioned_negative_labels=None, + partitioned_node_labels=None, + ) + dataset = DistDataset( + rank=0, + world_size=1, + edge_dir="in", + edge_quantization_metadata={ + _USER_TO_STORY: FeatureQuantizationMetadata( + bits=2, + feature_dim=4, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + }, + ) + dataset.build(partition_output=partition_output) + + mp.spawn( + fn=_run_incoming_heterogeneous_quantized_edge_feature_neighbor_loader, + args=(dataset, torch.tensor([[0.0, 10.0, 3.0, 20.0]])), + ) + if __name__ == "__main__": absltest.main() diff --git a/tests/unit/distributed/distributed_partitioner_test.py b/tests/unit/distributed/distributed_partitioner_test.py index 71845d444..f3aaea091 100644 --- a/tests/unit/distributed/distributed_partitioner_test.py +++ b/tests/unit/distributed/distributed_partitioner_test.py @@ -815,6 +815,17 @@ def test_partitioning_correctness( 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/utils/neighborloader_test.py b/tests/unit/distributed/utils/neighborloader_test.py index de0b5ab31..43884b495 100644 --- a/tests/unit/distributed/utils/neighborloader_test.py +++ b/tests/unit/distributed/utils/neighborloader_test.py @@ -187,6 +187,30 @@ def test_materialize_quantized_edge_features_uses_effective_edge_type( ) self.assertEqual(remaining_metadata, {}) + def test_materialize_quantized_edge_features_rejects_missing_packed_features_for_sampled_edge_type( + self, + ) -> None: + data = HeteroData() + data[_U2I_EDGE_TYPE].edge_index = torch.tensor([[0], [1]]) + data[_U2I_EDGE_TYPE].edge_attr = torch.tensor([[10.0]]) + + with self.assertRaisesRegex( + ValueError, "Missing packed quantized edge features" + ): + materialize_quantized_edge_features( + data=data, + metadata={}, + edge_quantization_metadata={ + _U2I_EDGE_TYPE: FeatureQuantizationMetadata( + bits=2, + feature_dim=3, + quantized_feature_indices=(0, 2), + clip_min=0.0, + clip_max=3.0, + ) + }, + ) + def test_materialize_quantized_node_features_uses_per_node_type_metadata( self, ) -> None: From 9b82deb0d17dd6f27ffc83d6f9dae2de4a02ebbd Mon Sep 17 00:00:00 2001 From: jchmura Date: Thu, 13 Aug 2026 22:23:03 +0000 Subject: [PATCH 78/78] Remove private dep from test --- gigl/distributed/base_sampler.py | 9 ++++---- .../feature_quantization_transform_test.py | 22 ++++++++++++++----- 2 files changed, 20 insertions(+), 11 deletions(-) diff --git a/gigl/distributed/base_sampler.py b/gigl/distributed/base_sampler.py index 30d3569ec..e92f2cf3d 100644 --- a/gigl/distributed/base_sampler.py +++ b/gigl/distributed/base_sampler.py @@ -484,11 +484,10 @@ async def _collate_fn( # GLT maps incoming wire edge types back to the dataset edge # type during collation. Metadata bypasses that mapping, so its # transport key must already match the final output store. - metadata_edge_type = etype - futs[ - f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{metadata_edge_type}" - ] = wrap_torch_future( - self.dist_edge_quantized_feature.async_get(eids, etype) + futs[f"#META.{EDGE_PACKED_FEATURES_METADATA_KEY}.{etype}"] = ( + wrap_torch_future( + self.dist_edge_quantized_feature.async_get(eids, etype) + ) ) if output.batch is not None: for ntype, batch in output.batch.items(): 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 3bd43425d..d2ef8c309 100644 --- a/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py +++ b/tests/integration/pipeline/data_preprocessor/feature_quantization_transform_test.py @@ -27,9 +27,6 @@ from gigl.src.common.types.pb_wrappers.preprocessed_metadata import ( PreprocessedMetadataPbWrapper, ) -from gigl.src.data_preprocessor.data_preprocessor import ( - _load_feature_quantization_metadata_pb, -) from gigl.src.data_preprocessor.lib.transform.feature_quantization import ( EDGE_PACKED_FEATURE_KEY, NODE_PACKED_FEATURE_KEY, @@ -87,10 +84,23 @@ def test_edge_quantization_round_trips_through_storage_and_loading(self) -> None ) tfdv.write_schema_text(logical_metadata.schema, schema_path) - quantization_metadata_pb = _load_feature_quantization_metadata_pb( - metadata_path=metadata_path, - entity_description="test edge type", + with tf.io.gfile.GFile(metadata_path) as metadata_file: + quantization_metadata = json.loads(metadata_file.read()) + quantization_metadata_pb = preprocessed_metadata_pb2.PreprocessedMetadata.FeatureQuantizationMetadata( + packed_feature_key=quantization_metadata["packed_feature_key"], + quantized_feature_indices=quantization_metadata[ + "quantized_feature_indices" + ], ) + quantization_metadata_pb.multi_bit_state.bits = quantization_metadata[ + "bits" + ] + quantization_metadata_pb.multi_bit_state.clip_min = quantization_metadata[ + "clip_min" + ] + quantization_metadata_pb.multi_bit_state.clip_max = quantization_metadata[ + "clip_max" + ] self.assertEqual( quantization_metadata_pb.packed_feature_key, "edge_packed_features" )