From ee5be73a3f984519884ae373a0d7b3d18cf51fce Mon Sep 17 00:00:00 2001 From: Lars Malmqvist <12750146+lmlearning@users.noreply.github.com> Date: Fri, 25 Sep 2026 13:40:29 +0200 Subject: [PATCH] Preserve extension labels and handle singleton training graphs --- README.md | 12 +++++++++++ graph_io.py | 22 ++++++++++++++++++++ tests/test_dataset_runtime.py | 38 +++++++++++++++++++++++++++++++++++ tests/test_extensions.py | 19 ++++++++++++++++++ train_consolidated.py | 6 +++--- train_refined_afgcn.py | 17 +++++++++------- 6 files changed, 104 insertions(+), 10 deletions(-) create mode 100644 tests/test_dataset_runtime.py create mode 100644 tests/test_extensions.py diff --git a/README.md b/README.md index 8e96ee4..0441c92 100644 --- a/README.md +++ b/README.md @@ -47,6 +47,18 @@ python -m pytest tests These component tests need no graph-learning libraries or checkpoints. They validate input handling; they do not reproduce training or evaluate model accuracy. +## Training-label format + +The dataset loader takes the union of identifiers in an ICCMA extension list such +as `[[a,b],[c]]` as its **credulous** acceptance target. A flat single extension +`[a,b]`, empty extension `[[]]` and empty extension list `[]` are also accepted. +Undeclared solution identifiers are rejected. Singleton and edgeless frameworks +retain a one-dimensional node output through both refined model implementations. + +Earlier parsing could omit the first argument of nested extension lists. Retrain +and evaluate checkpoints before attributing results to the corrected loader; +saved historical checkpoints and metrics have not been regenerated by this fix. + ## License See [LICENSE](LICENSE). diff --git a/graph_io.py b/graph_io.py index b3a6921..df8de99 100644 --- a/graph_io.py +++ b/graph_io.py @@ -1,5 +1,27 @@ """Dependency-free input parsing shared by the graph-training entry points.""" +import re + + +def parse_extension_union(text): + """Read accepted identifiers from an ICCMA extension list. + + Both ``[[a,b],[c]]`` (enumerated extensions) and ``[a,b]`` (one extension) + are accepted. The union supplies the existing credulous training target. + """ + text = re.sub(r"\s+", "", text) + if not re.fullmatch(r"\[(?:\[[^\[\]]*\](?:,\[[^\[\]]*\])*|[^\[\]]*)\]", text): + raise ValueError("expected an extension list such as [[a,b],[c]]") + accepted = set() + for extension in re.findall(r"\[([^\[\]]*)\]", text): + if not extension: + continue + identifiers = extension.split(",") + if any(not identifier for identifier in identifiers): + raise ValueError("empty identifier in extension list") + accepted.update(identifiers) + return accepted + def parse_tgf(path): """Read the unlabelled TGF subset used by the argumentation datasets. diff --git a/tests/test_dataset_runtime.py b/tests/test_dataset_runtime.py new file mode 100644 index 0000000..e0e8835 --- /dev/null +++ b/tests/test_dataset_runtime.py @@ -0,0 +1,38 @@ +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("torch_geometric") +pytest.importorskip("torch_scatter") + +from train_refined_afgcn import AFGraphDataset, RefinedAFGCN as OriginalModel +from train_consolidated import RefinedAFGCN + + +@pytest.mark.parametrize("graph,solution,expected", [ + ("a\n#\n", "[[a]]", [1]), + ("a\nb\n#\n", "[[a,b]]", [1, 1]), + ("a\nb\n#\na b\nb a\n", "[[a],[b]]", [1, 1]), + ("a\n#\na a\n", "[]", [0]), +]) +def test_dataset_labels_edges_and_real_model_backward(tmp_path, graph, solution, expected): + torch.set_num_threads(1) + (tmp_path / "tiny.tgf").write_text(graph) + (tmp_path / "tiny.EE-PR").write_text(solution) + data = AFGraphDataset(str(tmp_path))[0] + assert data.y.tolist() == expected + assert data.edge_index.shape[0] == 2 + for model in (OriginalModel(hidden=8), RefinedAFGCN(hidden=8)): + output = model(data) + logits = output[0] if isinstance(output, tuple) else output + assert logits.shape == (len(expected),) + loss = torch.nn.functional.binary_cross_entropy_with_logits(logits, data.y.float()) + loss.backward() + assert torch.isfinite(loss) + assert all(torch.isfinite(p.grad).all() for p in model.parameters() if p.grad is not None) + + +def test_undeclared_solution_argument_is_rejected(tmp_path): + (tmp_path / "tiny.tgf").write_text("a\n#\n") + (tmp_path / "tiny.EE-PR").write_text("[[missing]]") + with pytest.raises(ValueError, match="undeclared solution arguments"): + AFGraphDataset(str(tmp_path))[0] diff --git a/tests/test_extensions.py b/tests/test_extensions.py new file mode 100644 index 0000000..ac2edae --- /dev/null +++ b/tests/test_extensions.py @@ -0,0 +1,19 @@ +import pytest + +from graph_io import parse_extension_union + + +@pytest.mark.parametrize("text,expected", [ + ("[[a,b],[c]]", {"a", "b", "c"}), + (" [ [ 1, 2 ], [ 2, 3 ] ] \n", {"1", "2", "3"}), + ("[a,b]", {"a", "b"}), + ("[[]]", set()), ("[]", set()), ("[[first]]", {"first"}), +]) +def test_extension_union_preserves_first_argument(text, expected): + assert parse_extension_union(text) == expected + + +@pytest.mark.parametrize("text", ["", "a,b", "[[a]", "[[a],garbage]", "[[a,,b]]"]) +def test_malformed_solution_is_rejected(text): + with pytest.raises(ValueError): + parse_extension_union(text) diff --git a/train_consolidated.py b/train_consolidated.py index 121e6b4..413a82e 100644 --- a/train_consolidated.py +++ b/train_consolidated.py @@ -136,8 +136,8 @@ def forward(self, data): h = torch.cat([fixed_s, h], 1) h = l(edge_index, h) - class_logits = self.classification_head(h).squeeze() - rank_scores = self.ranking_head(h).squeeze() + class_logits = self.classification_head(h).squeeze(-1) + rank_scores = self.ranking_head(h).squeeze(-1) return class_logits, rank_scores # ╰───────────────────────────────────────────────────────────────────────────╯ @@ -309,4 +309,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/train_refined_afgcn.py b/train_refined_afgcn.py index 8bcf3a0..9f09919 100644 --- a/train_refined_afgcn.py +++ b/train_refined_afgcn.py @@ -16,7 +16,7 @@ """ import os, random, argparse, pickle, math, time -from graph_io import parse_tgf +from graph_io import parse_tgf, parse_extension_union import numpy as np, networkx as nx from tqdm import tqdm from sklearn.preprocessing import StandardScaler @@ -136,6 +136,8 @@ def __getitem__(self, idx): G = nx.DiGraph() G.add_nodes_from(args_list) G.add_edges_from(atts) + if not G.number_of_nodes(): + raise ValueError(f"{path}: training graphs must contain at least one argument") G, mapping = nx.convert_node_labels_to_integers(G, label_attribute="orig"), {n:i for i,n in enumerate(G.nodes())} # grounded flags @@ -164,10 +166,11 @@ def __getitem__(self, idx): # labels lbl_p = os.path.join(self.root_dir, os.path.splitext(gfile)[0] + self.label_ext) - lab_nodes = [] - with open(lbl_p) as f: - line = f.readline().strip()[1:-1].replace("]]", "") - lab_nodes = [z.strip("] ") for sub in line.split("],") for z in sub.split(',')] + with open(lbl_p, encoding="utf-8") as f: + lab_nodes = parse_extension_union(f.read()) + unknown = lab_nodes.difference(mapping) + if unknown: + raise ValueError(f"{lbl_p}: undeclared solution arguments: {sorted(unknown)}") y = np.zeros(G.number_of_nodes(), np.int64) for n in lab_nodes: if n in mapping: @@ -177,7 +180,7 @@ def __getitem__(self, idx): rank = categoriser_ranking(adj) # edge index with self-loops - edge_index = torch.tensor(list(G.edges()), dtype=torch.long).t().contiguous() + edge_index = torch.tensor(list(G.edges()), dtype=torch.long).reshape(-1, 2).t().contiguous() edge_index, _ = add_self_loops(edge_index, num_nodes=G.number_of_nodes()) data = Data( @@ -271,7 +274,7 @@ def forward(self, data): for l in self.layers[1:]: h = torch.cat([s_fixed, h], 1) h = l(edge_index, h) - return self.out(h).squeeze() + return self.out(h).squeeze(-1) # ─────────────────────────────────────────