@@ -23,16 +23,16 @@ import com.google.privacy.differentialprivacy.pipelinedp4j.core.DpEngine
2323import com.google.privacy.differentialprivacy.pipelinedp4j.core.DpEngineBudgetSpec
2424import com.google.privacy.differentialprivacy.pipelinedp4j.core.Encoder
2525import com.google.privacy.differentialprivacy.pipelinedp4j.core.EncoderFactory
26+ import com.google.privacy.differentialprivacy.pipelinedp4j.core.FeatureSpec
2627import com.google.privacy.differentialprivacy.pipelinedp4j.core.FeatureValuesExtractor
2728import com.google.privacy.differentialprivacy.pipelinedp4j.core.FrameworkCollection
2829import com.google.privacy.differentialprivacy.pipelinedp4j.core.FrameworkTable
2930import com.google.privacy.differentialprivacy.pipelinedp4j.core.MetricType
31+ import com.google.privacy.differentialprivacy.pipelinedp4j.core.ScalarFeatureSpec
3032import com.google.privacy.differentialprivacy.pipelinedp4j.core.SelectPartitionsParams
33+ import com.google.privacy.differentialprivacy.pipelinedp4j.core.VectorFeatureSpec
3134import com.google.privacy.differentialprivacy.pipelinedp4j.proto.DpAggregates
32- import com.google.privacy.differentialprivacy.pipelinedp4j.proto.PerFeature
3335import com.google.privacy.differentialprivacy.pipelinedp4j.proto.copy
34- import com.google.privacy.differentialprivacy.pipelinedp4j.proto.dpAggregates
35- import com.google.privacy.differentialprivacy.pipelinedp4j.proto.perFeature
3636
3737sealed interface Query <ReturnT > {
3838 /* * Executes the query (in production mode). */
@@ -142,21 +142,6 @@ protected constructor(
142142 valueAndVectorAggs.map { it.getFeatureId() }
143143 }
144144 return aggResults
145- .zip(featureIdPerRun)
146- .map { (table, featureId) ->
147- table.mapValues(" TagWithFeatureId" , encoderFactory.protos(DpAggregates ::class )) { _, agg ->
148- if (featureId == null ) {
149- agg
150- } else {
151- val perFeature = constructPerFeature(agg, featureId)
152- dpAggregates {
153- count = agg.count
154- privacyIdCount = agg.privacyIdCount
155- this .perFeature + = perFeature
156- }
157- }
158- }
159- }
160145 .reduce {
161146 acc: FrameworkTable <GroupKeysT , DpAggregates >,
162147 table: FrameworkTable <GroupKeysT , DpAggregates > ->
@@ -494,43 +479,59 @@ protected constructor(
494479 valueAggregations : ValueAggregations <* >? ,
495480 vectorAggregations : VectorAggregations <* >? ,
496481 ): AggregationParams {
497- val valueContributionBounds = valueAggregations?.contributionBounds
498- val vectorContributionBounds = vectorAggregations?.vectorContributionBounds
482+ val nonFeatureMetrics =
483+ aggregationSpecs
484+ .filter { it is Count || it is PrivacyIdCount }
485+ .map { it.toNonFeatureMetricDefinition() }
486+ val features =
487+ buildList<FeatureSpec > {
488+ if (valueAggregations != null ) {
489+ val valueContributionBounds = valueAggregations.contributionBounds
490+ add(
491+ ScalarFeatureSpec (
492+ featureId = valueAggregations.getFeatureId(),
493+ metrics =
494+ valueAggregations.valueAggregationSpecs
495+ .map { it.toMetricDefinition() }
496+ .toImmutableList(),
497+ minValue = valueContributionBounds.valueBounds?.minValue,
498+ maxValue = valueContributionBounds.valueBounds?.maxValue,
499+ minTotalValue = valueContributionBounds.totalValueBounds?.minValue,
500+ maxTotalValue = valueContributionBounds.totalValueBounds?.maxValue,
501+ )
502+ )
503+ }
504+ if (vectorAggregations != null ) {
505+ val vectorContributionBounds = vectorAggregations.vectorContributionBounds
506+ add(
507+ VectorFeatureSpec (
508+ featureId = vectorAggregations.getFeatureId(),
509+ metrics =
510+ vectorAggregations.vectorAggregationSpecs
511+ .map { it.toMetricDefinition() }
512+ .toImmutableList(),
513+ vectorSize = vectorAggregations.vectorSize,
514+ normKind = vectorContributionBounds.maxVectorTotalNorm.normKind.toInternalNormKind(),
515+ vectorMaxTotalNorm = vectorContributionBounds.maxVectorTotalNorm.value,
516+ )
517+ )
518+ }
519+ }
520+
499521 return AggregationParams (
500- metrics = ImmutableList .copyOf(aggregationSpecs.metrics()),
522+ nonFeatureMetrics = nonFeatureMetrics.toImmutableList(),
523+ features = features.toImmutableList(),
501524 noiseKind =
502525 checkNotNull(noiseKind) { " noiseKind cannot be null if there are aggregations." }
503526 .toInternalNoiseKind(),
504527 maxPartitionsContributed = contributionBoundingLevel.getMaxPartitionsContributed(),
505528 maxContributionsPerPartition = contributionBoundingLevel.getMaxContributionsPerPartition(),
506- minValue = valueContributionBounds?.valueBounds?.minValue,
507- maxValue = valueContributionBounds?.valueBounds?.maxValue,
508- minTotalValue = valueContributionBounds?.totalValueBounds?.minValue,
509- maxTotalValue = valueContributionBounds?.totalValueBounds?.maxValue,
510- vectorNormKind = vectorContributionBounds?.maxVectorTotalNorm?.normKind?.toInternalNormKind(),
511- vectorMaxTotalNorm = vectorContributionBounds?.maxVectorTotalNorm?.value,
512- vectorSize = vectorAggregations?.vectorSize,
513529 partitionSelectionBudget = groupsType.getBudget()?.toInternalBudgetPerOpSpec(),
514530 preThreshold = groupsType.getPreThreshold(),
515531 contributionBoundingLevel = contributionBoundingLevel.toInternalContributionBoundingLevel(),
516532 partitionsBalance = groupByAdditionalParameters.groupsBalance.toPartitionsBalance(),
517533 )
518534 }
519-
520- companion object {
521- private fun constructPerFeature (dpAggregates : DpAggregates , featureId : String ): PerFeature {
522- return perFeature {
523- this .featureId = featureId
524- sum = dpAggregates.sum
525- mean = dpAggregates.mean
526- variance = dpAggregates.variance
527- if (dpAggregates.quantilesList.isNotEmpty()) {
528- quantiles + = dpAggregates.quantilesList
529- }
530- if (dpAggregates.vectorSumList.isNotEmpty()) {
531- vectorSum + = dpAggregates.vectorSumList
532- }
533- }
534- }
535- }
536535}
536+
537+ private fun <T : Any > Iterable<T>.toImmutableList (): ImmutableList <T > = ImmutableList .copyOf(this )
0 commit comments