Skip to content

Learning the 2nd smallest number with LCM Aggregation

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora
Description: Can a network find the second smallest number in a set?

View in Colab   GitHub source


Learning the 2nd smallest number with LCM Aggregation

Can a network find the second smallest number in a set? Each example is a set of 2 to 16 random 8-bit numbers (each number written as 8 bits); the target is the bits of the second smallest one. Simple aggregations like sum or max cannot do this, but the Learnable Commutative Monoid (LCM) aggregation (Ong & Veličković, 2022) can learn it.

Same model as PyG's examples/lcm_aggr_2nd_min.py; with 8,192 training sets instead of 65,536 and 20 epochs.

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

Every set becomes a small graph without edges: its nodes are the numbers.

import numpy as np
import keras
from keras import ops
from k3_node.data import Data
from k3_node.layers import LCMAggregation
from k3_node.loader import DataLoader

NUM_BITS = 8
rng = np.random.default_rng(0)


def random_set(min_size, max_size):
    bits = rng.integers(0, 2, size=(rng.integers(min_size, max_size + 1), NUM_BITS))
    numbers = bits @ (2 ** np.arange(NUM_BITS)[::-1])
    second_smallest = bits[np.argsort(numbers, kind="stable")[1]]
    return Data(x=bits.astype("float32"), y=second_smallest[None].astype("float32"))


train_loader = DataLoader([random_set(2, 16) for _ in range(8192)], batch_size=128, shuffle=True)
val_loader = DataLoader([random_set(32, 32) for _ in range(1024)], batch_size=128)  # larger sets than in training

Define the model

An encoder embeds each number, LCM aggregates the numbers of each set, and a decoder outputs the 8 bits. (PyG's per-bit embeddings summed together equal one Dense layer on the bit vector.)

class LCM(keras.Model):
    def __init__(self, emb_dim, dropout=0.25):
        super().__init__()
        self.encoder = keras.Sequential([
            keras.layers.Dense(emb_dim),  # sum of per-bit embeddings
            keras.layers.Dense(emb_dim),
            keras.layers.Dropout(0.5),
            keras.layers.Activation("gelu"),
        ])
        self.aggr = LCMAggregation(emb_dim, emb_dim, project=False)
        self.decoder = keras.Sequential([
            keras.layers.Dense(emb_dim),
            keras.layers.Dropout(dropout),
            keras.layers.Activation("gelu"),
            keras.layers.Dense(NUM_BITS),
        ])

    def call(self, data, training=False):
        x = self.encoder(data.x, training=training)
        x = self.aggr(x, index=data.batch, dim_size=data.num_graphs)
        return self.decoder(x, training=training)


model = LCM(emb_dim=128)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=1e-4),
    loss=keras.losses.BinaryCrossentropy(from_logits=True),
    metrics=[keras.metrics.BinaryAccuracy(threshold=0.0, name="bit_accuracy")],  # logits > 0 means bit = 1
)
model.fit(train_loader, validation_data=val_loader, epochs=20, verbose=2)

Evaluate

loss, bit_accuracy = model.evaluate(val_loader, verbose=0)
print(f"Bit accuracy on larger sets: {bit_accuracy:.4f}")