molprim 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
molprim/__init__.py ADDED
@@ -0,0 +1,24 @@
1
+ """scikit-learn interface for primitive structure extraction and graph kernels.
2
+
3
+ Wongsriphisant et al., "A classification of biochemical compounds based on
4
+ their primitive structures and graph kernels", 2020.
5
+ """
6
+ from importlib.metadata import PackageNotFoundError, version
7
+
8
+ from .datasets import fetch_tudataset, read_tudataset
9
+ from .extractor import PrimitiveStructureExtractor
10
+ from .features import LabelPairEdgeCounter
11
+ from .kernels import GraphKernelTransformer
12
+
13
+ try:
14
+ __version__ = version("molprim")
15
+ except PackageNotFoundError: # running from a source checkout without installing
16
+ __version__ = "unknown"
17
+
18
+ __all__ = [
19
+ "GraphKernelTransformer",
20
+ "LabelPairEdgeCounter",
21
+ "PrimitiveStructureExtractor",
22
+ "fetch_tudataset",
23
+ "read_tudataset",
24
+ ]
molprim/_algorithms.py ADDED
@@ -0,0 +1,198 @@
1
+ """Primitive structure selection (Algorithm 1) and extraction (Algorithm 2).
2
+
3
+ Candidates are labelled cycles, stars and single edges, so instead of generic
4
+ subgraph isomorphism (VF2) each family is enumerated directly and identified
5
+ by a canonical key built from its labels:
6
+
7
+ - cycle: ``("C", labels)``, the smallest rotation or reflection of the labels
8
+ around the cycle
9
+ - star: ``("S", centre_label, sorted_leaf_labels)``
10
+ - edge: ``("P2", sorted_end_labels)``
11
+
12
+ Two labelled cycles (stars, edges) are isomorphic exactly when their keys are
13
+ equal, so this gives the same result as the VF2 code in the original
14
+ notebooks.
15
+ """
16
+ from collections import defaultdict
17
+ from itertools import combinations
18
+ from math import factorial
19
+
20
+ import networkx as nx
21
+ from joblib import Parallel, delayed
22
+
23
+ EDGE_PREFIX = "P2"
24
+
25
+
26
+ def _order(labels):
27
+ """Sort key that works for any hashable labels, including mixed types."""
28
+ return tuple(map(repr, labels))
29
+
30
+
31
+ def _sorted(labels):
32
+ return tuple(sorted(labels, key=lambda label: repr(label)))
33
+
34
+
35
+ def _cycle_key(labels):
36
+ labels = list(labels)
37
+ n = len(labels)
38
+ reflected = labels[::-1]
39
+ variants = [tuple(seq[i:] + seq[:i]) for seq in (labels, reflected) for i in range(n)]
40
+ return ("C", min(variants, key=_order))
41
+
42
+
43
+ def check_sizes(cycle_sizes, star_sizes):
44
+ for kind, sizes in (("cycle", cycle_sizes), ("star", star_sizes)):
45
+ for n in sizes:
46
+ if n < 3:
47
+ raise ValueError(f"{kind} sizes must be >= 3, got {n}")
48
+
49
+
50
+ def iter_structures(graph, cycle_sizes, star_sizes, include_edge, node_label, induced):
51
+ """Yield ``(prefix, key, nodes)`` for every cycle, star and edge of ``graph``.
52
+
53
+ With ``induced=True`` only induced structures are produced (chordless
54
+ cycles, stars whose leaves are pairwise non-adjacent), as in candidate
55
+ selection. With ``induced=False`` every occurrence is produced, as in
56
+ extraction.
57
+ """
58
+ label = {n: data[node_label] for n, data in graph.nodes(data=True)}
59
+ if cycle_sizes:
60
+ find_cycles = nx.chordless_cycles if induced else nx.simple_cycles
61
+ for cycle in find_cycles(graph, max(cycle_sizes)):
62
+ if len(cycle) in cycle_sizes:
63
+ yield f"C{len(cycle)}", _cycle_key(label[n] for n in cycle), cycle
64
+ for centre in graph:
65
+ neighbours = [n for n in graph[centre] if n != centre]
66
+ for size in star_sizes:
67
+ for leaves in combinations(neighbours, size - 1):
68
+ if induced and any(graph.has_edge(a, b) for a, b in combinations(leaves, 2)):
69
+ continue
70
+ yield f"S{size}", ("S", label[centre], _sorted(label[n] for n in leaves)), (centre, *leaves)
71
+ if include_edge:
72
+ for u, v in graph.edges:
73
+ if u != v:
74
+ yield EDGE_PREFIX, (EDGE_PREFIX, _sorted((label[u], label[v]))), (u, v)
75
+
76
+
77
+ def _graph_counts(graph, cycle_sizes, star_sizes, include_edge, node_label):
78
+ counts = defaultdict(int)
79
+ prefixes = {}
80
+ for prefix, key, _ in iter_structures(graph, cycle_sizes, star_sizes, include_edge, node_label, induced=True):
81
+ counts[key] += 1
82
+ prefixes[key] = prefix
83
+ return counts, prefixes
84
+
85
+
86
+ def structure_graph(key, node_label):
87
+ """The labelled cycle, star or edge described by ``key``."""
88
+ if key[0] == "C":
89
+ graph = nx.cycle_graph(len(key[1]))
90
+ labels = key[1]
91
+ elif key[0] == "S":
92
+ graph = nx.star_graph(len(key[2]))
93
+ labels = (key[1], *key[2])
94
+ else:
95
+ graph = nx.path_graph(2)
96
+ labels = key[1]
97
+ nx.set_node_attributes(graph, dict(enumerate(labels)), node_label)
98
+ return graph
99
+
100
+
101
+ def enumerate_candidates(graphs, cycle_sizes, star_sizes, include_edge, node_label, n_jobs=None):
102
+ """Collect every distinct labelled cycle, star and edge that occurs (induced) in ``graphs``.
103
+
104
+ Returns ``(keys, labels, counts)`` in extraction priority order: cycles
105
+ (largest first), stars (largest first), then edges; within a shape the most
106
+ frequent comes first and is labelled ``-1`` (e.g. ``"C6-1"``).
107
+
108
+ Counts follow the original notebooks, which counted VF2 mappings and
109
+ divided by the cycle length, the number of star leaves, or 2 for an edge:
110
+ 2 per cycle occurrence, (leaves - 1)! per star occurrence, 1 per edge.
111
+ """
112
+ check_sizes(cycle_sizes, star_sizes)
113
+ cycle_sizes, star_sizes = set(cycle_sizes), set(star_sizes)
114
+ per_graph = Parallel(n_jobs=n_jobs)(
115
+ delayed(_graph_counts)(g, cycle_sizes, star_sizes, include_edge, node_label) for g in graphs
116
+ )
117
+ occurrences = defaultdict(int)
118
+ prefix_of = {}
119
+ for counts, prefixes in per_graph:
120
+ for key, n in counts.items():
121
+ occurrences[key] += n
122
+ prefix_of.update(prefixes)
123
+
124
+ def weight(key):
125
+ if key[0] == "C":
126
+ return 2
127
+ if key[0] == "S":
128
+ return factorial(len(key[2]) - 1)
129
+ return 1
130
+
131
+ def priority(prefix):
132
+ kind, size = prefix[0], int(prefix[1:])
133
+ return ("CSP".index(kind), -size)
134
+
135
+ by_prefix = defaultdict(list)
136
+ for key, prefix in prefix_of.items():
137
+ by_prefix[prefix].append(key)
138
+
139
+ keys, labels, counts = [], [], []
140
+ for prefix in sorted(by_prefix, key=priority):
141
+ ranked = sorted(by_prefix[prefix], key=lambda k: (-occurrences[k], _order(k)))
142
+ for i, key in enumerate(ranked, start=1):
143
+ keys.append(key)
144
+ labels.append(f"{prefix}-{i}")
145
+ counts.append(occurrences[key] * weight(key))
146
+ return keys, labels, counts
147
+
148
+
149
+ def find_occurrences(graph, keys, node_label):
150
+ """Vertex sets (sorted tuples of node positions) where each primitive in ``keys`` occurs in ``graph``.
151
+
152
+ Occurrences are label-preserving, not necessarily induced, subgraphs, as
153
+ with ``graph_tool.subgraph_isomorphism`` in the original code.
154
+ """
155
+ index = {key: i for i, key in enumerate(keys)}
156
+ cycle_sizes = {len(k[1]) for k in keys if k[0] == "C"}
157
+ star_sizes = {len(k[2]) + 1 for k in keys if k[0] == "S"}
158
+ include_edge = any(k[0] == EDGE_PREFIX for k in keys)
159
+ position = {n: i for i, n in enumerate(graph)}
160
+
161
+ found = [set() for _ in keys]
162
+ for _, key, nodes in iter_structures(graph, cycle_sizes, star_sizes, include_edge, node_label, induced=False):
163
+ i = index.get(key)
164
+ if i is not None:
165
+ found[i].add(tuple(sorted(position[n] for n in nodes)))
166
+ return [sorted(occurrences) for occurrences in found]
167
+
168
+
169
+ def match_graphs(graphs, keys, node_label, n_jobs=None):
170
+ """:func:`find_occurrences` for every graph."""
171
+ return Parallel(n_jobs=n_jobs)(delayed(find_occurrences)(g, keys, node_label) for g in graphs)
172
+
173
+
174
+ def assemble_primitive_graph(occurrences, labels, node_label):
175
+ """Build the primitive graph from per-primitive occurrences (Algorithm 2).
176
+
177
+ Occurrences are taken greedily in priority order and dropped when they
178
+ share more than half of their vertices with an occurrence already taken.
179
+ Each kept occurrence becomes a vertex labelled with its primitive, and two
180
+ vertices are joined when their occurrences share an original vertex.
181
+ """
182
+ taken = []
183
+ owners = defaultdict(list) # original vertex -> indices into taken
184
+ for occurrence_list, label in zip(occurrences, labels):
185
+ for occurrence in occurrence_list:
186
+ occurrence = frozenset(occurrence)
187
+ overlapping = {j for v in occurrence for j in owners[v]}
188
+ if any(len(occurrence & taken[j][0]) > len(occurrence) / 2 for j in overlapping):
189
+ continue
190
+ for v in occurrence:
191
+ owners[v].append(len(taken))
192
+ taken.append((occurrence, label))
193
+
194
+ primitive_graph = nx.Graph()
195
+ primitive_graph.add_nodes_from((i, {node_label: label}) for i, (_, label) in enumerate(taken))
196
+ for indices in owners.values():
197
+ primitive_graph.add_edges_from(combinations(indices, 2))
198
+ return primitive_graph
molprim/_utils.py ADDED
@@ -0,0 +1,16 @@
1
+ import networkx as nx
2
+
3
+
4
+ def check_graphs(X, node_label):
5
+ """Validate ``X`` as a non-empty sequence of undirected graphs whose nodes all carry ``node_label``."""
6
+ graphs = list(X)
7
+ if not graphs:
8
+ raise ValueError("X must contain at least one graph")
9
+ for i, g in enumerate(graphs):
10
+ if not isinstance(g, nx.Graph) or g.is_directed():
11
+ raise TypeError(f"X[{i}] is {type(g).__name__}; expected an undirected networkx.Graph")
12
+ for n, data in g.nodes(data=True):
13
+ if node_label not in data:
14
+ raise ValueError(f"node {n!r} of X[{i}] has no {node_label!r} attribute")
15
+ return graphs
16
+
molprim/datasets.py ADDED
@@ -0,0 +1,59 @@
1
+ """Loader for the TU Dortmund graph benchmark datasets (MUTAG, NCI1, BZR, COX2, ...)."""
2
+ import io
3
+ import zipfile
4
+ from pathlib import Path
5
+ from urllib.request import urlopen
6
+
7
+ import networkx as nx
8
+ import numpy as np
9
+ from sklearn.datasets import get_data_home
10
+
11
+ TU_URL = "https://www.chrsmrrs.com/graphkerneldatasets/{name}.zip"
12
+
13
+
14
+ def fetch_tudataset(name, data_home=None):
15
+ """Download (once) and load a TU dataset as ``(graphs, y)``.
16
+
17
+ Files are cached under ``<data_home>/tudataset/<name>``; ``data_home``
18
+ defaults to scikit-learn's data directory.
19
+ """
20
+ root = Path(get_data_home(data_home)) / "tudataset"
21
+ folder = root / name
22
+ if not folder.exists():
23
+ with urlopen(TU_URL.format(name=name)) as response:
24
+ zipfile.ZipFile(io.BytesIO(response.read())).extractall(root)
25
+ return read_tudataset(folder)
26
+
27
+
28
+ def read_tudataset(folder):
29
+ """Load a TU dataset from an extracted folder as ``(graphs, y)``.
30
+
31
+ Each graph is an undirected ``networkx.Graph`` with nodes ``0..n-1``. Node
32
+ and edge labels, when the dataset has them, are stored in the ``"label"``
33
+ attribute.
34
+ """
35
+ folder = Path(folder)
36
+ name = folder.name
37
+
38
+ def read(suffix, optional=False):
39
+ path = folder / f"{name}_{suffix}.txt"
40
+ if optional and not path.exists():
41
+ return None
42
+ return np.loadtxt(path, delimiter=",", dtype=int, ndmin=2)
43
+
44
+ graph_of_node = read("graph_indicator")[:, 0] - 1
45
+ edges = read("A") - 1
46
+ y = read("graph_labels")[:, 0]
47
+ node_labels = read("node_labels", optional=True)
48
+ edge_labels = read("edge_labels", optional=True)
49
+
50
+ graphs = [nx.Graph() for _ in range(len(y))]
51
+ first_node = np.searchsorted(graph_of_node, np.arange(len(y)))
52
+ for node, g in enumerate(graph_of_node):
53
+ attrs = {} if node_labels is None else {"label": int(node_labels[node, 0])}
54
+ graphs[g].add_node(node - first_node[g], **attrs)
55
+ for k, (u, v) in enumerate(edges):
56
+ g = graph_of_node[u]
57
+ attrs = {} if edge_labels is None else {"label": int(edge_labels[k, 0])}
58
+ graphs[g].add_edge(u - first_node[g], v - first_node[g], **attrs)
59
+ return graphs, y
molprim/extractor.py ADDED
@@ -0,0 +1,104 @@
1
+ import numpy as np
2
+ from sklearn.base import BaseEstimator, TransformerMixin
3
+ from sklearn.utils.validation import check_is_fitted
4
+
5
+ from ._algorithms import EDGE_PREFIX, assemble_primitive_graph, enumerate_candidates, match_graphs, structure_graph
6
+ from ._utils import check_graphs
7
+
8
+
9
+ class PrimitiveStructureExtractor(TransformerMixin, BaseEstimator):
10
+ """Rewrite molecular graphs as graphs of their primitive structures.
11
+
12
+ ``fit`` collects the labelled cycles, stars and edges that occur in the
13
+ training graphs and keeps the frequent ones (Algorithm 1 of the paper).
14
+ ``transform`` turns each graph into a primitive graph: one vertex per
15
+ occurrence of a primitive, labelled with that primitive (e.g. ``"C6-1"``),
16
+ and an edge wherever two occurrences share an atom (Algorithm 2).
17
+
18
+ The output is a list of ``networkx.Graph``, meant to be followed by a graph
19
+ kernel (:class:`GraphKernelTransformer`) or :class:`LabelPairEdgeCounter`.
20
+
21
+ Parameters
22
+ ----------
23
+ threshold : float, default=0
24
+ Percentile (0-100) of candidate counts below which a cycle or star is
25
+ discarded. Edges (P2) are always kept. 0 keeps every candidate.
26
+ cycle_sizes : tuple of int, default=(3, ..., 10)
27
+ Numbers of vertices of the candidate cycles.
28
+ star_sizes : tuple of int, default=(3, ..., 7)
29
+ Numbers of vertices (centre included) of the candidate stars.
30
+ include_edge : bool, default=True
31
+ Whether the single edge P2 is a candidate.
32
+ node_label : str, default="label"
33
+ Node attribute holding the atom type. The primitive graphs use the
34
+ same attribute name for their labels.
35
+ n_jobs : int, default=None
36
+ Number of parallel jobs over graphs (joblib semantics).
37
+
38
+ Attributes
39
+ ----------
40
+ candidates_ : list of networkx.Graph
41
+ Every labelled candidate found in the training graphs, in extraction
42
+ priority order: cycles, then stars (largest first), then edges.
43
+ candidate_labels_ : list of str
44
+ Names such as ``"C6-1"``: shape, then rank by count within the shape.
45
+ candidate_counts_ : ndarray of int
46
+ Frequency of each candidate in the training graphs, normalised as in
47
+ the original notebooks.
48
+ candidate_keys_ : list of tuple
49
+ Canonical form of each candidate (see ``_algorithms``).
50
+ selected_ : ndarray of int
51
+ Indices into the candidate lists of the primitives kept by ``threshold``.
52
+ primitive_labels_ : list of str
53
+ Labels of the selected primitives.
54
+ """
55
+
56
+ def __init__(
57
+ self,
58
+ threshold=0,
59
+ cycle_sizes=(3, 4, 5, 6, 7, 8, 9, 10),
60
+ star_sizes=(3, 4, 5, 6, 7),
61
+ include_edge=True,
62
+ node_label="label",
63
+ n_jobs=None,
64
+ ):
65
+ self.threshold = threshold
66
+ self.cycle_sizes = cycle_sizes
67
+ self.star_sizes = star_sizes
68
+ self.include_edge = include_edge
69
+ self.node_label = node_label
70
+ self.n_jobs = n_jobs
71
+
72
+ def fit(self, X, y=None):
73
+ if not 0 <= self.threshold <= 100:
74
+ raise ValueError(f"threshold must be in [0, 100], got {self.threshold}")
75
+ graphs = check_graphs(X, self.node_label)
76
+ keys, labels, counts = enumerate_candidates(
77
+ graphs, self.cycle_sizes, self.star_sizes, self.include_edge, self.node_label, self.n_jobs
78
+ )
79
+ if not keys:
80
+ raise ValueError("no candidate primitive structure occurs in the training graphs")
81
+
82
+ counts = np.asarray(counts)
83
+ cut = np.percentile(counts, self.threshold)
84
+ is_edge = np.array([label.startswith(EDGE_PREFIX) for label in labels])
85
+
86
+ self.candidate_keys_ = keys
87
+ self.candidate_labels_ = labels
88
+ self.candidate_counts_ = counts
89
+ self.candidates_ = [structure_graph(key, self.node_label) for key in keys]
90
+ self.selected_ = np.flatnonzero(is_edge | (counts >= cut))
91
+ self.primitive_labels_ = [labels[i] for i in self.selected_]
92
+ return self
93
+
94
+ @property
95
+ def primitives_(self):
96
+ """The selected primitives as labelled graphs."""
97
+ return [self.candidates_[i] for i in self.selected_]
98
+
99
+ def transform(self, X):
100
+ check_is_fitted(self)
101
+ graphs = check_graphs(X, self.node_label)
102
+ keys = [self.candidate_keys_[i] for i in self.selected_]
103
+ occurrences = match_graphs(graphs, keys, self.node_label, self.n_jobs)
104
+ return [assemble_primitive_graph(occ, self.primitive_labels_, self.node_label) for occ in occurrences]
molprim/features.py ADDED
@@ -0,0 +1,52 @@
1
+ from itertools import combinations_with_replacement
2
+
3
+ import numpy as np
4
+ from sklearn.base import BaseEstimator, TransformerMixin
5
+ from sklearn.utils.validation import check_is_fitted
6
+
7
+ from ._utils import check_graphs
8
+
9
+
10
+ class LabelPairEdgeCounter(TransformerMixin, BaseEstimator):
11
+ """Count each graph's edges by the unordered pair of labels at their ends.
12
+
13
+ This is the "Edges" similarity measure of the paper: one feature per pair
14
+ of labels seen in ``fit`` (``a|b`` with ``a <= b``). Edges touching a label
15
+ unseen in ``fit`` are ignored. The paper follows it with
16
+ ``VarianceThreshold()`` and ``SVC(kernel="rbf", gamma="auto")``.
17
+
18
+ Parameters
19
+ ----------
20
+ node_label : str, default="label"
21
+ Node attribute holding the label.
22
+
23
+ Attributes
24
+ ----------
25
+ labels_ : list
26
+ Labels seen in ``fit``, sorted by their string form.
27
+ """
28
+
29
+ def __init__(self, node_label="label"):
30
+ self.node_label = node_label
31
+
32
+ def fit(self, X, y=None):
33
+ graphs = check_graphs(X, self.node_label)
34
+ self.labels_ = sorted({label for g in graphs for _, label in g.nodes(data=self.node_label)}, key=str)
35
+ self._pairs = {pair: k for k, pair in enumerate(combinations_with_replacement(range(len(self.labels_)), 2))}
36
+ return self
37
+
38
+ def transform(self, X):
39
+ check_is_fitted(self, "labels_")
40
+ graphs = check_graphs(X, self.node_label)
41
+ index = {label: i for i, label in enumerate(self.labels_)}
42
+ out = np.zeros((len(graphs), len(self._pairs)))
43
+ for row, g in enumerate(graphs):
44
+ for u, v in g.edges:
45
+ a, b = index.get(g.nodes[u][self.node_label]), index.get(g.nodes[v][self.node_label])
46
+ if a is not None and b is not None:
47
+ out[row, self._pairs[min(a, b), max(a, b)]] += 1
48
+ return out
49
+
50
+ def get_feature_names_out(self, input_features=None):
51
+ check_is_fitted(self, "labels_")
52
+ return np.array([f"{self.labels_[a]}|{self.labels_[b]}" for a, b in self._pairs], dtype=object)
molprim/kernels.py ADDED
@@ -0,0 +1,135 @@
1
+ from collections import Counter
2
+
3
+ import networkx as nx
4
+ import numpy as np
5
+ from scipy import sparse
6
+ from sklearn.base import BaseEstimator, TransformerMixin
7
+ from sklearn.utils.validation import check_is_fitted
8
+
9
+ from ._utils import check_graphs
10
+
11
+ KERNELS = ("wl_subtree", "wl_sp", "sp")
12
+
13
+
14
+ class GraphKernelTransformer(TransformerMixin, BaseEstimator):
15
+ """Graph kernel as a transformer, for ``SVC(kernel="precomputed")``.
16
+
17
+ ``fit_transform`` returns the kernel matrix of the training graphs and
18
+ ``transform`` returns the kernel between new graphs and the training
19
+ graphs, so ``make_pipeline(GraphKernelTransformer(), SVC(kernel="precomputed"))``
20
+ trains and predicts directly on lists of graphs.
21
+
22
+ Kernel values match grakel's ``GraphKernel`` with the same settings:
23
+
24
+ - ``"wl_subtree"``: Weisfeiler-Lehman subtree kernel,
25
+ ``[{"name": "weisfeiler_lehman", "n_iter": n_iter}, {"name": "subtree_wl"}]``
26
+ - ``"wl_sp"``: Weisfeiler-Lehman shortest-path kernel,
27
+ ``[{"name": "weisfeiler_lehman", "n_iter": n_iter}, {"name": "SP"}]``
28
+ - ``"sp"``: shortest-path kernel, ``[{"name": "SP"}]`` (ignores ``n_iter``)
29
+
30
+ Parameters
31
+ ----------
32
+ kernel : {"wl_subtree", "wl_sp", "sp"}, default="wl_subtree"
33
+ n_iter : int, default=3
34
+ Number of WL relabelling iterations h; base kernels are summed over
35
+ iterations 0..h.
36
+ normalize : bool, default=True
37
+ Cosine-normalise the kernel: K(x, y) / sqrt(K(x, x) K(y, y)).
38
+ node_label : str, default="label"
39
+ Node attribute holding the label.
40
+ """
41
+
42
+ def __init__(self, kernel="wl_subtree", n_iter=3, normalize=True, node_label="label"):
43
+ self.kernel = kernel
44
+ self.n_iter = n_iter
45
+ self.normalize = normalize
46
+ self.node_label = node_label
47
+
48
+ def fit(self, X, y=None):
49
+ self._fit(X)
50
+ return self
51
+
52
+ def fit_transform(self, X, y=None):
53
+ features = self._fit(X)
54
+ return self._kernel(features, self._fit_diag)
55
+
56
+ def transform(self, X):
57
+ check_is_fitted(self, "features_")
58
+ graphs = check_graphs(X, self.node_label)
59
+ counts = self._count_features(graphs, self._relabel_tables, learn=False)
60
+ diag = np.array([sum(c * c for c in counter.values()) for counter in counts], dtype=float)
61
+ return self._kernel(self._to_matrix(counts), diag)
62
+
63
+ def _fit(self, X):
64
+ if self.kernel not in KERNELS:
65
+ raise ValueError(f"kernel must be one of {KERNELS}, got {self.kernel!r}")
66
+ if self.n_iter < 0:
67
+ raise ValueError(f"n_iter must be >= 0, got {self.n_iter}")
68
+ graphs = check_graphs(X, self.node_label)
69
+ self._relabel_tables = [{} for _ in range(self._n_levels())]
70
+ counts = self._count_features(graphs, self._relabel_tables, learn=True)
71
+ self.vocabulary_ = {}
72
+ for counter in counts:
73
+ for key in counter:
74
+ self.vocabulary_.setdefault(key, len(self.vocabulary_))
75
+ self.features_ = self._to_matrix(counts)
76
+ self._fit_diag = np.asarray(self.features_.multiply(self.features_).sum(axis=1)).ravel()
77
+ return self.features_
78
+
79
+ def _n_levels(self):
80
+ return 1 if self.kernel == "sp" else self.n_iter + 1
81
+
82
+ def _kernel(self, features, diag):
83
+ K = (features @ self.features_.T).toarray()
84
+ if self.normalize:
85
+ scale = np.sqrt(np.outer(diag, self._fit_diag))
86
+ K = np.divide(K, scale, out=np.zeros_like(K), where=scale > 0)
87
+ return K
88
+
89
+ def _to_matrix(self, counts):
90
+ rows, cols, values = [], [], []
91
+ for i, counter in enumerate(counts):
92
+ for key, value in counter.items():
93
+ j = self.vocabulary_.get(key)
94
+ if j is not None: # features unseen in fit only affect the normalisation
95
+ rows.append(i)
96
+ cols.append(j)
97
+ values.append(value)
98
+ return sparse.csr_matrix((values, (rows, cols)), shape=(len(counts), len(self.vocabulary_)), dtype=float)
99
+
100
+ def _count_features(self, graphs, tables, learn):
101
+ if not learn:
102
+ # Unseen WL signatures get fresh ids in a scratch copy so the fitted tables stay unchanged.
103
+ tables = [dict(t) for t in tables]
104
+ return [self._graph_features(g, tables) for g in graphs]
105
+
106
+ def _graph_features(self, graph, tables):
107
+ position = {n: i for i, n in enumerate(graph)}
108
+ labels = [_compress(tables[0], data[self.node_label]) for _, data in graph.nodes(data=True)]
109
+ neighbours = [[position[m] for m in graph[n]] for n in graph]
110
+ if self.kernel != "wl_subtree":
111
+ # Ordered pairs of distinct, connected vertices, as grakel's shortest-path kernel counts them.
112
+ distances = [
113
+ (position[u], position[v], d)
114
+ for u, lengths in nx.all_pairs_shortest_path_length(graph)
115
+ for v, d in lengths.items()
116
+ if u != v
117
+ ]
118
+
119
+ counter = Counter()
120
+ for level in range(len(tables)):
121
+ if level > 0:
122
+ labels = [
123
+ _compress(tables[level], (label, tuple(sorted(labels[j] for j in neighbours[i]))))
124
+ for i, label in enumerate(labels)
125
+ ]
126
+ if self.kernel == "wl_subtree":
127
+ counter.update((level, label) for label in labels)
128
+ else:
129
+ counter.update((level, labels[i], labels[j], d) for i, j, d in distances)
130
+ return counter
131
+
132
+
133
+ def _compress(table, signature):
134
+ """Integer id of ``signature`` in ``table``, adding it when new."""
135
+ return table.setdefault(signature, len(table))
@@ -0,0 +1,133 @@
1
+ Metadata-Version: 2.4
2
+ Name: molprim
3
+ Version: 0.1.0
4
+ Summary: Primitive structure extraction and graph kernels for molecular graph classification, with a scikit-learn API
5
+ Author: Peemapat Wongsriphisant
6
+ License-Expression: BSD-3-Clause
7
+ Project-URL: Homepage, https://github.com/PeemapatW/primitive-structure-extraction
8
+ Project-URL: Repository, https://github.com/PeemapatW/primitive-structure-extraction
9
+ Project-URL: Issues, https://github.com/PeemapatW/primitive-structure-extraction/issues
10
+ Keywords: graph classification,graph kernels,cheminformatics,molecular graphs,scikit-learn
11
+ Classifier: Development Status :: 3 - Alpha
12
+ Classifier: Intended Audience :: Science/Research
13
+ Classifier: Operating System :: OS Independent
14
+ Classifier: Programming Language :: Python :: 3
15
+ Classifier: Topic :: Scientific/Engineering :: Bio-Informatics
16
+ Classifier: Topic :: Scientific/Engineering :: Chemistry
17
+ Requires-Python: >=3.9
18
+ Description-Content-Type: text/markdown
19
+ License-File: LICENSE
20
+ Requires-Dist: joblib
21
+ Requires-Dist: networkx>=3.1
22
+ Requires-Dist: numpy
23
+ Requires-Dist: scikit-learn>=1.0
24
+ Requires-Dist: scipy
25
+ Provides-Extra: test
26
+ Requires-Dist: pytest; extra == "test"
27
+ Dynamic: license-file
28
+
29
+ # molprim
30
+
31
+ **Mol**ecular **prim**itives: classify molecular graphs by their **primitive
32
+ structures**, that is, rings, branching points and bonds. The package rewrites
33
+ each molecule as a small graph whose vertices are these structures, then
34
+ compares molecules with graph kernels.
35
+ Everything is a scikit-learn transformer, so it plugs into `Pipeline`,
36
+ `GridSearchCV` and `cross_val_score`.
37
+
38
+ This is the implementation of:
39
+
40
+ > P. Wongsriphisant, C. Lursinsap, A. Suratanee and K. Plaimas,
41
+ > "A classification of biochemical compounds based on their primitive structures and graph kernels", IEEE, 2020, pp. 104–109.
42
+
43
+ ## Install
44
+
45
+ ```bash
46
+ pip install molprim
47
+ ```
48
+
49
+ Requires Python ≥ 3.9 with networkx, numpy, scipy and scikit-learn. It needs neither graph-tool nor grakel.
50
+
51
+ ## Quickstart
52
+
53
+ ```python
54
+ from sklearn.model_selection import cross_val_score
55
+ from sklearn.pipeline import make_pipeline
56
+ from sklearn.svm import SVC
57
+
58
+ from molprim import GraphKernelTransformer, PrimitiveStructureExtractor, fetch_tudataset
59
+
60
+ graphs, y = fetch_tudataset("MUTAG") # list of networkx.Graph with a "label" node attribute
61
+
62
+ model = make_pipeline(
63
+ PrimitiveStructureExtractor(),
64
+ GraphKernelTransformer(kernel="wl_sp", n_iter=2),
65
+ SVC(kernel="precomputed"),
66
+ )
67
+ print(cross_val_score(model, graphs, y, cv=5).mean())
68
+ ```
69
+
70
+ Inputs are lists of undirected `networkx.Graph` whose nodes carry the atom type
71
+ in a node attribute (`"label"` by default; change it with `node_label=`).
72
+ For a fuller example, see [`examples/quickstart.py`](examples/quickstart.py).
73
+
74
+ ## How it works
75
+
76
+ 1. **Selection** (`PrimitiveStructureExtractor.fit`): collects every labelled
77
+ cycle (3–10 atoms), star (3–7 atoms) and bond that occurs in the training
78
+ molecules, and keeps those whose frequency is at least the `threshold`
79
+ percentile. Bonds are always kept.
80
+ 2. **Extraction** (`PrimitiveStructureExtractor.transform`): finds the
81
+ primitives in each molecule, largest cycles first, then stars, then bonds.
82
+ An occurrence that shares more than half of its atoms with one already taken
83
+ is skipped. Each kept occurrence becomes a vertex labelled by its primitive
84
+ (e.g. `"C6-1"`, the most frequent 6-ring). Two vertices are joined when
85
+ their occurrences share an atom.
86
+ 3. **Similarity**: either a graph kernel on the primitive graphs, or a count of
87
+ adjacent primitive pairs fed to an RBF SVM.
88
+
89
+ | Component | Purpose |
90
+ |---|---|
91
+ | `PrimitiveStructureExtractor` | molecules → primitive graphs |
92
+ | `GraphKernelTransformer` | Weisfeiler-Lehman subtree (`"wl_subtree"`), WL shortest-path (`"wl_sp"`) or shortest-path (`"sp"`) kernel. `fit_transform` returns the training kernel matrix and `transform` the kernel against the training graphs, for `SVC(kernel="precomputed")`. Values match grakel's `GraphKernel`. |
93
+ | `LabelPairEdgeCounter` | counts edges per pair of end labels (the paper's "Edges" measure) |
94
+ | `fetch_tudataset`, `read_tudataset` | load TU Dortmund benchmark datasets (MUTAG, NCI1, BZR, COX2, …) as networkx graphs |
95
+
96
+ Tune the extraction like any other hyper-parameter, for example
97
+ `GridSearchCV(model, {"primitivestructureextractor__threshold": [0, 50], "graphkerneltransformer__n_iter": [1, 2, 3]})`.
98
+
99
+ ## Implementation notes
100
+
101
+ - Primitives are learned in `fit`, from the training graphs only.
102
+ - Candidate cycles, stars and bonds are enumerated directly and identified by
103
+ a canonical form of their labels, rather than with general subgraph
104
+ isomorphism. Extracting the primitive graphs of all 3,865 NCI1 molecules
105
+ takes a few seconds.
106
+ - When occurrences overlap, the one kept is chosen deterministically: larger
107
+ shapes first, then more frequent primitives, then vertex order.
108
+
109
+ ## Benchmarks
110
+
111
+ Scripts that evaluate the package on TU datasets (MUTAG, BZR, COX2, NCI1),
112
+ with their per-split results and runtimes, are in [`benchmarks/`](benchmarks/).
113
+
114
+ ## Repository layout
115
+
116
+ ```
117
+ src/molprim/ the package
118
+ tests/ pytest suite
119
+ examples/ usage examples
120
+ benchmarks/ benchmark scripts and results of this implementation
121
+ legacy/ original notebooks and results from the 2020 paper (see legacy/README.md)
122
+ ```
123
+
124
+ ## Development
125
+
126
+ ```bash
127
+ pip install -e ".[test]"
128
+ pytest
129
+ ```
130
+
131
+ ## License
132
+
133
+ BSD-3-Clause, the same as scikit-learn. See [LICENSE](LICENSE).
@@ -0,0 +1,12 @@
1
+ molprim/__init__.py,sha256=Xty7eqhnDTOpr0uL-YtQKB0Jr4Suqk5wSIejeQ30Tzk,773
2
+ molprim/_algorithms.py,sha256=Buc-R8y4nvkCLnnqS2zgrYNVMlJjUjGxkWyebnPfLn4,7909
3
+ molprim/_utils.py,sha256=O_jDlaMWY6nDsXHDcxEhDrbwp7JDS9XZJ0QNUjq4Xlo,652
4
+ molprim/datasets.py,sha256=77FCxDniiLDP0ErAc2kwrunJwYSImdwn_1WGwr8R2qY,2165
5
+ molprim/extractor.py,sha256=QMyDl-CW14uHqbAzZgudGXaBkWAFaPiBsJsszRxCSZc,4504
6
+ molprim/features.py,sha256=K3iRY0rNEynf2qLluUSzFYPRl1MaplY9EmzGIm0Jjcg,2050
7
+ molprim/kernels.py,sha256=4fo64sJVspy3fjJ0zScY1edPWZErN7klC52by4Ovj-I,5659
8
+ molprim-0.1.0.dist-info/licenses/LICENSE,sha256=Bpqz6JvJdb2Rx0IPH4XolloXbVZkOwkRzfgdncF5b5Y,1532
9
+ molprim-0.1.0.dist-info/METADATA,sha256=PUgyKftakyXbTgCP7kMKSAHT6H74GgU9f-QyMDsE_Co,5506
10
+ molprim-0.1.0.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
11
+ molprim-0.1.0.dist-info/top_level.txt,sha256=xnod4x6Ekh-o0oXCTpphW27Dfd-f05BXvsxXqHpqE6A,8
12
+ molprim-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (84.0.0)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -0,0 +1,28 @@
1
+ BSD 3-Clause License
2
+
3
+ Copyright (c) 2020-2026, Peemapat Wongsriphisant and contributors
4
+
5
+ Redistribution and use in source and binary forms, with or without
6
+ modification, are permitted provided that the following conditions are met:
7
+
8
+ 1. Redistributions of source code must retain the above copyright notice, this
9
+ list of conditions and the following disclaimer.
10
+
11
+ 2. Redistributions in binary form must reproduce the above copyright notice,
12
+ this list of conditions and the following disclaimer in the documentation
13
+ and/or other materials provided with the distribution.
14
+
15
+ 3. Neither the name of the copyright holder nor the names of its
16
+ contributors may be used to endorse or promote products derived from
17
+ this software without specific prior written permission.
18
+
19
+ THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
20
+ AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
21
+ IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
22
+ DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
23
+ FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
24
+ DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
25
+ SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
26
+ CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
27
+ OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
28
+ OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
@@ -0,0 +1 @@
1
+ molprim