Skip to content

Learning the median with Equilibrium Aggregation

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora
Description: Can a neural network learn to compute the median of a set of numbers?

View in Colab   GitHub source


Learning the median with Equilibrium Aggregation

Can a neural network learn to compute the median of a set of numbers? Equilibrium Aggregation (Bartunov et al., 2022) aggregates a set by solving a small optimization problem, which makes order statistics like the median learnable. Each training example is a set of 100 numbers drawn from a normal, gamma or uniform distribution, and the target is its median.

Same model as PyG's examples/equilibrium_median.py; trained on a fixed set of 2,000 examples for 20 epochs instead of 10 million freshly sampled sets.

Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

Create the data

2,000 random sets of 100 numbers each, with their medians as targets.

import numpy as np
import keras
from keras import ops
from k3_node.layers import EquilibriumAggregation

rng = np.random.default_rng(0)


def random_set(size=100):
    kind = rng.integers(3)
    if kind == 0:
        return rng.normal(0.5, 0.4, size)
    if kind == 1:
        return rng.gamma(0.2, 1 / 0.5, size)
    return rng.uniform(0, 1, size)


sets = np.stack([random_set() for _ in range(2000)])[..., None].astype("float32")  # [2000, 100, 1]
medians = np.median(sets, axis=1)  # [2000, 1]

Define the model

The aggregation reduces each set of 100 numbers to one number.

class MedianModel(keras.Model):
    def __init__(self):
        super().__init__()
        self.aggr = EquilibriumAggregation(1, 10, [256, 256], 1)

    def call(self, sets):
        num_sets, set_size = ops.shape(sets)[0], ops.shape(sets)[1]
        x = ops.reshape(sets, (-1, 1))  # all numbers in one list...
        index = ops.repeat(ops.arange(num_sets), set_size)  # ...with the set each one belongs to
        return self.aggr(x, index=index, dim_size=num_sets)


model = MedianModel()

Train

model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.001), loss="mae")
model.fit(sets, medians, batch_size=32, epochs=20, validation_split=0.1, verbose=2)

Evaluate

test_sets = np.stack([random_set() for _ in range(200)])[..., None].astype("float32")
error = np.abs(model.predict(test_sets, verbose=0) - np.median(test_sets, axis=1)).mean()
print(f"Mean absolute error on new sets: {error:.4f}")