Molecules as primitive structures

How molprim rewrites a molecular graph before a graph kernel compares it

A molecule can be read as a graph: atoms are vertices labelled by element, and bonds are edges. Graph kernels compare molecules by counting shared substructures, but on the raw atom graph those substructures are mostly single atoms and their immediate neighbours.

The JCSSE 2020 paper first rewrites each molecule as a graph of its primitive structures: rings (cycles of 3–10 atoms), branching points (stars of 3–7 atoms) and single bonds. Each occurrence of a primitive becomes one vertex, and two vertices are joined when their occurrences share an atom. A graph kernel then compares these smaller, more chemical graphs.

Everything on this page comes from molprim 0.1.0, the implementation of that paper (pip install molprim), run on MUTAG: 188 nitroaromatic compounds labelled by whether they are mutagenic.

Extract the primitive graphs with molprim
import json, urllib.request
import networkx as nx
import numpy as np
import molprim
from molprim import PrimitiveStructureExtractor, fetch_tudataset
from molprim._algorithms import find_occurrences

ELEMENTS = {0: "C", 1: "N", 2: "O", 3: "F", 4: "I", 5: "Cl", 6: "Br"}   # MUTAG node labels

graphs, y = fetch_tudataset("MUTAG")
extractor = PrimitiveStructureExtractor().fit(graphs)                    # learn primitives from all 188 molecules
primitive_graphs = extractor.transform(graphs)
keys = [extractor.candidate_keys_[i] for i in extractor.selected_]
labels = extractor.primitive_labels_

def atoms_of_primitives(graph):
    """Atoms covered by each vertex of the primitive graph.

    Occurrences come from molprim's own `find_occurrences`; they are taken in the same order and
    with the same overlap rule as molprim's `assemble_primitive_graph`, and the result is checked
    against molprim's output below.
    """
    taken, owners = [], {}
    for occurrence_list, label in zip(find_occurrences(graph, keys, "label"), labels):
        for occurrence in map(frozenset, occurrence_list):
            overlapping = {j for v in occurrence for j in owners.get(v, [])}
            if any(len(occurrence & taken[j][0]) > len(occurrence) / 2 for j in overlapping):
                continue
            for v in occurrence:
                owners.setdefault(v, []).append(len(taken))
            taken.append((occurrence, label))
    return taken

def describe(index):
    g, p = graphs[index], primitive_graphs[index]
    taken = atoms_of_primitives(g)
    # the reconstruction must agree with molprim exactly: same vertices, labels and edges
    assert [lab for _, lab in taken] == [p.nodes[i]["label"] for i in p.nodes]
    shared = {tuple(sorted((i, j))) for i in range(len(taken)) for j in range(i + 1, len(taken))
              if taken[i][0] & taken[j][0]}
    assert shared == {tuple(sorted(e)) for e in p.edges}
    pos = nx.kamada_kawai_layout(g)
    return {
        "index": int(index),
        "mutagenic": bool(y[index] == 1),
        "atoms": [[round(float(pos[n][0]), 3), round(float(pos[n][1]), 3), ELEMENTS[g.nodes[n]["label"]]] for n in g],
        "bonds": [[int(u), int(v)] for u, v in g.edges],
        "primitives": [{"label": lab, "kind": lab[0],
                        "x": round(float(np.mean([pos[a][0] for a in atoms])), 3),
                        "y": round(float(np.mean([pos[a][1] for a in atoms])), 3),
                        "atoms": sorted(int(a) for a in atoms)} for atoms, lab in taken],
        "links": [[int(u), int(v)] for u, v in p.edges],
    }

# a few varied examples from each class: the molecules with the most distinct primitive types
def variety(i):
    return len(set(nx.get_node_attributes(primitive_graphs[i], "label").values()))
examples = []
for label in (1, -1):
    pool = sorted((i for i in range(len(graphs)) if y[i] == label and len(graphs[i]) <= 26),
                  key=lambda i: (-variety(i), i))
    examples += pool[:4]

# benchmark results published in the molprim repository (10 random 80/20 splits, 5-fold grid search)
RESULTS = "https://raw.githubusercontent.com/PeemapatW/primitive-structure-extraction/main/benchmarks/results/{}_{}_{}.json"
bench = []
for dataset in ("MUTAG", "BZR", "COX2", "NCI1"):
    for method in ("wl_subtree", "wl_sp", "sp", "edges"):
        for primitive in (1, 0):
            rows = json.load(urllib.request.urlopen(RESULTS.format(dataset, method, primitive)))["rows"]
            acc = [r["acc"] for r in rows]
            bench.append({"dataset": dataset, "method": method, "graphs": "primitive" if primitive else "original",
                          "mean": float(np.mean(acc)), "sd": float(np.std(acc)), "n": len(acc)})

ojs_define(mp_examples=[describe(i) for i in examples], mp_bench=bench,
           mp_meta={"n_primitives": len(labels), "version": molprim.__version__,
                    "atoms": float(np.mean([len(g) for g in graphs])),
                    "prims": float(np.mean([len(p) for p in primitive_graphs]))})

From molecule to primitive graph

Left: the atom graph, coloured by element. Right: the primitive graph molprim builds from it, with each primitive drawn at the centre of its atoms: rings in orange, stars in blue, bonds in grey. Edges join primitives that share an atom. Click a primitive (or choose it above) to see which atoms it covers; click it again to clear.

What it helps with

  • Comparing molecules by chemically meaningful parts. A graph kernel on the atom graph mostly sees single atoms and their neighbours. On the primitive graph each vertex is a whole ring, branching point or bond, so the kernel compares larger units, and rings that are fused together appear as neighbouring vertices.
  • Working with any graph kernel. The rewriting is a separate step in front of the kernel. In molprim it is a scikit-learn transformer, so it can be placed in front of the Weisfeiler–Lehman or shortest-path kernels, or a simple edge count, without changing them.
  • Better accuracy in several settings. In the benchmark below, primitive graphs improve every graph kernel on MUTAG, and the shortest-path kernel and the edge-pair counts on NCI1 and BZR.
  • Readable features. Every primitive has a chemical meaning, such as C6-1, a six-carbon ring, or S4-3, a nitrogen bonded to one carbon and two oxygens (the nitro group present in every MUTAG compound). Results can therefore be discussed in terms of the structures involved.

Does it help?

The repository’s benchmark scripts compare the same graph kernels on the primitive graphs and on the original atom graphs. Each number below is the mean accuracy over 10 random 80/20 splits, with the kernel and SVM hyperparameters tuned by 5-fold grid search on the training part of each split. The values are read directly from the published result files.

Accuracy (%) on primitive graphs, then on the original graphs, then the difference in points.

Reading the table. The gain from primitive graphs depends on the dataset and the kernel:

  • On MUTAG the primitive graphs are better with every graph kernel, by about 4 to 8 points.
  • With the shortest-path kernel and the edge-pair counts, primitive graphs are never worse by more than half a point, and they gain clearly on NCI1 (+7.8 and +6.4 points), on MUTAG with the shortest-path kernel (+7.9) and on BZR with the edge pairs (+5.3).
  • With the Weisfeiler–Lehman kernels on BZR, COX2 and NCI1, the two representations are within about two points of each other, sometimes slightly in favour of the original graphs.
  • On COX2 every kernel gives nearly the same accuracy either way.

The splits are small for MUTAG, BZR and COX2, and the standard deviations across splits (3–6 points) are of the same order as several of these differences, so only the larger gaps should be read as real effects.