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?
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"
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)