Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).
22 changes: 22 additions & 0 deletions graph_io.py
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
38 changes: 38 additions & 0 deletions tests/test_dataset_runtime.py
Original file line number Diff line number Diff line change
@@ -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]
19 changes: 19 additions & 0 deletions tests/test_extensions.py
Original file line number Diff line number Diff line change
@@ -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)
6 changes: 3 additions & 3 deletions train_consolidated.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
# ╰───────────────────────────────────────────────────────────────────────────╯

Expand Down Expand Up @@ -309,4 +309,4 @@ def main():


if __name__ == "__main__":
main()
main()
17 changes: 10 additions & 7 deletions train_refined_afgcn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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(
Expand Down Expand Up @@ -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)


# ─────────────────────────────────────────
Expand Down
Loading