Skip to content
Léonel Vodounou
Chapters

Logic Tensor Networks · Chapter 5

Case study: adding digits without ever labeling them

An LTN learns to recognize handwritten MNIST digits from the sum of two digits alone, and stays at 95% accuracy on adding two-digit numbers, where a purely supervised network drops to 45%.

Léonel VODOUNOU

September 25, 2026 · 42 min read

The first four parts presented LTN’s mechanisms on toy examples. This case study puts them to the test on a harder problem: recognizing handwritten digits without ever seeing the label of a single digit.

The model is given two images of digits, and a single piece of information: their sum. From there, it has to learn to recognize each digit. We then compare this LTN with a purely supervised baseline, which learns to predict the sum directly, with no logical knowledge at all.

Why this example?

This problem was not picked at random. It is a classic case in the neuro-symbolic literature, introduced by the DeepProbLog paper (Manhaeve et al., 2018) and taken up in the original LTN paper. Three qualities make it a good testbed for comparison:

  • it isolates the difficulty we care about. Recognizing a handwritten digit is a task neural networks handle well. The challenge here is to combine several pieces of information to produce an answer, which lets us observe what logic brings;
  • the rule is known: the answer is the sum of the digits. But the two approaches use it differently: the baseline has to learn it from the examples, whereas the LTN states it explicitly as a logical rule;
  • the difficulty is easy to tune, by increasing the number of digits to combine. We can then see how both approaches behave as the problem gets harder.

The problem

We consider two tasks, built from the MNIST dataset of handwritten digits.

Adding two digits. Consider the predicate addition(X,Y,n)\mathrm{addition}(X, Y, n), where XX and YY are images of digits and nn an integer equal to their sum. It must estimate whether the addition is valid: addition(img(8),img(3),11)\mathrm{addition}(\mathrm{img}(8), \mathrm{img}(3), 11) is valid, addition(img(3),img(3),5)\mathrm{addition}(\mathrm{img}(3), \mathrm{img}(3), 5) is not. Here, img(x)\mathrm{img}(x) simply means “an image of the digit xx”, not a logical function.

Adding two-digit numbers. The predicate becomes:

addition([img(X1),img(X2)], [img(Y1),img(Y2)], n)\mathrm{addition}([\mathrm{img}(X_1), \mathrm{img}(X_2)],\ [\mathrm{img}(Y_1), \mathrm{img}(Y_2)],\ n)

where each list represents a two-digit number. The first addition below is valid, the second is not:

addition([img(2),img(0)], [img(1),img(7)], 37)\mathrm{addition}([\mathrm{img}(2), \mathrm{img}(0)],\ [\mathrm{img}(1), \mathrm{img}(7)],\ 37) addition([img(5),img(4)], [img(9),img(0)], 50)\mathrm{addition}([\mathrm{img}(5), \mathrm{img}(4)],\ [\mathrm{img}(9), \mathrm{img}(0)],\ 50)

The neuro-symbolic approach is to learn a single-digit classifier, and to rely on what we already know about addition. If a predicate digit(x,d)\mathrm{digit}(x, d) gives the likelihood that image xx represents digit dd, the addition addition(img(3),img(8),11)\mathrm{addition}(\mathrm{img}(3), \mathrm{img}(8), 11) is written in LTN as:

∃d1,d2:d1+d2=11  (digit(img(3),d1)∧digit(img(8),d2))\exists d_1, d_2 : d_1 + d_2 = 11\ \ \big(\mathrm{digit}(\mathrm{img}(3), d_1) \land \mathrm{digit}(\mathrm{img}(8), d_2)\big)

The difficulty is that no label of an individual digit is given during training: only the sum of each pair is known. The classifier is therefore never trained directly on its own task. Its output remains latent information, used only by the logic. This is not a problem for LTN, since the gradient flows through the whole logical structure, all the way to the classifier’s weights.

The LTN theory

Domains:

  • images\mathit{images}: the MNIST digit images;
  • results\mathit{results}: the integers that are results of additions;
  • digits\mathit{digits}: the digits from 0 to 9.

Variables:

  • x,yx, y range over images, nn over results and d1,d2d_1, d_2 over digits;
  • D(x)=D(y)=imagesD(x) = D(y) = \mathit{images}, D(n)=resultsD(n) = \mathit{results}, D(d1)=D(d2)=digitsD(d_1) = D(d_2) = \mathit{digits}.

Predicate: digit(x,d)\mathrm{digit}(x, d), the single-digit classifier, which returns the probability that image xx represents digit dd. Its input domain is Din(digit)=images,digitsD_{in}(\mathrm{digit}) = \mathit{images}, \mathit{digits}.

Axioms. For adding two digits:

∀ Diag(x,y,n)  (∃d1,d2:d1+d2=n  (digit(x,d1)∧digit(y,d2)))\forall\, \mathrm{Diag}(x, y, n)\ \ \Big(\exists d_1, d_2 : d_1 + d_2 = n\ \ \big(\mathrm{digit}(x, d_1) \land \mathrm{digit}(y, d_2)\big)\Big)

and for adding two-digit numbers:

∀ Diag(x1,x2,y1,y2,n)  (∃d1,d2,d3,d4:10d1+d2+10d3+d4=n  (digit(x1,d1)∧digit(x2,d2)∧digit(y1,d3)∧digit(y2,d4)))\forall\, \mathrm{Diag}(x_1, x_2, y_1, y_2, n)\ \ \Big(\exists d_1, d_2, d_3, d_4 : 10d_1 + d_2 + 10d_3 + d_4 = n\ \ \big(\mathrm{digit}(x_1, d_1) \land \mathrm{digit}(x_2, d_2) \land \mathrm{digit}(y_1, d_3) \land \mathrm{digit}(y_2, d_4)\big)\Big)

Two mechanisms from part 2 come back:

  • diagonal quantification Diag(x,y,n)\mathrm{Diag}(x, y, n) ensures that the ii-th elements of xx, yy and nn match: (G(x)i,G(y)i,G(n)i)(\mathcal{G}(x)_i, \mathcal{G}(y)_i, \mathcal{G}(n)_i) is a triple from the dataset. LTN aggregates each pair of images with its own result, not with any result;
  • guarded quantification keeps, among the possible digits, only those whose sum can give the result. This is how the symbolic information enters the system.

Grounding:

  • G(images)=[0,1]28×28×1\mathcal{G}(\mathit{images}) = [0, 1]^{28 \times 28 \times 1}: MNIST images are 28 × 28 grayscale pixels, and their values, from 0 to 255 initially, are scaled to [0,1][0, 1];
  • G(results)=N\mathcal{G}(\mathit{results}) = \mathbb{N} and G(digits)={0,1,…,9}\mathcal{G}(\mathit{digits}) = \{0, 1, \dots, 9\};
  • G(x),G(y)∈[0,1]m×28×28×1\mathcal{G}(x), \mathcal{G}(y) \in [0, 1]^{m \times 28 \times 28 \times 1} and G(n)∈Nm\mathcal{G}(n) \in \mathbb{N}^m: the three variables have the same number mm of examples, since Diag\mathrm{Diag} matches them one to one;
  • G(d1)=G(d2)=⟨0,1,…,9⟩\mathcal{G}(d_1) = \mathcal{G}(d_2) = \langle 0, 1, \dots, 9 \rangle;
  • G(digit∣θ):x,d↦onehot(d)⊤⋅softmax(CNNθ(x))\mathcal{G}(\mathrm{digit} \mid \theta) : x, d \mapsto \mathrm{onehot}(d)^\top \cdot \mathrm{softmax}(\mathrm{CNN}_\theta(x)), where CNN\mathrm{CNN} is a convolutional network with 10 outputs, one per digit. Unlike in the previous parts, dd is a raw integer here, which onehot(d)\mathrm{onehot}(d) converts at evaluation time.

The figure from the original paper sums up the full computation graph for adding two digits:

LTN computation graph for adding two digits

The computation graph for adding two digits. The same convolutional network produces a distribution over digits for xx and for yy; the conjunction TPT_P combines them for each pair (d1,d2)(d_1, d_2); the mask d1+d2=nd_1 + d_2 = n keeps only the pairs compatible with the result; finally the existential (ApMA_{pM}) aggregates over the pairs and the universal (ApMEA_{pME}) over the examples, with Diag(x,y,n)\mathrm{Diag}(x, y, n). Source: Badreddine et al. (2022).

The dataset

MNIST contains 70,000 images: 60,000 for training and 10,000 for testing. We need to derive an addition dataset from it.

  • For two digits, the first 30,000 training images serve as left operands, the next 30,000 as right operands, and the sum of their labels is the target. The same is done for the test set.
  • For two two-digit numbers, the training set is split into four groups of 15,000 images: the first two form the digits of the first number, the last two those of the second.
import torch
import pandas as pd
import torchvision

def get_mnist_dataset_for_digits_addition(single_digit=True):
    """
    Prepares the dataset for the MNIST digit addition example (single digit or multiple digits)
    de l'article LTN.

    :param single_digit: whether the dataset has to be generated for the single digit case
        or the multi digit case (see the specification above to understand the difference between the two).
    :return: a pair of two elements. The first is the training set, the second the test set.
        Each one is a list containing:
        1. a list [left_operands, right_operands], where left_operands is a list of MNIST images
           used as the left operand of the addition, and right_operands as the right operand;
        2. a list containing the sum of the labels of the images of point 1. The label of the left operand
           is added to that of the right operand, forming the target of the addition task.
    This describes the result for the single digit case. In the multi digit case, the list of point 1
    has 4 elements, since four digits are involved in each addition (two to represent
    the first operand, two for the second).
    """
    if single_digit:
        n_train_examples = 30000
        n_test_examples = 5000
        n_operands = 2
    else:
        n_train_examples = 15000
        n_test_examples = 2500
        n_operands = 4

    mnist_train = torchvision.datasets.MNIST("./datasets/", train=True, download=True,
                                             transform=torchvision.transforms.ToTensor())
    mnist_test = torchvision.datasets.MNIST("./datasets/", train=False, download=True,
                                            transform=torchvision.transforms.ToTensor())

    train_imgs, train_labels, test_imgs, test_labels = mnist_train.data, mnist_train.targets, \
                                                       mnist_test.data, mnist_test.targets

    train_imgs, test_imgs = train_imgs / 255.0, test_imgs / 255.0

    train_imgs, test_imgs = torch.unsqueeze(train_imgs, 1), torch.unsqueeze(test_imgs, 1)

    imgs_operand_train = [train_imgs[i * n_train_examples:i * n_train_examples + n_train_examples]
                          for i in range(n_operands)]
    labels_operand_train = [train_labels[i * n_train_examples:i * n_train_examples + n_train_examples]
                            for i in range(n_operands)]

    imgs_operand_test = [test_imgs[i * n_test_examples:i * n_test_examples + n_test_examples]
                         for i in range(n_operands)]
    labels_operand_test = [test_labels[i * n_test_examples:i * n_test_examples + n_test_examples]
                           for i in range(n_operands)]

    if single_digit:
        label_addition_train = labels_operand_train[0] + labels_operand_train[1]
        label_addition_test = labels_operand_test[0] + labels_operand_test[1]
    else:
        label_addition_train = 10 * labels_operand_train[0] + labels_operand_train[1] + \
                               10 * labels_operand_train[2] + labels_operand_train[3]

        label_addition_test = 10 * labels_operand_test[0] + labels_operand_test[1] + \
                              10 * labels_operand_test[2] + labels_operand_test[3]

    train_set = [torch.stack(imgs_operand_train, dim=1), label_addition_train]
    test_set = [torch.stack(imgs_operand_test, dim=1), label_addition_test]

    return train_set, test_set

# dataset for the single digit case
single_d_train_set, single_d_test_set = get_mnist_dataset_for_digits_addition(single_digit=True)
# dataset for the multi digit case
multi_d_train_set, multi_d_test_set = get_mnist_dataset_for_digits_addition(single_digit=False)
Understanding the structure of the data

There are three levels to tell apart: an example, the dataset and a batch.

An example is one addition. For 3+7=103 + 7 = 10, it contains two images and one label:

Example
├── image of the digit 3
├── image of the digit 7
└── label = 10

Each MNIST image is a tensor of shape (1, 28, 28): one channel (grayscale), 28 pixels high, 28 wide. The two images of an example are grouped in a tensor of shape (2, 1, 28, 28).

The dataset gathers mm examples. The images form a single tensor of shape (m, 2, 1, 28, 28), and the labels a tensor of shape (m,), with a direct correspondence: images[i] and labels[i] describe the same addition.

(m, 2, 1, 28, 28)
 ↑  ↑
 │  └── number of images per example
 └───── number of examples

This is exactly what single_d_train_set holds: a list of two elements, the images at [0] and the labels at [1]. single_d_train_set[0][0] first selects the images, then the first example: a (2, 1, 28, 28) tensor. single_d_train_set[1][0] selects the label of that same example, a scalar.

A batch gathers BB examples: its images have shape (B, 2, 1, 28, 28) and its labels (B,). The first two dimensions must not be confused: (32, 2, 1, 28, 28) represents 32 additions, each described by 2 images.

For adding two-digit numbers, the principle is the same with 4 images per example: (m, 4, 1, 28, 28) for the dataset, (B, 4, 1, 28, 28) for a batch.

Let’s display the first training example for adding two digits:

import matplotlib.pyplot as plt
first_example_images = single_d_train_set[0][0]
first_example_label = single_d_train_set[1][0]

print("The operands are shown in the following images:")
fig = plt.figure()
ax = fig.add_subplot(1, 2, 1)
imgplot = plt.imshow(first_example_images[0].permute(1, 2, 0))
ax.set_title('First operand')
ax = fig.add_subplot(1, 2, 2)
imgplot = plt.imshow(first_example_images[1].permute(1, 2, 0))
ax.set_title('Second operand')
plt.show()
print("The target label (the sum) for these operands is: %d" % first_example_label.item())

The two images of the first training example: a 5 and a 3

The operands are shown in the following images:
The target label (the sum) for these operands is: 8

A 5 and a 3, adding up to 8: this is the only information the model will get for this example. (The plot titles were generated by the notebook, which is in French.)

The LTN model

We need to define the digit\mathrm{digit} predicate, the variables d1d_1 to d4d_4, the connectives, the quantifiers and SatAgg. Connectives and quantifiers follow the product configuration seen in part 3.

The digit\mathrm{digit} predicate relies on two models:

  • the first is a CNN that returns the logits of the ten classes for an image xx;
  • the second takes a pair (x,d)(x, d), computes the logits with the first, applies a softmax, and returns the probability of class dd: the likelihood that image xx represents digit dd.

We keep the two separate because we need both outputs: the logits are used to measure classification accuracy, and the probabilities are read as degrees of truth to compute the satisfaction of the knowledge base.

The variables d1d_1 to d4d_4 range over the ten digits, ⟨0,1,…,9⟩\langle 0, 1, \dots, 9 \rangle. Unlike the one-hot labels of the previous parts, these are raw integer indices: torch.gather directly selects the probability of index dd.

from torch.nn.init import xavier_uniform_, normal_, kaiming_uniform_
import ltn

# we define the variables
d_1 = ltn.Variable("d_1", torch.tensor(range(10)))
d_2 = ltn.Variable("d_2", torch.tensor(range(10)))
# only used in the multi digit case
d_3 = ltn.Variable("d_3", torch.tensor(range(10)))
d_4 = ltn.Variable("d_4", torch.tensor(range(10)))

# we define the digit predicate
class MNISTConv(torch.nn.Module):
    """
    CNN that returns embeddings for MNIST images.
    Arguments:
        conv_channels_sizes: tuple containing the number of channels of the convolutional layers of the model. The first
        element must be the number of input channels of the first layer, the last one the number of output
        channels of the last layer. The number of layers built is `len(conv_channels_sizes) - 1`;
        
        kernel_sizes: tuple containing the sizes of the kernels used in the convolutional layers;
        
        linear_layers_sizes: tuple containing the sizes of the final dense layers of the architecture. The first
        element must be the number of input features of the first dense layer, the last one the number of
        output features of the last one. The number of layers built is `len(linear_layers_sizes) - 1`.
    """
    def __init__(self, conv_channels_sizes=(1, 6, 16), kernel_sizes=(5, 5), linear_layers_sizes=(256, 100)):
        super(MNISTConv, self).__init__()
        self.conv_layers = torch.nn.ModuleList([torch.nn.Conv2d(conv_channels_sizes[i - 1], conv_channels_sizes[i],
                                                                kernel_sizes[i - 1])
                                                  for i in range(1, len(conv_channels_sizes))])
        self.elu = torch.nn.ELU()  # activation used everywhere, as in baselines.py (Keras, activation="elu")
        self.maxpool = torch.nn.MaxPool2d((2, 2))
        self.linear_layers = torch.nn.ModuleList([torch.nn.Linear(linear_layers_sizes[i - 1], linear_layers_sizes[i])
                                                  for i in range(1, len(linear_layers_sizes))])

        self.init_weights()

    def forward(self, x):
        for conv in self.conv_layers:
            x = self.elu(conv(x))
            x = self.maxpool(x)
        x = torch.flatten(x, start_dim=1)
        for linear in self.linear_layers:
            x = self.elu(linear(x))
        return x

    def init_weights(self):
        r"""Initializes the weights of the network.
        All weights (convolutional and dense) are initialized with :py:func:`torch.nn.init.xavier_uniform_`
        (equivalent to Keras's default "glorot_uniform" initializer, used in baselines.py),
        biases with :py:func:`torch.nn.init.zeros_`.
        """
        for layer in self.conv_layers:
            xavier_uniform_(layer.weight)
            torch.nn.init.zeros_(layer.bias)

        for layer in self.linear_layers:
            xavier_uniform_(layer.weight)
            torch.nn.init.zeros_(layer.bias)


class SingleDigitClassifier(torch.nn.Module):
    """
    Model that classifies an MNIST digit image among 10 possible classes. It returns the logits, an
    unnormalized output. Architecture faithful to baselines.SingleDigit: convolutional part (MNISTConv, ELU), then a
    hidden dense layer (ELU), then the final classification layer
    Arguments:
        layers_sizes: tuple containing the sizes of the final dense layers of the architecture. The first
        element must be the number of input features of the first layer, the last one the number of
        output features of the last one. The number of layers built is `len(layers_sizes) - 1`.
    """
    def __init__(self, layers_sizes=(100, 84, 10)):
        super(SingleDigitClassifier, self).__init__()
        self.mnistconv = MNISTConv()  # convolutional part of the architecture
        self.elu = torch.nn.ELU()  # activation of the hidden dense layers
        self.linear_layers = torch.nn.ModuleList([torch.nn.Linear(layers_sizes[i - 1], layers_sizes[i])
                                                  for i in range(1, len(layers_sizes))])
        self.init_weights()

    def forward(self, x):
        x = self.mnistconv(x)
        for i in range(len(self.linear_layers) - 1):
            x = self.elu(self.linear_layers[i](x))
        return self.linear_layers[-1](x)  # a sigmoid or a softmax must be applied on the last layer

    def init_weights(self):
        """Initializes the weights of the dense layers of the network.
        Weights are initialized with :py:func:`torch.nn.init.xavier_uniform_`,
        biases with :py:func:`torch.nn.init.zeros_`.
        """
        for layer in self.linear_layers:
            xavier_uniform_(layer.weight)
            torch.nn.init.zeros_(layer.bias)


class LogitsToPredicate(torch.nn.Module):
    """
    This model wraps a logits model, which computes the class logits for an input image x.
    The idea is to keep logits and probabilities separate: the logits model returns the logits for an example,
    while this model returns the probabilities computed from those logits.

    Concretely, it takes an image x and a class label d as input. It applies the logits model
    to x to get the logits, then a softmax function to get the probability of each class. Finally, it
    only returns the probability of class d, selected directly by indexing rather than
    by a product with a one-hot vector, since d is a raw integer index here.
    """
    def __init__(self, logits_model):
        super(LogitsToPredicate, self).__init__()
        self.logits_model = logits_model
        self.softmax = torch.nn.Softmax(dim=1)

    def forward(self, x, d):
        logits = self.logits_model(x)
        probs = self.softmax(logits)
        out = torch.gather(probs, 1, d)
        return out

# single digit
cnn_s_d = SingleDigitClassifier()
Digit_s_d = ltn.Predicate(LogitsToPredicate(cnn_s_d)).to(ltn.device)
# multiple digits
cnn_m_d = SingleDigitClassifier()
Digit_m_d = ltn.Predicate(LogitsToPredicate(cnn_m_d)).to(ltn.device)

# we define the connectives, quantifiers, and SatAgg
And = ltn.Connective(ltn.fuzzy_ops.AndProd())
Exists = ltn.Quantifier(ltn.fuzzy_ops.AggregPMean(p=2), quantifier="e")
Forall = ltn.Quantifier(ltn.fuzzy_ops.AggregPMeanError(p=2), quantifier="f")
SatAgg = ltn.fuzzy_ops.SatAgg()

The baseline

For comparison, we also train a standard network, with no logic. It encodes each image with the same convolutional trunk as the LTN (weights shared across operands), concatenates the representations, then predicts the sum directly. It is a 19-class classification for two digits (0 to 18), and a 199-class one for two two-digit numbers (0 to 198).

# ===============================================
# BASELINE (standard CNN, no logic)
# ===============================================

class BaselineNet(torch.nn.Module):
    """
    "Baseline" network: encodes each operand image with the SAME convolutional trunk
    (MNISTConv, shared weights, siamese architecture), concatenates the embeddings, then a
    hidden dense layer (ELU) and a final classification layer (raw logits) directly predict
    the class of the result (the sum), with no intermediate logical predicates.
    """
    def __init__(self, n_operands, n_classes, embedding_dim=100, hidden_dim=100):
        super(BaselineNet, self).__init__()
        self.encoder = MNISTConv()  # convolutional trunk shared by the operands
        self.fc1 = torch.nn.Linear(embedding_dim * n_operands, hidden_dim)
        self.fc2 = torch.nn.Linear(hidden_dim, n_classes)
        self.elu = torch.nn.ELU()
        self.init_weights()

    def init_weights(self):
        xavier_uniform_(self.fc1.weight)
        torch.nn.init.zeros_(self.fc1.bias)
        xavier_uniform_(self.fc2.weight)
        torch.nn.init.zeros_(self.fc2.bias)

    def forward(self, images_list):
        embeddings = [self.encoder(img) for img in images_list]
        x = torch.cat(embeddings, dim=1)
        x = self.elu(self.fc1(x))
        return self.fc2(x)  # raw logits


# baseline for the single digit case: sum of 2 digits -> 19 classes (0 to 18)
baseline_s_d = BaselineNet(n_operands=2, n_classes=19, hidden_dim=84).to(ltn.device)
# baseline for the multi digit case: sum of two 2-digit numbers -> 199 classes (0 to 198)
baseline_m_d = BaselineNet(n_operands=4, n_classes=199, hidden_dim=128).to(ltn.device)

Finally, we define a data loader, for training and testing on both tasks:

import numpy as np

# standard PyTorch data loader, for training and testing the model
class DataLoader(object):
    def __init__(self,
                 fold,
                 batch_size=1,
                 shuffle=True):
        self.fold = fold
        self.batch_size = batch_size
        self.shuffle = shuffle

    def __len__(self):
        return int(np.ceil(self.fold[0].shape[0] / self.batch_size))

    def __iter__(self):
        n = self.fold[0].shape[0]
        idxlist = list(range(n))
        if self.shuffle:
            np.random.shuffle(idxlist)

        for _, start_idx in enumerate(range(0, n, self.batch_size)):
            end_idx = min(start_idx + self.batch_size, n)
            digits = self.fold[0][idxlist[start_idx:end_idx]]
            addition_labels = self.fold[1][idxlist[start_idx:end_idx]]

            yield digits, addition_labels

# create the training and test loaders
single_train_loader = DataLoader(single_d_train_set, 32, shuffle=True)
single_test_loader = DataLoader(single_d_test_set, 32, shuffle=False)
multi_d_train_loader = DataLoader(multi_d_train_set, 32, shuffle=True)
multi_d_test_loader = DataLoader(multi_d_test_set, 32, shuffle=False)

The two approaches therefore share everything, except how they learn:

LTNBaseline
Input dataidenticalidentical
Convolutional trunkidentical architectureidentical architecture
Batch size3232
OptimizerAdam, lr=0.001lr = 0.001Adam, lr=0.001lr = 0.001
Objectivemaximize satisfactionminimize the error on the sum
Loss function1−SatAgg1 - \mathrm{SatAgg}cross-entropy
Logical knowledge∃d1,d2:d1+d2=n\exists d_1, d_2 : d_1 + d_2 = nnone
Epochs2030

The baseline is trained longer (30 epochs versus 20), following the protocol of the paper.

Learning

Let DD be the set of examples. The objective is:

SatAggϕ∈K Gθ, x←D(ϕ)\mathrm{SatAgg}_{\phi \in \mathcal{K}}\ \mathcal{G}_{\theta,\, x \leftarrow D}(\phi)

that is, each formula of K\mathcal{K} is evaluated on the examples of DD with the model of parameters θ\theta, then SatAgg\mathrm{SatAgg} aggregates these satisfactions. In practice, the optimizer minimizes, on a mini-batch BB drawn from DD:

L=1−SatAggϕ∈K Gθ, x←B(ϕ)L = 1 - \mathrm{SatAgg}_{\phi \in \mathcal{K}}\ \mathcal{G}_{\theta,\, x \leftarrow B}(\phi)

Since the knowledge base here contains a single axiom, SatAgg is not used explicitly: applied to a single value, it would return it unchanged. The result of Forall directly serves as the satisfaction level.

The schedule of pp. The parameter pp of the existential quantifier increases in steps during training: p=1p = 1, then 2, then 4, then 6. At the start, the classifier is close to random. A high pp would focus the whole gradient on the best hypothesis of the moment, and could reinforce a wrong combination of digits just because it has the highest confidence at initialization. A low pp spreads the gradient over all plausible combinations and lets the right signal emerge; the focus is then tightened once the classifier becomes reliable. This is a direct application of the trade-off seen in part 3.

Accuracy is measured by predicting each digit with the CNN, then checking that the reconstructed addition gives the right result.

Here is the core of training, the formula evaluated on each batch:

sat_agg = Forall(
    ltn.diag(images_x, images_y, labels_z),
    Exists(
        [d_1, d_2],
        And(Digit_s_d(images_x, d_1), Digit_s_d(images_y, d_2)),
        cond_vars=[d_1, d_2, labels_z],
        cond_fn=lambda d1, d2, z: torch.eq(d1.value + d2.value, z.value),
        p=p
    )).value

It reads: for each example, there must be two digits that match the two images and whose sum equals the observed result.

How this formula learns the digits, step by step

Take a batch of two examples: 2+3=52 + 3 = 5 and 1+4=51 + 4 = 5.

1. The data. The first images go into images_x, the second ones into images_y, the results into labels_z, each with a batch dimension of 2. ltn.diag(images_x, images_y, labels_z) keeps the triples aligned: (image of 2, image of 3, 5) and (image of 1, image of 4, 5), without crossing the images of one example with the result of another.

2. The Digit predicate. For an image, the CNN produces logits, then a softmax gives a distribution over the ten digits, for example:

d      0     1     2     3     4     5   ...
p_d   0.01  0.02  0.90  0.01  0.01  0.01 ...

Digit_s_d(x, d) is the probability of digit dd: here Digit(x,2)=0.90\mathrm{Digit}(x, 2) = 0.90, a high degree of truth for “image xx represents a 2”.

3. The admissible pairs. d1d_1 and d2d_2 each range over {0,…,9}\{0, \dots, 9\}, giving 100 possible pairs. The cond_fn mask only keeps those whose sum equals the result. For z=5z = 5: (0,5),(1,4),(2,3),(3,2),(4,1),(5,0)(0, 5), (1, 4), (2, 3), (3, 2), (4, 1), (5, 0). This mask recognizes nothing in the pixels: it only says which combinations of digits are compatible with the label.

4. The conjunction. For each admissible pair, And multiplies the two probabilities. Suppose the network gives 0.90 to the 2 for the first image and 0.85 to the 3 for the second:

(d1,d2)(d_1, d_2)Digit(x,d1)\mathrm{Digit}(x, d_1)Digit(y,d2)\mathrm{Digit}(y, d_2)AndProd
(0, 5)0.010.020.0002
(1, 4)0.020.010.0002
(2, 3)0.900.850.7650
(3, 2)0.010.030.0003
(4, 1)0.010.020.0002
(5, 0)0.010.010.0001

Several pairs are mathematically correct, and the constraint alone is not enough to decide. The Digit predicate brings the visual information: the logic says d1+d2=5d_1 + d_2 = 5, and the CNN says which values are the most plausible given the images.

5. The existential. Exists aggregates the six values with pMean: for p=2p = 2, (16∑ivi2)1/2\big(\frac{1}{6} \sum_i v_i^2\big)^{1/2}. The result is higher when at least one explanation is very plausible, here the pair (2, 3).

6. The universal. Exists produces one satisfaction per example, then Forall aggregates them with pMeanError into an overall satisfaction for the batch. A badly explained example weighs more than one that is already well satisfied.

7. The gradient. This whole chain is differentiable: CNN weights → logits → softmax → Digit → And → Exists → Forall. The gradient flows back the other way. For the pair (2, 3), the satisfaction is v=abv = ab, with a=Px(2)a = P_x(2) and b=Py(3)b = P_y(3), so ∂v∂a=b\frac{\partial v}{\partial a} = b and ∂v∂b=a\frac{\partial v}{\partial b} = a: the gradient acts on the probabilities, then on the logits, then on the CNN weights.

8. Why the digits end up being learned. A single example like 2+3=52 + 3 = 5 is not enough to identify both digits. The information comes from all the examples together: a 2 also appears in 2+4=62 + 4 = 6, 2+1=32 + 1 = 3, 2+5=72 + 5 = 7… The hypothesis “this is a 2” stays consistent with all these additions, while a wrong hypothesis becomes incompatible with a growing number of examples. Since the same CNN is shared by every example, the gradients of all these constraints add up, and push it toward distributions where P(2∣x)≈1P(2 \mid x) \approx 1 for images of 2. The individual digits act as latent variables: the CNN provides the predictions, the logic provides the constraints.

Adding two digits

We train the LTN for 20 epochs:

history_single = {"train_loss": [], "train_sat": [], "train_acc": [],
                   "test_loss": [], "test_sat": [], "test_acc": []}

optimizer = torch.optim.Adam(Digit_s_d.parameters(), lr=0.001)

for epoch in range(20):
    # schedule of the parameter p for the existential quantifier, as described in the LTN paper
    if epoch in range(0, 4):
        p = 1
    if epoch in range(4, 8):
        p = 2
    if epoch in range(8, 12):
        p = 4
    if epoch in range(12, 20):
        p = 6
    train_loss, test_loss = 0.0, 0.0
    train_sat, test_sat = 0.0, 0.0
    train_acc, test_acc = 0.0, 0.0
    # training step
    cnn_s_d.train()
    for batch_idx, (operand_images, addition_label) in enumerate(single_train_loader):
        optimizer.zero_grad()
        # ground the variables with the current batch data
        images_x = ltn.Variable("x", operand_images[:, 0])
        images_y = ltn.Variable("y", operand_images[:, 1])
        labels_z = ltn.Variable("z", addition_label)
        sat_agg = Forall(
            ltn.diag(images_x, images_y, labels_z),
            Exists(
                [d_1, d_2],
                And(Digit_s_d(images_x, d_1), Digit_s_d(images_y, d_2)),
                cond_vars=[d_1, d_2, labels_z],
                cond_fn=lambda d1, d2, z: torch.eq(d1.value + d2.value, z.value),
                p=p
            )).value
        loss = 1. - sat_agg
        loss.backward()
        optimizer.step()
        train_loss += loss.item()
        train_sat += sat_agg.item()
        # compute the training accuracy
        predictions_x = torch.argmax(cnn_s_d(operand_images[:, 0].to(ltn.device)), dim=1)
        predictions_y = torch.argmax(cnn_s_d(operand_images[:, 1].to(ltn.device)), dim=1)
        predictions = predictions_x + predictions_y
        train_acc += torch.count_nonzero(torch.eq(addition_label.to(ltn.device), predictions)) / predictions.shape[0]
    train_loss = train_loss / len(single_train_loader)
    train_sat = train_sat / len(single_train_loader)
    train_acc = train_acc / len(single_train_loader)

    # test step
    cnn_s_d.eval()
    with torch.no_grad():
        for batch_idx, (operand_images, addition_label) in enumerate(single_test_loader):
            # ground the variables with the current batch data
            images_x = ltn.Variable("x", operand_images[:, 0])
            images_y = ltn.Variable("y", operand_images[:, 1])
            labels_z = ltn.Variable("z", addition_label)
            sat_agg = Forall(
                ltn.diag(images_x, images_y, labels_z),
                Exists(
                    [d_1, d_2],
                    And(Digit_s_d(images_x, d_1), Digit_s_d(images_y, d_2)),
                    cond_vars=[d_1, d_2, labels_z],
                    cond_fn=lambda d1, d2, z: torch.eq(d1.value + d2.value, z.value),
                    p=p
                )).value
            loss = 1. - sat_agg
            test_loss += loss.item()
            test_sat += sat_agg.item()
            # compute the test accuracy
            predictions_x = torch.argmax(cnn_s_d(operand_images[:, 0].to(ltn.device)), dim=1)
            predictions_y = torch.argmax(cnn_s_d(operand_images[:, 1].to(ltn.device)), dim=1)
            predictions = predictions_x + predictions_y
            test_acc += torch.count_nonzero(torch.eq(addition_label.to(ltn.device), predictions)) / predictions.shape[0]
    test_loss = test_loss / len(single_test_loader)
    test_sat = test_sat / len(single_test_loader)
    test_acc = test_acc / len(single_test_loader)

    # record the history
    history_single["train_loss"].append(float(train_loss))
    history_single["train_sat"].append(float(train_sat))
    history_single["train_acc"].append(float(train_acc))
    history_single["test_loss"].append(float(test_loss))
    history_single["test_sat"].append(float(test_sat))
    history_single["test_acc"].append(float(test_acc))

    # print the metrics at each epoch
    print(" epoch %d | Train loss %.4f | Train Sat %.4f | Train Acc %.4f | Test loss %.4f | Test Sat %.4f |"
                " Test Acc %.4f " %
          (epoch, train_loss, train_sat, train_acc, test_loss, test_sat, test_acc))
 epoch 0 | Train loss 0.8594 | Train Sat 0.1406 | Train Acc 0.8124 | Test loss 0.8325 | Test Sat 0.1675 | Test Acc 0.9357
 epoch 1 | Train loss 0.8349 | Train Sat 0.1651 | Train Acc 0.9415 | Test loss 0.8293 | Test Sat 0.1707 | Test Acc 0.9496
 epoch 2 | Train loss 0.8313 | Train Sat 0.1687 | Train Acc 0.9568 | Test loss 0.8282 | Test Sat 0.1718 | Test Acc 0.9554
 epoch 3 | Train loss 0.8298 | Train Sat 0.1702 | Train Acc 0.9624 | Test loss 0.8275 | Test Sat 0.1725 | Test Acc 0.9616
 epoch 4 | Train loss 0.6184 | Train Sat 0.3816 | Train Acc 0.9614 | Test loss 0.6125 | Test Sat 0.3875 | Test Acc 0.9630
 epoch 5 | Train loss 0.6125 | Train Sat 0.3875 | Train Acc 0.9699 | Test loss 0.6102 | Test Sat 0.3898 | Test Acc 0.9662
 epoch 6 | Train loss 0.6090 | Train Sat 0.3910 | Train Acc 0.9751 | Test loss 0.6097 | Test Sat 0.3903 | Test Acc 0.9668
 epoch 7 | Train loss 0.6085 | Train Sat 0.3915 | Train Acc 0.9761 | Test loss 0.6068 | Test Sat 0.3932 | Test Acc 0.9739
 epoch 8 | Train loss 0.4046 | Train Sat 0.5954 | Train Acc 0.9711 | Test loss 0.3988 | Test Sat 0.6012 | Test Acc 0.9703
 epoch 9 | Train loss 0.3984 | Train Sat 0.6016 | Train Acc 0.9756 | Test loss 0.3984 | Test Sat 0.6016 | Test Acc 0.9703
 epoch 10 | Train loss 0.3970 | Train Sat 0.6030 | Train Acc 0.9765 | Test loss 0.4024 | Test Sat 0.5976 | Test Acc 0.9668
 epoch 11 | Train loss 0.3956 | Train Sat 0.6044 | Train Acc 0.9774 | Test loss 0.3956 | Test Sat 0.6044 | Test Acc 0.9731
 epoch 12 | Train loss 0.3032 | Train Sat 0.6968 | Train Acc 0.9782 | Test loss 0.3022 | Test Sat 0.6978 | Test Acc 0.9749
 epoch 13 | Train loss 0.3030 | Train Sat 0.6970 | Train Acc 0.9777 | Test loss 0.3082 | Test Sat 0.6918 | Test Acc 0.9707
 epoch 14 | Train loss 0.3016 | Train Sat 0.6984 | Train Acc 0.9789 | Test loss 0.3093 | Test Sat 0.6907 | Test Acc 0.9691
 epoch 15 | Train loss 0.2990 | Train Sat 0.7010 | Train Acc 0.9799 | Test loss 0.2994 | Test Sat 0.7006 | Test Acc 0.9765
 epoch 16 | Train loss 0.2984 | Train Sat 0.7016 | Train Acc 0.9806 | Test loss 0.3218 | Test Sat 0.6782 | Test Acc 0.9598
 epoch 17 | Train loss 0.2977 | Train Sat 0.7023 | Train Acc 0.9807 | Test loss 0.3004 | Test Sat 0.6996 | Test Acc 0.9755
 epoch 18 | Train loss 0.2982 | Train Sat 0.7018 | Train Acc 0.9807 | Test loss 0.3043 | Test Sat 0.6957 | Test Acc 0.9723
 epoch 19 | Train loss 0.2961 | Train Sat 0.7039 | Train Acc 0.9818 | Test loss 0.3090 | Test Sat 0.6910 | Test Acc 0.9701

Then the baseline, for 30 epochs:

history_baseline_single = {"train_loss": [], "train_acc": [], "test_loss": [], "test_acc": []}

optimizer_baseline_s = torch.optim.Adam(baseline_s_d.parameters(), lr=0.001)
criterion = torch.nn.CrossEntropyLoss()

for epoch in range(30):
    train_loss, train_acc = 0.0, 0.0
    # training step
    baseline_s_d.train()
    for batch_idx, (operand_images, addition_label) in enumerate(single_train_loader):
        images_x = operand_images[:, 0].to(ltn.device)
        images_y = operand_images[:, 1].to(ltn.device)
        labels = addition_label.to(ltn.device).long()

        optimizer_baseline_s.zero_grad()
        logits = baseline_s_d([images_x, images_y])
        loss = criterion(logits, labels)
        loss.backward()
        optimizer_baseline_s.step()

        train_loss += loss.item()
        predictions = torch.argmax(logits, dim=1)
        train_acc += torch.count_nonzero(torch.eq(labels, predictions)) / predictions.shape[0]
    train_loss = train_loss / len(single_train_loader)
    train_acc = train_acc / len(single_train_loader)

    # test step
    test_loss, test_acc = 0.0, 0.0
    baseline_s_d.eval()
    with torch.no_grad():
        for batch_idx, (operand_images, addition_label) in enumerate(single_test_loader):
            images_x = operand_images[:, 0].to(ltn.device)
            images_y = operand_images[:, 1].to(ltn.device)
            labels = addition_label.to(ltn.device).long()

            logits = baseline_s_d([images_x, images_y])
            loss = criterion(logits, labels)

            test_loss += loss.item()
            predictions = torch.argmax(logits, dim=1)
            test_acc += torch.count_nonzero(torch.eq(labels, predictions)) / predictions.shape[0]
    test_loss = test_loss / len(single_test_loader)
    test_acc = test_acc / len(single_test_loader)

    history_baseline_single["train_loss"].append(float(train_loss))
    history_baseline_single["train_acc"].append(float(train_acc))
    history_baseline_single["test_loss"].append(float(test_loss))
    history_baseline_single["test_acc"].append(float(test_acc))

    print(" epoch %d | [Baseline] Train loss %.4f | Train Acc %.4f | Test loss %.4f | Test Acc %.4f " %
          (epoch, train_loss, train_acc, test_loss, test_acc))
 epoch 0 | [Baseline] Train loss 1.6665 | Train Acc 0.4358 | Test loss 0.5332 | Test Acc 0.8404
 epoch 1 | [Baseline] Train loss 0.3587 | Train Acc 0.8948 | Test loss 0.2420 | Test Acc 0.9270
 epoch 2 | [Baseline] Train loss 0.1966 | Train Acc 0.9409 | Test loss 0.1737 | Test Acc 0.9479
 epoch 3 | [Baseline] Train loss 0.1388 | Train Acc 0.9589 | Test loss 0.1846 | Test Acc 0.9467
 epoch 4 | [Baseline] Train loss 0.1040 | Train Acc 0.9700 | Test loss 0.1476 | Test Acc 0.9538
 epoch 5 | [Baseline] Train loss 0.0799 | Train Acc 0.9753 | Test loss 0.1356 | Test Acc 0.9618
 epoch 6 | [Baseline] Train loss 0.0633 | Train Acc 0.9800 | Test loss 0.1477 | Test Acc 0.9622
 epoch 7 | [Baseline] Train loss 0.0508 | Train Acc 0.9836 | Test loss 0.1554 | Test Acc 0.9582
 epoch 8 | [Baseline] Train loss 0.0468 | Train Acc 0.9844 | Test loss 0.1357 | Test Acc 0.9646
 epoch 9 | [Baseline] Train loss 0.0361 | Train Acc 0.9883 | Test loss 0.1743 | Test Acc 0.9586
 epoch 10 | [Baseline] Train loss 0.0329 | Train Acc 0.9890 | Test loss 0.1625 | Test Acc 0.9642
 epoch 11 | [Baseline] Train loss 0.0253 | Train Acc 0.9919 | Test loss 0.1950 | Test Acc 0.9596
 epoch 12 | [Baseline] Train loss 0.0290 | Train Acc 0.9903 | Test loss 0.1727 | Test Acc 0.9646
 epoch 13 | [Baseline] Train loss 0.0270 | Train Acc 0.9917 | Test loss 0.2054 | Test Acc 0.9592
 epoch 14 | [Baseline] Train loss 0.0307 | Train Acc 0.9905 | Test loss 0.2011 | Test Acc 0.9640
 epoch 15 | [Baseline] Train loss 0.0186 | Train Acc 0.9939 | Test loss 0.2010 | Test Acc 0.9668
 epoch 16 | [Baseline] Train loss 0.0187 | Train Acc 0.9942 | Test loss 0.1747 | Test Acc 0.9636
 epoch 17 | [Baseline] Train loss 0.0255 | Train Acc 0.9924 | Test loss 0.2053 | Test Acc 0.9654
 epoch 18 | [Baseline] Train loss 0.0243 | Train Acc 0.9926 | Test loss 0.1583 | Test Acc 0.9703
 epoch 19 | [Baseline] Train loss 0.0175 | Train Acc 0.9942 | Test loss 0.1953 | Test Acc 0.9670
 epoch 20 | [Baseline] Train loss 0.0226 | Train Acc 0.9931 | Test loss 0.2413 | Test Acc 0.9648
 epoch 21 | [Baseline] Train loss 0.0204 | Train Acc 0.9945 | Test loss 0.2181 | Test Acc 0.9668
 epoch 22 | [Baseline] Train loss 0.0218 | Train Acc 0.9934 | Test loss 0.2141 | Test Acc 0.9672
 epoch 23 | [Baseline] Train loss 0.0166 | Train Acc 0.9946 | Test loss 0.2378 | Test Acc 0.9638
 epoch 24 | [Baseline] Train loss 0.0187 | Train Acc 0.9943 | Test loss 0.2273 | Test Acc 0.9693
 epoch 25 | [Baseline] Train loss 0.0148 | Train Acc 0.9955 | Test loss 0.2393 | Test Acc 0.9674
 epoch 26 | [Baseline] Train loss 0.0194 | Train Acc 0.9940 | Test loss 0.2325 | Test Acc 0.9684
 epoch 27 | [Baseline] Train loss 0.0205 | Train Acc 0.9944 | Test loss 0.2574 | Test Acc 0.9628
 epoch 28 | [Baseline] Train loss 0.0177 | Train Acc 0.9950 | Test loss 0.2633 | Test Acc 0.9705
 epoch 29 | [Baseline] Train loss 0.0189 | Train Acc 0.9948 | Test loss 0.2083 | Test Acc 0.9709

Both approaches end up at the same level: about 97% test accuracy. On this simple case, knowing the rule brings no measurable advantage: with 30,000 examples and 19 possible sums, a supervised network learns the task just as well.

A difference already shows in the dynamics: the LTN is above 93% test accuracy after the very first epoch, while the baseline starts at 84%.

Adding two-digit numbers

Training is the same; only the grounding of the variables changes, to account for the four digits:

history_multi = {"train_loss": [], "train_sat": [], "train_acc": [],
                  "test_loss": [], "test_sat": [], "test_acc": []}

optimizer = torch.optim.Adam(Digit_m_d.parameters(), lr=0.001)

for epoch in range(20):
    # schedule of the parameter p for the existential quantifier, as described in the LTN paper
    if epoch in range(0, 4):
        p = 1
    if epoch in range(4, 8):
        p = 2
    if epoch in range(8, 12):
        p = 4
    if epoch in range(12, 20):
        p = 6
    train_loss, test_loss = 0.0, 0.0
    train_sat, test_sat = 0.0, 0.0
    train_acc, test_acc = 0.0, 0.0
    # training step
    cnn_m_d.train()
    for batch_idx, (operand_images, addition_label) in enumerate(multi_d_train_loader):
        optimizer.zero_grad()
        # ground the variables with the current batch data
        images_x1 = ltn.Variable("x1", operand_images[:, 0])
        images_x2 = ltn.Variable("x2", operand_images[:, 1])
        images_y1 = ltn.Variable("y1", operand_images[:, 2])
        images_y2 = ltn.Variable("y2", operand_images[:, 3])
        labels_z = ltn.Variable("z", addition_label)
        sat_agg = Forall(
            ltn.diag(images_x1, images_x2, images_y1, images_y2, labels_z),
            Exists(
                [d_1, d_2, d_3, d_4],
                And(
                    And(Digit_m_d(images_x1, d_1), Digit_m_d(images_x2, d_2)),
                    And(Digit_m_d(images_y1, d_3), Digit_m_d(images_y2, d_4))
                ),
                cond_vars=[d_1, d_2, d_3, d_4, labels_z],
                cond_fn=lambda d1, d2, d3, d4, z: torch.eq(10 * d1.value + d2.value + 10 * d3.value + d4.value, z.value),
                p=p
            )).value
        loss = 1. - sat_agg
        loss.backward()
        optimizer.step()
        train_loss += loss.item()
        train_sat += sat_agg.item()
        # compute the training accuracy
        predictions_x1 = torch.argmax(cnn_m_d(operand_images[:, 0].to(ltn.device)), dim=1)
        predictions_x2 = torch.argmax(cnn_m_d(operand_images[:, 1].to(ltn.device)), dim=1)
        predictions_y1 = torch.argmax(cnn_m_d(operand_images[:, 2].to(ltn.device)), dim=1)
        predictions_y2 = torch.argmax(cnn_m_d(operand_images[:, 3].to(ltn.device)), dim=1)
        predictions = 10 * predictions_x1 + predictions_x2 + 10 * predictions_y1 + predictions_y2
        train_acc += torch.count_nonzero(torch.eq(addition_label.to(ltn.device), predictions)) / predictions.shape[0]
    train_loss = train_loss / len(multi_d_train_loader)
    train_sat = train_sat / len(multi_d_train_loader)
    train_acc = train_acc / len(multi_d_train_loader)

    # test step
    cnn_m_d.eval()
    with torch.no_grad():
        for batch_idx, (operand_images, addition_label) in enumerate(multi_d_test_loader):
            # ground the variables with the current batch data
            images_x1 = ltn.Variable("x1", operand_images[:, 0])
            images_x2 = ltn.Variable("x2", operand_images[:, 1])
            images_y1 = ltn.Variable("y1", operand_images[:, 2])
            images_y2 = ltn.Variable("y2", operand_images[:, 3])
            labels_z = ltn.Variable("z", addition_label)
            sat_agg = Forall(
                ltn.diag(images_x1, images_x2, images_y1, images_y2, labels_z),
                Exists(
                    [d_1, d_2, d_3, d_4],
                    And(
                        And(Digit_m_d(images_x1, d_1), Digit_m_d(images_x2, d_2)),
                        And(Digit_m_d(images_y1, d_3), Digit_m_d(images_y2, d_4))
                    ),
                    cond_vars=[d_1, d_2, d_3, d_4, labels_z],
                    cond_fn=lambda d1, d2, d3, d4, z: torch.eq(10 * d1.value + d2.value + 10 * d3.value + d4.value, z.value),
                    p=p
                )).value
            loss = 1. - sat_agg
            test_loss += loss.item()
            test_sat += sat_agg.item()
            # compute the test accuracy
            predictions_x1 = torch.argmax(cnn_m_d(operand_images[:, 0].to(ltn.device)), dim=1)
            predictions_x2 = torch.argmax(cnn_m_d(operand_images[:, 1].to(ltn.device)), dim=1)
            predictions_y1 = torch.argmax(cnn_m_d(operand_images[:, 2].to(ltn.device)), dim=1)
            predictions_y2 = torch.argmax(cnn_m_d(operand_images[:, 3].to(ltn.device)), dim=1)
            predictions = 10 * predictions_x1 + predictions_x2 + 10 * predictions_y1 + predictions_y2
            test_acc += torch.count_nonzero(torch.eq(addition_label.to(ltn.device), predictions)) / predictions.shape[0]
    test_loss = test_loss / len(multi_d_test_loader)
    test_sat = test_sat / len(multi_d_test_loader)
    test_acc = test_acc / len(multi_d_test_loader)

    # record the history
    history_multi["train_loss"].append(float(train_loss))
    history_multi["train_sat"].append(float(train_sat))
    history_multi["train_acc"].append(float(train_acc))
    history_multi["test_loss"].append(float(test_loss))
    history_multi["test_sat"].append(float(test_sat))
    history_multi["test_acc"].append(float(test_acc))

    # print the metrics at each epoch
    print(" epoch %d | Train loss %.4f | Train Sat %.4f | Train Acc %.4f | Test loss %.4f | Test Sat %.4f |"
                " Test Acc %.4f " %
          (epoch, train_loss, train_sat, train_acc, test_loss, test_sat, test_acc))
 epoch 0 | Train loss 0.9897 | Train Sat 0.0103 | Train Acc 0.5497 | Test loss 0.9832 | Test Sat 0.0168 | Test Acc 0.8180
 epoch 1 | Train loss 0.9833 | Train Sat 0.0167 | Train Acc 0.8576 | Test loss 0.9822 | Test Sat 0.0178 | Test Acc 0.8699
 epoch 2 | Train loss 0.9824 | Train Sat 0.0176 | Train Acc 0.8962 | Test loss 0.9821 | Test Sat 0.0179 | Test Acc 0.8778
 epoch 3 | Train loss 0.9826 | Train Sat 0.0174 | Train Acc 0.8917 | Test loss 0.9826 | Test Sat 0.0174 | Test Acc 0.8382
 epoch 4 | Train loss 0.8840 | Train Sat 0.1160 | Train Acc 0.8997 | Test loss 0.8773 | Test Sat 0.1227 | Test Acc 0.9260
 epoch 5 | Train loss 0.8785 | Train Sat 0.1215 | Train Acc 0.9329 | Test loss 0.8775 | Test Sat 0.1225 | Test Acc 0.9252
 epoch 6 | Train loss 0.8764 | Train Sat 0.1236 | Train Acc 0.9444 | Test loss 0.8749 | Test Sat 0.1251 | Test Acc 0.9415
 epoch 7 | Train loss 0.8757 | Train Sat 0.1243 | Train Acc 0.9489 | Test loss 0.8756 | Test Sat 0.1244 | Test Acc 0.9375
 epoch 8 | Train loss 0.6739 | Train Sat 0.3261 | Train Acc 0.9356 | Test loss 0.6760 | Test Sat 0.3240 | Test Acc 0.9173
 epoch 9 | Train loss 0.6675 | Train Sat 0.3325 | Train Acc 0.9474 | Test loss 0.6676 | Test Sat 0.3324 | Test Acc 0.9363
 epoch 10 | Train loss 0.6613 | Train Sat 0.3387 | Train Acc 0.9593 | Test loss 0.6636 | Test Sat 0.3364 | Test Acc 0.9450
 epoch 11 | Train loss 0.6597 | Train Sat 0.3403 | Train Acc 0.9617 | Test loss 0.6631 | Test Sat 0.3369 | Test Acc 0.9466
 epoch 12 | Train loss 0.5290 | Train Sat 0.4710 | Train Acc 0.9610 | Test loss 0.5342 | Test Sat 0.4658 | Test Acc 0.9438
 epoch 13 | Train loss 0.5255 | Train Sat 0.4745 | Train Acc 0.9633 | Test loss 0.5264 | Test Sat 0.4736 | Test Acc 0.9561
 epoch 14 | Train loss 0.5240 | Train Sat 0.4760 | Train Acc 0.9651 | Test loss 0.5314 | Test Sat 0.4686 | Test Acc 0.9478
 epoch 15 | Train loss 0.5213 | Train Sat 0.4787 | Train Acc 0.9692 | Test loss 0.5298 | Test Sat 0.4702 | Test Acc 0.9494
 epoch 16 | Train loss 0.5197 | Train Sat 0.4803 | Train Acc 0.9705 | Test loss 0.5289 | Test Sat 0.4711 | Test Acc 0.9521
 epoch 17 | Train loss 0.5191 | Train Sat 0.4809 | Train Acc 0.9717 | Test loss 0.5324 | Test Sat 0.4676 | Test Acc 0.9470
 epoch 18 | Train loss 0.5195 | Train Sat 0.4805 | Train Acc 0.9706 | Test loss 0.5288 | Test Sat 0.4712 | Test Acc 0.9529
 epoch 19 | Train loss 0.5189 | Train Sat 0.4811 | Train Acc 0.9711 | Test loss 0.5336 | Test Sat 0.4664 | Test Acc 0.9454

The ∧\land connective is binary, so the conjunctions have to be nested to combine the four degrees of truth. Since the product is associative and commutative, the order of the grouping does not change the result.

The baseline follows the same principle, with 4 input images and 199 output classes:

history_baseline_multi = {"train_loss": [], "train_acc": [], "test_loss": [], "test_acc": []}

optimizer_baseline_m = torch.optim.Adam(baseline_m_d.parameters(), lr=0.001)
criterion = torch.nn.CrossEntropyLoss()

for epoch in range(30):
    train_loss, train_acc = 0.0, 0.0
    # training step
    baseline_m_d.train()
    for batch_idx, (operand_images, addition_label) in enumerate(multi_d_train_loader):
        images_x1 = operand_images[:, 0].to(ltn.device)
        images_x2 = operand_images[:, 1].to(ltn.device)
        images_y1 = operand_images[:, 2].to(ltn.device)
        images_y2 = operand_images[:, 3].to(ltn.device)
        labels = addition_label.to(ltn.device).long()

        optimizer_baseline_m.zero_grad()
        logits = baseline_m_d([images_x1, images_x2, images_y1, images_y2])
        loss = criterion(logits, labels)
        loss.backward()
        optimizer_baseline_m.step()

        train_loss += loss.item()
        predictions = torch.argmax(logits, dim=1)
        train_acc += torch.count_nonzero(torch.eq(labels, predictions)) / predictions.shape[0]
    train_loss = train_loss / len(multi_d_train_loader)
    train_acc = train_acc / len(multi_d_train_loader)

    # test step
    test_loss, test_acc = 0.0, 0.0
    baseline_m_d.eval()
    with torch.no_grad():
        for batch_idx, (operand_images, addition_label) in enumerate(multi_d_test_loader):
            images_x1 = operand_images[:, 0].to(ltn.device)
            images_x2 = operand_images[:, 1].to(ltn.device)
            images_y1 = operand_images[:, 2].to(ltn.device)
            images_y2 = operand_images[:, 3].to(ltn.device)
            labels = addition_label.to(ltn.device).long()

            logits = baseline_m_d([images_x1, images_x2, images_y1, images_y2])
            loss = criterion(logits, labels)

            test_loss += loss.item()
            predictions = torch.argmax(logits, dim=1)
            test_acc += torch.count_nonzero(torch.eq(labels, predictions)) / predictions.shape[0]
    test_loss = test_loss / len(multi_d_test_loader)
    test_acc = test_acc / len(multi_d_test_loader)

    history_baseline_multi["train_loss"].append(float(train_loss))
    history_baseline_multi["train_acc"].append(float(train_acc))
    history_baseline_multi["test_loss"].append(float(test_loss))
    history_baseline_multi["test_acc"].append(float(test_acc))

    print(" epoch %d | [Baseline] Train loss %.4f | Train Acc %.4f | Test loss %.4f | Test Acc %.4f " %
          (epoch, train_loss, train_acc, test_loss, test_acc))
 epoch 0 | [Baseline] Train loss 4.9616 | Train Acc 0.0147 | Test loss 4.6097 | Test Acc 0.0206
 epoch 1 | [Baseline] Train loss 4.3830 | Train Acc 0.0333 | Test loss 4.3773 | Test Acc 0.0320
 epoch 2 | [Baseline] Train loss 4.0406 | Train Acc 0.0649 | Test loss 4.2692 | Test Acc 0.0388
 epoch 3 | [Baseline] Train loss 3.6883 | Train Acc 0.1100 | Test loss 4.1066 | Test Acc 0.0676
 epoch 4 | [Baseline] Train loss 3.1502 | Train Acc 0.1911 | Test loss 3.7083 | Test Acc 0.1040
 epoch 5 | [Baseline] Train loss 2.5091 | Train Acc 0.3078 | Test loss 3.3443 | Test Acc 0.1487
 epoch 6 | [Baseline] Train loss 1.9065 | Train Acc 0.4393 | Test loss 3.0455 | Test Acc 0.2065
 epoch 7 | [Baseline] Train loss 1.4043 | Train Acc 0.5754 | Test loss 2.8107 | Test Acc 0.2627
 epoch 8 | [Baseline] Train loss 1.0039 | Train Acc 0.6915 | Test loss 2.7919 | Test Acc 0.3074
 epoch 9 | [Baseline] Train loss 0.7097 | Train Acc 0.7794 | Test loss 2.7845 | Test Acc 0.3212
 epoch 10 | [Baseline] Train loss 0.4998 | Train Acc 0.8419 | Test loss 2.9036 | Test Acc 0.3402
 epoch 11 | [Baseline] Train loss 0.3643 | Train Acc 0.8850 | Test loss 3.0704 | Test Acc 0.3552
 epoch 12 | [Baseline] Train loss 0.2681 | Train Acc 0.9129 | Test loss 3.1775 | Test Acc 0.3675
 epoch 13 | [Baseline] Train loss 0.2111 | Train Acc 0.9322 | Test loss 3.2700 | Test Acc 0.3908
 epoch 14 | [Baseline] Train loss 0.2028 | Train Acc 0.9321 | Test loss 3.3315 | Test Acc 0.3952
 epoch 15 | [Baseline] Train loss 0.1517 | Train Acc 0.9515 | Test loss 3.6011 | Test Acc 0.3821
 epoch 16 | [Baseline] Train loss 0.1467 | Train Acc 0.9514 | Test loss 3.5471 | Test Acc 0.3916
 epoch 17 | [Baseline] Train loss 0.1561 | Train Acc 0.9474 | Test loss 3.7553 | Test Acc 0.3833
 epoch 18 | [Baseline] Train loss 0.1457 | Train Acc 0.9504 | Test loss 3.6462 | Test Acc 0.4110
 epoch 19 | [Baseline] Train loss 0.1029 | Train Acc 0.9671 | Test loss 3.6674 | Test Acc 0.4229
 epoch 20 | [Baseline] Train loss 0.1139 | Train Acc 0.9619 | Test loss 3.9440 | Test Acc 0.4102
 epoch 21 | [Baseline] Train loss 0.1303 | Train Acc 0.9573 | Test loss 3.7222 | Test Acc 0.4260
 epoch 22 | [Baseline] Train loss 0.0985 | Train Acc 0.9679 | Test loss 3.9490 | Test Acc 0.4118
 epoch 23 | [Baseline] Train loss 0.1044 | Train Acc 0.9641 | Test loss 3.9776 | Test Acc 0.4233
 epoch 24 | [Baseline] Train loss 0.0848 | Train Acc 0.9718 | Test loss 3.9414 | Test Acc 0.4213
 epoch 25 | [Baseline] Train loss 0.0752 | Train Acc 0.9753 | Test loss 3.9455 | Test Acc 0.4355
 epoch 26 | [Baseline] Train loss 0.1106 | Train Acc 0.9649 | Test loss 3.9096 | Test Acc 0.4407
 epoch 27 | [Baseline] Train loss 0.1121 | Train Acc 0.9621 | Test loss 4.0718 | Test Acc 0.4557
 epoch 28 | [Baseline] Train loss 0.0972 | Train Acc 0.9681 | Test loss 3.8898 | Test Acc 0.4415
 epoch 29 | [Baseline] Train loss 0.0640 | Train Acc 0.9788 | Test loss 4.0434 | Test Acc 0.4490

This time, the gap is clear: the LTN reaches about 95% on the test set, while the baseline plateaus around 45%, even though its training accuracy is above 97%.

Results

Test accuracy at the last epoch:

LTNBaseline
Adding two digits97%97%
Adding two-digit numbers95%45%

To visualize training, we plot for each task the accuracy of both approaches and the satisfaction of the LTN, with the changes of pp. This is a single training run, not averaged over several random seeds for lack of compute: the curves therefore have no uncertainty band.

import matplotlib.pyplot as plt

def plot_task_with_baseline(history_ltn, history_baseline, title,
                             p_transitions=(4, 8, 12), p_labels=(1, 2, 4, 6),
                             max_epochs=20):
    epochs = range(max_epochs)
    fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))

    colors = {
        "baseline_train": "#8fc7e8", "baseline_test": "#1f6fb2",
        "ltn_train": "#a8d98a", "ltn_test": "#2e8b3d",
    }

    # --- Accuracy ---
    ax = axes[0]
    ax.plot(epochs, history_baseline["train_acc"][:max_epochs],
             color=colors["baseline_train"], lw=1, label="Baseline (train)")
    ax.plot(epochs, history_baseline["test_acc"][:max_epochs],
             color=colors["baseline_test"], lw=1.5, ls="--", label="Baseline (test)")
    ax.plot(epochs, history_ltn["train_acc"][:max_epochs],
             color=colors["ltn_train"], lw=1, label="LTN (train)")
    ax.plot(epochs, history_ltn["test_acc"][:max_epochs],
             color=colors["ltn_test"], lw=1.5, ls="--", label="LTN (test)")
    ax.set_ylim(0, 1.02); ax.set_xlim(0, max_epochs - 1)
    ax.set_xlabel("Epoch"); ax.set_ylabel("Accuracy")
    ax.set_title(f"{title} examples")

    # --- Sat ---
    ax = axes[1]
    ax.plot(epochs, history_ltn["train_sat"][:max_epochs],
             color=colors["ltn_train"], lw=1, label="LTN (train)")
    ax.plot(epochs, history_ltn["test_sat"][:max_epochs],
             color=colors["ltn_test"], lw=1.5, ls="--", label="LTN (test)")
    for t in p_transitions:
        ax.axvline(t, color="grey", ls=":", lw=0.8)
    boundaries = [0] + list(p_transitions) + [max_epochs - 1]
    for i, p in enumerate(p_labels):
        mid = (boundaries[i] + boundaries[i + 1]) / 2
        ax.text(mid, 0.95, f"$p_2={p}$", ha="center", fontsize=8, color="grey")
    ax.set_ylim(0, 1.02); ax.set_xlim(0, max_epochs - 1)
    ax.set_xlabel("Epoch"); ax.set_ylabel("Sat")
    ax.set_title(f"{title} examples")

    # remove the top/right spines on both axes
    for a in axes:
        a.spines["top"].set_visible(False)
        a.spines["right"].set_visible(False)

    handles, labels_ = axes[0].get_legend_handles_labels()
    fig.legend(handles, labels_, loc="center right", bbox_to_anchor=(1.18, 0.5))
    plt.tight_layout()
    plt.show()

plot_task_with_baseline(history_single, history_baseline_single, "30000")
plot_task_with_baseline(history_multi, history_baseline_multi, "15000")

Accuracy and satisfaction for adding two digits

Adding two digits (30,000 examples). Left, the accuracy of the LTN and the baseline; right, the satisfaction of the LTN, with the steps of pp.

Accuracy and satisfaction for adding two-digit numbers

Adding two-digit numbers (15,000 examples). The baseline’s training curve keeps rising, but its test curve plateaus much lower; the LTN’s two curves stay close to each other.

The curves only show the first 20 epochs of the baseline, to line them up with the LTN’s: its final test accuracy, at epoch 30, is 45%.

On adding two digits, the test curves of both approaches almost overlap from start to finish.

On adding two-digit numbers, the baseline overfits: its training accuracy keeps rising, but its test accuracy plateaus. It memorizes the examples without generalizing. The LTN keeps its training and test curves close to each other.

The reason lies in the structure of the problem. The baseline has to predict the sum directly among 199 classes, with 15,000 examples, that is, few examples per class, and extreme sums are very rare. The LTN only learns to tell 10 digits apart, and each training image gives it a signal about one of them; the logical rule then takes care of combining the digits into the sum. It never needs to learn the 199 possible sums.

Satisfaction rises in steps that coincide exactly with the changes of pp. These jumps need careful reading: for fixed degrees of truth, pMean mechanically increases with pp, since it gets closer to the maximum. A large part of each jump therefore comes from the change of the measure itself, not from a sudden improvement of the classifier, as confirmed by the accuracy, which does not jump at the same moments. Satisfaction levels are only comparable at equal pp.

Key takeaways

  • An LTN can learn a classifier without any individual label, from a single logical constraint on the sum: the digits are latent variables, linked to the data by the logic alone.
  • Diagonal quantification pairs each couple of images with its own result, and guarded quantification injects the rule d1+d2=nd_1 + d_2 = n.
  • On adding two digits, LTN and baseline are on par (97%). On adding two-digit numbers, the LTN stays at 95% while the baseline drops to 45%: by breaking the problem down into digits, the logic avoids having to learn 199 classes from few examples.
  • The schedule of pp applies the lesson of part 3: a low pp at first to spread the gradient, then a higher one once the classifier becomes reliable.

This case study concludes the series. It illustrates LTN’s central idea: logic fixes the form of the reasoning, here “the sum of the digits equals the result”, and the content, recognizing each digit, is learned by gradient descent.

The full notebook for this case study is available in my repository.

Weekly Notes

Every Sunday, I share what I've been learning: papers, ideas, experiments, and questions that stayed with me.

You can unsubscribe at any time with a single click.

0 Likes • 0 Comments

Discussion about this post0

Join the discussion

A secure sign-in link will be sent to your email address.

Loading discussion...